混合精度与通信——BF16 为什么是 LLM 训练的主流选择 / Mixed Precision and Communication with BF16
📅 创建时间:2026-06-02 🏷️ 标签:#混合精度 #BF16 #FP8 #FP16 #NCCL #通信优化 📚 前置知识:[[01-gpu-hardware]](GPU 硬件基础) [[02-distributed-training]](分布式训练) 📚 相关知识:[[03-memory-optimization]](显存优化) [[05-pretraining]](预训练稳定性)
场景:FP16 训练 Loss 溢出,换 BF16 后稳定了
┌─────────────────────────────────────────────────────────────────┐
│ │
│ 凌晨 3 点,你的训练脚本报错了。 │
│ │
│ 第一次尝试:FP16 训练 │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ Loss: 3.45 → 12.78 → nan │ │
│ │ 梯度爆炸了!Loss 直接溢出变成 NaN。 │ │
│ │ 排查:learning rate、初始化、FP16 动态范围 │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ 第二次尝试:换成 BF16 训练 │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ Loss: 3.45 → 3.21 → 2.98 → 2.75 │ │
│ │ 稳定下降,没有溢出。 │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ 为什么同一个模型,FP16 会溢出,BF16 就不会? │
│ 这就是本章要回答的问题。 │
│ │
└─────────────────────────────────────────────────────────┘第1节:浮点数基础——为什么精度用 16 位就够了
浮点数的结构
┌─────────────────────────────────────────────────────────────┐
│ 浮点数的三部分构成 │
├─────────────────────────────────────────────────────────────┤
│ │
│ 一个浮点数由三部分组成: │
│ │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ (-1)^sign × (1 + fraction) × 2^exponent │ │
│ │ 1 bit 10/8/7 bits 5/8/7 bits │ │
│ │ │ │
│ │ sign = 符号位(正 or 负) │ │
│ │ exponent = 指数位(表示数量级,10^x 中的 x) │ │
│ │ fraction = 尾数位(表示精度,1.xxxx 中的 xxxx) │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ 动态范围(Dynamic Range)= 能表示的最大数和最小数之间的比值 │
│ 精度(Precision)= 两个相邻可表示数之间的最小间隔 │
│ │
└─────────────────────────────────────────────────────────────┘FP16 vs BF16 vs FP32——结构对比
┌─────────────────────────────────────────────────────────────┐
│ FP16 vs BF16 vs FP32 内存布局对比 │
├─────────────────────────────────────────────────────────────┤
│ │
│ FP32(单精度,4 字节 = 32 bits): │
│ ┌────────┬──────────┬───────────────────┐ │
│ │ Sign │ Exponent │ Fraction(尾数) │ │
│ │ 1 bit │ 8 bits │ 23 bits │ │
│ │ 正/负 │ 2^e │ 精度位 │ │
│ └────────┴──────────┴───────────────────┘ │
│ 动态范围:~10^38 (和 BF16 一样大) │
│ 精度:~7 位十进制小数 │
│ │
│ FP16(半精度,2 字节 = 16 bits): │
│ ┌────────┬──────────┬───────────────────┐ │
│ │ Sign │ Exponent │ Fraction(尾数) │ │
│ │ 1 bit │ 5 bits │ 10 bits │ │
│ └────────┴──────────┴───────────────────┘ │
│ 动态范围:~65504 (最大约 6.5×10^4) │
│ 精度:~4 位十进制小数 │
│ │
│ BF16(Brain Float,2 字节 = 16 bits): │
│ ┌────────┬──────────┬───────────────────┐ │
│ │ Sign │ Exponent │ Fraction(尾数) │ │
│ │ 1 bit │ 8 bits │ 7 bits │ │
│ └────────┴──────────┴───────────────────┘ │
│ 动态范围:~10^38 (和 FP32 一样大!) │
│ 精度:~3 位十进制小数 │
│ │
│ 关键差异: │
│ FP16 把 16 bits 分给指数(5)和尾数(10) │
│ BF16 把 16 bits 中 8 bits 给指数,7 bits 给尾数 │
│ → BF16 的指数和 FP32 一样多,动态范围相同 │
│ → BF16 只比 FP16 少了 3 bits 的精度,但动态范围大得多 │
│ │
└─────────────────────────────────────────────────────────────┘为什么 FP16 会溢出
┌─────────────────────────────────────────────────────────────┐
│ FP16 溢出 vs BF16 稳定的根本原因 │
├─────────────────────────────────────────────────────────────┤
│ │
│ 问题出在动态范围: │
│ │
│ FP16 的最大表示数: │
│ (1 + 0.1111111111)₂ × 2^30 = 1.999 × 2^30 ≈ 65504 │
│ → 任何超过 65504 的数字,在 FP16 中就会变成 Infinity │
│ │
│ LLM 训练中的大数字来源: │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ 1. 损失值大:Cross-Entropy Loss 早期可能是 10-20 │ │
│ │ 2. 梯度累积:多层链式乘法,梯度可能非常大 │ │
│ │ 3. 权重数值:LayerNorm、Attention 等的中间值 │ │
│ │ 4. Loss Scaling:动态损失缩放翻车时更容易溢出 │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ BF16 的最大表示数: │
│ (1 + 0.1111111)₂ × 2^254 ≈ 3.4 × 10^38 │
│ → 和 FP32 一样大,训练中的任何中间值都不会溢出 │
│ │
│ 结论: │
│ BF16 = FP32 的动态范围 + FP16 的显存效率 │
│ 这就是为什么 2020 年后,所有 LLM 训练都迁移到了 BF16 │
│ │
└─────────────────────────────────────────────────────────────┘第2节:混合精度训练原理——什么时候用 FP32
什么是混合精度训练
┌─────────────────────────────────────────────────────────────┐
│ 混合精度训练:不同步骤用不同精度 │
├─────────────────────────────────────────────────────────────┤
│ │
│ 核心观察: │
│ → Forward 和 Backward 用 BF16 足够(动态范围大,不溢出) │
│ → Optimizer State 用 FP32 更好(精度高,更新更准确) │
│ │
│ 混合精度训练的工作流程: │
│ │
│ Forward(BF16): │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ Input (FP32) │ │
│ │ ↓ │ │
│ │ 模型权重转 BF16 │ │
│ │ ↓ │ │
│ │ Forward 计算(全部 BF16)→ Output │ │
│ │ ↓ │ │
│ │ Loss 计算(BF16) │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ Backward(BF16): │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ Loss (BF16) │ │
│ │ ↓ │ │
│ │ 反向传播(全部 BF16)→ 梯度(BF16) │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ 参数更新(FP32): │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ 梯度 BF16 → 转 FP32 │ │
│ │ ↓ │ │
│ │ Optimizer State (FP32) + 梯度 (FP32) → 更新 │ │
│ │ ↓ │ │
│ │ 更新后的权重(FP32)→ 转 BF16(用于下一步 Forward)│ │
│ └─────────────────────────────────────────────────────┘ │
│ │
└─────────────────────────────────────────────────────────────┘PyTorch 混合精度实现
python
from torch.cuda.amp import autocast, GradScaler
# GradScaler:管理 loss scale,防止 BF16 下溢出
scaler = GradScaler()
for data, labels in dataloader:
data = data.cuda()
labels = labels.cuda()
# Forward:用 autocast 自动把支持的操作转为 BF16/FP16
with autocast(dtype=torch.bfloat16):
outputs = model(data)
loss = loss_fn(outputs, labels)
# Backward:用 scaler.scale 防止梯度溢出
scaler.scale(loss).backward()
# 参数更新:用 scaler.step(内部做了 unscaling)
scaler.step(optimizer)
scaler.update() # 动态调整 loss scaleLoss Scaling——FP16 时代的补救方案
┌─────────────────────────────────────────────────────────────┐
│ Loss Scaling 的原理(FP16 时代的技术) │
├─────────────────────────────────────────────────────────────┤
│ │
│ 背景:FP16 动态范围小,容易溢出 │
│ │
│ 核心思想: │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ 梯度值太小(小于 FP16 最小精度)? │ │
│ │ → 乘以一个大的 scale(S),变成 FP16 可表示的数 │ │
│ │ → 正常更新 │ │
│ │ → 除以 S,恢复到原始尺度 │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ 动态 Loss Scaling: │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ 如果当前 batch 的梯度全是 0(溢出变 Inf): │ │
│ │ → Loss scale 翻倍(2× S) │ │
│ │ → 重算 │ │
│ │ 如果连续 N 个 batch 没有溢出: │ │
│ │ → Loss scale 减半(S/2) │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ BF16 不需要 Loss Scaling(动态范围够大) │
│ 但 GradScaler 在 BF16 训练中仍然有用(极端情况保护) │
│ │
└─────────────────────────────────────────────────────────────┘第3节:FP8——H100 的下一代精度
FP8 的两种格式
┌─────────────────────────────────────────────────────────────┐
│ FP8:H100 的 8 位浮点精度 │
├─────────────────────────────────────────────────────────────┤
│ │
│ H100 支持两种 FP8 格式(NVIDIA 发明): │
│ │
│ E4M3(4 位指数 + 3 位尾数): │
│ ┌────────┬──────────┬───────────────────┐ │
│ │ Sign │ Exponent │ Fraction │ │
│ │ 1 bit │ 4 bits │ 3 bits │ │
│ └────────┴──────────┴───────────────────┘ │
│ 范围:±448(最大),精度约 2 位小数 │
│ → 用于:Forward 计算(精度要求相对较低) │
│ │
│ E5M2(5 位指数 + 2 位尾数): │
│ ┌────────┬──────────┬───────────────────┐ │
│ │ Sign │ Exponent │ Fraction │ │
│ │ 1 bit │ 5 bits │ 2 bits │ │
│ └────────┴──────────┴───────────────────┘ │
│ 范围:±57344(最大),精度约 1 位小数 │
│ → 用于:Backward 计算(动态范围更重要) │
│ │
│ 为什么 Forward 和 Backward 用不同精度? │
│ Forward:需要高精度(模型权重对精度敏感) │
│ Backward:梯度动态范围大,精度要求较低 │
│ │
└─────────────────────────────────────────────────────────────┘FP8 在 LLM 训练中的应用
┌─────────────────────────────────────────────────────────────┐
│ FP8 训练的实际应用 │
├─────────────────────────────────────────────────────────────┤
│ │
│ FP8 的速度优势: │
│ H100 FP8 算力:989 TFLOPS(稀疏后 1978 TFLOPS) │
│ H100 BF16 算力:989 TFLOPS │
│ H100 FP16 算力:989 TFLOPS │
│ → FP8 峰值算力和 BF16 一样,但显存减半! │
│ │
│ 但 FP8 的挑战: │
│ 1. 精度损失:动态范围更小,需要精细的 scaling │
│ 2. 实现复杂度:Transformer Engine 是关键 │
│ 3. 不是所有算子都支持 FP8(Softmax、LayerNorm 等仍需 BF16)│
│ │
│ 典型配置(2024-2025 年): │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ Forward:FP8(矩阵乘法)+ BF16(Softmax/LayerNorm)│ │
│ │ Backward:FP8(梯度计算)+ BF16(通信) │ │
│ │ Optimizer:FP32(保持高精度) │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ Transformer Engine(TE): │
│ → 自动处理精度转换 + 动态 scaling │
│ → DeepSpeed 和 Megatron 都集成了 TE │
│ │
└─────────────────────────────────────────────────────────────┘FP16 / BF16 / FP8 精度选择对比
┌─────────────────────────────────────────────────────────────┐
│ 训练精度选择全景图 │
├─────────────────────────────────────────────────────────────┤
│ │
│ │ 精度 │ 显存占用 │ 动态范围 │ 训练稳定性 │ 速度 │
│ ├─────────┼──────────┼──────────┼────────────┼─────────┤
│ │ FP32 │ 100% │ 最大 │ 最稳定 │ 最慢 │
│ ├─────────┼──────────┼──────────┼────────────┼─────────┤
│ │ BF16 │ 50% │ 最大 │ 稳定 │ 快 │
│ ├─────────┼──────────┼──────────┼────────────┼─────────┤
│ │ FP16 │ 50% │ 小 │ 不稳定(需LS)│ 快 │
│ ├─────────┼──────────┼──────────┼────────────┼─────────┤
│ │ FP8 │ 25% │ 中 │ 较稳定 │ 最快 │
│ │
│ 2020-2023:BF16 是主流选择(GPT-4 之前) │
│ 2024-2025:FP8 开始在 H100 上普及(训练速度优势) │
│ FP16 在现代 LLM 训练中已基本被淘汰 │
│ │
└─────────────────────────────────────────────────────────────┘第4节:NCCL 通信原语——分布式训练的通信基础
为什么分布式训练需要通信
┌─────────────────────────────────────────────────────────────┐
│ 分布式训练中的通信需求 │
├─────────────────────────────────────────────────────────────┤
│ │
│ 数据并行(DP): │
│ 每张卡独立计算梯度 → 需要同步 → AllReduce │
│ │
│ 张量并行(TP): │
│ 每张卡算一部分矩阵乘法 → 需要聚合 → AllReduce │
│ │
│ 流水线并行(PP): │
│ GPU 之间的激活值和梯度传递 → 点对点通信(P2P) │
│ │
│ FSDP/ZeRO-3: │
│ 参数分片 → 需要时广播 → AllGather │
│ │
│ 通信 = 分布式训练的隐形瓶颈 │
│ 通信慢 → GPU 等待 → MFU 下降 │
│ │
└─────────────────────────────────────────────────────────────┘三大通信原语
┌─────────────────────────────────────────────────────────────┐
│ AllReduce / AllGather / ReduceScatter │
├─────────────────────────────────────────────────────────────┤
│ │
│ AllReduce(最常用): │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ │ │
│ │ GPU 0: [1, 2, 3, 4] AllReduce → [10, 20, 30, 40] │
│ │ GPU 1: [2, 4, 6, 8] ────────────────→ 每张卡都一样 │
│ │ GPU 2: [3, 6, 9, 12] │ │
│ │ GPU 3: [4, 8, 12, 16] │ │
│ │ │ │
│ │ 操作:对应位置求和/平均/最大值等 │ │
│ │ 用途:数据并行梯度同步 │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ AllGather(广播): │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ │ │
│ │ GPU 0: [1, 2] │ │
│ │ GPU 1: [3, 4] AllGather → [1, 2, 3, 4] │ │
│ │ GPU 2: [5, 6] ──────────────→ 每张卡都一样 │ │
│ │ GPU 3: [7, 8] │ │
│ │ │ │
│ │ 操作:收集所有 GPU 的数据,拼成完整结果 │ │
│ │ 用途:ZeRO-3 参数分片广播 │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ ReduceScatter(分片归约): │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ │ │
│ │ GPU 0: [1, 2, 3, 4] │ │
│ │ GPU 1: [2, 4, 6, 8] ReduceScatter → GPU0:[3,6,9,12] │
│ │ GPU 2: [3, 6, 9, 12] ──────────────→ GPU1:[6,12,18,24] │
│ │ GPU 3: [4, 8, 12, 16] → GPU2:[9,18,27,36] │
│ │ → GPU3:[12,24,36,48] │
│ │ │ │
│ │ 操作:先求和,然后分片发给各 GPU │ │
│ │ 用途:ZeRO 梯度分片收集 │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
└─────────────────────────────────────────────────────────────┘NCCL——GPU 集合通信的标准库
┌─────────────────────────────────────────────────────────────┐
│ NCCL:NVIDIA Collective Communications Library │
├─────────────────────────────────────────────────────────────┤
│ │
│ NCCL 是 NVIDIA 开发的 GPU 间通信库: │
│ → 所有主流分布式训练框架都在用(DeepSpeed / Megatron / FSDP)│
│ → 支持 GPUDirect RDMA(绕过 CPU 直接通信) │
│ → 自动选择最优通信路径(NVLink / IB / PCIe) │
│ │
│ NCCL Ring vs Tree: │
│ │
│ Ring AllReduce(链路环形): │
│ GPU0 → GPU1 → GPU2 → GPU3 → GPU0 │
│ → 通信量:2×(N-1)/N × data(次优但稳定) │
│ → 适合:跨节点通信(IB) │
│ │
│ Tree AllReduce(树形): │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ GPU0 │ │
│ │ / \ │ │
│ │ GPU1 GPU2 │ │
│ │ | | │ │
│ │ GPU3 GPU4 ... │ │
│ └─────────────────────────────────────────────────────┘ │
│ → 通信量:(N-1)/N × data(最优) │
│ → 适合:节点内通信(NVLink) │
│ │
│ NCCL 自动选择:节点内用 Tree,跨节点用 Ring │
│ │
└─────────────────────────────────────────────────────────────┘通信与计算的重叠——隐藏延迟
┌─────────────────────────────────────────────────────────────┐
│ 通信与计算重叠:让 GPU 不等待 │
├─────────────────────────────────────────────────────────────┤
│ │
│ 朴素做法(串行): │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ Forward ──────────────────▶ Backward ──────────▶ │ │
│ │ 全部完成后 → AllReduce ─────────────────────────▶ │ │
│ │ ↑ GPU 在这里空等 │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ 重叠做法(并行): │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ Micro-batch 1: F1 ──▶ B1 ──▶ AllReduce │ │
│ │ Micro-batch 2: F2 ──▶ B2 (和上一步 AllReduce 并行) │ │
│ │ Micro-batch 3: F3 ──▶ B3 (和上一步 AllReduce 并行) │ │
│ │ │ │
│ │ 通信时间被计算隐藏 → 总时间大幅减少 │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ 在 FSDP 中的应用: │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ 当前 micro-batch: AllGather 参数 │ │
│ │ 上一个 micro-batch: Backward 计算(并行) │ │
│ │ 下一个 micro-batch: Forward 计算(并行) │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ 实现关键:CUDA Stream(见 [[01-gpu-hardware]]) │
│ │
└─────────────────────────────────────────────────────────────┘升华:精度与效率的永恒博弈
┌─────────────────────────────────────────────────────────────┐
│ 从 FP16 到 BF16 到 FP8:精度选择的工程哲学 │
├─────────────────────────────────────────────────────────────┤
│ │
│ 1. 精度损失的接受度在提高 │
│ → FP32 → BF16:动态范围换取稳定性(可接受) │
│ → BF16 → FP8:精度换取速度(可接受) │
│ → 趋势:能用低精度就不用高精度,能省就省 │
│ │
│ 2. 通信永远是瓶颈 │
│ → GPU 算力增长 > 显存带宽增长 > 互联带宽增长 │
│ → 越往后,通信越可能成为瓶颈 │
│ → AllReduce 的优化(Ring vs Tree vs 拓扑感知)是工程重点│
│ │
│ 3. BF16 是当前训练的主流,FP8 是未来 │
│ → H100 的 FP8 硬件支持已经成熟 │
│ → 但框架支持(Transformer Engine)还在完善中 │
│ → 消费级显卡(4090)不支持 BF16 以外的新精度 │
│ │
│ 一句话总结: │
│ 混合精度训练的本质是在"精度够用"的边界上,尽可能节省显存和计算。│
│ │
└─────────────────────────────────────────────────────────────┘"AI 可查 vs 必须理解"清单
AI 可查:
✅ Transformer Engine 的具体配置参数
✅ NCCL 通信的 benchmark 数据
✅ 不同精度在不同 GPU 型号上的支持情况
必须理解:
🔴 FP16 / BF16 / FP32 的内存布局差异,以及为什么 BF16 比 FP16 更适合 LLM 训练
🔴 混合精度训练的流程:Forward BF16 + Optimizer FP32
🔴 FP8 的 E4M3 / E5M2 两种格式及其适用场景
🔴 NCCL 三大通信原语:AllReduce / AllGather / ReduceScatter 的使用场景
🔴 为什么通信与计算重叠能提升 MFU
🔴 为什么张量并行需要 NVLink(带宽需求),而数据并行可以用 InfiniBand学习状态:🟡 开始学习