chapter 06 / transformer · 全课程最重要的一章 · 预计学习时间 180-240 分钟
第 5 章结尾我们把注意力重述为「可微分软检索」:query 找 key,按相关度取 value。2017 年的《Attention Is All You Need》做了三步看似简单的改造,然后把 RNN 整个扔掉了:
收益是结构性的:RNN 中位置 1 的信息要走 $n$ 步才能影响位置 $n$(第 5 章的指数衰减就发生在路上),self-attention 中任意两位置一步直达;RNN 必须串行等 $\mathbf{h}_{t-1}$,self-attention 的所有位置同时计算。代价是 $O(n^2)$——第 5 章思考题的答案,第 9 章 KV cache 的起点。
那个 $\sqrt{d_k}$ 是论文里唯一的「魔法常数」,值得把推导走完。设 $\mathbf{q}, \mathbf{k}$ 的各分量独立、均值 0、方差 1,点积 $\mathbf{q}^\top\mathbf{k} = \sum_{i=1}^{d_k} q_i k_i$ 是 $d_k$ 个独立项之和:
$$\mathbb{E}[\mathbf{q}^\top\mathbf{k}] = 0, \qquad \text{Var}(\mathbf{q}^\top\mathbf{k}) = \sum_{i=1}^{d_k} \text{Var}(q_i k_i) = d_k$$即点积的典型幅度是 $\sqrt{d_k}$。$d_k = 128$ 时分数动辄 ±11,softmax 里 $e^{11}$ 与 $e^{-11}$ 相差 $10^9$ 倍——输出饱和成 one-hot。回忆第 2 章:softmax 饱和处梯度趋零,注意力学不动。除以 $\sqrt{d_k}$ 把方差归一回 1,softmax 工作在梯度健康的区间。整个 Transformer 的设计哲学缩影:每个组件都为「梯度能顺畅流动」服务。
语言模型的训练目标(第 7 章主角)是预测下一个词,那么位置 $i$ 绝不能看见位置 $j > i$——否则就是抄答案。实现得粗暴而高效:在 softmax 之前把上三角全部设为 $-\infty$,softmax 后这些位置权重精确为 0。
5 个 token、$d{=}4$ 的玩具尺寸,但计算是真的(矩阵就是真 $QK^\top$)。必做实验:① 按「下一步」走完五个阶段,每一步对照 §2 的公式;② 在阶段③④来回切换「因果掩码」开关——看右上三角从 −∞ 到有值、softmax 行分布如何重排(关掉掩码就是 BERT 式双向注意力);③ 在阶段④数一数:第一行为什么永远是 1.00?(「猫」只能看自己)。
真实计算:E(5×4) → Q,K,V → 五个阶段。绿色深浅=数值热度,红色=被掩码。
一次注意力只能学一种「相关性」——但「垫子」需要同时关注语法主语(猫)、空间关系(上)、修饰结构。解法:把 $d$ 维空间切成 $h$ 份,每份独立做注意力,最后拼接再投影:
$$\text{head}_i = \text{Attention}(XW_Q^{(i)},\, XW_K^{(i)},\, XW_V^{(i)}), \qquad \text{MHA}(X) = \text{Concat}(\text{head}_1, \dots, \text{head}_h)\, W_O$$每个头在 $d_k = d/h$ 的子空间工作(GPT-2:$d{=}768, h{=}12, d_k{=}64$),总计算量与单头相同——免费的多视角。可解释性研究(Anthropic 的工作,附录信息源有链接)确实找到了分工明确的头:专盯前一个词的、追踪括号配对的、把宾语信息搬给动词的「归纳头」……当然也有大量冗余头(可被剪枝)。
纯注意力有个致命盲点:它是集合运算——把输入顺序打乱,输出跟着同样打乱(置换等变),「猫坐垫子」和「垫子坐猫」在它眼里没有区别。必须人为注入位置信息:
注意力负责搬运信息(token 之间),还需要一个组件加工信息(token 内部):前馈网络 FFN,对每个位置独立地做两层 MLP,中间扩到 $4d$:
$$\text{FFN}(\mathbf{x}) = W_2\, \text{GELU}(W_1 \mathbf{x}), \qquad W_1: d \to 4d, \quad W_2: 4d \to d$$它占了 Transformer 约 2/3 的参数,可解释性研究视其为模型的「键值存储」——事实知识主要存在这里。现代 LLM 用 SwiGLU 变体(第 3 章激活函数表的最后一行)。把全部零件装进一个块(Pre-LN 结构,先归一化再进子层——比原论文的 Post-LN 训练稳定得多,深层不需要 warmup 魔法):
x = x + MHA(LayerNorm(x)) # 子层1:注意力 + 残差
x = x + FFN(LayerNorm(x)) # 子层2:前馈 + 残差
两条残差连接正是第 4 章的 $y = F(x) + x$——GPT 就是一摞(修了梯度高速公路的)残差块,加上底部的 token embedding 和顶部的 lm_head(GPT-2 里两者权重共享,省一半嵌入参数)。整体架构在你眼前完整了:
tokens → Embedding(+位置) → [GPT Block] × L → LayerNorm → lm_head → 下一词概率
读懂任何 LLM 配置表的基本功。每个块:注意力 $W_Q, W_K, W_V, W_O$ 各 $d^2$ → $4d^2$;FFN 两个矩阵 $d \times 4d$ → $8d^2$。合计每块 $12d^2$,于是著名的速算公式:
$$\text{非嵌入参数} \approx 12\, L\, d^2$$验证 GPT-2 small($L{=}12, d{=}768$,词表 50257):块参数 $12 \times 12 \times 768^2 \approx 85\text{M}$,嵌入 $50257 \times 768 \approx 38.6\text{M}$(与 lm_head 共享),合计 $\approx 124\text{M}$ ✓。同一个公式套 LLaMA-2 7B($L{=}32, d{=}4096$):$12 \times 32 \times 4096^2 \approx 6.4\text{B}$,加嵌入正是 ~7B。以后看到任何「xB 模型」,你都能反推它的骨架。
import torch, torch.nn as nn, torch.nn.functional as F
class CausalSelfAttention(nn.Module):
def __init__(self, d, h, n_max=1024):
super().__init__()
self.h, self.dk = h, d // h
self.qkv = nn.Linear(d, 3 * d) # Q,K,V 一次投影
self.proj = nn.Linear(d, d) # W_O
mask = torch.tril(torch.ones(n_max, n_max))
self.register_buffer('mask', mask) # 下三角=可见
def forward(self, x):
B, N, d = x.shape
q, k, v = self.qkv(x).chunk(3, dim=-1)
# 拆多头:(B, N, d) → (B, h, N, dk)
q, k, v = (t.view(B, N, self.h, self.dk).transpose(1, 2) for t in (q, k, v))
att = (q @ k.transpose(-2, -1)) / self.dk ** 0.5 # ② 缩放点积
att = att.masked_fill(self.mask[:N, :N] == 0, float('-inf')) # ③ 因果掩码
att = F.softmax(att, dim=-1) # ④ 注意力权重
y = att @ v # ⑤ 加权搬运
y = y.transpose(1, 2).contiguous().view(B, N, d) # 拼回 (B, N, d)
return self.proj(y)
class Block(nn.Module):
"""一个完整的 GPT 块:你正在读的所有 LLM 都是它的 L 次重复"""
def __init__(self, d=768, h=12):
super().__init__()
self.ln1, self.ln2 = nn.LayerNorm(d), nn.LayerNorm(d)
self.attn = CausalSelfAttention(d, h)
self.ffn = nn.Sequential(
nn.Linear(d, 4 * d), nn.GELU(), nn.Linear(4 * d, d))
def forward(self, x):
x = x + self.attn(self.ln1(x)) # Pre-LN + 残差(第 4 章的高速公路)
x = x + self.ffn(self.ln2(x))
return x
x = torch.randn(2, 10, 768) # batch=2, 10 个 token
print(Block()(x).shape) # → (2, 10, 768)
print(sum(p.numel() for p in Block().parameters()) / 1e6, 'M') # ≈ 7.1M ≈ 12d²/1e6 ✓
压轴大餐:Karpathy 用 3.5 小时从空文件写出完整可训练的 GPT(含数据加载、训练循环、采样生成)。这是全课程唯一标注「必须跟写」的视频——写完它,你对 LLM 的理解将从「读过推导」变成「亲手造过」: