AI 工程基础体系 · 第 71/100 篇。内容覆盖机器学习、深度学习与生成式 AI;模型、数据、评测、权限和成本会作为同一生产系统处理。

Attention 数学细解:QKV、缩放点积、Mask、多头与复杂度

Attention(注意力)是一种根据查询内容,从一组候选信息中计算“应关注多少”的加权聚合机制。Transformer 将它作为核心计算单元:输入序列中的每个位置都可以根据与其他位置的相关性,动态读取其他位置的信息。

本文中的符号约定如下:

  • nn:序列长度;
  • dmodeld_{\text{model}}:模型隐藏维度;
  • dkd_k:Query、Key 的维度;
  • dvd_v:Value 的维度;
  • Q,K,VQ,K,V:Query、Key、Value 矩阵;
  • SS:未归一化的注意力分数,也叫 logits;
  • AA:经过 Mask 和 Softmax 后的注意力权重矩阵;
  • OO:Attention 输出。

1. 从加权平均开始理解 Attention

假设有三个候选信息:

v1=[10,0],v2=[0,10],v3=[5,5]v_1=[10,0],\quad v_2=[0,10],\quad v_3=[5,5]

如果某个查询认为三个信息的重要性分别是:

a1=0.6,a2=0.3,a3=0.1a_1=0.6,\quad a_2=0.3,\quad a_3=0.1

那么输出就是加权平均:

o=0.6v1+0.3v2+0.1v3o=0.6v_1+0.3v_2+0.1v_3

计算得到:

o=0.6[10,0]+0.3[0,10]+0.1[5,5]=[6.5,3.5]o= 0.6[10,0]+0.3[0,10]+0.1[5,5] =[6.5,3.5]

Attention 的关键不在于“加权平均”本身,而在于权重 aia_i 不是固定参数,而是由当前查询与候选信息之间的匹配关系动态计算出来的:

Query与各个 Key 比较得到权重加权读取 Value\text{Query} \longrightarrow \text{与各个 Key 比较} \longrightarrow \text{得到权重} \longrightarrow \text{加权读取 Value}

这也解释了为什么需要三组向量,而不是只使用一组向量:

  • Query:当前请求“想找什么”;
  • Key:每个候选项“具有什么索引特征”;
  • Value:真正被读取和聚合的内容。

Key 用于匹配,Value 用于传递信息。Key 和 Value 可以来自同一个输入,也可以来自不同输入。


2. Q、K、V 的数学定义

设输入序列表示为矩阵:

XRn×dmodelX\in\mathbb{R}^{n\times d_{\text{model}}}

其中第 iixix_i 是第 ii 个位置的隐藏向量。

通过三个独立的线性投影得到:

Q=XWQQ=XW_Q

K=XWKK=XW_K

V=XWVV=XW_V

如果是单头 Attention,通常有:

WQRdmodel×dkW_Q\in\mathbb{R}^{d_{\text{model}}\times d_k}

WKRdmodel×dkW_K\in\mathbb{R}^{d_{\text{model}}\times d_k}

WVRdmodel×dvW_V\in\mathbb{R}^{d_{\text{model}}\times d_v}

因此:

QRn×dk,KRn×dk,VRn×dvQ\in\mathbb{R}^{n\times d_k},\quad K\in\mathbb{R}^{n\times d_k},\quad V\in\mathbb{R}^{n\times d_v}

2.1 Self-Attention

Self-Attention 中,Q,K,VQ,K,V 都来自同一个输入 XX

Q=XWQ,K=XWK,V=XWVQ=XW_Q,\quad K=XW_K,\quad V=XW_V

因此序列中的每个位置都可以读取同一序列中其他位置的信息。

例如句子:

小明把书放在桌子上,因为它很重。

“它”所在位置可以通过 Self-Attention 读取“小明”“书”“桌子”等位置的表示,模型据此判断指代关系。

2.2 Cross-Attention

Cross-Attention 中,Query 和 Key/Value 来自不同序列:

Q=XqueryWQQ=X_{\text{query}}W_Q

K=XsourceWK,V=XsourceWVK=X_{\text{source}}W_K,\quad V=X_{\text{source}}W_V

典型场景是编码器—解码器模型:

  • 编码器输出源语言序列;
  • 解码器当前状态生成 Query;
  • 编码器输出生成 Key 和 Value;
  • 解码器根据当前生成需求读取源序列。

如果目标序列长度为 nqn_q,源序列长度为 nkn_k,则:

QRnq×dkQ\in\mathbb{R}^{n_q\times d_k}

KRnk×dk,VRnk×dvK\in\mathbb{R}^{n_k\times d_k},\quad V\in\mathbb{R}^{n_k\times d_v}

最终注意力矩阵大小是:

nq×nkn_q\times n_k

它不必是方阵。


3. 点积如何产生匹配分数

对于第 ii 个 Query 和第 jj 个 Key,最基本的匹配分数是点积:

sij=qikjs_{ij}=q_i\cdot k_j

展开为:

sij=r=1dkqirkjrs_{ij}=\sum_{r=1}^{d_k}q_{ir}k_{jr}

