Transformer 最小闭环
把注意力、FFN、残差、归一化和位置编码组成真正可训练的 decoder-only Transformer。
- 建议时长
- 9–12 小时
- 难度
- 核心实现
- 课程位置
- 3 / 12
学完这一课,你应该能:
- 能解释 Transformer block 每个部件的职责
- 能比较 Pre-Norm 与 Post-Norm 的梯度路径
- 能实现可训练的 decoder-only Tiny Transformer
- 能通过小批次过拟合与消融定位实现错误
1. 一个 Decoder Block 包含什么
现代 decoder-only 模型的基本块可抽象为:
x = x + Attention(Norm(x))
x = x + FFN(Norm(x))
这是 Pre-Norm 写法。注意力负责 token mixing,FFN 负责 channel mixing:
| 部件 | 职责 | 去掉后的典型现象 |
|---|---|---|
| causal self-attention | 在可见上下文中聚合信息 | token 无法交互 |
| FFN | 每个位置独立做非线性变换 | 容量明显下降 |
| residual | 保留恒等路径,改善梯度传播 | 深层训练困难 |
| normalization | 控制特征尺度与优化条件 | 激活和梯度不稳定 |
| position | 表达顺序与距离 | 难以区分排列 |
2. Pre-Norm 与 Post-Norm
原始 Transformer 常用 Post-Norm:
x = Norm(x + Attention(x))
x = Norm(x + FFN(x))
Pre-Norm 把 Norm 放在子层前,残差主干有更直接的恒等梯度路径,通常更容易训练深层网络。但“更稳定”不是绝对优劣。公平比较要固定参数量、初始化、学习率、warmup、深度、batch 与精度。
3. FFN 为什么扩张再压缩
FFN(x) = W₂ · activation(W₁x + b₁) + b₂
若模型维度为 D,中间宽度常为约 4D。FFN 对每个 token 使用同一组参数,不混合序列位置,却通过大矩阵承担大量参数和计算。现代模型常用 SwiGLU:
SwiGLU(x) = [SiLU(xW_gate) ⊙ (xW_up)] W_down
替换激活函数时必须对齐总参数量;否则“效果提升”可能只是隐藏层更宽。
4. 位置编码决定顺序如何进入模型
没有位置信息,自注意力对输入排列具有置换等变性。常见方案:
- 学习式绝对位置:实现简单,但最大位置固定。
- 正弦位置:无需训练,可计算任意位置,外推质量不保证。
- 相对位置 bias:直接修改注意力 logits。
- RoPE:旋转 Q/K,让点积带相对位置信息,现代 LLM 常用。
第一版 Tiny Transformer 先用学习式位置 embedding,减少变量。基线正确后再替换 RoPE。
5. 实现 Decoder Block
import torch
from torch import nn
from torch.nn import functional as F
class CausalSelfAttention(nn.Module):
def __init__(self, d_model, n_heads, dropout=0.0):
super().__init__()
assert d_model % n_heads == 0
self.n_heads = n_heads
self.d_head = d_model // n_heads
self.qkv = nn.Linear(d_model, 3 * d_model, bias=False)
self.out = nn.Linear(d_model, d_model, bias=False)
self.dropout = dropout
def forward(self, x):
b, t, d = x.shape
q, k, v = self.qkv(x).chunk(3, dim=-1)
def heads(z):
return z.view(b, t, self.n_heads, self.d_head).transpose(1, 2)
q, k, v = map(heads, (q, k, v))
y = F.scaled_dot_product_attention(
q, k, v,
dropout_p=self.dropout if self.training else 0.0,
is_causal=True,
)
y = y.transpose(1, 2).contiguous().view(b, t, d)
return self.out(y)
class MLP(nn.Module):
def __init__(self, d_model, ratio=4, dropout=0.0):
super().__init__()
hidden = ratio * d_model
self.net = nn.Sequential(
nn.Linear(d_model, hidden),
nn.GELU(),
nn.Linear(hidden, d_model),
nn.Dropout(dropout),
)
def forward(self, x):
return self.net(x)
class Block(nn.Module):
def __init__(self, d_model, n_heads, dropout=0.0):
super().__init__()
self.norm1 = nn.LayerNorm(d_model)
self.attn = CausalSelfAttention(d_model, n_heads, dropout)
self.norm2 = nn.LayerNorm(d_model)
self.mlp = MLP(d_model, dropout=dropout)
def forward(self, x):
x = x + self.attn(self.norm1(x))
x = x + self.mlp(self.norm2(x))
return x
函数式 SDPA 不会自动读取模块的 training 状态,所以 eval 时必须显式把 dropout_p 设为 0。
6. 组成语言模型
class TinyTransformerLM(nn.Module):
def __init__(self, vocab_size, max_seq, d_model=128, n_heads=4, n_layers=4):
super().__init__()
self.max_seq = max_seq
self.token_emb = nn.Embedding(vocab_size, d_model)
self.pos_emb = nn.Embedding(max_seq, d_model)
self.blocks = nn.ModuleList([
Block(d_model, n_heads, dropout=0.1)
for _ in range(n_layers)
])
self.norm = nn.LayerNorm(d_model)
self.lm_head = nn.Linear(d_model, vocab_size, bias=False)
self.lm_head.weight = self.token_emb.weight # weight tying
def forward(self, input_ids, labels=None):
b, t = input_ids.shape
if t > self.max_seq:
raise ValueError("sequence too long")
pos = torch.arange(t, device=input_ids.device)
x = self.token_emb(input_ids) + self.pos_emb(pos)[None]
for block in self.blocks:
x = block(x)
logits = self.lm_head(self.norm(x))
loss = None
if labels is not None:
loss = F.cross_entropy(
logits.reshape(-1, logits.size(-1)),
labels.reshape(-1),
ignore_index=-100,
)
return logits, loss
语言模型标签必须错位:
input : [BOS, 我, 喜, 欢, AI]
label : [我, 喜, 欢, AI, EOS]
输入和标签若未 shift,模型只需复制当前 token,就会得到虚假的低 Loss。
7. 参数预算与训练观测
每层粗略参数:
Attention QKV + output ≈ 4D²
FFN(4D hidden)≈ 8D²
每层总计 ≈ 12D²
词表很大时,embedding/lm_head 也占大量参数。Weight tying 共享输入 embedding 与输出分类矩阵,减少参数。
训练至少记录:
- train/eval loss 与 perplexity
- learning rate 与 global gradient norm
- 各层 activation RMS
- tokens/s 与峰值显存
只看平均 Loss 无法发现某一层已输出 NaN,或数据加载让 GPU 长时间空等。
8. 最强的正确性测试:过拟合小批次
固定 8–32 条短序列,关闭 dropout,重复训练同一批。正确的小模型应能把 Loss 压得非常低。失败时按顺序检查:
- 标签是否 shift。
- causal mask 是否正确。
- padding label 是否为
-100。 - optimizer 是否包含所有参数。
train()/eval()状态。- 学习率是否极端。
小批次过拟合先证明模型“有能力记住”,再扩到完整语料。它能把实现错误与数据噪声分开。
9. 三个有价值的消融
去掉残差
比较 2 层与 8 层,记录每层 gradient norm。深度增加后,无残差版本通常更难优化。
Pre-Norm 改为 Post-Norm
固定参数量与数据,扫描学习率,比较早期 loss、梯度波动与最终验证集。
去掉位置 embedding
构造 token 集合相同但顺序不同的序列,例如 A B C D 与 D C B A,直接检测顺序区分能力。
10. 本课交付物
tiny-transformer/
├── model.py
├── train.py
├── generate.py
├── tests/
│ ├── test_attention.py
│ ├── test_causal_mask.py
│ └── test_overfit_batch.py
├── configs/tiny.yaml
└── report.md
报告必须回答:参数量如何分布、小批次能否过拟合、消融发生了什么、理论复杂度与 tokens/s 是否一致、最先遇到的错误及证据是什么。
11. 练习与答案提示
- FFN 不混合 token,为何参数仍多?**答:**隐藏宽度常约 4D,两个大矩阵约 8D²。
- Weight tying 后改
token_emb.weight会怎样?**答:**输出分类矩阵同步改变。 - 固定长度样本还需 padding mask 吗?**答:**不需要;自回归仍需 causal mask。
- Pre-Norm 为什么还需要 output norm?**答:**多个残差累积后需控制最终表示尺度。
- Perplexity 和交叉熵关系?**答:**自然对数 Loss 下通常为
exp(loss)。