跳到正文
Learn Everything
返回

大模型基础(二):Self-Attention 从张量到多头机制

更新于:
编辑文章

Self-Attention 从张量投影到上下文聚合

上一篇建立了 Transformer 的整体坐标,这一篇只深入一件事:一组 Token 表示进入 Attention 后,究竟经历了哪些张量变换?

Attention 经常被解释成“给重要的词更高权重”,但这句话隐藏了很多关键细节:

本文会从张量形状开始,完整走过一次 Scaled Dot-Product Attention,再扩展到 Multi-Head、Cross-Attention 与自回归推理。

符号约定:BB 表示 Batch Size,NN 表示序列长度,DD 表示模型维度,HH 表示 Head 数量,Dh=D/HD_h=D/H 表示每个 Head 的维度。

一、Attention 不是一个分数,而是一条数据流水线

一次标准 Self-Attention 可以写成:

Q=XWQ,K=XWK,V=XWV,S=QK⊤Dh,A=softmax⁡(S+M),Z=AV.\begin{aligned} Q&=XW_Q, & K&=XW_K,\\ V&=XW_V, & S&=\frac{QK^{\top}}{\sqrt{D_h}},\\ A&=\operatorname{softmax}(S+M), & Z&=AV. \end{aligned}

其中:

Self-Attention 的完整张量流水线

图 1:单个 Attention Head 的完整数据流。矩阵形状明确说明了二次项产生在 N×NN\times N 的 Score 和 Weight,而不是所有步骤都具有二次复杂度。

这条流水线可以分成三个问题:

  1. 用什么空间描述“我要找什么”和“我有什么”?
  2. 怎样把匹配分数变成合法权重?
  3. 怎样根据权重取回真正的信息?

Q、K、V 分别回答这三个问题。

二、为什么需要 Query、Key 和 Value

假设 Token xix_i 同时包含词义、位置、句法和上下文特征。如果直接用 ⟨xi,xj⟩\langle x_i,x_j\rangle 计算相似度,匹配条件与被取回内容会被绑定在同一个空间里。

Transformer 使用三套可学习投影:

qi=xiWQ,ki=xiWK,vi=xiWVq_i=x_iW_Q,\quad k_i=x_iW_K,\quad v_i=x_iW_V

可以用检索系统类比:

名称检索类比在 Attention 中的职责
Query搜索条件当前 Token 想寻找什么
Key索引字段当前 Token 可以怎样被匹配
Value文档内容被匹配后真正返回什么

关键在于:匹配空间与内容空间被解耦。

两个 Token 可以在 Q/K 空间中高度匹配,但返回的 V 不必等于用于匹配的 K。这让模型能够学习类似“用代词特征寻找名词,但取回名词的语义表示”这样的变换。

Self-Attention 的“Self”是什么意思

Self-Attention 中三者来自同一个输入:

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

“Self”不是说某个 Token 只关注自己,而是说 Query、Key 和 Value 都来自同一组序列表示。

Cross-Attention 有什么不同

在 Encoder–Decoder Transformer 中:

Q=YWQ,K=XencWK,V=XencWV.\begin{aligned} Q&=YW_Q,\\ K&=X_{\mathrm{enc}}W_K,\quad V=X_{\mathrm{enc}}W_V. \end{aligned}

Decoder 用自己的状态提出 Query,到 Encoder Memory 中匹配 Key 并取回 Value。

因此 Self-Attention 和 Cross-Attention 使用同一个计算公式,区别在于 Q、K、V 的来源。

三、从输入 X 到 Q、K、V 的张量形状

先忽略 Batch 和 Multi-Head,假设:

X∈RN×D,WQ,WK∈RD×Dh,WV∈RD×Dv.\begin{aligned} X&\in\mathbb{R}^{N\times D},\\ W_Q,W_K&\in\mathbb{R}^{D\times D_h},\\ W_V&\in\mathbb{R}^{D\times D_v}. \end{aligned}

投影后:

Q,K∈RN×Dh,V∈RN×DvQ,K\in\mathbb{R}^{N\times D_h},\qquad V\in\mathbb{R}^{N\times D_v}