如果 qiq_ikjk_j 方向相近,点积较大;方向相反,点积较小;接近正交时,点积接近零。

把所有 Query 和 Key 的两两点积写成矩阵形式:

S=QKTS=QK^\mathsf{T}

其中:

SRnq×nkS\in\mathbb{R}^{n_q\times n_k}

矩阵元素满足:

Sij=qikjS_{ij}=q_i\cdot k_j

因此第 ii 行表示第 ii 个 Query 对所有 Key 的匹配分数。

3.1 点积不是概率

SijS_{ij} 可以是任意实数,包括负数,也不满足行和为 1。因此它只是 logits,不是注意力权重。

需要对每一行应用 Softmax:

Aij=exp(Sij)t=1nkexp(Sit)A_{ij}= \frac{\exp(S_{ij})} {\sum_{t=1}^{n_k}\exp(S_{it})}

Softmax 后:

Aij0A_{ij}\ge 0

并且对每个未屏蔽的 Query 行:

jAij=1\sum_j A_{ij}=1

于是每个输出位置都是 Value 的加权平均:

oi=j=1nkAijvjo_i=\sum_{j=1}^{n_k}A_{ij}v_j

矩阵形式是:

O=AVO=AV


4. 为什么需要缩放点积

Transformer 使用的核心公式是:

Attention(Q,K,V)=softmax(QKTdk)V\operatorname{Attention}(Q,K,V) = \operatorname{softmax} \left( \frac{QK^\mathsf{T}}{\sqrt{d_k}} \right)V

相比普通点积,多了一个缩放因子:

1dk\frac{1}{\sqrt{d_k}}

4.1 方差推导

假设 Query 和 Key 的每个分量独立,满足:

E[qr]=0,E[kr]=0\mathbb{E}[q_r]=0,\quad \mathbb{E}[k_r]=0

Var(qr)=1,Var(kr)=1\operatorname{Var}(q_r)=1,\quad \operatorname{Var}(k_r)=1

点积为:

qk=r=1dkqrkrq\cdot k=\sum_{r=1}^{d_k}q_rk_r

每一项 qrkrq_rk_r 的期望为 0,方差约为 1。独立相加后:

Var(qk)dk\operatorname{Var}(q\cdot k)\approx d_k

因此标准差约为:

dk\sqrt{d_k}

dkd_k 增大时,未缩放点积的绝对值通常也会增大。

4.2 大 logits 会使 Softmax 饱和

考虑两个 logits:

[1,2][1,2]

其 Softmax 约为:

[0.269,0.731][0.269,0.731]

如果 logits 被放大为:

[10,20][10,20]

Softmax 约为:

[0.000045,0.999955][0.000045,0.999955]

这时几乎所有权重都集中到一个位置。输出虽然看似“明确”,但 Softmax 的梯度在极端区域会变得很小,导致训练信号变弱。

缩放后,点积的方差约为:

Var(qkdk)=dkdk=1\operatorname{Var}\left(\frac{q\cdot k}{\sqrt{d_k}}\right) = \frac{d_k}{d_k}=1

这不会保证 logits 永远处于理想范围,但能使不同维度规模下的初始化更稳定。

4.3 缩放不是归一化概率

QKTdk\frac{QK^\mathsf{T}}{\sqrt{d_k}}

只是对 logits 做尺度调整,之后仍然必须执行 Softmax。它不是除以行和,也不是 L2 归一化。


5. 一个完整的数值算例

设:

Q=[1001],K=[1001]Q= \begin{bmatrix} 1&0\\ 0&1 \end{bmatrix}, \quad K= \begin{bmatrix} 1&0\\ 0&1 \end{bmatrix}

V=[100020]V= \begin{bmatrix} 10&0\\ 0&20 \end{bmatrix}

这里:

n=2,dk=2,dv=2n=2,\quad d_k=2,\quad d_v=2

5.1 计算未缩放点积

QKT=[1001][1001]=[1001]QK^\mathsf{T} = \begin{bmatrix} 1&0\\ 0&1 \end{bmatrix} \begin{bmatrix} 1&0\\ 0&1 \end{bmatrix} = \begin{bmatrix} 1&0\\ 0&1 \end{bmatrix}

5.2 缩放

因为:

dk=2\sqrt{d_k}=\sqrt{2}

所以:

S=QKT2=[0.7071000.7071]S= \frac{QK^\mathsf{T}}{\sqrt{2}} = \begin{bmatrix} 0.7071&0\\ 0&0.7071 \end{bmatrix}

5.3 对每一行应用 Softmax

第一行:

softmax([0.7071,0])[0.6698,0.3302]\operatorname{softmax}([0.7071,0]) \approx [0.6698,0.3302]

第二行:

softmax([0,0.7071])[0.3302,0.6698]\operatorname{softmax}([0,0.7071]) \approx [0.3302,0.6698]

所以:

A[0.66980.33020.33020.6698]A\approx \begin{bmatrix} 0.6698&0.3302\\ 0.3302&0.6698 \end{bmatrix}

