跳到主要内容

L3.3 主流框架与工具实践

三维坐标 layer: L3(训练)level: Engineer / Seniorpillar: 训推框架

L3.1 讲清了「为什么要分布式」(显存墙 + 算力墙),L3.2 讲清了「并行策略的数学」(DP/TP/PP/ZeRO)。本文回答最现实的问题:这些策略到底用哪个框架落地、怎么调、怎么把训了一半的模型安全存下来再恢复。目标是让你从「会跑 demo」走到「能在多卡上选对框架、调对 config、扛住中断」。

学习目标

  • 前置知识:读过 L3.1(分布式训练的动机:显存墙 + 算力墙、单卡装不下的根因)与 L3.2(并行策略的数学:DP / TP / PP / SP / ZeRO 三阶段如何切参数/梯度/优化器状态);写过基础 PyTorch 训练循环;知道「多卡要通信」即可。无需手写过 CUDA kernel。
  • 学完产出:① 能用一张表说清 DDP / FSDP / DeepSpeed 在「单个 step 里搬了什么」(参数/梯度/优化器状态的分片粒度与通信量),并据此估算单卡显存;② 能照「选型决策树」走完一次真实决策——给定模型规模与硬件,判断该用 DDP / FSDP2 / DeepSpeed ZeRO-3 / Megatron TP,并说清各自代价;③ 能读懂并调一份 DeepSpeed ZeRO-3 + offload 的 JSON config,知道 stage、offload、bucket size 各自换的是显存还是吞吐;④ 能讲清 BF16 / FP16 / FP8 的数值取舍,说出 FP8 训练「哪些张量绝不能降精度」的铁律;⑤ 亲手用 FSDP + torch.distributed.checkpoint 落一份分片 checkpoint 并重启恢复,验证权重/优化器/step 一致。
  • 阅读姿势:盯住一条主线——「框架选型与调优,本质都是在『显存』『通信』『工程易用性』这三者之间做权衡」。无论是 ZeRO 分片、CPU offload、3D 并行还是混合精度,每一个旋钮拧动时都在问同一个问题:我愿意用多少通信/复杂度,去换多少显存或吞吐?

背景与现状

大模型训练框架的竞争,本质是 「显存利用率」与「工程易用性」的权衡之争。同一份 ZeRO-3 风格的全分片思想,在不同框架里有截然不同的工程取舍:

  • DDP(DistributedDataParallel):PyTorch 原生的纯数据并行。每张卡持有完整的模型副本,只对梯度做 All-Reduce。简单、稳定、通信模式可预测——但模型必须装得进单卡,是「小模型多卡加速」的事实标准。
  • FSDP(FullyShardedDataParallel):PyTorch 原生的 ZeRO-3 风格全分片(FSDP1 API)。2026 新基线FSDP2(fully_shard + DTensor):per-parameter 分片、与 torch.compile/TP 组合更顺,官方推荐新项目优先 FSDP2;FSDP1 仍广泛存在于存量代码。
  • DeepSpeed:微软出品,以 ZeRO(Zero Redundancy Optimizer) 闻名。除分片外提供极其丰富的特性:CPU/NVMe offload、ZeRO-Infinity、稀疏注意力、MoE、量化训练、流水线并行——是「特性最全、配置最重」的一极。
  • Megatron-Core:NVIDIA 出品,张量并行(TP)+ 流水线并行(PP)+ 序列并行(SP) 的工业级实现,为「单卡装不下一层」的超大模型(百亿~万亿)而生,常与 DeepSpeed/FSDP 组合成 3D 并行。

从产业演进看:

  • 2020–2022:DeepSpeed/ZeRO 一家独大,是「让大模型训得动」的破局者。
  • 2022–2024:Megatron-LM 的 TP/PP 成为千卡级训练底座;FSDP 作为原生方案快速追赶。
  • 2024 至今:FSDP2 + DTensor + 分布式 checkpoint 让「原生路径」工程体验逼近 DeepSpeed,社区出现明显的「去重依赖、回归 PyTorch 原生」趋势;而 FP8 训练(H100 Transformer Engine)成为新的性能分水岭。

