GOCLAWLLM ENGINEERING
GoClaw 首页

6. Transformer 逐部件讲透

核心原理4~6 小时
学习目标
  1. 从 Q/K/V 推导因果注意力
  2. 解释多头、位置编码、归一化与 MLP
  3. 用形状和断言验证 Decoder Block
前置知识
  • 矩阵乘法与概率
  • 语言模型训练样本

本章产物手写注意力输出,并完成与 PyTorch 实现的数值核对。

6.1 为什么从 RNN 走向 Transformer

RNN 按时间顺序更新隐藏状态,第 t 步依赖第 t-1 步,训练难以在序列维度充分并行。长距离信息还要经过很多递归步骤。

Self-Attention 让任意两个位置直接建立联系;训练时,同一层所有位置可以并行计算。代价是标准注意力矩阵随序列长度 T 呈 增长。

6.2 Q、K、V 的直觉

把每个 token 想成图书馆中的一张卡片:

它们来自同一输入 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) V

6.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

减少 K/V 头能显著减小 KV Cache 和解码内存带宽。GQA 是质量与效率的折中。假设 Q 有 32 头、KV 有 8 头,则每 4 个 Query 头共享一组 K/V。

6.7 位置编码

Attention 本身若不加入位置信息,对 token 顺序不敏感。常见方案:

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,按以下顺序执行:

  1. 构造形状为 [B, H, T, D] 的 Q、K、V。
  2. 计算 Q @ K.transpose(-2, -1) / sqrt(D)
  3. 在 softmax 之前把未来位置填成负无穷。
  4. 检查每行概率和为 1,未来位置概率为 0。
  5. 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)

做两个故意失败的实验:

验收标准:能从张量形状定位错误发生在 head 拆分、矩阵转置、mask 广播还是输出拼接,而不是通过随机修改维度让代码“碰巧运行”。


本章依据

原理性结论以原始论文、官方文档或公开教材为依据。论文中的实验结果只适用于其声明的模型、数据、硬件和评估设置。

  1. Scaled dot-product attention、多头注意力、位置编码和 Transformer 架构。

  2. RoPE 的旋转表示及相对位置信息进入注意力点积的方式。

  3. RMSNorm 的定义、缩放不变性和与 LayerNorm 的差异。

  4. SwiGLU 等门控前馈网络变体。

  5. MHA、MQA 与 GQA 的 KV Head 组织和质量—效率折中。