训练框架 — mHC 的工程脚手架
Muon 与 ZeRO 为什么天然冲突、V4 怎么"稠密 / MoE 二分 + FP32→BF16 SR 通信"折中,mHC 怎么靠融合 kernel + DualPipe overlap 把 6.7% 开销压下去,以及 CSA / HCA 压缩边界让传统 CP 失效后的 two-stage CP 怎么救场。
训练框架要解决三件事:(1) Muon × ZeRO 的"完整梯度 vs 切分梯度"冲突;(2) mHC 残差路径的 6.7% wall-time 开销;(3) CSA/HCA 把序列压成 1/m 后传统 CP 边界不对齐
一句话:把梯度按"稠密 / MoE"二分别切、把通信精度从 FP32 stochastic round 到 BF16 减半带宽、把 mHC forward/backward 融成单 kernel 嵌进 DualPipe 1F1B 流水、把 CP 切成两阶段先在压缩边界对齐再 all-gather。 每条都是为了让 V4 的架构层创新(Muon、mHC、CSA/HCA)在 1.6T 规模下真能跑得动。
- ZeRO(Zero Redundancy Optimizer)
- 把 optimizer state、grad、param 沿数据维切到不同 rank,每 rank 只存一片。显存压力 ↓ N×,但每步训练需要 all-gather param(forward 用)+ reduce-scatter grad(backward 后)。ZeRO-3 是最激进版本,权重也切。
- Muon × ZeRO 冲突
- Muon 的 NS 迭代要看完整的更新矩阵 $M$(要做 $MM^T$)。ZeRO 把 $M$ 沿某一维切到各 rank,每个 rank 只看自己那片 → $MM^T$ 算不出来。这是架构与并行的根本冲突。
- 稠密 / MoE 二分切
- V4 的折中方案:稠密参数(attention 投影、共享 FFN)限制 ZeRO 并行度 $\le P_{\max}$,用背包均衡把不同 shape 的稠密层打包到 $\le P_{\max}$ 个 rank;MoE 参数每个 expert 独立优化,flatten 所有 down/up/gate 后跨 rank 平均切,不限 ZeRO 并行度(因为 expert 之间彼此独立做 NS)。
- 背包均衡(knapsack-balanced sharding)
- 稠密参数矩阵的尺寸不同,简单按矩阵数量分配容易失衡。V4 用背包算法把矩阵分配到受限数量的 ZeRO ranks,使各 rank 负载大致均衡,再把 bucket padding 到最大 bucket 的尺寸。报告称这种 padding 的内存开销通常低于 10%。
- FP32 → BF16 Stochastic Rounding(SR)
- 跨 rank 同步梯度时的带宽减半方案:发送方把 FP32 梯度随机舍入成 BF16 再发。$E[\text{round}] = $ 原 FP32 值(无偏),跨多步统计平均不引入偏差。比 round-to-nearest(有偏)好得多。
- 两阶段 all-to-all + 本地 FP32 求和
- SR 的搭档:BF16 over wire → 接收方 BF16 升回 FP32 后 本地 求和(保 FP32 精度)→ 第二阶段 all-to-all 同步。这样通信带宽减半 + 数值精度保持。
- DualPipe 1F1B
- 双向流水线(V3 引入,V4 沿用):每个 micro-batch 走"前向 → 后向"循环,1 个 forward 紧跟 1 个 backward 排队。可以把不同 micro-batch 的 forward 与 backward 重叠。给 mHC 的融合 kernel 提供 overlap 窗口。
- mHC 工程优化(融合 kernel + 选择性 ckpt)
- V4 为 mHC 的训练与推理设计融合 kernel,选择性 checkpoint 中间张量,并调整 DualPipe 1F1B 的重叠方式。三项优化合计把 mHC 的 wall-time 开销限制在重叠流水阶段的 6.7%;报告没有披露优化前的比例。
- Tensor 级 Activation Checkpointing
- V4 选择性保存中间张量:重算大部分层间 hidden state 和所有 normalized layer input,同时避免重算计算密集的操作,以平衡显存与计算开销。报告没有提到 TorchFX 或“无额外开销”。
- Two-stage CP(Context Parallelism)
- 长上下文沿序列维切到多 rank。CSA/HCA 把序列每 $m$ token 压成 1 个 → CP 边界经常落在压缩块中间,传统 CP 无法直接切。V4 的方案:Stage 1: 每 rank 把自己 $m$ 个未压缩 KV 发给下一 rank 合并压缩,固定输出长度 $s/m + 1$;Stage 2: 跨 rank all-gather + select-and-pad 重组完整压缩序列。
1. Muon × ZeRO:完整梯度 vs 切分梯度
回顾 Ch06:Muon 每步要把更新矩阵 $M$ polar 化,用 NS 迭代 $M \leftarrow aM + b\,MM^TM + c(MM^T)^2 M$。问题:
- $MM^T$ 是 $n \times n$ 的方阵,要 $M$ 完整才能算;
- ZeRO-3 把 $M$ 沿某一维切到 $P$ 个 rank,每 rank 只有 $M$ 的一片;
- 各 rank 算自己的 $M_i M_i^T$ 后不能直接相加 —— 因为切的是行还是列决定矩阵乘的语义。
设 $M \in \mathbb{R}^{n\times m}$ 沿 $m$ 维切到 $P$ 个 rank:$M = [M_1 | M_2 | \cdots | M_P]$。要算的是:
这条等式里 $M_i M_i^T$ 仍是 $n \times n$,所以简单切分并不会把 Newton–Schulz 所需的矩阵计算变成逐元素优化。V4 因而限制稠密参数的最大 ZeRO 并行规模;对彼此独立的 MoE expert,则允许使用全部 ranks。
1.1 稠密:背包均衡
稠密参数矩阵尺寸不同,按矩阵数量平均分配会造成负载不均。V4 用背包算法把矩阵分配到受限数量的 ZeRO ranks,并把各 bucket padding 到最大 bucket 的尺寸,便于执行 reduce-scatter。报告没有给出 $P_{\max}$ 的具体值,也没有承诺 rank 间误差不超过 5%;明确披露的是 padding 通常带来不足 10% 的内存开销,且每个 rank 最多管理 5 个参数矩阵。
1.2 MoE:每 expert 独立切
对 MoE 参数,V4 分别展平所有层、所有 experts 的 down、up 和 gate 投影矩阵,再进行 padding,使向量可以在不拆分逻辑独立矩阵的前提下均匀分到各 ranks。由于 expert 数量多,MoE 参数不限制 ZeRO 并行规模,padding 开销也可以忽略。
1.3 通信精度:SR FP32→BF16
两个最大的通信:reduce-scatter 梯度(每步)、all-gather param(每步)。原本 FP32 = 4 字节/数。SR 把发送方先 stochastic round 到 BF16 = 2 字节/数。
设 FP32 值 $x$ 落在两个 BF16 邻近码点 $\lfloor x \rfloor$ 和 $\lceil x \rceil$ 之间,距离比例 $p = (x - \lfloor x\rfloor) / (\lceil x \rceil - \lfloor x \rfloor)$:
- 无偏:$E[\mathrm{SR}(x)] = p \cdot \lceil x \rceil + (1-p) \cdot \lfloor x \rfloor = x$;
- vs round-to-nearest:RTN 是有偏的(小数总是往最近偶数靠),多步累积会漂移;SR 不漂移,跨多步取平均逼近真值;
- 方差:$\mathrm{Var}[\mathrm{SR}(x)] = p(1-p) \cdot (\lceil x \rceil - \lfloor x \rfloor)^2 \le \frac{1}{4} \cdot \mathrm{ULP}^2$,与 RTN 同量级。
SR 的期望值无偏,但单次舍入仍会引入方差。报告称 Newton–Schulz 的 BF16 矩阵乘保持稳定,并通过本地 FP32 求和避免低精度累加误差;它没有给出“质量几乎无损”的独立消融数据。
2. mHC 的 6.7%:融合、重算与流水重叠
mHC 相比普通残差连接会增加 activation memory 和流水阶段间的通信量。报告给出三项优化:
- 融合 kernel:为 mHC 的训练和推理实现专用融合 kernel;
- 选择性 checkpoint:重算大部分层间 hidden state 和全部 normalized layer input,同时避开计算密集操作的重复执行;
- DualPipe overlap:调整 1F1B 重叠方案,以容纳更多流水通信,并让部分 mHC 操作并发执行。
这是上述优化共同作用后,mHC 相对于 overlapped 1F1B pipeline stage 的 wall-time overhead。报告没有给出 eager 基线、单 kernel 微秒数、重算比例或显存节省比例,因此无法把 6.7% 拆成更细的时延账。
报告只说明稠密参数限制最大 ZeRO 并行规模,超出的 data-parallel groups 会冗余计算 Muon update;它没有披露 $P_{\max}$、稠密参数总量或每步通信量,因此只能说明机制,不能绘制带固定配置数字的显存图。
3. Two-stage CP:让 CSA/HCA 压缩边界与 CP 边界和谐共处
长上下文 CP 的意思:把 1M token 沿序列维切到 $P_{\text{cp}}$ 个 rank,每 rank 处理 $1\text{M}/P_{\text{cp}}$ token。每层 attention 时跨 rank 互通需要 KV 的 ring all-reduce。 问题:CSA / HCA 要把每 $m$ 个 token 压成 1 个,CP 边界经常落在某个压缩块中间 —— 那个边界 token 需要前后 $m$ 个 token 的信息,但前一段在邻居 rank 上。
V4 的两阶段方案:
- Stage 1: 每个 rank 把自己最末尾 $m$ 个未压缩 KV 发给下一个 rank。下一个 rank 收到后与本地的 $m$ 个 KV 合并,压缩成 1 个 "桥接 token"。固定输出长度 $s/m + 1$ 每 rank。
- Stage 2: 所有 rank all-gather 自己的压缩 KV 序列 + bridging token,select-and-pad 重组成完整压缩序列,attention 在这个完整序列上做。
每个 rank 输入 $s$ 个 token,本地能压成 $s/m$ 个完整压缩 token 加上 1 个 bridging(含本地末尾 + 邻居首端 $m$ 个 KV 合并而成):
"+1" 是 bridging token,它代表跨 rank 边界的 $m$ 个 token 联合压缩。这条精算让所有 rank 输出长度一致,方便 stage 2 的 all-gather。
4. Tensor 级 Activation Checkpointing
传统 module 级 checkpointing 通常按模块边界保存输入,反向时重算模块内部。mHC 引入了更多细粒度中间张量,因此 V4 改用选择性的 tensor-level checkpointing,在显存与重算之间取舍。报告没有给出“重算占约 10%”的数字。
V4 使用选择性 checkpoint,在保存与重算之间取平衡。报告明确给出的策略是:
- 重算大部分层间 hidden state;
- 重算所有 normalized layer input;
- 避免重算计算密集的操作。
报告没有披露总体显存节省比例、重算 wall-time 或是否由图变换系统自动选择边界。
本章小结
- Muon 要完整梯度,ZeRO 切分梯度 → V4 把"稠密限制 $P_{\max}$ + MoE 不限"二分;通信精度 BF16 SR 减半带宽。
- 融合 kernel、选择性 checkpoint 与 DualPipe overlap 共同把 mHC 的 wall-time overhead 限制在 6.7%。
- CSA/HCA 压缩边界让传统 CP 无法切,two-stage CP(邻居桥接 + all-gather)固定输出 $s/m + 1$ 解决。
- 选择性 checkpoint 重算大部分 hidden state 与 normalized input,同时避开计算密集操作。
- 这些机制分别处理优化器切分、残差路径开销和长上下文并行问题。