业界信号:HuggingFace accelerate / transformers 同时把 FSDP 与 DeepSpeed 作为一等公民后端,torchtitan(PyTorch 官方大模型训练参考)全面押注 FSDP2 + DTensor——这说明「选哪个框架」已不再有唯一答案,而是取决于模型规模、硬件代际与团队工程能力的工程决策。

原理与架构

2.1 三大框架的分片粒度与通信模式

理解选型,先看三者在**「单个训练 step 里搬了什么」**:

维度DDPFSDP (ZeRO-3 风格)DeepSpeed ZeRO-1/2/3
参数 (P)每卡全量副本沿 rank 全分片,用时 All-GatherStage 3 分片,否则全量
梯度 (G)All-Reduce 后全量分片,Reduce-ScatterStage 2/3 分片
优化器状态 (OS)每卡全量全分片Stage 1/2/3 分片
单卡显存高(3 份冗余)最低(≈1/N)随 stage 递减
通信量/step1× All-ReduceAll-Gather + Reduce-Scatter(≈1.5×)同 FSDP 量级
外部依赖无(原生)无(原生)deepspeed 库 + config
offload有限(CPU offload)强(CPU/NVMe,ZeRO-Infinity)
典型场景模型可入单卡模型超单卡、想留在原生栈极致显存压榨 / 特性需求多

核心洞察:DDP 用「3 倍显存冗余」换「最简单的通信」;FSDP/ZeRO-3 用「1.5 倍通信」换「N 倍显存节省」。当模型装不进单卡,分片是唯一出路;通信开销则靠 overlap(计算与通信重叠) 和高速互联(NVLink)来摊薄。

2.2 FSDP 的分片与按需还原机制

逐层读这张图:FSDP 把模型按「FSDP unit」(通常是 transformer block)切分。前向走到某个 unit 时 All-Gather 还原它的完整参数 → 算完该 unit 的 forward → 立即释放非本 rank 的参数分片(峰值显存只多一个 unit 的量,而非整个模型)。反向同理,梯度算完用 Reduce-Scatter 让每个 rank 只保留自己负责的那一片。这正是 ZeRO-3 的精髓:用通信换显存,且通信与计算重叠

2.3 选型决策树

2.4 DeepSpeed config 调优 playbook

DeepSpeed 的威力全在 JSON config。一份典型 ZeRO-3 + offload 配置与调优要点:

{
"train_micro_batch_size_per_gpu": 4,
"gradient_accumulation_steps": 8,
"gradient_clipping": 1.0,
"bf16": { "enabled": true },
"zero_optimization": {
"stage": 3,
"offload_optimizer": { "device": "cpu", "pin_memory": true },
"offload_param": { "device": "cpu", "pin_memory": true },
"overlap_comm": true,
"contiguous_gradients": true,
"reduce_bucket_size": 5e8,
"stage3_prefetch_bucket_size": 5e8,
"stage3_param_persistence_threshold": 1e6,
"stage3_max_live_parameters": 1e9,
"stage3_gather_16bit_weights_on_model_save": true
}
}

调优心法(按优先级)

  1. 先选对 stage:显存够 → stage 1(只分优化器状态,通信最省);不够 → stage 2(再分梯度);还不够 → stage 3(连参数也分)。每升一级换更多显存、付更多通信
  2. offload 是最后的救命稻草,不是首选offload_optimizer/param 到 CPU/NVMe 能再省一大块显存,但 PCIe 带宽会成为瓶颈,吞吐可能腰斩。显存够就别开
  3. gradient_accumulation_steps:用「小 micro-batch × 累积步数」凑大 global batch,是单卡显存不足时撑大 batch 的标准手段。global_batch = micro_bsz × accum × world_size。
  4. overlap_comm: true + contiguous_gradients: true:让通信与反向计算重叠、梯度内存连续,几乎是免费的吞吐提升,默认就该开。
  5. reduce_bucket_size / stage3_prefetch_bucket_size:调大减少通信次数(更高吞吐)、调小省显存——典型显存与速度的旋钮。
  6. 保存权重务必开 stage3_gather_16bit_weights_on_model_save,否则导出的是分片碎权重。

2.5 Megatron-Core 并行实践要点

