大模型训练 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% 空转。
这类知识不写下来就会丢。所以这份记录同时是三样东西:
- 一套可以照着学的课:8 节,循序渐进,原理 → 为什么 → 真实数字 → 常见坑。
- 一份 SFT/RL infra 的调参依据:为什么是 TP1/PP1/EP8/DP40,为什么 micro batch 只能是 1,为什么梯度 buffer 从 fp32 改 bf16。
- 一份事故档案:六次排障的现象、误判、定位手法和最终结论。
教材项目:真实参数
集群
| 项目 | 参数 |
|---|---|
| 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 节:硬件层——算力、显存、带宽的三角,与一张集群带宽地图
- 第 2 节:显存账——把 85.6GB 峰值逐项拆开算
- 第 3 节:并行策略全景——为什么是 TP1 / PP1 / EP8 / DP40
- 第 4 节:通信层与死锁诊断——5 小时静默死锁该怎么在 5 分钟内定位
- 第 5 节:框架景观——从 2.5 天到一夜跑完 2 epoch
- 第 6 节:手把手算 MFU——8-10% 其实被低估了
- 第 7 节:稳定性工程——失败模式、监控、checkpoint 与脏环境治理
- 第 8 节:从 SFT 到 RL——一套 infra 两个入口
事故档案索引
六次真实排障,会在对应章节展开。这里先列出来,因为它们本身就是本课程最硬的部分:
| # | 现象 | 根因 | 所在章节 |
|---|---|---|---|
| 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,每降一级带宽掉一个数量级;并行策略设计的全部本质,就是让高频通信发生在高带宽层级。 后面七节都是这句话的展开。