解析 FlashAttention(1):从标准 Attention 讲起

1. 标准 Attention 的定义

在正式分析 FlashAttention 之前,要先彻底地了解标准 Attention 的做法和含义。

给定单个 Query 向量 qRN×d\mathbf{q} \in \mathbb{R}^{N \times d},以及 Key 矩阵 KRNkv×d\mathbf{K} \in \mathbb{R}^{N_{kv} \times d} 和 Value 矩阵 VRNkv×d\mathbf{V} \in \mathbb{R}^{N_{kv} \times d},其中 NkvN_{kv} 是 KV Cache 的序列长度(可能很长),dd 是注意力头维度。标准的缩放点积注意力(Scaled Dot-Product Attention)定义为:

Attention(q,K,V)=softmax(qKd)V\mathrm{Attention}(\mathbf{q}, \mathbf{K}, \mathbf{V}) = \mathrm{softmax}\left(\frac{\mathbf{q}\mathbf{K}^{\top}}{\sqrt{d}}\right) \mathbf{V}


2. 以简化的矩阵为例进行推导

n=2n=2(两个词)、d=3d=3(向量维度)为例,Q,K,VR2×3Q, K, V \in \mathbb{R}^{2 \times 3},完整推导 Attention 计算过程。

Q=[q1q2]=[q11q12q13q21q22q23]Q = \begin{bmatrix} q_1 \\ q_2 \end{bmatrix} = \begin{bmatrix} q_{11} & q_{12} & q_{13} \\ q_{21} & q_{22} & q_{23} \end{bmatrix}

K=[k1k2]=[k11k12k13k21k22k23]\quad K = \begin{bmatrix} k_1 \\ k_2 \end{bmatrix} = \begin{bmatrix} k_{11} & k_{12} & k_{13} \\ k_{21} & k_{22} & k_{23} \end{bmatrix}

V=[v1v2]=[v11v12v13v21v22v23]\quad V = \begin{bmatrix} v_1 \\ v_2 \end{bmatrix} = \begin{bmatrix} v_{11} & v_{12} & v_{13} \\ v_{21} & v_{22} & v_{23} \end{bmatrix}

其中 qi,ki,viR1×3q_i, k_i, v_i \in \mathbb{R}^{1 \times 3} 为第 ii 个词对应的 query、key、value 行向量。


2. S=QKS = QK^\top:注意力得分矩阵

S=QK=[q1k1q1k2q2k1q2k2]=[l=13q1lk1ll=13q1lk2ll=13q2lk1ll=13q2lk2l]R2×2S = QK^\top = \begin{bmatrix} q_1 k_1^\top & q_1 k_2^\top \\ q_2 k_1^\top & q_2 k_2^\top \end{bmatrix} = \begin{bmatrix} \sum_{l=1}^{3} q_{1l}k_{1l} & \sum_{l=1}^{3} q_{1l}k_{2l} \\[6pt] \sum_{l=1}^{3} q_{2l}k_{1l} & \sum_{l=1}^{3} q_{2l}k_{2l} \end{bmatrix} \in \mathbb{R}^{2 \times 2}

含义:

  • 元素 Sij=qikj=l=13qilkjlS_{ij} = q_i k_j^\top = \sum_{l=1}^{3} q_{il} k_{jl}:第 ii 个词的 query 向量 qiq_i 与第 jj 个词的 key 向量 kjk_j内积
  • 行向量 Si=[Si1,  Si2]S_i = [S_{i1},\; S_{i2}]:第 ii 个词的 query 与序列中所有词的 key 计算得到的注意力得分向量。

3. Softmax:注意力权重矩阵

SS 逐行做 Softmax。先减去行最大值 mi=maxjSijm_i = \max_j S_{ij} 保证数值稳定,令 S~ij=Sijmi\tilde{S}_{ij} = S_{ij} - m_i,再归一化:

Pij=eS~ijj=12eS~ijP_{ij} = \frac{e^{\tilde{S}_{ij}}}{\sum_{j'=1}^{2} e^{\tilde{S}_{ij'}}}

P=[P11P12P21P22]R2×2,满足 P11+P12=1,  P21+P22=1P = \begin{bmatrix} P_{11} & P_{12} \\ P_{21} & P_{22} \end{bmatrix} \in \mathbb{R}^{2 \times 2}, \quad \text{满足 } P_{11}+P_{12}=1,\; P_{21}+P_{22}=1

含义:

  • 元素 Pij[0,1]P_{ij} \in [0,1]:第 ii 个词的 query 对第 jj 个词的 key 的注意力权重,即第 ii 个词关注第 jj 个词的概率
  • 行向量 Pi=[Pi1,  Pi2]P_i = [P_{i1},\; P_{i2}]:第 ii 个词的 query 与所有 key 经 Softmax 后的注意力分布(概率向量,和为 1)。

4. O=PVO = PV:输出矩阵

O=PV=[P11P12P21P22][v1v2]=[P11v1+P12v2P21v1+P22v2]=[O1O2]R2×3O = PV = \begin{bmatrix} P_{11} & P_{12} \\ P_{21} & P_{22} \end{bmatrix} \begin{bmatrix} v_1 \\ v_2 \end{bmatrix} = \begin{bmatrix} P_{11}v_1 + P_{12}v_2 \\ P_{21}v_1 + P_{22}v_2 \end{bmatrix} = \begin{bmatrix} O_1 \\ O_2 \end{bmatrix} \in \mathbb{R}^{2 \times 3}