Megatron-Core 解决的是「单层就装不进单卡」的极端情形,核心是 张量并行(TP):把一个 Linear 的权重矩阵沿列/行切到多卡,前向各算一片、用 All-Reduce/All-Gather 拼回。配合 流水线并行(PP) 把不同 layer 放到不同卡、序列并行(SP) 切 LayerNorm/Dropout 的激活,组成 3D 并行。实践三原则:

  • TP 只在单机内(NVLink 域)开:TP 通信密集(每层都通信),跨机 PCIe/IB 会拖垮性能,tp_size 一般 ≤ 单机 GPU 数(如 8)。
  • PP 跨机:PP 通信稀疏(只在 stage 边界传激活),适合跨节点,但要靠 interleaved 1F1B 调度压缩流水线 bubble。
  • DP 在最外层:TP×PP 确定单副本布局后,用 DP 复制多份扩展 batch。total_gpus = TP × PP × DP

2.6 混合精度:BF16 / FP16 / FP8 策略

格式位宽动态范围精度关键特性适用
FP1616窄(易上溢/下溢)高(10 位尾数)必须配 loss scaling 防梯度下溢老硬件 (V100)
BF1616与 FP32 同宽低(7 位尾数)无需 loss scaling,数值更稳A100/H100 首选
FP88E4M3/E5M2 两种很低per-tensor scaling + Transformer EngineH100+ 极致吞吐

数值稳定要点

  • 优先 BF16:动态范围与 FP32 一致,几乎不会溢出,省掉 FP16 那套 loss scaling 的调试地狱,是现代训练默认。
  • FP16 必开动态 loss scaling:梯度乘一个大 scale 防下溢,溢出则回退缩小——DeepSpeed/AMP 已内置,但仍需监控 skipped step
  • FP8 是混合的混合:通常只对 GEMM 的输入用 FP8,累加(accumulate)仍在 FP32master weight 与优化器状态保留 BF16/FP32,否则训练发散。FP8 收益主要在 H100 的算力翻倍,但对 scaling factor 校准极敏感。
  • 永远保留 FP32 master weight:混合精度的铁律——计算用低精度,权重更新累加在高精度副本上,否则小梯度被「吃掉」。

2.7 分布式 checkpoint 与恢复

千卡训练动辄数周,节点故障是必然事件,checkpoint 是 L3 的生命线。朴素做法(rank0 收集全量权重再存)在大模型上会 OOM 且极慢。现代方案是 torch.distributed.checkpoint(DCP,sharded state dict)

  • 每个 rank 只存自己持有的那片,并行写盘,无单点收集 → 又快又不 OOM。
  • 存储格式与并行布局解耦:DCP 存的是「逻辑张量 + 分片元数据」,恢复时可在不同 world_size / 不同并行度下重新切分加载(resharding),这是它最强的特性——今天 8 卡训的,明天能 16 卡接着训。
  • 恢复必须同时存取 model + optimizer + 训练进度(step/epoch/RNG state),否则 LR scheduler 错位、数据重复。

2.8 Scaling Laws 与超参 Sweeps

Scaling Laws(缩放定律) 是 L3 调参的「先验地图」:

  • Chinchilla 最优:在固定算力预算 C 下,模型参数量 N 与训练 token 数 D 应约 1:20 配比D ≈ 20N),而非一味堆大模型。它直接决定「给定卡时,该练多大模型、喂多少数据」。
  • Loss ∝ 幂律L(N) ≈ (Nc/N)^α,可用小规模实验外推大模型的最终 loss,避免盲目烧卡。
  • 超参可迁移(µP / µTransfer):在小模型上 sweep 出的最优 LR 等超参,经合适参数化后可零成本迁移到大模型,把 sweep 成本降几个数量级。

Sweeps 实践:用 W&B Sweeps / Optuna 等做贝叶斯/网格搜索,优先 sweep LR、warmup、batch size、weight decay;务必在小规模 + 短步数上 sweep,再把赢家放大——这正是 scaling laws 给的底气。

动手实践:极简代码实操

实验目标:用 PyTorch 原生 FSDP 包一个小 transformer,微调几个 step,然后用 torch.distributed.checkpoint 落一份分片 checkpoint,重启进程加载恢复并验证权重/优化器/step 一致——亲手走完「分布式训练 → 容错存档 → 安全恢复」的完整闭环。产出物:两段日志(训练后 loss + 恢复后对比)证明 checkpoint 可恢复。

3.1 环境准备

