← 返回首页

先说结论

Transformer 最重要的变化,不是简单地“加入注意力”,而是把序列中不同位置之间的信息交换写成一次可并行的矩阵计算。它解决了循环网络必须逐步传递隐状态的问题,却没有消除序列长度带来的成本:训练时注意力矩阵通常随长度平方增长,自回归推理时则需要逐 token 生成并维护 KV Cache。

因此,理解 Transformer 至少要同时回答三个问题:

  1. 一个 token 如何从其他位置读取信息?
  2. 训练时的并行计算为什么到了生成时又变成串行?
  3. 模型结构、学习目标和最终可靠性之间是什么关系?

一层 Transformer 在做什么

以下采用常见的 Pre-Norm 结构。设输入为 (X\in\mathbb{R}^{B\times T\times d}),其中 (B) 是批量大小,(T) 是序列长度,(d) 是隐藏维度:

\[ \begin{aligned} H &= X + \operatorname{MHA}(\operatorname{LN}(X)),\\ Y &= H + \operatorname{MLP}(\operatorname{LN}(H)). \end{aligned} \]

这两行包含四种不同作用:

注意力负责“位置之间”的混合,MLP 负责“特征维度内”的变换。把两者都称为注意力,会掩盖它们完全不同的计算角色。

从张量形状理解自注意力

将输入分别投影为 Query、Key 和 Value:

\[ Q=XW_Q,\qquad K=XW_K,\qquad V=XW_V. \]

对单个注意力头,设头维度为 (d_h)。加入掩码 (M) 后:

\[ S=\frac{QK^\top}{\sqrt{d_h}}+M,\qquad A=\operatorname{softmax}(S),\qquad O=AV. \]
张量 典型形状 含义
(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 位置。自回归语言模型则使用因果掩码:

\[ M_{ij}= \begin{cases} 0,& j\le i,\\ -\infty,& j>i. \end{cases} \]

softmax 后,未来位置的权重为零。训练时虽然整段序列一次进入模型,但第 (i) 个位置仍只能利用 (1,\ldots,i) 的信息;这就是 teacher forcing 能并行、又不泄露未来 token 的原因。

多头、MQA 与 GQA

标准多头注意力为每个头分别计算 Q、K、V。多个头不是多个独立模型,而是让同一层在不同投影子空间中读取信息:

\[ \operatorname{MHA}(X)= \operatorname{Concat}(\operatorname{head}_1,\ldots,\operatorname{head}_h)W_O. \]

推理时,K 和 V 需要被缓存。为了减少缓存:

因此,估算 KV Cache 时不能只知道总参数量;还需要层数、头维度和 K/V 头数。

位置从哪里来

纯注意力对输入排列本身没有方向感,模型需要额外的位置机制。常见做法包括绝对位置嵌入和旋转位置编码(RoPE)。它们的共同目的,是让注意力分数能够区分“内容相同但位置不同”的 token。

位置机制并不保证模型能无限外推到更长上下文。训练长度、缩放方式、注意力模式和数据分布都会影响长度外推;“支持 128K 输入”也不等于模型能在整个窗口内同等可靠地利用信息。

训练并行,生成仍然串行

训练语言模型时,已知目标序列,所有位置的损失可以同时计算:

\[ \mathcal{L}(\theta) =-\sum_{t=1}^{T}\log p_\theta(x_t\mid x_{<t}). \]

生成时,第 (t+1) 个 token 的输入依赖第 (t) 个 token 的实际采样结果,所以无法跨生成步完全并行。KV Cache 避免重复计算历史 token 的 Key 和 Value,但不能消除逐步生成。

设层数为 (L)、批量为 (B)、已缓存长度为 (T)、K/V 头数为 (h_{kv})、头维度为 (d_h)、每个元素占 (b) 字节,则 KV Cache 近似为:

\[ M_{\mathrm{KV}} =2LBT\,h_{kv}d_hb. \]

前面的 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 预测;强化学习后训练又会受奖励定义、采样策略和优势估计影响。

如果训练目标只强调常见、高概率路径,困难样本和低概率但有信息的轨迹可能贡献不足。反过来,强调尾部或困难区域也可能增加方差和训练不稳定。可靠的改进需要同时说明:

常见误区

  1. “注意力就是模型的解释。”
    注意力权重是计算路径的一部分,不是完整的因果解释。
  2. “训练能并行,所以生成也能并行。”
    训练已知整段目标,生成必须等待前一步的实际输出。
  3. “参数量决定全部显存。”
    推理还受 KV Cache、临时张量和框架开销影响;训练还包含梯度、优化器状态和激活值。
  4. “上下文窗口越长,模型记得越好。”
    可接收、可检索和可可靠利用是三个不同层次。
  5. “模拟基准提升就等于真实任务更可靠。”
    需要检查数据分布、评估口径、真实约束和样本外表现。

阅读实现时的检查顺序

面对一段 Transformer 代码,可以按以下顺序核对:

  1. 输入、Q/K/V 和输出的形状;
  2. mask 的广播方向以及是否真的阻断未来位置;
  3. num_headsnum_kv_headshead_dim 的关系;
  4. RoPE 或其他位置机制应用在 Q/K 的哪个阶段;
  5. KV Cache 的追加维度和生命周期;
  6. 残差、LayerNorm 是 Pre-Norm 还是 Post-Norm;
  7. 训练损失、采样方法与最终评估指标是否一致。

这套顺序比只看模块名称更可靠,因为许多实现错误都发生在形状、广播、缓存和目标口径之间。