5.4 加权读取 Value

O=AVO=AV

第一行:

o1=0.6698[10,0]+0.3302[0,20][6.698,6.604]o_1= 0.6698[10,0]+0.3302[0,20] \approx [6.698,6.604]

第二行:

o2=0.3302[10,0]+0.6698[0,20][3.302,13.396]o_2= 0.3302[10,0]+0.6698[0,20] \approx [3.302,13.396]

最终:

O[6.6986.6043.30213.396]O\approx \begin{bmatrix} 6.698&6.604\\ 3.302&13.396 \end{bmatrix}

注意力权重决定“读多少”,Value 决定“读到什么”。即使两个位置的 Key 相同,Value 不同,输出仍然可以不同。


6. Softmax 的数值稳定性

直接计算:

softmax(xi)=exijexj\operatorname{softmax}(x_i) = \frac{e^{x_i}}{\sum_j e^{x_j}}

可能发生溢出。例如:

e1000e^{1000}

在浮点数中可能变成无穷大。

利用 Softmax 对整体平移不敏感的性质:

softmax(x)=softmax(xc)\operatorname{softmax}(x) = \operatorname{softmax}(x-c)

通常取:

c=maxixic=\max_i x_i

于是:

softmax(xi)=eximax(x)jexjmax(x)\operatorname{softmax}(x_i) = \frac{e^{x_i-\max(x)}}{\sum_j e^{x_j-\max(x)}}

所有指数的输入都不大于 0,因此最大指数为 1。

常见深度学习框架会在底层实现稳定版本,但自定义 Attention、量化内核或 NumPy 原型时不能忽略这一点。


7. Mask:在 Softmax 前限制可见范围

Mask 用于禁止某些 Query 读取某些 Key。最稳妥的数学形式是先对 logits 添加一个加性掩码:

Sij=Sij+MijS'_{ij}=S_{ij}+M_{ij}

其中:

Mij={0,允许关注,禁止关注M_{ij}= \begin{cases} 0,&\text{允许关注}\\ -\infty,&\text{禁止关注} \end{cases}

然后:

A=softmax(S)A=\operatorname{softmax}(S')

禁止位置的指数项为:

e=0e^{-\infty}=0

因此它们的注意力权重为 0。

实际代码中不一定真的存储 -\infty,也可能使用一个足够小的数,例如 104-10^4。但具体数值需要结合数据类型:

  • FP32 中 109-10^9 通常足够小;
  • FP16、BF16 中应避免不必要的溢出或类型转换;
  • 底层 fused kernel 可能使用专门的布尔 Mask 逻辑。

7.1 Padding Mask

不同样本长度不一致时,通常将 Batch 补齐到同一长度。例如:

样本 A:我 喜欢 猫
样本 B:我 喜欢 狗 很

补齐后:

样本 A:我 喜欢 猫 PAD
样本 B:我 喜欢 狗 很

Padding 位置不是有效输入,因此其他位置不应读取它对应的 Key/Value。Padding Mask 通常屏蔽 Key 维度上的 PAD 列。

如果某个 Query 本身是 PAD,是否还需要计算它的输出取决于后续实现:

  • 后续损失是否忽略 PAD;
  • 是否在残差连接或池化前清除 PAD 位置;
  • 模型是否使用专门的 packed sequence。

只屏蔽 Key 不等于自动让 PAD Query 的输出无效。

7.2 Causal Mask

自回归生成要求位置 ii 不能读取未来位置 j>ij>i。Causal Mask 定义为:

Mij={0,ji,j>iM_{ij}= \begin{cases} 0,&j\le i\\ -\infty,&j>i \end{cases}

长度为 4 时,允许矩阵为:

C=[1000110011101111]C= \begin{bmatrix} 1&0&0&0\\ 1&1&0&0\\ 1&1&1&0\\ 1&1&1&1 \end{bmatrix}

第 1 个位置只能读取自己;第 2 个位置可以读取第 1、2 个位置;依此类推。

如果把方向写反,模型会读取未来 Token,训练时可能表现异常好,但生成时出现严重的信息泄漏。

7.3 将数值算例改为 Causal Attention

前面的未缩放分数为:

S=[0.7071000.7071]S= \begin{bmatrix} 0.7071&0\\ 0&0.7071 \end{bmatrix}

加入长度为 2 的 Causal Mask:

S=[0.707100.7071]S'= \begin{bmatrix} 0.7071&-\infty\\ 0&0.7071 \end{bmatrix}

Softmax 后:

A=[100.33020.6698]A= \begin{bmatrix} 1&0\\ 0.3302&0.6698 \end{bmatrix}

于是:

O=[1003.30213.396]O= \begin{bmatrix} 10&0\\ 3.302&13.396 \end{bmatrix}

第一个位置不再读取第二个位置的信息。

7.4 为什么不能先 Softmax 再乘 Mask

错误做法是:

A=softmax(S)A=\operatorname{softmax}(S)

然后:

A=ACA'=A\odot C

对于:

S=[1,2]S=[1,2]

Softmax 是:

A=[0.269,0.731]A=[0.269,0.731]

如果屏蔽第二列,直接乘 Mask 得到:

A=[0.269,0]A'=[0.269,0]

它的和是 0.269,不再是概率分布。正确做法是先屏蔽:

S=[1,]S'=[1,-\infty]

然后:

softmax(S)=[1,0]\operatorname{softmax}(S')=[1,0]

这表示“在可见候选集合中重新归一化”。

7.5 全部位置被屏蔽的边界

如果某个 Query 行全部是 -\infty,那么:

softmax([,])\operatorname{softmax}([-\infty,-\infty])

会出现:

00\frac{0}{0}

结果通常是 NaN。

因此实现中必须保证:

  • 每个需要输出的 Query 至少有一个可见 Key;
  • 或者对全屏蔽行做特殊处理;
  • 不要仅通过“把所有无效位置都设为 -\infty”就假设结果一定安全。

这类问题经常出现在 Padding Mask 与 Causal Mask 合并、变长批处理以及自定义 Attention 内核中。


8. Attention 的完整公式与张量形状

单头缩放点积 Attention 为:

Attention(Q,K,V)=softmax(QKTdk+M)V\operatorname{Attention}(Q,K,V) = \operatorname{softmax} \left( \frac{QK^\mathsf{T}}{\sqrt{d_k}}+M \right)V

逐步对应:

  1. QKTQK^\mathsf{T}:计算 Query 与 Key 的两两匹配;
  2. 除以 dk\sqrt{d_k}:控制 logits 的尺度;
  3. 加上 MM:删除不可见连接;
  4. Softmax:将每个 Query 的分数转成权重;
  5. 乘以 VV:聚合 Value。

设:

QRnq×dkQ\in\mathbb{R}^{n_q\times d_k}

KRnk×dkK\in\mathbb{R}^{n_k\times d_k}

VRnk×dvV\in\mathbb{R}^{n_k\times d_v}

则:

QKTRnq×nkQK^\mathsf{T}\in\mathbb{R}^{n_q\times n_k}

ARnq×nkA\in\mathbb{R}^{n_q\times n_k}

O=AVRnq×dvO=AV\in\mathbb{R}^{n_q\times d_v}

矩阵乘法要求 Key 的数量和 Value 的行数相同,因为每个 Key 对应一个 Value。


9. Attention 的反向传播结构

训练时,Attention 不只是前向的加权平均,还需要对 Q,K,VQ,K,V 传播梯度。

令:

S=QKTdk+MS=\frac{QK^\mathsf{T}}{\sqrt{d_k}}+M

A=softmax(S)A=\operatorname{softmax}(S)

O=AVO=AV

如果上游梯度为:

GO=LOG_O=\frac{\partial L}{\partial O}

那么对 Value 的梯度为:

LV=ATGO\frac{\partial L}{\partial V} = A^\mathsf{T}G_O

这符合直觉:某个 Value 被关注得越多,输出对它的梯度通常越大。

对注意力权重的梯度为:

LA=GOVT\frac{\partial L}{\partial A} = G_OV^\mathsf{T}

Softmax 的逐行 Jacobian 为:

aisj=ai(δijaj)\frac{\partial a_i}{\partial s_j} = a_i(\delta_{ij}-a_j)

其中 δij\delta_{ij} 是 Kronecker delta:

δij={1,i=j0,ij\delta_{ij}= \begin{cases} 1,&i=j\\ 0,&i\ne j \end{cases}

用向量形式表示,若某一行 Softmax 输出为 aa,上游梯度为 gag_a,则:

gs=a(ga(gaa)1)g_s = a\odot \left( g_a-(g_a\cdot a)\mathbf{1} \right)

这里 \odot 是逐元素乘法。

再考虑缩放点积:

S=QKTdk+MS=\frac{QK^\mathsf{T}}{\sqrt{d_k}}+M

令:

GS=LSG_S=\frac{\partial L}{\partial S}

则:

LQ=GSKdk\frac{\partial L}{\partial Q} = \frac{G_SK}{\sqrt{d_k}}

LK=GSTQdk\frac{\partial L}{\partial K} = \frac{G_S^\mathsf{T}Q}{\sqrt{d_k}}

Mask 位置不是可训练连接。实现中通常通过 -\infty 或布尔逻辑使其梯度为零。


10. 多头 Attention

单头 Attention 只能在一个投影空间中建立匹配。Multi-Head Attention(多头注意力)让模型并行使用多个独立的投影空间。

设头数为 hh,第 rr 个头的投影为:

Qr=QWQ(r)Q_r=QW_Q^{(r)}

Kr=KWK(r)K_r=KW_K^{(r)}

Vr=VWV(r)V_r=VW_V^{(r)}

rr 个头的输出为:

Hr=softmax(QrKrTdk+M)VrH_r= \operatorname{softmax} \left( \frac{Q_rK_r^\mathsf{T}}{\sqrt{d_k}} +M \right)V_r

将所有头拼接:

H=Concat(H1,H2,,Hh)H=\operatorname{Concat}(H_1,H_2,\ldots,H_h)

最后经过输出投影:

O=HWOO=HW_O

完整公式为:

MultiHead(Q,K,V)=Concat(head1,,headh)WO\operatorname{MultiHead}(Q,K,V) = \operatorname{Concat}(\text{head}_1,\ldots,\text{head}_h)W_O

其中:

headr=Attention(QWQ(r),KWK(r),VWV(r))\text{head}_r= \operatorname{Attention} (QW_Q^{(r)},KW_K^{(r)},VW_V^{(r)})

10.1 为什么需要多头

不同头可以学习不同类型的关系,例如:

  • 局部相邻关系;
  • 主谓关系;
  • 指代关系;
  • 句法结构;
  • 长距离依赖;
  • 特定语义特征。

这不是规范保证,而是模型训练后可能形成的行为。不能简单地把某个头固定解释为某种语法功能;头的语义通常需要通过可视化、消融实验或探针任务验证。

10.2 维度如何划分

Transformer 中常见设置是:

dk=dv=dheadd_k=d_v=d_{\text{head}}

并且:

hdhead=dmodelh\cdot d_{\text{head}}=d_{\text{model}}

例如:

dmodel=512,h=8d_{\text{model}}=512,\quad h=8

则:

dhead=64d_{\text{head}}=64

每个头单独在 64 维空间中进行 Attention,最后拼回 512 维。

注意,除数应当是每个头的 Key 维度:

dhead\sqrt{d_{\text{head}}}

而不是:

dmodel\sqrt{d_{\text{model}}}

这是实现多头 Attention 时常见的错误。

10.3 多头不是简单复制单头

如果所有头共享完全相同的 WQ,WK,WVW_Q,W_K,W_V,再复制结果,那么多头不会获得真正的表示多样性。多头的作用来自每个头拥有不同的投影参数和独立的注意力矩阵。

多头输出还需要 WOW_O 做跨头混合。没有输出投影时,各头只能在拼接后的固定分块中传递信息,表达能力和后续层的接口都会不同。


11. 从输入到多头输出的形状变化

以 Batch 输入为例:

XRB×n×dmodelX\in\mathbb{R}^{B\times n\times d_{\text{model}}}

经过线性层:

Q,K,VRB×n×dmodelQ,K,V\in\mathbb{R}^{B\times n\times d_{\text{model}}}

然后拆成头:

Q,K,VRB×h×n×dheadQ,K,V\in\mathbb{R}^{B\times h\times n\times d_{\text{head}}}

进行转置后,矩阵乘法得到:

QKTRB×h×n×nQK^\mathsf{T} \in \mathbb{R}^{B\times h\times n\times n}

Softmax 必须沿最后一个 Key 位置维度进行:

A=softmax(,dim=1)A=\operatorname{softmax}(\cdots,\text{dim}=-1)

而不是沿头维度或 Query 维度。最终:

AVRB×h×n×dheadAV\in \mathbb{R}^{B\times h\times n\times d_{\text{head}}}

转回:

B×n×h×dheadB\times n\times h\times d_{\text{head}}

再合并头:

B×n×dmodelB\times n\times d_{\text{model}}


12. 可执行的 NumPy 实现

下面的实现包含:

  • 缩放点积;
  • 稳定 Softmax;
  • 布尔 Mask;
  • Causal Mask;
  • 全屏蔽行检查。
import numpy as np


def stable_softmax(x, axis=-1):
    """数值稳定的 Softmax。"""
    x_max = np.max(x, axis=axis, keepdims=True)
    exp_x = np.exp(x - x_max)
    return exp_x / np.sum(exp_x, axis=axis, keepdims=True)


def scaled_dot_product_attention(q, k, v, allow=None):
    """
    q: [n_q, d_k]
    k: [n_k, d_k]
    v: [n_k, d_v]
    allow: [n_q, n_k] 的 bool 矩阵。
           True 表示允许关注,False 表示屏蔽。
    """
    if q.ndim != 2 or k.ndim != 2 or v.ndim != 2:
        raise ValueError("q、k、v 必须都是二维矩阵")

    n_q, d_k = q.shape
    n_k, d_k_k = k.shape

    if d_k != d_k_k:
        raise ValueError("q 和 k 的最后一维必须相同")

    if v.shape[0] != n_k:
        raise ValueError("k 的行数必须等于 v 的行数")

    scores = q @ k.T / np.sqrt(d_k)

    if allow is not None:
        if allow.shape != (n_q, n_k):
            raise ValueError("allow 的形状必须是 [n_q, n_k]")

        # 每个 Query 至少需要一个可见 Key。
        if np.any(np.sum(allow, axis=-1) == 0):
            raise ValueError("存在完全被屏蔽的 Query 行")

        scores = np.where(allow, scores, -np.inf)

    weights = stable_softmax(scores, axis=-1)
    output = weights @ v
    return output, weights


q = np.array([
    [1.0, 0.0],
    [0.0, 1.0],
])

k = np.array([
    [1.0, 0.0],
    [0.0, 1.0],
])

v = np.array([
    [10.0, 0.0],
    [0.0, 20.0],
])

# 每个位置都可以看到全部位置
output_full, weights_full = scaled_dot_product_attention(q, k, v)

# Causal Mask:位置 i 只能看到位置 j <= i
causal = np.tril(np.ones((2, 2), dtype=bool))
output_causal, weights_causal = scaled_dot_product_attention(
    q, k, v, allow=causal
)

np.set_printoptions(precision=4, suppress=True)

print("full weights:")
print(weights_full)
print("full output:")
print(output_full)

print("causal weights:")
print(weights_causal)
print("causal output:")
print(output_causal)

预期输出近似为:

full weights:
[[0.6698 0.3302]
 [0.3302 0.6698]]

full output:
[[ 6.6976  6.6041]
 [ 3.302  13.3959]]

causal weights:
[[1.     0.    ]
 [0.3302 0.6698]]

causal output:
[[10.      0.    ]
 [ 3.302  13.3959]]

这里的 allow 使用“允许矩阵”语义,和某些框架中使用的 1 表示保留、0 表示屏蔽的约定相似,但不同 API 对 Mask 的语义可能相反。调用框架接口时必须确认:

  • True 是允许还是屏蔽;
  • 数值 1 是保留还是屏蔽;
  • Mask 是加到 logits 上,还是传给底层布尔内核;
  • Mask 的形状是否支持广播;
  • 是否需要显式提供 Causal Mask。

不能仅凭参数名 attention_mask 推断语义。


13. 复杂度:时间、空间与真正的瓶颈

设:

  • Batch 大小为 BB
  • Query 长度为 nqn_q
  • Key 长度为 nkn_k
  • 模型维度为 dmodeld_{\text{model}}
  • 头数为 hh
  • 每头维度为 dheadd_{\text{head}}

13.1 投影复杂度

计算 Q,K,VQ,K,V 的线性投影大致需要:

O(Bndmodel2)O(Bn d_{\text{model}}^2)

如果 Q,K,VQ,K,V 都来自长度为 nn 的序列,三组投影的常数大约是三次矩阵乘法;实际实现通常将它们合并为一个大的线性层以提高硬件利用率。

输出投影还需要:

O(Bndmodel2)O(Bn d_{\text{model}}^2)

13.2 分数矩阵复杂度

每个头需要计算:

QKTQK^\mathsf{T}

其复杂度为:

O(Bhnqnkdhead)O(Bh n_q n_k d_{\text{head}})

由于:

hdhead=dmodelh d_{\text{head}}=d_{\text{model}}

可以写成:

O(Bnqnkdmodel)O(B n_q n_k d_{\text{model}})

Self-Attention 中 nq=nk=nn_q=n_k=n,因此是:

O(Bn2dmodel)O(Bn^2d_{\text{model}})

这就是通常所说的 Attention 对序列长度呈二次复杂度。

13.3 加权聚合复杂度

计算:

AVAV

复杂度同样为:

O(Bhnqnkdhead)=O(Bnqnkdmodel)O(Bh n_q n_k d_{\text{head}}) = O(B n_q n_k d_{\text{model}})

因此完整 Attention 的主要计算量可以概括为:

O(Bndmodel2)+O(Bn2dmodel)O(Bn d_{\text{model}}^2) + O(Bn^2d_{\text{model}})

其中:

  • nn 较小时,投影层的 dmodel2d_{\text{model}}^2 项可能占主导;
  • nn 很大时,n2n^2 项通常成为瓶颈。

13.4 注意力矩阵的内存复杂度

分数矩阵和权重矩阵形状为:

B×h×n×nB\times h\times n\times n

因此显式保存它们的空间复杂度为:

O(Bhn2)O(Bhn^2)

训练时还可能需要保存:

  • Softmax 前的 logits;
  • Softmax 后的权重;
  • Dropout 掩码;
  • Q、K、V;
  • 反向传播所需的中间结果。

所以实际显存压力不只来自一次矩阵乘法。

FlashAttention 等实现的核心思想之一,是通过分块和在线 Softmax 避免显式保存完整的 n×nn\times n 注意力矩阵,同时保持精确 Attention 的数学结果或在浮点误差范围内一致。它降低的是中间内存访问和显存占用,不代表理论上的全连接 Attention 已经从二次复杂度变成线性复杂度。


14. 自回归生成与 KV Cache

训练时,因果 Self-Attention 通常一次处理整个序列:

Q,K,VRn×dQ,K,V\in\mathbb{R}^{n\times d}

虽然未来位置被 Mask,但所有位置的计算仍然可以并行。

生成时,每一步只需要为新 Token 计算一个新的 Query。历史 Token 的 Key 和 Value 不再变化,因此可以缓存:

KcacheRL×dkK_{\text{cache}}\in\mathbb{R}^{L\times d_k}

VcacheRL×dvV_{\text{cache}}\in\mathbb{R}^{L\times d_v}

当前步:

qnewR1×dkq_{\text{new}}\in\mathbb{R}^{1\times d_k}

只计算:

qnewKcacheTq_{\text{new}}K_{\text{cache}}^\mathsf{T}

其复杂度从单步的全序列计算变为:

O(Ldk)O(Ld_k)

再加上当前 Token 的线性投影和输出层计算。

但是,如果生成 TT 个 Token,缓存 Attention 的总读取量仍大致为:

O(T2dmodel)O(T^2d_{\text{model}})

因为第 tt 步需要读取长度约为 tt 的历史。KV Cache 主要避免了重复计算历史 Token 的 K,VK,V,并使单步延迟可控;它没有消除长序列生成的全部成本。

KV Cache 的显存大致与以下因素成正比:

O(BL层数hdhead2)O(B\cdot L\cdot \text{层数}\cdot h\cdot d_{\text{head}}\cdot 2)

最后的 2 来自 Key 和 Value 两份缓存。实际大小还受数据类型、是否使用 MQA/GQA、张量并行布局等影响。


15. 多头 Attention 的复杂度是否增加

如果保持:

hdhead=dmodelh d_{\text{head}}=d_{\text{model}}

那么所有头的分数计算总复杂度仍然近似为:

O(Bn2dmodel)O(Bn^2d_{\text{model}})

不是简单地乘上一个额外的 hh,因为每个头的维度同时缩小了。

但头数会影响:

  • 中间张量的形状;
  • Kernel 是否适合硬件;
  • 每个头的并行粒度;
  • 注意力权重的存储布局;
  • KV Cache 的组织方式;
  • 小头维度下的数值和表达能力。

因此理论 FLOPs 相近,不意味着不同头数的实际延迟、显存和吞吐完全相同。


16. Attention 的排列等变性与位置信息

如果只对输入 Token 做排列,而没有任何位置编码,Self-Attention 本身不会知道原始顺序。

PP 是一个排列矩阵,输入变为:

X=PXX'=PX

则:

Q=PXWQ=PQQ'=PXW_Q=PQ

K=PK,V=PVK'=PK,\quad V'=PV

分数矩阵变为:

QKT=(PQ)(PK)T=PQKTPTQ'K'^\mathsf{T} = (PQ)(PK)^\mathsf{T} = PQK^\mathsf{T}P^\mathsf{T}

Softmax 按行应用后,注意力权重也相应排列,最终输出满足:

O=POO'=PO

这叫排列等变性:输入顺序变化,输出按相同方式变化,但 Attention 没有独立能力判断“谁在前、谁在后”。

因此 Transformer 需要加入位置信息,例如:

  • 绝对位置编码;
  • 相对位置偏置;
  • RoPE;
  • ALiBi;
  • 其他位置表示机制。

位置机制与 Causal Mask 解决的是不同问题:

  • Causal Mask 限制“能否看见未来”;
  • 位置机制告诉模型“当前 Token 位于什么位置、两个 Token 相距多远”。

有 Causal Mask 并不意味着模型自动获得了完整的位置信息。


17. 常见错误与失败表现

17.1 忘记除以 dk\sqrt{d_k}

表现可能包括:

  • Softmax 权重过早接近 one-hot;
  • 注意力分布熵很低;
  • 梯度变小;
  • 训练初期不稳定;
  • 增大隐藏维度后问题更明显。

但并非所有模型都显式使用同一种缩放方式;如果使用了余弦相似度、特殊归一化或专用内核,公式可能不同,不能机械叠加缩放。

17.2 在错误维度上做 Softmax

对于:

ARB×h×nq×nkA\in\mathbb{R}^{B\times h\times n_q\times n_k}

Softmax 应沿 Key 位置维度 nkn_k 进行,使每个 Query 对所有 Key 的权重之和为 1。

如果错误地沿 Query 维度做 Softmax,得到的含义变成“某个 Key 被不同 Query 关注的相对程度”,不再是标准 Attention。

诊断方法是检查:

weights.sum(axis=-1)

对于没有异常 Mask 的有效 Query 行,应接近 1。

17.3 Mask 方向写反

Causal Mask 的正确条件是:

jij\le i

若误写成:

jij\ge i

模型就会看到未来信息。训练损失可能异常低,但真实逐 Token 生成时性能明显下降。

可以用一个长度为 3 的人工输入验证第一个 Query:

  • 正确因果 Mask:第一行只有第一个 Key 可见;
  • 若第一行能看到第二、第三个 Key,说明存在泄漏。

17.4 只屏蔽 Query,不屏蔽 Key

Padding Mask 主要阻止有效 Query 读取 PAD Key。若只把 PAD Query 的输出置零,却没有屏蔽 PAD Key,那么有效 Token 仍可能把 PAD 的 Value 聚合进来。

17.5 用 Mask 实现权限控制

Attention Mask 只能控制一次前向计算中的信息可见性,不能替代权限系统。

如果一个用户无权访问某条文档,不能仅依赖:

  • Prompt 中的提示;
  • 生成阶段的 Attention Mask;
  • 模型“应该不会关注”的经验。

权限必须在数据检索、服务端授权和上下文组装阶段完成。模型内的 Mask 是计算约束,不是安全边界。


18. Dropout 在 Attention 中的位置

原始 Transformer 论文在注意力权重上使用 Dropout。若 AA 是 Softmax 后的权重,训练时可以得到:

A~=Dropout(A)\widetilde{A}=\operatorname{Dropout}(A)

然后:

O=A~VO=\widetilde{A}V

这与 Causal Mask 不同:

  • Mask 是结构性禁止连接;
  • Dropout 是训练时随机丢弃部分连接;
  • Mask 在推理时仍然存在;
  • Dropout 通常在推理时关闭。

Dropout 后的权重在实现中可能采用 inverted dropout 缩放,因此其行和不一定仍然严格等于 1。不能看到注意力权重经过 Dropout 后行和变化,就判断实现错误。


19. 长上下文的真实边界

理论复杂度为:

O(n2)O(n^2)

意味着序列长度翻倍时,两两交互相关的计算量和显式注意力矩阵大小约变为四倍。

这会影响:

  • 训练显存;
  • 推理延迟;
  • KV Cache 体积;
  • 带宽和显存访问;
  • Batch size;
  • 单请求成本;
  • 并发能力。

工程上可以使用以下方向,但它们改变了计算或存储方式,不能自动视为与标准全连接 Attention 等价:

  • FlashAttention:优化精确 Attention 的内存访问;
  • 滑动窗口或局部 Attention:只连接邻近位置;
  • 稀疏 Attention:只计算部分连接;
  • MQA/GQA:减少 Key/Value 头数,降低 KV Cache;
  • 线性 Attention:改变核函数或归一化形式,试图避免显式 n2n^2
  • 分块、检索或层次化上下文:减少一次性送入模型的 Token 数。

选择方案时应同时测量:

  • 相同质量目标下的效果;
  • 端到端延迟,而不是只看单个 Kernel;
  • 峰值显存;
  • 长度增长时的退化;
  • 训练和推理是否使用不同实现;
  • 异常长输入、空输入和全 Mask 输入的行为。

20. 评测时如何识别 Attention 泄漏

如果任务要求模型只能使用历史信息,评测不能只检查最终准确率。还应构造能暴露未来信息泄漏的测试。

一种简单方法是对未来 Token 做扰动:

  1. 固定当前位置及其历史;
  2. 修改未来位置的 Token;
  3. 重新计算当前位置输出;
  4. 比较两次输出。

对于严格的 Causal Attention,当前位置的输出不应依赖未来 Token。实际浮点计算可能存在极小误差,但不应出现有意义的系统变化。

训练数据也可能产生“伪泄漏”:

  • 标签被拼入输入;
  • 预处理使用了未来统计量;
  • Padding 或特殊 Token 处理错误;
  • 训练和评测的 Mask 方向不一致。

因此 Mask 正确只是必要条件,不是完整的数据泄漏检查。


21. 与生产系统的关系

Attention 的数学正确性最终会影响生产系统中的多个维度:

  • 模型:缩放、Mask、头维度和位置机制决定模型实际计算;
  • 数据:Padding、特殊 Token、序列截断会改变可见范围;
  • 评测:训练时并行的 Causal Mask 必须与推理时逐 Token 生成一致;
  • 权限:数据授权应在模型输入前完成,不能把 Attention Mask 当作访问控制;
  • 成本:序列长度同时影响二次交互计算、显存和 KV Cache;
  • 并发:长序列请求会占用更多显存,可能降低 Batch 合并和服务吞吐;
  • 故障处理:全 Mask 行、错误广播、NaN logits 和 KV Cache 形状不匹配都应在运行时尽早检测。

一个可复现的 Attention 实现,至少应验证以下不变量:

shape(QKT)=[nq,nk]\operatorname{shape}(QK^\mathsf{T}) = [n_q,n_k]

Aij=0当位置 (i,j) 被屏蔽A_{ij}=0 \quad\text{当位置 }(i,j)\text{ 被屏蔽}

jAij1对于有效 Query 行\sum_j A_{ij}\approx 1 \quad\text{对于有效 Query 行}

O=AVO=AV

dk(Q)=dk(K)d_k(Q)=d_k(K)

nk(K)=nk(V)n_k(K)=n_k(V)

这些条件分别对应形状、Mask、归一化、聚合和矩阵乘法的基本正确性。实现优化、混合精度或替换内核后,应重新检查这些数学性质,而不能只根据最终 loss 或吞吐判断正确。

Attention 的核心可以归结为一条链路:

Q,K,VQKTdkMaskSoftmaxAVQ,K,V \rightarrow \frac{QK^\mathsf{T}}{\sqrt{d_k}} \rightarrow \text{Mask} \rightarrow \text{Softmax} \rightarrow AV

Q 决定查询需求,K 决定匹配方式,V 决定被传递的内容;缩放控制 logits 的统计尺度,Mask 控制可见连接,多头提供多个投影空间,复杂度则决定了这种全连接交互在长序列和生产推理中的成本边界。


系列导航与关联阅读

官方资料

本文依据研究论文、标准组织与主流框架官方文档重新梳理;正文、示例与工程清单由 WR BLOG 编写。