# 推荐 Python 3.11;用 uv 或 venv 隔离
python3 -m venv .venv && source .venv/bin/activate
pip install torch --index-url https://download.pytorch.org/whl/cpu # CPU/Mac 默认路径
# NVIDIA GPU 路径:pip install torch (自动选 CUDA wheel)

路径说明:本实验CPU(gloo 后端)即可完整跑通 FSDP + 分布式 checkpoint,无需 GPU。有 NVIDIA GPU 时把后端切到 nccl、设备切到 cuda 即享受真实加速,代码逻辑完全一致。

3.2 代码:FSDP 微调 + 落分片 checkpoint + 恢复验证

# fsdp_ckpt.py —— torchrun --nproc_per_node=2 fsdp_ckpt.py
# 单卡环境:torchrun --nproc_per_node=1 fsdp_ckpt.py(PyTorch 2.x FSDP 需 GPU)
import os, torch, torch.nn as nn, torch.distributed as dist
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
import torch.distributed.checkpoint as dcp
from torch.distributed.checkpoint.state_dict import (
get_state_dict, set_state_dict, StateDictOptions,
)

CKPT_DIR = "ckpt_fsdp"

def setup():
dist.init_process_group(backend="gloo") # GPU 改 "nccl"
# CPU 路径无需 set_device;GPU 路径:torch.cuda.set_device(local_rank)

class TinyTransformer(nn.Module):
def __init__(self, d=64, layers=4):
super().__init__()
self.emb = nn.Linear(16, d)
self.blocks = nn.ModuleList([
nn.TransformerEncoderLayer(d, nhead=4, dim_feedforward=128, batch_first=True)
for _ in range(layers)
])
self.head = nn.Linear(d, 1)
def forward(self, x):
x = self.emb(x)
for b in self.blocks:
x = b(x)
return self.head(x).mean(dim=1)

def train_steps(model, opt, steps=5):
rank = dist.get_rank()
for s in range(steps):
x = torch.randn(8, 10, 16)
y = torch.randn(8, 1)
opt.zero_grad()
loss = ((model(x) - y) ** 2).mean()
loss.backward()
opt.step()
if rank == 0:
print(f"[train] step {s} loss={loss.item():.4f}")

def main():
setup()
rank = dist.get_rank()
torch.manual_seed(0) # 所有 rank 同初始化

model = FSDP(TinyTransformer())
opt = torch.optim.AdamW(model.parameters(), lr=1e-3)

# ---- 训练几步 ----
train_steps(model, opt, steps=5)

# ---- 落分片 checkpoint(每个 rank 只存自己那片)----
msd, osd = get_state_dict(model, opt)
dcp.save({"model": msd, "optim": osd}, checkpoint_id=CKPT_DIR)
dist.barrier()
if rank == 0:
print(f"[save] sharded checkpoint -> {CKPT_DIR}/")

# ---- 模拟重启:新建模型/优化器,从 checkpoint 恢复 ----
model2 = FSDP(TinyTransformer())
opt2 = torch.optim.AdamW(model2.parameters(), lr=1e-3)
msd2, osd2 = get_state_dict(model2, opt2)
dcp.load({"model": msd2, "optim": osd2}, checkpoint_id=CKPT_DIR)
set_state_dict(model2, opt2, model_state_dict=msd2, optim_state_dict=osd2)

# ---- 验证:恢复后的本地分片参数应与原模型逐元素一致 ----
p1 = next(model.parameters()).detach()
p2 = next(model2.parameters()).detach()
same = torch.allclose(p1, p2, atol=1e-6)
if rank == 0:
print(f"[verify] params identical after restore: {same}")
dist.destroy_process_group()

if __name__ == "__main__":
main()

3.3 运行与观察

# CPU / Mac 路径:2 个进程模拟 2 卡,gloo 后端
torchrun --nproc_per_node=2 fsdp_ckpt.py

# 单卡 GPU(6GB 等):FSDP 需 GPU,用 1 进程即可验证 checkpoint 流程
# torchrun --nproc_per_node=1 fsdp_ckpt.py