计算:

QK⊤:[N,Dh] [Dh,N]⟶[N,N]QK^{\top}: [N,D_h]\,[D_h,N]\longrightarrow[N,N]

Score Matrix 的含义是:

Sij=⟨qi,kj⟩DhS_{ij}=\frac{\langle q_i,k_j\rangle}{\sqrt{D_h}}

其中行索引 ii 对应 Query,列索引 jj 对应 Key。

这也是理解 Softmax 轴的关键:对每一个 Query,应该在它可以访问的所有 Key 上归一化,所以 Softmax 沿最后一个维度,也就是 Key 维度执行。

weights = torch.softmax(scores, dim=-1)

每一行满足:

∑jAij=1\sum_j A_{ij}=1

但不同 Query 的两行之间不需要归一化,也不要求注意力矩阵对称。

即使 Q=KQ=K,Softmax 的逐行归一化也可能让最终 Attention Weight 不对称;实际模型通常还有不同的 Q、K 投影。

四、为什么除以 sqrt(Dh)

如果 Query 和 Key 每个分量近似独立、均值为 0、方差为 1,那么点积:

⟨q,k⟩=∑r=1Dhqrkr\langle q,k\rangle=\sum_{r=1}^{D_h}q_rk_r

由 DhD_h 个项相加,方差会随 DhD_h 增长到约 DhD_h,标准差约为 Dh\sqrt{D_h}。

维度越大,未经缩放的点积绝对值越容易变大。大幅度 Logit 进入 Softmax 后,输出会非常接近 one-hot:

Softmax 进入饱和区域后,除最大项外的梯度会很小。

因此标准 Attention 使用:

S=QK⊤DhS=\frac{QK^{\top}}{\sqrt{D_h}}

缩放的目标不是改变哪个 Key 最大,而是让分数尺度在不同 Head Dimension 下更稳定。

五、Mask 必须在 Softmax 之前应用

Mask 的作用不是把输出权重“看起来变成零”,而是让被禁止位置不参与概率归一化。

通常做法是:记 Ai\mathcal{A}_i 为 Query ii 允许访问的 Key 集合,则

