这篇做完,你能拿到什么
一篇代码跟下来,你会得到一个能在 CPU 上跑通前向 + 反向的小型 LLM:4 层、隐藏维度 128~256 的 decoder-only 结构,里面包含现代大模型最核心的五个零件——AI 词典:RMSNorm">RMSNorm、RoPE、GQA 注意力、SwiGLU 前馈、MoE 路由,外加配套的负载均衡辅助损失。
它当然生成不出像样的文章,但它能干三件有价值的事:
1. 验证形状与梯度:打印出 logits 形状、损失值、router 的梯度范数,确认整条链路是活的。
2. 看清架构演进:你可以把 GQA 换回 MHA、把 RoPE 换回可学习位置编码、把 MoE 换回单个 FFN,跑同一份数据,对比参数量和 loss 曲线的差别。这种"控制变量"的体感,比读十篇论文摘要都来得直接。
3. 写出可过拟合的 Sanity Check:把一小段文本背下来。能背下来,说明你的反向传播、位置编码、注意力掩码都是对的。
本文代码是自包含的 PyTorch 实现,思路参照 OpenArch 这类"把现代 LLM 架构拆开讲"的开源参考实现。如果你想对照阅读,建议 clone 一份 OpenArch,把它的模块文件和下面每一节一一对应;具体的模块命名、目录结构和接口以 OpenArch 官方仓库和文档为准。
前置条件清单
- Python 3.9 以上,PyTorch 2.x(安装命令、版本要求以 PyTorch 官方页面为准)
- 会用
nn.Module写最普通的网络层 - 不需要 GPU,不需要数据集,不需要下载任何权重
- 一点点耐心:下面的代码是按"零件 → 总装"的顺序写的,请按顺序敲,每写一段就 import 跑一次
环境准备:
```bash
pip install torch
```
先看清全局骨架
一次前向传播长这样:
```text
token id
→ N × Block {
RMSNorm → Attention(GQA + RoPE) → 残差相加
RMSNorm → FFN(SwiGLU 或 MoE) → 残差相加
}
→ RMSNorm
→ LM Head(与 Embedding 共享权重)
→ logits
```
记住这个骨架,下面每个模块都是在往里填一块积木。
第 1 步:RMSNorm —— 把 LayerNorm 的"减均值"去掉
LayerNorm 干两件事:减去均值、除以标准差,再乘一个缩放系数 γ、加一个偏置 β。
RMSNorm 说:减均值这一步可以省。它只保留"除以均方根"和缩放 γ:
```python
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
class RMSNorm(nn.Module):
"""去掉均值减法、去掉 bias 的归一化层。"""
def __init__(self, dim: int, eps: float = 1e-6):
super().__init__()
self.eps = eps
self.weight = nn.Parameter(torch.ones(dim))
def forward(self, x):
dtype = x.dtype
统计量用 fp32 计算:bf16 下直接算平方会掉精度
x = x.float()
x = x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps)
return self.weight * x.to(dtype)
```
省掉一次 reduce 和一组 bias 参数,在大模型里就是实打实的显存和带宽。注意 keepdim=True,否则广播会把后面的维度吃掉。
第 2 步:RoPE —— 用旋转表示位置
早期模型用"可学习的位置 embedding",把它加到 token embedding 上。问题有两个:一是序列一长,位置表就不够用了;二是模型很难学到"相对距离"这个概念。
RoPE 换了个思路:不往向量里加东西,而是把 Query 和 Key 向量旋转一个角度。角度由位置决定,转完之后两个向量做点积,结果天然只跟它们的相对距离有关。
具体做法是把 head_dim 维度两两分组,每一对看作二维平面上的一根向量,位置 t 就按 t * θᵢ 旋转,θᵢ 随维度索引递减——低频管长距离,高频管短距离。
```python
def build_rope_cache(seq_len: int, head_dim: int, base: float = 10000.0,
device=None, dtype=torch.float32):
"""预计算每个位置的 cos / sin,形状都是 (seq_len, head_dim // 2)。"""
assert head_dim % 2 == 0, "head_dim 必须是偶数,否则没法两两配对"
第 i 组的角频率:theta_i = base ** (-2i / head_dim)
inv_freq = 1.0 / (base ** (torch.arange(0, head_dim, 2, device=device).float() / head_dim))
pos = torch.arange(seq_len, device=device).float()
freqs = torch.outer(pos, inv_freq) # (seq_len, head_dim // 2)
return freqs.cos().to(dtype), freqs.sin().to(dtype)
def apply_rope(x, cos, sin):
"""x: (B, H, T, D);cos / sin: (T, D // 2)。"""
B, H, T, D = x.shape
cos = cos[:T].view(1, 1, T, D // 2)
sin = sin[:T].view(1, 1, T, D // 2)
x1, x2 = x[..., : D // 2], x[..., D // 2:]
二维旋转矩阵:[cos -sin; sin cos] 作用在 (x1, x2) 上
return torch.cat([x1 * cos - x2 * sin,
x1 * sin + x2 * cos], dim=-1)
```
关键点:cos / sin 一定要预先算好、注册成 buffer,不要在每个 forward 里重算。旋转的"配对方式"(前后各半还是奇偶交错)必须全局一致,否则位置信息会错乱。
第 3 步:GQA —— 让 KV 头比 Q 头少
标准的 MHA 里,Q、K、V 的头数一样。推理时每生成一个 token,都要把新的 K、V 追加到缓存里,缓存大小正比于头数。
- MQA:K、V 只剩 1 个头。缓存最小,但质量掉得明显。
- GQA:K、V 的头数取 Q 头数的一个约数,比如 Q 有 32 个头、K/V 有 8 个,那么每 4 个 Q 头共享一组 K/V。
代码上就是"投影矩阵输出维度变小 + 注意力计算前把 K/V 复制对齐":
```python
class GQAAttention(nn.Module):
def __init__(self, d_model, n_heads, n_kv_heads, dropout=0.0):
super().__init__()
assert d_model % n_heads == 0
assert n_heads % n_kv_heads == 0, "Q 头数必须能被 KV 头数整除"
self.n_heads = n_heads
self.n_kv_heads = n_kv_heads
self.n_rep = n_heads // n_kv_heads # 每组 KV 被几个 Q 头共享
self.head_dim = d_model // n_heads
self.q_proj = nn.Linear(d_model, n_heads * self.head_dim, bias=False)
self.k_proj = nn.Linear(d_model, n_kv_heads * self.head_dim, bias=False)
self.v_proj = nn.Linear(d_model, n_kv_heads * self.head_dim, bias=False)
self.o_proj = nn.Linear(n_heads * self.head_dim, d_model, bias=False)
self.dropout = dropout
def forward(self, x, cos, sin):
B, T, _ = x.shape
q = self.q_proj(x).view(B, T, self.n_heads, self.head_dim).transpose(1, 2)
k = self.k_proj(x).view(B, T, self.n_kv_heads, self.head_dim).transpose(1, 2)
v = self.v_proj(x).view(B, T, self.n_kv_heads, self.head_dim).transpose(1, 2)
q = apply_rope(q, cos, sin)
k = apply_rope(k, cos, sin)
关键:用 repeat_interleave 按顺序复制,保证 KV 分组和 Q 头对齐
k = k.repeat_interleave(self.n_rep, dim=1)
v = v.repeat_interleave(self.n_rep, dim=1)
out = F.scaled_dot_product_attention(
q, k, v,
is_causal=True, # 因果掩码,别漏
dropout_p=self.dropout if self.training else 0.0,
) # (B, H, T, head_dim)
out = out.transpose(1, 2).reshape(B, T, -1)
return self.o_proj(out)
```
算一笔账:d_model=512、n_heads=8(head_dim=64)时,MHA 的 q/k/v/o 四个投影一共约 105 万参数;换成 n_kv_heads=2 的 GQA,k/v 投影各缩到四分之一,总数降到约 66 万。推理时 KV cache 更是直接小到四分之一。
第 4 步:SwiGLU —— 带门控的前馈层
原始 Transformer 的 FFN 是 Linear → ReLU → Linear。SwiGLU 换成"门控"结构:一路算出候选值,另一路算出门控系数,逐元素相乘后再投影回原维度。
```python
class SwiGLU(nn.Module):
def __init__(self, d_model, hidden):
super().__init__()
self.w_gate = nn.Linear(d_model, hidden, bias=False)
self.w_up = nn.Linear(d_model, hidden, bias=False)
self.w_down = nn.Linear(hidden, d_model, bias=False)
def forward(self, x):
return self.w_down(F.silu(self.w_gate(x)) * self.w_up(x))
```
多了一个矩阵,所以 hidden 一般取 8/3 × d_model 而不是原来的 4 × d_model,让总参数量大致持平。实践里还会把 hidden 对齐到 128 的整数倍,做矩阵乘时对硬件更友好。
第 5 步:MoE —— 把 FFN 换成"专家团 + 门卫"
MoE 的核心想法:把上面那个 FFN 复制成 E 份(叫"专家"),再加一个小的 router 网络。每个 token 只走其中 top-k 个专家,其余专家完全不参与这个 token 的计算。
于是参数量变成 E 倍,但每个 token 的算力只增加 k 倍。这就是"稀疏激活"。
附带一个问题:router 很可能"赢家通吃"——所有 token 都涌向同一个专家,其他专家永远学不到东西。所以要加一个负载均衡辅助损失,惩罚路由分布过于集中。
```python
class MoE(nn.Module):
def __init__(self, d_model, hidden, n_experts=8, top_k=2, aux_coef=0.01):
super().__init__()
assert 0 < top_k <= n_experts
self.n_experts = n_experts
self.top_k = top_k
self.aux_coef = aux_coef
self.router = nn.Linear(d_model, n_experts, bias=False)
self.experts = nn.ModuleList(SwiGLU(d_model, hidden) for _ in range(n_experts))
def forward(self, x):
B, T, D = x.shape
flat = x.reshape(-1, D) # (N, D)
N = flat.shape[0]
logits = self.router(flat) # (N, E)
probs = F.softmax(logits, dim=-1)
top_p, top_idx = probs.topk(self.top_k, dim=-1) # (N, k)
重新归一化,让被选中的 k 个权重加起来等于 1
top_p = top_p / top_p.sum(dim=-1, keepdim=True).clamp_min(1e-9)
教学版 combine:每个专家都算一遍,再按掩码加权求和。
注意 out 用 torch.zeros 显式创建(而不是 zeros_like 输入的张量),
然后全程用 out-of-place 的加法累加,避免原地操作踩自动求导的坑。
out = torch.zeros(N, D, dtype=flat.dtype, device=flat.device)
for e in range(self.n_experts):
mask = (top_idx == e).to(probs.dtype) # (N, k)
weight = (top_p * mask).sum(dim=-1, keepdim=True) # (N, 1)
if float(weight.abs().sum()) == 0.0:
continue
out = out + self.expertse * weight
---- 负载均衡损失:Switch Transformer 的形式 ----
counts = torch.bincount(top_idx.reshape(-1),
minlength=self.n_experts).to(probs.dtype)
f = counts / (N * self.top_k) # 每个专家实际分到的 token 比例
P = probs.mean(dim=0) # 每个专家的平均路由概率
aux = self.aux_coef * self.n_experts * (f * P).sum()
return out.reshape(B, T, D), aux
```
这段的 combine 是故意写成教学版的:为了短、好读、绝对能跑通,它把每个专家都在全部 token 上算了一遍,等于白烧了 E 倍算力。生产实现会先按专家把 token 分桶(dispatch),只对分到该专家的 token 做计算,再用 scatter / index_add 把结果拼回去——思路一样,只是多了排序和索引对齐的工程活儿。想深入的话搜 "MoE token dispatch combine" 看官方实现,注意原地累加类算子对自动求导的要求,细节以 PyTorch 官方文档为准。
第 6 步:总装 —— Block 和 MiniLLM
```python
class Block(nn.Module):
def __init__(self, d_model, n_heads, n_kv_heads, ffn_hidden,
use_moe=True, n_experts=8, top_k=2, dropout=0.0):
super().__init__()
self.norm1 = RMSNorm(d_model)
self.attn = GQAAttention(d_model, n_heads, n_kv_heads, dropout)
self.norm2 = RMSNorm(d_model)
self.use_moe = use_moe
if use_moe:
self.ffn = MoE(d_model, ffn_hidden, n_experts, top_k)
else:
self.ffn = SwiGLU(d_model, ffn_hidden)
def forward(self, x, cos, sin):
pre-norm:归一化放在残差分支里,主干是一条干净的直通路
x = x + self.attn(self.norm1(x), cos, sin)
h = self.norm2(x)
if self.use_moe:
h, aux = self.ffn(h)
else:
h, aux = self.ffn(h), torch.zeros((), device=x.device, dtype=x.dtype)
return x + h, aux
class MiniLLM(nn.Module):
def __init__(self, vocab_size, d_model=256, n_layers=4, n_heads=8,
n_kv_heads=2, n_experts=8, top_k=2, max_seq_len=256,
use_moe=True, dropout=0.0):
super().__init__()
self.head_dim = d_model // n_heads
hidden = int(8 * d_model / 3)
hidden = 128 * ((hidden + 127) // 128) # 向上对齐到 128 的倍数
self.embed = nn.Embedding(vocab_size, d_model)
self.blocks = nn.ModuleList([
Block(d_model, n_heads, n_kv_heads, hidden,
use_moe, n_experts, top_k, dropout)
for _ in range(n_layers)
])
self.norm_f = RMSNorm(d_model)
self.lm_head = nn.Linear(d_model, vocab_size, bias=False)
self.lm_head.weight = self.embed.weight # 权重共享,省一份 embedding
self.max_seq_len = max_seq_len
cos, sin = build_rope_cache(max_seq_len, self.head_dim)
self.register_buffer("cos", cos, persistent=False)
self.register_buffer("sin", sin, persistent=False)
def forward(self, idx):
B, T = idx.shape
assert T <= self.max_seq_len, "序列长度超过了预计算的 RoPE 长度"
x = self.embed(idx)
aux_total = torch.zeros((), device=idx.device, dtype=x.dtype)
for blk in self.blocks:
x, aux = blk(x, self.cos, self.sin)
aux_total = aux_total + aux
x = self.norm_f(x)
return self.lm_head(x), aux_total
```
第 7 步:跑通它 —— 三步冒烟测试
测试 A:形状和梯度
```python
torch.manual_seed(0)
model = MiniLLM(vocab_size=1000, d_model=256, n_layers=4, n_heads=8,
n_kv_heads=2, n_experts=8, top_k=2, max_seq_len=128)
opt = torch.optim.AdamW(model.parameters(), lr=3e-4)
idx = torch.randint(0, 1000, (2, 64))
target = torch.randint(0, 1000, (2, 64))
logits, aux = model(idx)
loss = F.cross_entropy(logits.reshape(-1, logits.size(-1)), target.reshape(-1)) + aux
loss.backward()
opt.step()
print("logits 形状:", tuple(logits.shape)) # 期望 (2, 64, 1000)
print("总损失:", float(loss))
print("辅助损失:", float(aux))
print("router 梯度范数:",
model.blocks[0].ffn.router.weight.grad.norm().item()) # 必须是正数
```
随机数据下 loss 大约在 ln(vocab_size) 附近(约 6.9),这是对的——模型还没学任何东西。如果 router 梯度范数是 0,说明辅助损失没接上,或者某个专家一次都没被选中。
测试 B:对比 GQA 和 MHA 的参数量
```python
attn_mha = GQAAttention(512, n_heads=8, n_kv_heads=8)
attn_gqa = GQAAttention(512, n_heads=8, n_kv_heads=2)
print(sum(p.numel() for p in attn_mha.parameters()))
print(sum(p.numel() for p in attn_gqa.parameters()))
```
测试 C:过拟合一小段文本(最重要的 sanity check)
```python
text = "人工智能正在改变世界。" * 6
chars = sorted(set(text))
stoi = {c: i for i, c in enumerate(chars)}
data = torch.tensor([stoi[c] for c in text]).unsqueeze(0) # (1, T)
model = MiniLLM(vocab_size=len(chars), d_model=128, n_layers=4, n_heads=4,
n_kv_heads=1, n_experts=4, top_k=2, max_seq_len=256)
opt = torch.optim.AdamW(model.parameters(), lr=3e-3)
for step in range(300):
logits, aux = model(data[:, :-1])
loss = F.cross_entropy(logits.reshape(-1, logits.size(-1)),
data[:, 1:].reshape(-1)) + aux
opt.zero_grad()
loss.backward()
opt.step()
if step % 60 == 0:
print(step, round(float(loss), 4))
```
loss 应该从 2 点多一路降到接近 0。降不下去就是有 bug,别急着调超参,先去下面的排错清单里找。
常见坑与排错
1. head_dim 不是偶数。 RoPE 要把维度两两配对,奇数直接报错。选 d_model 和 n_heads 时先除一下。
2. n_heads 不能被 n_kv_heads 整除。 GQA 的 n_rep 必须是整数,否则复制后头数对不上。
3. 用 repeat 而不是 repeat_interleave。 k.repeat(1, 4, 1, 1) 会把 [k0, k1] 变成 [k0, k1, k0, k1],而 repeat_interleave 给出 [k0, k0, k0, k0, k1, ...]。后者才和 Q 的分组对齐。这个 bug 不会报错,只会让模型变笨。
4. 忘了 is_causal=True。 模型能"看见未来",训练 loss 掉得飞快,但推理时生成的全是乱码。这是最隐蔽的坑之一。
5. cos / sin 长度不够。 序列长度超过 max_seq_len 时,cos[:T] 会静默地给你一个更短的张量,然后广播出错或者结果错位。加个 assert 挡住。
6. 路由塌缩。 训练几百步后打印 top_idx 的分布,如果 90% 的 token 都去了专家 0,说明 aux_coef 太小或者 router 初始化不合适。可以先把 aux_coef 调大观察效果,收敛后再调回去。
7. 混合精度下数值不稳。 RMSNorm 里的平方和、softmax 里的指数,都建议在 fp32 下算完再转回去,上面的代码已经这么做了。
8. 原地操作 + 自动求导。 MoE 里把专家输出累加回一个张量时,尽量用 out-of-place 的加法或显式创建的零张量;直接在"需要梯度的叶子张量"上做原地写,很容易撞上 "a leaf Variable that requires grad is being used in an in-place operation"。具体哪些算子支持原地、梯度怎么传,以 PyTorch 官方文档为准。
9. 权重共享后优化器重复。 model.parameters() 会自动去重,所以 lm_head.weight = embed.weight 是安全的;但如果你手动拼了一份参数列表,记得去重。
下一步建议
按难度从低到高,建议这么往下走:
1. 写一个 generate 函数。 贪心解码就行:每步取 logits 最后一列,argmax 出下一个 token,拼回去再前向。跑通之后你会立刻感觉到"没有 KV cache 有多慢"。
2. 加 KV cache。 在 GQAAttention 里维护 past_k / past_v,新 token 只算自己的 Q/K/V 再拼接。GQA 的收益到这一步才真正体现出来。
3. 把 MoE 换成 dispatch 版。 排序 → 分桶 → 按专家并行计算 → scatter 回来。这是理解工业级 MoE 实现的必经之路。
4. 做控制变量实验。 固定其他一切,分别跑 MHA / GQA / MQA,看参数量、显存和 loss 曲线;把 RoPE 换成可学习位置编码再跑一遍。
5. 长上下文外推。 RoPE 的天然优势是可以通过"位置插值"或 NTK 缩放把训练长度外推出去,这也是现在很多长上下文方案的基础。
最后提醒一句:这套代码是为了理解而写的,不是为效率写的。真实训练里还有 FlashAttention、张量并行、专家并行、梯度检查点等一大堆工程优化,但那些都是在今天这些模块之上做的加法。骨架搞懂了,剩下的都是查文档的事。