# NVIDIA GPU 路径(多卡):把 backend 改成 nccl 后同样命令
# torchrun --nproc_per_node=<GPU数> fsdp_ckpt.py
  • 预期输出:先看到 5 行 [train] step N loss=...(loss 逐步下降)→ [save] sharded checkpoint -> ckpt_fsdp/ → 关键的 [verify] params identical after restore: True
  • ckpt_fsdp/ 目录看:会发现是多个分片文件 + 一个 .metadata,而非单个大权重文件——这正是 sharded state dict「每 rank 存一片、并行写盘」的实证。
  • resharding 验证(进阶):把 --nproc_per_node 从 2 改成 1 再 dcp.load,依然能加载成功——证明 DCP 存储格式与并行度解耦,今天 2 卡存的明天 1 卡能接着训。

踩坑预警 (Gotchas)

  • 不能用裸 model.state_dict() 存 FSDP:那拿到的是当前 rank 的分片碎片,直接存会损坏。必须走 get_state_dict(DCP 配套 API)或配置 FullStateDictConfig,否则恢复时张量形状对不上。
  • 恢复必须先 dcp.load 写回 dict、再 set_state_dict 灌回模型——两步缺一不可。只 load 不 set,模型参数纹丝不动。
  • 优化器状态别忘了:只存 model 不存 optim,恢复后 Adam 的一阶/二阶动量清零,等于「热启动变冷启动」,loss 会抖动甚至发散。生产中还要存 step / RNG state / LR scheduler
  • CPU 路径用 gloo,GPU 用 nccl:在无 GPU 机器上误设 nccl 会直接报错;GPU 上用 gloo 则慢到没意义。
  • 所有 rank 必须同初始化种子torch.manual_seed 要在建模型前对齐,否则各 rank 分片来自不同初始权重,分片语义错乱。
  • save 后加 dist.barrier():避免快的 rank 先去读还没写完的 checkpoint。

深入思考

下面三题每题先给题干,再用 <details> 折叠一份图文并茂的参考答案。建议先合上答案自己想 3 分钟,再展开对照。

思考题 1:13B 模型在 8×A100-40G 上的选型权衡

你的团队要训一个 13B 模型,硬件是 8×A100-40G 单机。模型 FP32 权重 ≈ 52GB、加梯度与 Adam 优化器状态总需 ≈ 200GB+。请用 2.3 节的选型决策树 论证:为什么 DDP 必然 OOM?在 FSDP(ZeRO-3)和 DeepSpeed ZeRO-3 + CPU offload 之间,你如何根据「吞吐 vs 显存」做取舍?什么情况下你会被迫引入 Megatron 的 TP?

展开参考答案(含显存预算-选型决策流程图 + 算一遍)

结论:DDP 因为每卡要装完整的 200GB+ 副本而必然 OOM;只要单卡装得下「整模型的 1/8 分片 + 一层临时还原」,FSDP 与 DeepSpeed ZeRO-3 都能不 offload 跑通且吞吐更高;只有当某一层本身就超过单卡显存时,才被迫引入 Megatron 的张量并行。

用具体数字算一遍(混合精度 + Adam,按每参数显存口径估算):

  1. 13B 参数,Adam 混合精度的经典口径约 16 字节/参数(FP32 master 权重 4 + 一阶动量 4 + 二阶动量 4 + BF16 权重 2 + BF16 梯度 2)≈ 208GB
  2. DDP:每卡各持一份完整 208GB → 单卡 40GB 直接爆,且没有分片机制可省,必然 OOM。这对应决策树第一问「模型能装进单卡吗?」答「不能」。
  3. FSDP / ZeRO-3 全分片:208GB / 8 卡 ≈ 26GB/卡,再加上前向时临时 All-Gather 还原「一个 transformer block」的参数(峰值只多一层,非整模型)+ 激活,单卡 40GB 大体可控(必要时配激活重算)。这是首选——不 offload,吞吐最高
  4. 何时上 offload:若 26GB + 激活 + 想要的大 batch 顶破 40GB,再开 DeepSpeed ZeRO-3 的 CPU offload,把优化器状态甩到主机内存——显存立刻松一大块,但 PCIe 带宽成为瓶颈,吞吐可能腰斩。offload 是救命稻草不是首选
  5. 何时上 Megatron TP:13B 的单个 Linear 权重远小于单卡,不需要 TP。只有当模型放大到「单层就装不进单卡」(百亿~万亿级的超宽层)时,才被迫把单层矩阵沿行/列切到多卡,且 TP 只在 NVLink 域内开。

