← 返回全部思考

大模型训练 Infra 实战课:用一次 35B MoE 后训练把原理讲透

以 40 卡 H20 集群上 Qwen3.6-35B-A3B(35B 总参 / 3B 激活 MoE)的 SFT+RL 后训练为教材的 8 节完整课程:从硬件带宽地图、显存账、并行选型、死锁诊断、框架对比、MFU 计算、稳定性工程一路讲到 RL infra 演进,所有公式都代入真实数字,附六次真实排障实录。

大模型训练 Infra 实战课

这不是一套从论文里抄来的课。它的教材是一次正在进行的真实后训练任务:40 张 H20 上跑 Qwen3.6-35B-A3B 的 SFT,接下来接 RL。

课程里每一个公式都会代入这套集群的真实数字算出来——不是 A100 的示例数字,是 120 秒/步、2.31M token/步、峰值 85.6GB、MFU 8-10% 这些从 log 里捞出来的数。每一个"坑"也都真踩过,包括一次让 39 张卡静默死锁 5 小时的事故。

为什么要留这份记录

后训练 infra 有个特点:它的知识几乎全部长在失败里。

论文告诉你 ZeRO 分三级、EP 要做 all-to-all,但不会告诉你 FusedAdam 的优化器状态是在第一次 optimizer.step() 才惰性分配的——于是你的显存预算算得再准,也可能在第一步之后才 OOM,而且只 OOM 一张卡,剩下 39 张卡在集体通信里安静地等它,一直等到你一小时后发现 GPU 利用率是 100% 空转。

这类知识不写下来就会丢。所以这份记录同时是三样东西:

  1. 一套可以照着学的课:8 节,循序渐进,原理 → 为什么 → 真实数字 → 常见坑。
  2. 一份 SFT/RL infra 的调参依据:为什么是 TP1/PP1/EP8/DP40,为什么 micro batch 只能是 1,为什么梯度 buffer 从 fp32 改 bf16。
  3. 一份事故档案:六次排障的现象、误判、定位手法和最终结论。

教材项目:真实参数

集群

项目 参数
GPU H20,96GB HBM3,bf16 峰值约 148 TFLOPs,带宽约 4.0 TB/s
规模 5 节点 × 8 卡 = 40 卡,集群峰值 5.92 PFLOPs
节点内 NVLink + NVSwitch,900 GB/s 双向全互联
节点间 eth0(管理面) + RoCE 多网卡(数据面)
硬约束 驱动上限 CUDA 12.8 → vllm ≤ 0.19、torch ≤ 2.11,SGLang 出局
共享存储 NFS(数据集、代码、checkpoint)

H20 的性格一句话讲完:大显存、弱算力。它是被出口管制削过的 H100——算力砍到约 15%,但显存反而更大(96 vs 80GB)、带宽反而更高(4.0 vs 3.35 TB/s)、NVLink 一刀未动。这个畸形的比例会贯穿整门课的所有取舍。

模型

Qwen3.6-35B-A3B:35B 总参 / 3B 激活的 MoE,混合 GDN 线性注意力(每 4 层 1 层全注意力),多模态。

并行与 batch

维度 取值 一句话理由
TP 1 96GB 显存装得下,不必付每层两次 all-reduce 的通信税
PP 1 40 卡规模不需要,且 PP 会引入 bubble 和额外调度复杂度
EP 8 恰好等于单节点卡数,把 all-to-all 全部关在 NVLink 域内
DP 40 每步才通信一次的低频流量,才允许跨出节点走 RoCE
micro batch 1 49K 序列已经把单卡显存吃满
global batch 80 梯度累积 2
序列长度 49K 长序列是 MFU 的主要拖累项之一
其他 full recompute + distributed optimizer 用算力换显存、用分片换显存

实测性能

指标 数值
单步耗时 120 秒
单步 token 2.31M(80 样本 × 平均 29K token)
集群吞吐 约 19,300 tok/s
单卡吞吐 482 tok/s
MFU 约 8-10%

显存账

项目 数值
每卡 bf16 权重 约 13.8GB(EP8 分完专家后每卡约 6.9B 参数)
fp32 梯度 buffer 曾需 27.7GB,改 bf16 后减半
distributed optimizer 分片有效(注意专家的 DP 组只有 5)
峰值显存 从 95GB 降到 85.6GB

框架

  • SFT:Megatron-SWIFT,底座 megatron-core 0.16
  • RL:计划用 verl,同 mcore 底座,safetensors 互通
  • 历史对照:v1 用 transformers + DeepSpeed,同数据 2.5 天;v2 换 Megatron 后一夜跑完 2 epoch

课程结构

事故档案索引

六次真实排障,会在对应章节展开。这里先列出来,因为它们本身就是本课程最硬的部分:

# 现象 根因 所在章节
1 个别 rank 在第一步之后 OOM,其余 39 rank 静默死锁 5 小时 FusedAdam 优化器状态在首次 optimizer.step() 才惰性分配 第 2 节
2 死锁一小时后才被"炸"出来 NCCL watchdog / ddp_timeout 的作用与局限 第 4 节
3 新任务一启动就 OOM 被 kill 的进程显存未释放,脏启动竞态 第 7 节
4 不知道谁死了、谁在等谁 py-spy 全集群扫栈,还原死锁拓扑 第 4 节
5 异常 traceback 打不出来 被卡死的 destroy_process_group 挡住,必须先打印再清理 第 4 节
6 集体通信无故死锁 多网卡下 NCCL_SOCKET_IFNAME 选口错误 + 节点全局 bashrc 里残留不一致的 NCCL 调试变量 第 1、7 节

怎么读这套课

每节固定四段结构:直觉 → 公式 → 代入真实数字 → 常见坑,末尾留 1-2 道思考题。

如果你只想拿走一个判断框架,那就是第 1 节最后那张带宽地图:从 HBM 到 NFS,每降一级带宽掉一个数量级;并行策略设计的全部本质,就是让高频通信发生在高带宽层级。 后面七节都是这句话的展开。