先说结论
Transformer 最重要的变化,不是简单地“加入注意力”,而是把序列中不同位置之间的信息交换写成一次可并行的矩阵计算。它解决了循环网络必须逐步传递隐状态的问题,却没有消除序列长度带来的成本:训练时注意力矩阵通常随长度平方增长,自回归推理时则需要逐 token 生成并维护 KV Cache。
因此,理解 Transformer 至少要同时回答三个问题:
- 一个 token 如何从其他位置读取信息?
- 训练时的并行计算为什么到了生成时又变成串行?
- 模型结构、学习目标和最终可靠性之间是什么关系?
一层 Transformer 在做什么
以下采用常见的 Pre-Norm 结构。设输入为 (X\in\mathbb{R}^{B\times T\times d}),其中 (B) 是批量大小,(T) 是序列长度,(d) 是隐藏维度:
这两行包含四种不同作用:
- LayerNorm 控制每个 token 表示的数值尺度;
- 多头注意力 在 token 之间交换信息;
- MLP 对每个位置独立地做非线性特征变换;
- 残差连接 保留原表示,并为深层网络提供更稳定的梯度路径。
注意力负责“位置之间”的混合,MLP 负责“特征维度内”的变换。把两者都称为注意力,会掩盖它们完全不同的计算角色。
从张量形状理解自注意力
将输入分别投影为 Query、Key 和 Value:
对单个注意力头,设头维度为 (d_h)。加入掩码 (M) 后:
| 张量 | 典型形状 | 含义 |
|---|---|---|
| (Q,K,V) | (B\times h\times T\times d_h) | 每个头的查询、索引和值 |
| (S) | (B\times h\times T\times T) | 每个查询位置对所有键位置的原始分数 |
| (A) | (B\times h\times T\times T) | 归一化后的读取权重 |
| (O) | (B\times h\times T\times d_h) | 从上下文聚合后的结果 |
这里的“相关”并不等于因果关系,也不自动等于人类可解释的重要性。注意力权重只是当前参数化计算中的中间量。
为什么除以 (\sqrt{d_h})
假设 (q_i) 和 (k_i) 独立、均值为零、方差约为一,则点积 (\sum_{i=1}^{d_h}q_i k_i) 的方差会随 (d_h) 增长。维度越大,未经缩放的分数绝对值越容易变大,softmax 更容易过早进入接近 one-hot 的区域,使梯度变得尖锐而不稳定。
除以 (\sqrt{d_h}) 后,分数的典型尺度不再随头维度快速增长。这不是为了“让结果更小”这么简单,而是为了让不同维度设置下的 softmax 处于相近的数值区间。
掩码决定模型能看到什么
在编码器式双向注意力中,一个位置通常可以读取同一序列的所有非 padding 位置。自回归语言模型则使用因果掩码:
softmax 后,未来位置的权重为零。训练时虽然整段序列一次进入模型,但第 (i) 个位置仍只能利用 (1,\ldots,i) 的信息;这就是 teacher forcing 能并行、又不泄露未来 token 的原因。
多头、MQA 与 GQA
标准多头注意力为每个头分别计算 Q、K、V。多个头不是多个独立模型,而是让同一层在不同投影子空间中读取信息:
推理时,K 和 V 需要被缓存。为了减少缓存:
- MHA:每个 Query 头都有独立的 K/V 头;
- MQA:所有 Query 头共享一组 K/V;
- GQA:若干 Query 头共享一组 K/V,是效果与内存之间的折中。
因此,估算 KV Cache 时不能只知道总参数量;还需要层数、头维度和 K/V 头数。
位置从哪里来
纯注意力对输入排列本身没有方向感,模型需要额外的位置机制。常见做法包括绝对位置嵌入和旋转位置编码(RoPE)。它们的共同目的,是让注意力分数能够区分“内容相同但位置不同”的 token。
位置机制并不保证模型能无限外推到更长上下文。训练长度、缩放方式、注意力模式和数据分布都会影响长度外推;“支持 128K 输入”也不等于模型能在整个窗口内同等可靠地利用信息。
训练并行,生成仍然串行
训练语言模型时,已知目标序列,所有位置的损失可以同时计算:
生成时,第 (t+1) 个 token 的输入依赖第 (t) 个 token 的实际采样结果,所以无法跨生成步完全并行。KV Cache 避免重复计算历史 token 的 Key 和 Value,但不能消除逐步生成。
设层数为 (L)、批量为 (B)、已缓存长度为 (T)、K/V 头数为 (h_{kv})、头维度为 (d_h)、每个元素占 (b) 字节,则 KV Cache 近似为:
前面的 2 分别对应 K 和 V。该公式解释了为什么长上下文、大批量和标准 MHA 会迅速增加推理内存,也解释了 GQA/MQA 的工程价值。
计算复杂度与真正的瓶颈
| 部分 | 主要复杂度 | 随序列长度的变化 |
|---|---|---|
| Q/K/V 与输出投影 | (O(Td^2)) | 线性 |
| 注意力分数与加权 | (O(T^2d)) | 平方 |
| MLP | (O(Td\,d_{ff})) | 线性 |
| KV Cache | (O(TLh_{kv}d_h)) | 线性 |
短序列、大隐藏维度时,线性层和 MLP 可能占主要计算;序列很长时,标准注意力的平方项才会越来越突出。不能脱离模型宽度、硬件和批量,仅凭 (O(T^2)) 判断实际速度。
模型结构不等于学习目标
Transformer 决定了信息如何流动,但没有决定模型最终优化什么。交叉熵倾向于拟合平均 token 预测;强化学习后训练又会受奖励定义、采样策略和优势估计影响。
如果训练目标只强调常见、高概率路径,困难样本和低概率但有信息的轨迹可能贡献不足。反过来,强调尾部或困难区域也可能增加方差和训练不稳定。可靠的改进需要同时说明:
- 优化的是平均、分位数、尾部风险还是其他分布函数;
- 训练信号如何分配给不同难度的样本;
- 提升来自更好的搜索、更多采样,还是能力边界真正变化;
- 结论是否在分布变化和独立测试上仍成立。
常见误区
- “注意力就是模型的解释。”
注意力权重是计算路径的一部分,不是完整的因果解释。 - “训练能并行,所以生成也能并行。”
训练已知整段目标,生成必须等待前一步的实际输出。 - “参数量决定全部显存。”
推理还受 KV Cache、临时张量和框架开销影响;训练还包含梯度、优化器状态和激活值。 - “上下文窗口越长,模型记得越好。”
可接收、可检索和可可靠利用是三个不同层次。 - “模拟基准提升就等于真实任务更可靠。”
需要检查数据分布、评估口径、真实约束和样本外表现。
阅读实现时的检查顺序
面对一段 Transformer 代码,可以按以下顺序核对:
- 输入、Q/K/V 和输出的形状;
- mask 的广播方向以及是否真的阻断未来位置;
num_heads、num_kv_heads与head_dim的关系;- RoPE 或其他位置机制应用在 Q/K 的哪个阶段;
- KV Cache 的追加维度和生命周期;
- 残差、LayerNorm 是 Pre-Norm 还是 Post-Norm;
- 训练损失、采样方法与最终评估指标是否一致。
这套顺序比只看模块名称更可靠,因为许多实现错误都发生在形状、广播、缓存和目标口径之间。