取舍一句话:显存够 → FSDP2/ZeRO-3 不 offload(图省事 + 原生栈选 FSDP2,要特性/极致省显存选 DeepSpeed);显存差一口气 → ZeRO-3 + CPU offload 换吞吐;单层爆卡 → 才动 Megatron TP。

思考题 2:千卡训练的 checkpoint 策略设计

一个千卡训练任务平均每 6 小时挂一个节点。请设计一套 checkpoint 策略:多久存一次(结合存档耗时与重算损失的权衡)、为什么必须用 sharded state dict 而非 rank0 收集、恢复时如何保证「数据不重复、LR scheduler 不错位」?如果故障后只剩一半节点可用,2.7 节 的 DCP resharding 能力如何救你?

展开参考答案(含故障-恢复时间线图 + 间隔权衡算一遍)

结论:checkpoint 间隔应取「单次存档耗时」与「平均重算损失」的平衡点——存太勤浪费 I/O、存太疏一次崩溃赔掉几小时;必须用 sharded state dict 才能并行写盘、不在 rank0 OOM;恢复时连同 step / 数据迭代位置 / RNG / LR scheduler 一起存取,才能不重复、不错位;而 DCP 的 resharding 让你在节点减半后仍能用新 world_size 直接续训。

间隔怎么定,算一遍(数量级估算):

取值(示例)说明
平均故障间隔 MTBF6 小时题设
单次 sharded 存档耗时~1 分钟每 rank 并行写自己那片,故很快
若每 30 分钟存一次存档开销 ≈ 1/30 ≈ 3.3%崩溃平均损失 ≈ 15 分钟重算
若每 3 小时存一次存档开销 ≈ 0.6%崩溃平均损失 ≈ 90 分钟重算

经验法则:让「存档总开销」与「期望重算损失」量级相当。MTBF 6h、单次存档仅 ~1 分钟时,每 15~30 分钟 存一次较合理——存档开销几个百分点,换来单次崩溃只赔十几分钟。存档越快(sharded 并行写就是为了这个),就越敢存得勤。

为什么必须 sharded state dict:rank0 收集全量的朴素做法要把千卡的所有分片汇聚到一张卡——大模型上 rank0 直接 OOM,且单点串行写盘慢到离谱(存一次几十分钟,把训练拖死)。sharded state dict 让每个 rank 只存自己持有的那片、并行写盘,无单点收集,又快又不 OOM(对应 2.7 节)。

恢复时不重复、不错位:checkpoint 必须是「model + optimizer + 训练进度」的原子三件套——

  • optimizer 状态:不存则 Adam 动量清零,热启动变冷启动,loss 抖动。
  • step / 数据迭代位置:不存则 dataloader 从头来,数据重复(已学过的样本再学一遍,污染分布)。
  • LR scheduler 进度 / RNG state:不存则学习率回到错误档位、随机性不一致,scheduler 错位导致 loss 突变。

节点减半时 DCP 如何救你:DCP 存的是「逻辑张量 + 分片元数据」,存储格式与并行布局解耦。原来 1000 卡存的 checkpoint,故障后只剩 500 卡可用时,直接用新的 world_size dcp.load 即可——DCP 会在加载时按新并行度重新切分(resharding),无需任何离线转换。这就是「今天 N 卡存的,明天 M 卡接着训」的容错底气。

思考题 3:FP16→BF16 不再 NaN 的归因与 FP8 铁律

同一份训练脚本,从 FP16 切到 BF16 后「莫名其妙不再 NaN 了」,但有人坚持用 FP16 + loss scaling。请结合 2.6 节 从「动态范围 vs 尾数精度」解释这个现象。再进一步:要在 H100 上用 FP8 把吞吐再翻倍,你必须守住哪几条数值稳定铁律(哪些张量绝不能降到 FP8)?

展开参考答案(含三种精度数值轴对比图 + 对比表)

结论:BF16 的指数位与 FP32 一样多,动态范围一致,梯度几乎不会下溢/上溢成 NaN,代价是尾数少、精度低;FP16 指数位少、范围窄,必须靠 loss scaling 把梯度搬进可表示窗口才不发散;而 FP8 范围更窄、可表示值更少,只能让 GEMM 输入吃 FP8,master 权重、优化器状态与累加器绝不能降到 FP8。