S~ij={Sij,j∈Ai,−∞,j∉Ai,A=softmax⁡(S~).\begin{aligned} \widetilde S_{ij}&= \begin{cases} S_{ij}, & j\in\mathcal{A}_i,\\ -\infty, & j\notin\mathcal{A}_i, \end{cases}\\ A&=\operatorname{softmax}(\widetilde S). \end{aligned}

因为:

exp⁡(−∞)=0\exp(-\infty)=0

所以被屏蔽位置的权重严格为零,其余合法位置重新归一化。

从 Score Matrix 到 Causal Attention Weight

图 2:Causal Mask 在 Softmax 前加入。每一行代表一个 Query,每一列代表一个 Key;上三角位置属于未来 Token,因此不参与当前行归一化。

Causal Mask

自回归生成中,第 ii 个位置只能看到:

j≤ij\le i

四个 Token 的可见性为:

      Key 1  Key 2  Key 3  Key 4
Q1      ✓      ×      ×      ×
Q2      ✓      ✓      ×      ×
Q3      ✓      ✓      ✓      ×
Q4      ✓      ✓      ✓      ✓

Padding Mask

Batch 中句子长度不同时,短句通常会补 Padding。Padding Key 不应被任何 Query 读取,因此需要在对应列上屏蔽。

两种 Mask 可以叠加

Decoder 训练时经常同时使用:

Mfinal=Mcausal∨MpaddingM_{\mathrm{final}}=M_{\mathrm{causal}}\lor M_{\mathrm{padding}}

为什么不能在 Softmax 后直接乘零

假设原权重是:

[0.2, 0.3, 0.5][0.2,\ 0.3,\ 0.5]

Softmax 后再把第三项乘零:

[0.2, 0.3, 0][0.2,\ 0.3,\ 0]

剩余权重和只有 0.5,不再是合法归一化分布。虽然可以再次除以总和,但这等价于重新归一化,也更容易在数值和实现上出错。

六、Softmax 后为什么乘 V

对于第 ii 个 Query,输出是:

zi=∑jAijvjz_i=\sum_j A_{ij}v_j

这表示 ziz_i 是所有可见 Value 的加权和。

需要特别注意:

矩阵形式:

A∈RN×N,V,Z∈RN×DvA\in\mathbb{R}^{N\times N},\qquad V,Z\in\mathbb{R}^{N\times D_v}

因此:

Z=AVZ=AV

Attention Matrix 决定跨位置的信息路由,Value Matrix 提供真正被路由的内容。

七、用一个可手算案例走完整流程

设序列只有三个 Token,每个 Head 的维度为 2。为了集中观察聚合过程,取:

Q=K=[100111],V=[100211]Q=K= \begin{bmatrix} 1&0\\ 0&1\\ 1&1 \end{bmatrix}, \qquad V= \begin{bmatrix} 1&0\\ 0&2\\ 1&1 \end{bmatrix}

现在只计算第三个 Query:

q3=[1,1]q_3=[1,1]

1. 与三个 Key 点积

⟨q3,k1⟩=1,⟨q3,k2⟩=1,⟨q3,k3⟩=2.\begin{aligned} \langle q_3,k_1\rangle&=1, & \langle q_3,k_2\rangle&=1,\\ \langle q_3,k_3\rangle&=2. \end{aligned}

2. 除以 sqrt(2)

s≈[0.707, 0.707, 1.414]s\approx[0.707,\ 0.707,\ 1.414]

3. 计算 Softmax

a=softmax⁡(s)≈[0.248, 0.248, 0.503]a=\operatorname{softmax}(s) \approx[0.248,\ 0.248,\ 0.503]

由于四舍五入,三项显示值之和约为 0.999;使用完整精度时总和为 1。

4. 聚合 Value

z3=0.248[1,0]+0.248[0,2]+0.503[1,1]≈[0.751, 0.999].\begin{aligned} z_3 &=0.248[1,0]+0.248[0,2]+0.503[1,1]\\ &\approx[0.751,\ 0.999]. \end{aligned}

第三个 Query 的可手算 Attention 案例

图 3:这是人工构造的教学案例,不是模型训练得到的权重。它展示了一个 Query 如何通过三个标量权重聚合三个 Value Vector。

这个案例说明,第三个 Query 最关注第三个 Token,但输出并不是复制 v3v_3,而是三个 Value 的混合。

八、Multi-Head Attention 怎样组织张量

真实 Transformer 不只运行一个 Head。

输入:

X∈RB×N×DX\in\mathbb{R}^{B\times N\times D}

经过一次大的线性投影,或者三次独立投影后,通常会 reshape 为:

Q,K,V∈RB×H×N×DhQ,K,V\in\mathbb{R}^{B\times H\times N\times D_h}

其中:

D=H DhD=H\,D_h

每个 Head 独立计算:

Ah=softmax⁡ ⁣(QhKh⊤Dh+M),Zh=AhVh.\begin{aligned} A_h&=\operatorname{softmax}\!\left( \frac{Q_hK_h^{\top}}{\sqrt{D_h}}+M \right),\\ Z_h&=A_hV_h. \end{aligned}

随后:

Concat⁡(Z1,Z2,…,ZH)∈RB×N×D\operatorname{Concat}(Z_1,Z_2,\ldots,Z_H) \in\mathbb{R}^{B\times N\times D}

最后经过输出投影:

Y=Concat⁡(Z1,…,ZH)WOY=\operatorname{Concat}(Z_1,\ldots,Z_H)W_O

Multi-Head Attention 的切分、并行和合并

图 4:Multi-Head 不是把完整维度无代价复制 H 次,而是通常把模型维度切分到 H 个 Head,再并行计算、拼接并经过输出投影。

为什么多头可能比单头更有表达力

不同 Head 拥有不同的投影参数:

WQ(h),WK(h),WV(h)W_Q^{(h)},\qquad W_K^{(h)},\qquad W_V^{(h)}

因此它们可以在不同表示子空间中建立路由关系。一个 Head 的高权重位置,不要求与另一个 Head 相同。

但需要避免过度解释:

九、MHA、MQA 和 GQA 有什么区别

标准 Multi-Head Attention(MHA)中,每个 Query Head 都有自己的 K 和 V Head:

HQ=HK=HV=HH_Q=H_K=H_V=H

Multi-Query Attention(MQA)让多个 Query Head 共享一组 K/V:

HQ=H,HK=HV=1H_Q=H,\qquad H_K=H_V=1

Grouped-Query Attention(GQA,分组查询注意力)位于两者之间。设 Query Head 数量为 HH,K/V Head 数量为 GG,并假设 HH 能被 GG 整除,则每组包含 R=H/GR=H/G 个 Query Head:

HQ=H,HK=HV=G,R=HG,1<G<H.\begin{aligned} H_Q&=H, & H_K=H_V&=G,\\ R&=\frac{H}{G}, & 1&<G<H. \end{aligned}

第 hh 个 Query Head 使用第 g(h)g(h) 组的 Key 和 Value:

g(h)=⌈hR⌉,head⁡h=Attention⁡ ⁣(Qh,Kg(h),Vg(h)),h∈{1,…,H}.\begin{aligned} g(h)&=\left\lceil\frac{h}{R}\right\rceil,\\ \operatorname{head}_h &=\operatorname{Attention}\!\left(Q_h,K_{g(h)},V_{g(h)}\right),\\ h&\in\{1,\ldots,H\}. \end{aligned}

例如 H=8H=8、G=2G=2 时,每四个 Query Head 共享一组 K/V:

Query Head使用的 K/V Head
1–41
5–82

从这个定义可以看出,MHA 对应 G=HG=H,MQA 对应 G=1G=1,而通常所说的 GQA 取 1<G<H1<G<H。

GQA 中的 Group 指“分组共享 K/V”,不是把组内 Query 或 Attention 输出做算术平均。全部 HH 个 Query Head 仍会分别产生输出,随后拼接并经过输出投影。

在 Head Dimension 和序列长度相同的近似下,GQA 的 K/V Cache 规模约为 MHA 的 G/HG/H。上面的 H=8H=8、G=2G=2 示例只需保留约四分之一的 K/V Head 状态,因此能够降低自回归推理中的 Cache 容量和读取带宽。代价是更多 Query Head 共享 K/V 表示,所以 GQA 常被用作模型质量与推理效率之间的折中。

MQA/GQA 改变的是 Query Head 与 K/V Head 的组织方式,不改变 Scaled Dot-Product Attention 的基本语义。

一个可运行的 GQA 教学实现

下面用 PyTorch 实现前面的 H=8H=8、G=2G=2 示例。为了让分组关系直观可见,代码先生成 2 组 K/V,再把每组显式分配给 4 个 Query Head:

import torch
from torch import nn

B, T, D = 2, 5, 512
Hq, Hkv, d = 8, 2, 64
assert Hq % Hkv == 0

x = torch.randn(B, T, D)

# nn.Linear(in_features, out_features)
q_proj = nn.Linear(D, Hq * d, bias=False)
k_proj = nn.Linear(D, Hkv * d, bias=False)
v_proj = nn.Linear(D, Hkv * d, bias=False)
o_proj = nn.Linear(Hq * d, D, bias=False)

# 线性投影 → 拆头 → 将 Head 维度移到前面
q = q_proj(x).reshape(B, T, Hq, d).transpose(1, 2)
k = k_proj(x).reshape(B, T, Hkv, d).transpose(1, 2)
v = v_proj(x).reshape(B, T, Hkv, d).transpose(1, 2)
# q: [2, 8, 5, 64];k、v: [2, 2, 5, 64]

# 教学实现:显式重复共享的 K/V Head
repeats = Hq // Hkv
k_exp = k.repeat_interleave(repeats, dim=1)
v_exp = v.repeat_interleave(repeats, dim=1)
# k_exp、v_exp: [2, 8, 5, 64]

scores = (q @ k_exp.transpose(-2, -1)) / (d**0.5)
# scores: [2, 8, 5, 5]

# Causal Mask:屏蔽未来位置
future_mask = torch.ones(T, T, dtype=torch.bool, device=x.device).triu(1)
scores = scores.masked_fill(future_mask, float("-inf"))

attn = torch.softmax(scores, dim=-1)  # [2, 8, 5, 5]
out = attn @ v_exp  # [2, 8, 5, 64]
out = out.transpose(1, 2).reshape(B, T, Hq * d)
y = o_proj(out)  # [2, 5, 512]

这段实现中,真正投影并需要长期缓存的 K/V 仍然只有 2 个 Head;repeat_interleave 只是为了用普通批量矩阵乘法清楚展示共享关系。

生产级 GQA Kernel 通常不会像这个教学版本一样物化重复后的 k_exp 和 v_exp,否则会增加临时内存与带宽开销,削弱分组共享带来的收益。

十、训练和推理为什么不一样

训练:一次计算整段序列

Teacher Forcing 下,完整目标序列已知,只需用 Causal Mask 阻止未来信息泄漏。

因此一层可以并行计算:

Q,K,V:[B,H,N,Dh],S:[B,H,N,N].\begin{aligned} Q,K,V&:[B,H,N,D_h],\\ S&:[B,H,N,N]. \end{aligned}

虽然具有因果约束,但不是必须像 RNN 那样按 Token 逐步训练。

推理:一次通常只产生一个新 Token

已经生成 NN 个 Token 后,下一个 Decode Step 只产生一条新 Query、Key 和 Value:

qnew:[B,H,1,Dh],knew,vnew:[B,Hkv,1,Dh].\begin{aligned} q_{\mathrm{new}}&:[B,H,1,D_h],\\ k_{\mathrm{new}},v_{\mathrm{new}}&:[B,H_{kv},1,D_h]. \end{aligned}

历史 Token 的 K/V 不需要重复计算,因此保存为 KV Cache:

Kcache,Vcache:[B,Hkv,N,Dh]K_{\mathrm{cache}},V_{\mathrm{cache}} :[B,H_{kv},N,D_h]

新 Query 与整个 KcacheK_{\mathrm{cache}} 匹配,再从 VcacheV_{\mathrm{cache}} 聚合信息。

这带来两个事实:

  1. KV Cache 避免重复计算历史 K/V;
  2. Cache 会随序列长度线性增长,并在长上下文推理中占据大量显存和内存带宽。

因此 MQA、GQA、PagedAttention 和 KV Cache 量化主要属于推理系统优化,而不是重新定义 Attention 的基本公式。

十一、Attention 的复杂度应该怎样看

一句“Attention 是 O(N2)O(N^2)”并不完整。设模型维度为 DD:

步骤主要时间复杂度主要中间形状
Q/K/V 投影O(ND2)O(ND^2)[N,D][N,D]
Score 计算O(N2D)O(N^2D)[H,N,N][H,N,N]
SoftmaxO(HN2)O(HN^2)[H,N,N][H,N,N]
Value 聚合O(N2D)O(N^2D)[H,N,Dh][H,N,D_h]
输出投影O(ND2)O(ND^2)[N,D][N,D]

当 NN 较短、DD 很大时,线性投影和 FFN 也可能占据大量计算;当上下文很长时,N2N^2 项会越来越突出。

还要区分:

FlashAttention 主要优化显存 IO 与中间激活存储,并保持标准 Attention 结果;Linear Attention 则通过改变代数形式避免显式 N×NN\times N 关系矩阵。两者不能只因为“都更快”而归为一类。

十二、一个最小但正确的实现

下面的 PyTorch 风格函数展示单次 Scaled Dot-Product Attention:

import math
import torch


def scaled_dot_product_attention(q, k, v, allowed=None):
    """
    q: [B, H, Nq, Dh]
    k: [B, H, Nk, Dh]
    v: [B, H, Nk, Dv]
    allowed: 可广播到 [B, H, Nq, Nk] 的布尔张量
             True 表示允许读取,False 表示屏蔽
    """
    scores = q @ k.transpose(-2, -1)
    scores = scores / math.sqrt(q.size(-1))

    if allowed is not None:
        scores = scores.masked_fill(~allowed, float("-inf"))

    weights = torch.softmax(scores, dim=-1)
    output = weights @ v
    return output, weights

生产实现还需要处理:

实际项目应优先使用框架提供的优化实现,例如 PyTorch 的 scaled_dot_product_attention,让运行时根据设备和输入选择可用 Kernel。

十三、数值稳定性容易出错的地方

1. Softmax 前减去最大值

稳定实现会使用:

softmax⁡(x)i=exp⁡(xi−max⁡jxj)∑kexp⁡(xk−max⁡jxj)\operatorname{softmax}(x)_i =\frac{\exp(x_i-\max_j x_j)} {\sum_k\exp(x_k-\max_j x_j)}

减去同一个常数不改变 Softmax 结果,却能降低指数溢出风险。

2. Mask 的负数必须足够小

概念上使用 −∞-\infty。实际低精度 Kernel 可能使用数据类型可表示的极小值,但必须确保被屏蔽位置的指数贡献为零。

3. 避免整行全部被屏蔽

如果一整行都是 −∞-\infty,Softmax 会出现未定义的 0/00/0,可能产生 NaN。Padding、序列切分和特殊 Token 逻辑需要保证每个有效 Query 至少存在一个可见 Key,或者对全 Mask 行单独处理。

4. 不要把 Attention Weight 当概率预测

每行权重和为 1,只说明它是 Value 聚合系数。它不是词表概率,也不直接等于模型对事实的置信度。

十四、怎样阅读一个新的 Attention 变体

遇到新方法时,可以按下面顺序检查:

  1. Q、K、V 来自哪里?
  2. 每个 Query 可以读取哪些 Key?
  3. 是否仍显式构造 N×NN\times N Score?
  4. 相似度仍是 Softmax Dot-Product 吗?
  5. 归一化是否改变?
  6. 因果场景怎样维护历史状态?
  7. 优化的是 FLOPs、激活内存、KV Cache,还是显存 IO?
  8. 结果是精确等价、受控近似,还是新的 Attention 定义?

这组问题可以区分很多容易混淆的名称:

Sparse Attention:改变连接图
Linear Attention:改变代数形式或相似度族
FlashAttention:改变精确 Attention 的执行方式
MQA / GQA:改变 Query Head 与 KV Head 的组织
Cross-Attention:改变 Q 与 K/V 的来源

十五、小结

Self-Attention 的完整逻辑可以压缩成一句话:

每个 Query 先与允许访问的 Key 计算缩放匹配分数,再经过 Mask 和逐行 Softmax 得到权重,最后用这些权重聚合 Value。

真正需要记住的不是一条孤立公式,而是它的张量语义:

X [B,N,D]
  ↓ 可学习投影与多头 reshape
Q/K/V [B,H,N,Dh]
  ↓ Query–Key 匹配
Scores [B,H,N,N]
  ↓ Mask + row-wise Softmax
Weights [B,H,N,N]
  ↓ 加权聚合 Value
Heads [B,H,N,Dh]
  ↓ Concat + W_O
Output [B,N,D]

下一篇将继续回答:没有递归和卷积时,Transformer 怎样通过正弦位置编码、相对位置、RoPE 与 ALiBi 表示顺序和距离?

参考资料

  1. Vaswani et al., Attention Is All You Need, 2017.
  2. Rush, The Annotated Transformer, Harvard NLP.
  3. PyTorch, torch.nn.functional.scaled_dot_product_attention.
  4. Shazeer, Fast Transformer Decoding: One Write-Head is All You Need, 2019.
  5. Ainslie et al., GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints, 2023.
  6. Dao et al., FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness, 2022.
  7. Jain and Wallace, Attention is not Explanation, 2019.

编辑文章
分享这篇文章:

上一篇
大模型基础(一):Transformer 到底改变了什么?
下一篇
高效 Attention 变体:从稀疏、低秩到线性注意力