1. 标准 Attention 的定义
在正式分析 FlashAttention 之前,要先彻底地了解标准 Attention 的做法和含义。
给定单个 Query 向量 q∈RN×d,以及 Key 矩阵 K∈RNkv×d 和 Value 矩阵 V∈RNkv×d,其中 Nkv 是 KV Cache 的序列长度(可能很长),d 是注意力头维度。标准的缩放点积注意力(Scaled Dot-Product Attention)定义为:
Attention(q,K,V)=softmax(dqK⊤)V
2. 以简化的矩阵为例进行推导
以 n=2(两个词)、d=3(向量维度)为例,Q,K,V∈R2×3,完整推导 Attention 计算过程。
Q=[q1q2]=[q11q21q12q22q13q23]
K=[k1k2]=[k11k21k12k22k13k23]
V=[v1v2]=[v11v21v12v22v13v23]
其中 qi,ki,vi∈R1×3 为第 i 个词对应的 query、key、value 行向量。
2. S=QK⊤:注意力得分矩阵
S=QK⊤=[q1k1⊤q2k1⊤q1k2⊤q2k2⊤]=[∑l=13q1lk1l∑l=13q2lk1l∑l=13q1lk2l∑l=13q2lk2l]∈R2×2
含义:
- 元素 Sij=qikj⊤=∑l=13qilkjl:第 i 个词的 query 向量 qi 与第 j 个词的 key 向量 kj 的内积。
- 行向量 Si=[Si1,Si2]:第 i 个词的 query 与序列中所有词的 key 计算得到的注意力得分向量。
3. Softmax:注意力权重矩阵
对 S 逐行做 Softmax。先减去行最大值 mi=maxjSij 保证数值稳定,令 S~ij=Sij−mi,再归一化:
Pij=∑j′=12eS~ij′eS~ij
P=[P11P21P12P22]∈R2×2,满足 P11+P12=1,P21+P22=1
含义:
- 元素 Pij∈[0,1]:第 i 个词的 query 对第 j 个词的 key 的注意力权重,即第 i 个词关注第 j 个词的概率。
- 行向量 Pi=[Pi1,Pi2]:第 i 个词的 query 与所有 key 经 Softmax 后的注意力分布(概率向量,和为 1)。
4. O=PV:输出矩阵
O=PV=[P11P21P12P22][v1v2]=[P11v1+P12v2P21v1+P22v2]=[O1O2]∈R2×3
逐行展开:
O1=P11[v11,v12,v13]+P12[v21,v22,v23]=[O11,O12,O13]
O2=P21[v11,v12,v13]+P22[v21,v22,v23]=[O21,O22,O23]
含义:
- O1:第 1 个词的 query 向量 q1 与所有 key 向量(k1,k2)计算注意力得分并归一化为注意力分布 P1=[P11,P12] 后,用该分布作为权重对所有词的 value 向量(v1,v2)加权求和的结果。即 O1=P11v1+P12v2。
- O2:同理,第 2 个词的 query 向量 q2 与所有 key 计算得分并归一化为 P2=[P21,P22] 后,对所有词的 value 向量加权求和的结果。即 O2=P21v1+P22v2。
其中分量形式为:
Oij=Pi1v1j+Pi2v2j=l=1∑2Pilvlj
元素 Oij 的详细解释:
Oij 是输出矩阵 O 的第 i 行第 j 列元素。具体计算过程为:
- 第 i 个词的 query 向量 qi 与所有词的 key 向量 k1,k2 计算注意力得分,经 Softmax 归一化得到注意力分布 Pi=[Pi1,Pi2];
- 用该分布对所有词的 value 向量 v1,v2 做加权求和,得到第 i 个词的输出向量 Oi=Pi1v1+Pi2v2;
- Oij 就是这个输出向量 Oi 的第 j 个分量(即嵌入维度的第 j 维)。
从矩阵乘法展开看,Oij 等于所有词的 value 向量第 j 维分量按第 i 个词对各词的关注概率加权求和:
Oij=Pi1⋅v1j+Pi2⋅v2j
即:第 1 个词 value 向量的第 j 维分量 v1j 乘以第 i 个词对第 1 个词的关注概率 Pi1,加上第 2 个词 value 向量的第 j 维分量 v2j 乘以第 i 个词对第 2 个词的关注概率 Pi2。
- 行向量 Oi=[Oi1,Oi2,Oi3]:第 i 个词的 query 经注意力机制后,按注意力分布 Pi 对所有词的 value 向量加权求和得到的输出向量。
5. 加入 Causal Mask(因果掩码)
Decoder 中需防止第 i 个词看到第 j>i 个词。对 n=2:
M=[00−∞0]
Masked 分数
Smask=S+M=[S11S21−∞S22]
Masked Softmax
第 1 行中 e−∞=0,故:
P11mask=eS~11+0eS~11=1,P12mask=0
第 2 行正常计算:
P21mask=eS~21+eS~22eS~21,P22mask=eS~21+eS~22eS~22
Masked 输出
Omask=PmaskV=[1⋅v1+0⋅v2P21maskv1+P22maskv2]=[v1P21maskv1+P22maskv2]
结果:第 1 个词只能看到自己(O1=v1),第 2 个词可以看到自己和第 1 个词。
评论