从 GPT-2 到 MoE:单卡 RTX 3090 上 8 天训练 446M 稀疏专家模型 · 干货攻略
- 链接: https://www.gilesthomas.com/2026/09/gpt-2-to-moe
- 分类: x-tips
- 来源: X @rasbt
- 作者: Jay
- 更新: 2026-10-07
这是什么
这是一篇由独立开发者 Giles Thomas 撰写的实操复盘(2026年9月10日发布),记录了如何从 Sebastian Raschka 的书《Build a Large Language Model (from Scratch)》中的 GPT-2 代码出发,添加 MoE(Mixture-of-Experts,稀疏专家混合)支持,并在单张 RTX 3090(24GB)上从零训练一个 446M 参数模型的完整过程。模型架构本质上是将 GPT-2 small 的前馈网络(FFN)层替换为 6 个独立的 FFN 专家,每次推理激活其中表现最优的 2 个(top-2 routing)。
为什么值得关注
MoE 是当前主流大模型的核心架构之一——Claude、ChatGPT、DeepSeek-V4、Kimi K3 等都被认为基于 MoE。它解决了大模型"参数多但推理慢"的矛盾:总参数量大(所以知识容量大),但每次只激活一部分(所以推理快)。然而主流教程大多只讲概念,缺乏从零实现+真实训练的完整闭环。
Giles Thomas 的这篇博客恰好填补了这个空白——他把 Raschka 书中的 GPT-2 代码一步步改造成 MoE,完整记录了路由器梯度消失的陷阱、auxiliary load-balancing loss 的原理与实现,以及在 RTX 3090 上跑 8 天的真实成本。他的结论值得注意:446M 总参数量、220M 激活参数,在测试集 loss 上超过了原始 GPT-2 small(124M),接近 GPT-2 medium(345M),但有更多总参数却 fewer active params。
📌 Raschka 的书配套代码库 rasbt/LLMs-from-scratch 已有 104.9k GitHub stars,第 4 章专门讲 MoE,也为本文提供了坚实的理论基础。
核验过程
官方来源 1:Giles Thomas 博客原文(https://www.gilesthomas.com/2026/09/gpt-2-to-moe) - 训练时间:不到 8 天(~90,823 global steps),用 Token 约 8,928,215,040(约 8.9B) - 模型规模:446M 总参数 / 220M 激活参数(top-2 / 6 experts),对比基准:GPT-2 small = 124M,GPT-2 medium = 345M - 硬件:RTX 3090 24GB - 训练结果:test loss 超过本人之前所有模型,原帖称「好于 GPT-2 small,接近 GPT-2 medium」
官方来源 2:Raschka 书代码库 rasbt/LLMs-from-scratch(https://github.com/rasbt/LLMs-from-scratch) - 确认第 4 章包含 MoE 相关代码(GPT-2 + MoE),以及 Qwen3 Dense & MoE 从零实现 - 该仓库是本文的工程起点
交叉验证(Switch Transformers 论文 + Mixtral 实现)
原帖明确提到实现参考了 4 篇关键论文,其中 Switch Transformers(Google, 2021)是 auxiliary load-balancing loss 的来源。将原帖描述的实现逻辑与 Hugging Face Mixtral 源码(https://github.com/huggingface/transformers/blob/main/src/transformers/models/mixtral/modeling_mixtral.py)对比,原帖作者确认"MoE 核心实现与 Mixtral 基本一致,仅 auxiliary loss 计算有细微差异"。
核验结论: - 训练数字(~8 天、8.9B tokens、90,823 steps):来自原帖,经与公开 RTX 3090 LLM 训练基准交叉核验,数量级合理。 - 模型参数(446M 总 / 220M 激活):来自原帖。 - GPT-2 small = 124M、GPT-2 medium = 345M:OpenAI 官方文档确认。 - Auxiliary loss 来自 Switch Transformers:原帖+交叉验证确认。 - 性能比较(原帖自称):为作者自测,无独立第三方基准,标注为「原帖主张,未核验」。
上手步骤
前提条件
# Python 3.10+, PyTorch 2.x, 约 30GB 磁盘空间(数据集+checkpoint)
pip install torch torchvision numpy tqdm
# 推荐 24GB+ 显存GPU;原帖 RTX 3090(24GB)
Step 1:理解 Raschka GPT-2 的 FFN 位置
在标准 Transformer block 中,FFN 占 GPT-2 总参数量的约 2/3(原帖引用 Raschka 的分析)。MoE 的核心改动就是把这一个 FFN 替换为多个独立的 FFN(专家)。
标准 Transformer block:
Input → LayerNorm → Self-Attention → ResidualAdd
→ LayerNorm → [单个 FFN] → ResidualAdd → Output
MoE Transformer block:
Input → LayerNorm → Self-Attention → ResidualAdd
→ LayerNorm → [Router → Top-K 选择专家 → 加权合并] → ResidualAdd → Output
Step 2:路由器的数学(Top-K + Softmax 稀疏化)
原帖给出了完整的逐步推导,这是最核心的知识点:
原始 router 输出(单层 linear,d_emb → n_experts):
logits = router(inputs) # shape: (batch, n_experts)
# e.g., tensor([0.0418, -0.1140, 0.4254, 0.1342, 0.5106, -0.1385])
关键问题:直接取 top-k → 无梯度 → 路由器无法训练!
如果直接 top_k_indices = torch.topk(logits, k),反向传播时路由选择这一步是离散的,不在计算图中,路由器权重永远得不到梯度更新。
解法(Outrageously Large Neural Networks, 2017):将 top-k 掩码注入 softmax
# Step 1: 取 top-k 的值,其余设为 -inf
top_k_logits, _ = torch.topk(logits, k)
# 将非 top-k 位置设为 -inf
masked_logits = logits.masked_fill(
logits < top_k_logits[:, -1].unsqueeze(-1), float('-inf')
)
# Step 2: softmax → top-k 变稀疏权重,其余为 0
weights = torch.softmax(masked_logits, dim=-1)
# 例: tensor([0.0000, 0.0000, 0.4787, 0.0000, 0.5213, 0.0000])
# Step 3: 跳过权重为 0 的专家,只跑 top-k → 节省计算
⚠️ Top-1 routing 陷阱(Switch Transformers 特别指出)
当 k=1 时,softmax 后权重永远为 [0,0,...,1,0,...](恒为 one-hot),梯度恒为 0,路由器无法训练。因此必须 k≥2。Giles Thomas 的模型用的是 top-2 routing(6 个专家中选 2 个)。
Step 3:Auxiliary Load-Balancing Loss(防止路由崩溃)
如果只优化语言建模损失,路由器会迅速将所有 token 路由到同一个专家(多数专家永远不被激活),导致模型退化。解决方案来自 Switch Transformers:添加 auxiliary loss 惩罚 expert 利用率不均。
Switch Transformers 的 load-balancing loss(原帖使用):
# 对每个 expert 计算"重要性分数"(router 输出概率的均值)
import torch
def load_balancing_loss(router_probs, top_k_idx, n_experts):
# router_probs: (n_tokens, n_experts) softmax 概率
# top_k_idx: (n_tokens, k) 被选中的 expert id
f_i = router_probs.mean(dim=0) # 每个 expert 的平均概率
P_i = torch.zeros(n_experts, device=f_i.device)
P_i.scatter_add_(0, top_k_idx.flatten(), torch.ones_like(top_k_idx, dtype=torch.float).flatten())
P_i = P_i / top_k_idx.numel() # 实际选中频率
loss = n_experts * (f_i * P_i).sum() # 越高说明越不均衡
return loss
💡 原帖提到 Mixtral 在 auxiliary loss 计算上有细微差异(但未详述具体差异内容)。实用角度:关键是让
f_i(软概率)和P_i(硬选中频率)尽量对齐。
Step 4:完整 MoE Block 示例代码
import torch
import torch.nn.functional as F
import torch.nn as nn
class MoEFeedForward(nn.Module):
def __init__(self, d_model, n_experts=6, top_k=2):
super().__init__()
self.n_experts = n_experts
self.top_k = top_k
# Router: 单层 linear,无 bias(参考 Switch Transformers)
self.router = nn.Linear(d_model, n_experts, bias=False)
# 6 个独立的 FFN 专家
self.experts = nn.ModuleList([
nn.Sequential(
nn.Linear(d_model, d_model * 4),
nn.GELU(),
nn.Linear(d_model * 4, d_model)
)
for _ in range(n_experts)
])
def forward(self, x):
B, T, C = x.shape # batch, seq_len, hidden
x_flat = x.view(-1, C) # (B*T, C)
# Router forward
router_logits = self.router(x_flat) # (B*T, n_experts)
router_weights = F.softmax(router_logits, dim=-1)
# Top-K 选择
top_k_weights, top_k_idx = torch.topk(router_weights, self.top_k, dim=-1)
top_k_weights = top_k_weights / top_k_weights.sum(dim=-1, keepdim=True) # 归一化
# 稀疏前向:只跑被选中的 expert
out = torch.zeros_like(x_flat)
for i in range(self.top_k):
expert_id = top_k_idx[:, i] # (B*T,)
expert_weight = top_k_weights[:, i] # (B*T,)
for e in range(self.n_experts):
mask = (expert_id == e)
if mask.any():
expert_out = self.experts[e](x_flat[mask])
out[mask] += expert_weight[mask].unsqueeze(-1) * expert_out
return out.view(B, T, C)
⚠️ 上述代码为简化示意。生产级实现需要考虑 expert capacity(专家容量上限)、grouped routing(批处理优化)等,Giles Thomas 博客中均有讨论。
Step 5:训练配置(来自原帖实测)
# 来自 Giles Thomas 原帖的配置摘要
config = {
"model": "GPT-2 small 架构 + MoE 替换",
"n_layers": 12, # 同 GPT-2 small
"n_heads": 12,
"n_experts": 6, # 总专家数
"top_k": 2, # 每次激活专家数
"d_model": 768, # 隐藏层维度
"total_params": "446M", # 总参数
"active_params": "220M", # 激活参数
"total_tokens": "8.9B (~90,823 steps)",
"hardware": "RTX 3090 24GB",
"training_time": "~8 天",
"dataset": "原帖未披露具体数据集(为作者个人数据集)",
}
Step 6:复现路径
- 买 Raschka 的《Build a Large Language Model (from Scratch)》,或直接在 github.com/rasbt/LLMs-from-scratch 看第 4 章代码
- 在 GPT-2 实现基础上按上文的
MoEFeedForward替换 FFN - 添加 auxiliary load-balancing loss 到训练循环
- 准备数据集(原帖用约 8.9B tokens,作者未披露数据集名)
- 开训,预计需要 7-10 天单卡
坑与适用边界
坑 1:auxiliary loss 权重必须调
原帖没有给出 auxiliary loss 的具体权重(alpha),这个超参数对训练稳定性影响很大。太小 → router 仍会崩溃;太大 → 主损失被稀释。Switch Transformers 原论文建议 0.01-0.1 量级,需要根据 loss 曲线调。
坑 2:RTX 3090 24GB 是最低可行配置 8.9B tokens 跑 8 天意味着 checkpoint 保存、梯度累积、batch size 都需要精心调配。如果显存更小,需要更多梯度累积步数,训练时间会显著延长。
坑 3:auxiliary loss 计算与 Mixtral 细节差异 原帖仅提及"有细微差异"但未详述。追求生产级准确度的话,务必对照 Mixtral HF 源码 核验 backward pass 的实现。
适用边界
- 适合已有 PyTorch LLM 训练经验的开发者(需要能看懂 Raschka 的 GPT-2 代码)
- 不适合纯新手:没有提供可直接 git clone 的配套代码仓库
- 训练数据集原帖未披露,无法精确复现
- 性能数字(原帖自称超过 GPT-2 small/接近 medium):未经独立第三方验证,仅供参考
一句话结论
从 Raschka GPT-2 到 MoE 的实操天花板——用 RTX 3090 单卡、8 天、8.9B tokens,从零训出 446M MoE(top-2/6 experts),核心难点是 router 梯度消失问题和 auxiliary load-balancing loss 的正确实现,推荐配合 rasbt/LLMs-from-scratch 第 4 章代码阅读。