为什么 BF16 不再 NaN:NaN 通常源于梯度下溢成 0 后参与除法上溢成 Inf。FP16 只有 5 位指数,最小正规数约 6e-5、最大约 65504——梯度量级一旦跌出这个窄窗口就下溢成 0、激活偏大就上溢成 Inf,反传出 NaN。BF16 有 8 位指数(与 FP32 完全相同),动态范围一致,梯度几乎不可能溢出,所以「莫名其妙不再 NaN」。代价是 BF16 只有 7 位尾数(FP16 是 10 位),精度更低——但训练对动态范围远比对尾数精度敏感,这笔买卖划算,所以 BF16 成现代默认。

为什么有人仍坚持 FP16 + loss scaling:FP16 的 10 位尾数精度更高,在数值范围可控的场景(如某些推理或老硬件 V100 无原生 BF16)仍有价值;loss scaling 把整个 loss/梯度乘一个大 scale,把小梯度「平移」进 FP16 可表示窗口,溢出则自动回退缩小——本质是用一个动态旋钮弥补 FP16 范围窄的先天缺陷。

精度指数位尾数位动态范围是否需 scaling不再 NaN 的根因
FP16510必须 loss scaling—(窄范围正是 NaN 来源)
BF1687与 FP32 同宽不需要范围足够宽,梯度不溢出
FP84 或 53 或 2更窄必须 per-tensor scaling仅作 GEMM 输入,配缩放因子

FP8 翻倍吞吐的数值铁律(哪些张量绝不能降到 FP8)

  1. 累加器(accumulate)必须留 FP32:FP8 只做 GEMM 的输入操作数,乘加的累加在 FP32 里进行,否则误差迅速累积。
  2. master 权重必须留 BF16/FP32:权重更新累加在高精度副本上,小梯度才不会被「吃掉」(混合精度的通用铁律)。
  3. 优化器状态(Adam 一阶/二阶动量)必须高精度:动量是长期累积量,降到 FP8 会失真导致发散。
  4. 必须配 per-tensor scaling + Transformer Engine:逐张量动态维护缩放因子(前向 E4M3、反向梯度 E5M2),把数值搬进 FP8 窗口,缺了它会大面积下溢成 0 或上溢成 Inf。

一句话:FP8 是「混合的混合」——只让最耗算力的 GEMM 输入吃 FP8,权重/动量/累加这些「记账用」的张量一律守住高精度,否则 H100 的吞吐翻倍会以训练发散为代价。

延伸阅读

1. 核心 Paper

  • ZeRO: Memory Optimizations Toward Training Trillion Parameter Models(2020)— DeepSpeed 分片三阶段的理论原点。
  • Megatron-LM: Training Multi-Billion Parameter Language Models Using Model Parallelism(2019)+ Reducing Activation Recomputation(SP,2022)— TP/PP/SP 工业级实现。
  • Training Compute-Optimal Large Language Models(Chinchilla,2022)— scaling law 的算力最优配比,决定「练多大、喂多少」。
  • FP8 Formats for Deep Learning(NVIDIA,2022)— FP8 训练的数值规范与 scaling 策略。

2. 相关高 Star 仓库与源码必读路径

  • pytorch/pytorch — 看 torch/distributed/fsdp/(FSDP 实现)与 torch/distributed/checkpoint/(DCP),理解原生分片与存档。
  • pytorch/torchtitan — PyTorch 官方大模型训练参考,FSDP2 + DTensor + DCP 的最佳实践范本,强烈建议精读。
  • microsoft/DeepSpeed — 入口看 deepspeed/runtime/zero/ 的 stage 实现与 deepspeed/runtime/engine.py
  • NVIDIA/Megatron-LM — 看 megatron/core/ 下 TP/PP/SP 的张量切分组织方式。

3. 优质博客 / 视频

  • HuggingFace「The Ultra-Scale Playbook」/ accelerate 文档中 FSDP 与 DeepSpeed 后端对照章节。
  • DeepSpeed 官方 ZeRO 系列教程与 config 参数详解。
  • Stanford CS336(Language Modeling from Scratch)的并行训练与混合精度章节。

下一篇L3.4 万卡集群弹性容错:把本文「单任务的 checkpoint 恢复」放大到千卡集群,讲清弹性训练(elastic)、故障自愈与 straggler 治理的工程体系。