6. Transformer 逐部件讲透
- 从 Q/K/V 推导因果注意力
- 解释多头、位置编码、归一化与 MLP
- 用形状和断言验证 Decoder Block
- 矩阵乘法与概率
- 语言模型训练样本
本章产物手写注意力输出,并完成与 PyTorch 实现的数值核对。
6.1 为什么从 RNN 走向 Transformer
RNN 按时间顺序更新隐藏状态,第 t 步依赖第 t-1 步,训练难以在序列维度充分并行。长距离信息还要经过很多递归步骤。
Self-Attention 让任意两个位置直接建立联系;训练时,同一层所有位置可以并行计算。代价是标准注意力矩阵随序列长度 T 呈 T² 增长。
6.2 Q、K、V 的直觉
把每个 token 想成图书馆中的一张卡片:
- Query:我当前想找什么?
- Key:我能被什么问题匹配?
- Value:如果选中我,要读取什么内容?
它们来自同一输入 X,但经过不同可学习矩阵:
Q = X W_Q
K = X W_K
V = X W_V注意力:
Attention(Q,K,V)
= softmax((Q K^T) / sqrt(d_k) + Mask) V6.3 用形状理解注意力
单头情况下:
Q: [T, D]
K: [T, D]
K^T: [D, T]
Q K^T: [T, T][i,j] 表示第 i 个位置对第 j 个位置的匹配分数。softmax 通常沿最后一维进行,使每个查询位置对所有可见 Key 的权重和为 1。
最后:
[T,T] @ [T,D] → [T,D]每个位置得到所有可见 Value 的加权和。
运行:
python code/02_attention_from_scratch.py输出矩阵上三角应为 0,因为未来位置被因果遮罩。
6.4 因果遮罩为什么必须存在
预训练把整段 target 一次送入模型以便并行计算。如果没有 mask,位置 t 能看到 t+1,而 t+1 正是它应该预测的答案,训练 loss 会虚假地迅速下降。
因果 mask 概念上是:
位置 0: 可看 [0]
位置 1: 可看 [0,1]
位置 2: 可看 [0,1,2]
位置 3: 可看 [0,1,2,3]未来分数加上负无穷,softmax 后概率为 0。
配套实验:从公式实现因果注意力 Notebook。实验逐项验证缩放、遮罩、概率归一化,并与 PyTorch 标准实现核对。
6.5 多头注意力
把 C 维拆成 H 个 D 维头:
D = C / H多个头允许模型在不同投影子空间建立不同关系。不要把它机械解释成“某个头一定负责语法”;头可能冗余、混合或随层次变化。
多头输出重新拼接:
[B,H,T,D] → [B,T,H×D] = [B,T,C]再经过输出投影 W_O。
6.6 MHA、MQA 和 GQA
- MHA:每个 Query 头都有独立 K/V 头。
- MQA:所有 Query 头共享一组 K/V。
- GQA:多个 Query 头分组共享 K/V。
减少 K/V 头能显著减小 KV Cache 和解码内存带宽。GQA 是质量与效率的折中。假设 Q 有 32 头、KV 有 8 头,则每 4 个 Query 头共享一组 K/V。
6.7 位置编码
Attention 本身若不加入位置信息,对 token 顺序不敏感。常见方案:
- 绝对位置 embedding:位置 0、1、2 各自有可学习向量。
- 正弦位置编码:使用不同频率的 sin/cos。
- RoPE:对 Q/K 的二维分量按位置旋转,使相对距离进入点积关系。
- ALiBi:在注意力分数上加入与距离相关的偏置。
RoPE 的直觉是“同一个内容向量在不同位置有不同旋转角度”;两个位置的 Q/K 点积会自然包含角度差,即相对位置。扩展上下文时不能只把最大长度数字改大,还需要考虑模型训练过的距离分布和 RoPE scaling。
6.8 前馈网络:每个 token 自己思考
Attention 负责位置之间交换信息;MLP 对每个位置独立应用相同非线性网络。
经典形式:
MLP(x) = W_2 GELU(W_1 x)许多现代 LLM 使用 SwiGLU:
SwiGLU(x)
= W_down( SiLU(W_gate x) ⊙ (W_up x) )门控分支决定哪些特征通过。MLP 参数通常占每层很大比例。
6.9 残差连接
y = x + sublayer(x)残差提供一条不经过复杂变换的直接路径:
- 信息可以原样向深层传播。
- 梯度有更直接的回传通道。
- 子层只需学习“在原表示上改多少”。
6.10 LayerNorm 与 RMSNorm
LayerNorm 对一个 token 的隐藏维度计算均值和方差,再做缩放和平移。RMSNorm 主要使用均方根缩放,不减均值,计算更简洁。
两者目的不是让所有 token 变成相同,而是稳定每个位置激活的数值尺度。
现代 decoder 常使用 Pre-Norm:
x = x + Attention(Norm(x))
x = x + MLP(Norm(x))相较原始 Post-Norm,Pre-Norm 往往更利于训练深层网络。
6.11 一个完整 Decoder Block
输入 X
│
├───────────────┐
▼ │
Norm → Causal Attention
│ │
└──── 相加 ◀────┘
│
├───────────────┐
▼ │
Norm → MLP │
│ │
└──── 相加 ◀────┘
│
输出堆叠 L 层后,最终 Norm 和 LM Head 产生对词表的 logits。
6.12 为什么 Attention 是 O(T²)
QK 矩阵形状是 [T,T]。序列从 8K 增加到 32K:
长度扩大 4 倍
注意力矩阵元素数扩大 16 倍Flash Attention 并不改变精确注意力的数学定义,而是通过分块、融合和减少中间矩阵的高成本读写来降低实际内存与时间。长上下文还会受 KV Cache、位置外推和训练数据长度分布制约。
6.13 动手:用断言验证注意力,而不是只看热力图
实验 03|Causal Attention 资源:CPU;时间:约 2 分钟;产物:形状记录、注意力矩阵与三项正确性断言。
运行独立实现:
python code/02_attention_from_scratch.py然后打开因果注意力 Notebook,按以下顺序执行:
- 构造形状为
[B, H, T, D]的 Q、K、V。 - 计算
Q @ K.transpose(-2, -1) / sqrt(D)。 - 在 softmax 之前把未来位置填成负无穷。
- 检查每行概率和为 1,未来位置概率为 0。
- 与
torch.nn.functional.scaled_dot_product_attention的输出比较。
必须保留的断言:
assert scores.shape == (batch, heads, sequence, sequence)
assert torch.allclose(weights.sum(dim=-1), torch.ones_like(weights[..., 0]))
assert torch.allclose(manual_output, torch_output, atol=1e-6)做两个故意失败的实验:
- 去掉
sqrt(D):随着 head dimension 增大,记录 softmax 最大概率与梯度变化。 - 去掉 causal mask:观察训练 loss 可能异常变好,但模型已经读取未来标签,属于信息泄漏。
验收标准:能从张量形状定位错误发生在 head 拆分、矩阵转置、mask 广播还是输出拼接,而不是通过随机修改维度让代码“碰巧运行”。
本章依据
原理性结论以原始论文、官方文档或公开教材为依据。论文中的实验结果只适用于其声明的模型、数据、硬件和评估设置。
Scaled dot-product attention、多头注意力、位置编码和 Transformer 架构。
RoPE 的旋转表示及相对位置信息进入注意力点积的方式。
RMSNorm 的定义、缩放不变性和与 LayerNorm 的差异。
SwiGLU 等门控前馈网络变体。
MHA、MQA 与 GQA 的 KV Head 组织和质量—效率折中。