从 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:复现路径

  1. 买 Raschka 的《Build a Large Language Model (from Scratch)》,或直接在 github.com/rasbt/LLMs-from-scratch 看第 4 章代码
  2. 在 GPT-2 实现基础上按上文的 MoEFeedForward 替换 FFN
  3. 添加 auxiliary load-balancing loss 到训练循环
  4. 准备数据集(原帖用约 8.9B tokens,作者未披露数据集名)
  5. 开训,预计需要 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 章代码阅读。