BotOf TechAI / IoT / Full-Stack / 植物养护知识分享
W0301 / 夯实底层理论与架构

Transformer 最小闭环

把注意力、FFN、残差、归一化和位置编码组成真正可训练的 decoder-only Transformer。

建议时长
9–12 小时
难度
核心实现
课程位置
3 / 12
LEARNING OBJECTIVES

学完这一课,你应该能:

  • 能解释 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 压得非常低。失败时按顺序检查:

  1. 标签是否 shift。
  2. causal mask 是否正确。
  3. padding label 是否为 -100
  4. optimizer 是否包含所有参数。
  5. train() / eval() 状态。
  6. 学习率是否极端。

小批次过拟合先证明模型“有能力记住”,再扩到完整语料。它能把实现错误与数据噪声分开。

9. 三个有价值的消融

去掉残差

比较 2 层与 8 层,记录每层 gradient norm。深度增加后,无残差版本通常更难优化。

Pre-Norm 改为 Post-Norm

固定参数量与数据,扫描学习率,比较早期 loss、梯度波动与最终验证集。

去掉位置 embedding

构造 token 集合相同但顺序不同的序列,例如 A B C DD 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. 练习与答案提示

  1. FFN 不混合 token,为何参数仍多?**答:**隐藏宽度常约 4D,两个大矩阵约 8D²。
  2. Weight tying 后改 token_emb.weight 会怎样?**答:**输出分类矩阵同步改变。
  3. 固定长度样本还需 padding mask 吗?**答:**不需要;自回归仍需 causal mask。
  4. Pre-Norm 为什么还需要 output norm?**答:**多个残差累积后需控制最终表示尺度。
  5. Perplexity 和交叉熵关系?**答:**自然对数 Loss 下通常为 exp(loss)

12. 延伸阅读