逐行展开:

O1=P11[v11,  v12,  v13]+P12[v21,  v22,  v23]=[O11,  O12,  O13]O_1 = P_{11}[v_{11},\; v_{12},\; v_{13}] + P_{12}[v_{21},\; v_{22},\; v_{23}] = [O_{11},\; O_{12},\; O_{13}]

O2=P21[v11,  v12,  v13]+P22[v21,  v22,  v23]=[O21,  O22,  O23]O_2 = P_{21}[v_{11},\; v_{12},\; v_{13}] + P_{22}[v_{21},\; v_{22},\; v_{23}] = [O_{21},\; O_{22},\; O_{23}]

含义:

  • O1O_1:第 1 个词的 query 向量 q1q_1 与所有 key 向量(k1,k2k_1, k_2)计算注意力得分并归一化为注意力分布 P1=[P11,  P12]P_1 = [P_{11},\; P_{12}] 后,用该分布作为权重对所有词的 value 向量v1,v2v_1, v_2加权求和的结果。即 O1=P11v1+P12v2O_1 = P_{11}v_1 + P_{12}v_2
  • O2O_2:同理,第 2 个词的 query 向量 q2q_2 与所有 key 计算得分并归一化为 P2=[P21,  P22]P_2 = [P_{21},\; P_{22}] 后,对所有词的 value 向量加权求和的结果。即 O2=P21v1+P22v2O_2 = P_{21}v_1 + P_{22}v_2

其中分量形式为:

Oij=Pi1v1j+Pi2v2j=l=12PilvljO_{ij} = P_{i1}v_{1j} + P_{i2}v_{2j} = \sum_{l=1}^{2} P_{il} v_{lj}

元素 OijO_{ij} 的详细解释:

OijO_{ij} 是输出矩阵 OO 的第 ii 行第 jj 列元素。具体计算过程为:

  1. ii 个词的 query 向量 qiq_i 与所有词的 key 向量 k1,k2k_1, k_2 计算注意力得分,经 Softmax 归一化得到注意力分布 Pi=[Pi1,  Pi2]P_i = [P_{i1},\; P_{i2}]
  2. 用该分布对所有词的 value 向量 v1,v2v_1, v_2 做加权求和,得到第 ii 个词的输出向量 Oi=Pi1v1+Pi2v2O_i = P_{i1}v_1 + P_{i2}v_2
  3. OijO_{ij} 就是这个输出向量 OiO_i 的第 jj 个分量(即嵌入维度的第 jj 维)。

从矩阵乘法展开看,OijO_{ij} 等于所有词的 value 向量第 jj 维分量按第 ii 个词对各词的关注概率加权求和:

Oij=Pi1v1j+Pi2v2jO_{ij} = P_{i1} \cdot v_{1j} + P_{i2} \cdot v_{2j}

即:第 1 个词 value 向量的第 jj 维分量 v1jv_{1j} 乘以第 ii 个词对第 1 个词的关注概率 Pi1P_{i1},加上第 2 个词 value 向量的第 jj 维分量 v2jv_{2j} 乘以第 ii 个词对第 2 个词的关注概率 Pi2P_{i2}

  • 行向量 Oi=[Oi1,  Oi2,  Oi3]O_i = [O_{i1},\; O_{i2},\; O_{i3}]:第 ii 个词的 query 经注意力机制后,按注意力分布 PiP_i 对所有词的 value 向量加权求和得到的输出向量

5. 加入 Causal Mask(因果掩码)

Decoder 中需防止第 ii 个词看到第 j>ij > i 个词。对 n=2n=2

M=[000]M = \begin{bmatrix} 0 & -\infty \\ 0 & 0 \end{bmatrix}

Masked 分数

Smask=S+M=[S11S21S22]S^{\text{mask}} = S + M = \begin{bmatrix} S_{11} & -\infty \\ S_{21} & S_{22} \end{bmatrix}

Masked Softmax

第 1 行中 e=0e^{-\infty}=0,故:

P11mask=eS~11eS~11+0=1,P12mask=0P_{11}^{\text{mask}} = \frac{e^{\tilde{S}_{11}}}{e^{\tilde{S}_{11}}+0} = 1, \quad P_{12}^{\text{mask}} = 0

第 2 行正常计算:

P21mask=eS~21eS~21+eS~22,P22mask=eS~22eS~21+eS~22P_{21}^{\text{mask}} = \frac{e^{\tilde{S}_{21}}}{e^{\tilde{S}_{21}}+e^{\tilde{S}_{22}}}, \quad P_{22}^{\text{mask}} = \frac{e^{\tilde{S}_{22}}}{e^{\tilde{S}_{21}}+e^{\tilde{S}_{22}}}

Masked 输出

Omask=PmaskV=[1v1+0v2P21maskv1+P22maskv2]=[v1P21maskv1+P22maskv2]O^{\text{mask}} = P^{\text{mask}} V = \begin{bmatrix} 1 \cdot v_1 + 0 \cdot v_2 \\ P_{21}^{\text{mask}} v_1 + P_{22}^{\text{mask}} v_2 \end{bmatrix} = \begin{bmatrix} v_1 \\ P_{21}^{\text{mask}} v_1 + P_{22}^{\text{mask}} v_2 \end{bmatrix}

结果:第 1 个词只能看到自己(O1=v1O_1 = v_1),第 2 个词可以看到自己和第 1 个词。

评论