跳到主内容
快讯直播
AI智模界
教程

从零手搭现代 LLM:GQA、RoPE、MoE 逐模块实现

这篇做完,你能拿到什么

一篇代码跟下来,你会得到一个能在 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

Embedding

→ 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=512n_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_modeln_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、张量并行、专家并行、梯度检查点等一大堆工程优化,但那些都是在今天这些模块之上做的加法。骨架搞懂了,剩下的都是查文档的事。

AI 生成本文由 AI 基于公开信息自动生成,仅供参考。