混合精度训练
混合精度训练
用 FP16/BF16 算前向反向(快、省显存),用 FP32 做累加和参数更新(准)。理解混合精度是理解为什么训练能吃这么多显存的关键。从 FP16 Loss Scaling 到黑科技 FP8 Transformer Engine,再到 “零成本加速” TF32,一文讲透。
1. Why Mixed Precision — 显存的故事
纯 FP32 训练:每个参数占 4 bytes,纯 FP16 占 2 bytes。对于 LLaMA-70B:
| 精度 | 参数显存 | 梯度显存 | 优化器显存(Adam) | 总显存(估算) |
|---|---|---|---|---|
| 纯 FP32 | 280 GB | 280 GB | 560 GB(m + v) | ~1120 GB |
| 纯 FP16 | 140 GB | 140 GB | 280 GB | ~560 GB |
| 混合精度(FP16 master + FP32 optimizer) | 140 GB | 140 GB | 560 GB(FP32) | ~840 GB |
结论:混合精度不是单纯省一半,而是在可接受精度损失下,让训练成为可能。纯 FP32 训练 70B 模型需要超过 1 TB 显存,即使用 8×A100-80GB 也无法容纳——必须用混合精度。
详见 显存计算详解
2. How It Actually Works — 训练循环拆解
混合精度的核心设计:FP16 权重副本用于前向/反向,FP32 主副本用于优化器更新。
权重存储与同步
FP32 Master Weights (W_32) ← 优化器更新在这里
│
▼ 每次 step 前转换
FP16 Working Copy (W_16) ← 前向/反向在这里
完整训练循环
for each batch:
┌─────────────────────────────────────────────┐
│ 1. W_16 = W_32.to(FP16) # 拷贝+转换 │
│ 2. L_16 = forward(X, W_16) # FP16 前向 │
│ 3. L_32 = L_16.to(FP32) # 提升精度 │
│ 4. L_scaled = L_32 * scale # Loss Scaling │
│ 5. G_16 = backward(L_scaled) # FP16 反向 │
│ 6. G_32 = G_16.to(FP32) # 梯度升精度 │
│ 7. G_32 = G_32 / scale # Unscale │
│ 8. W_32 = optimizer.step(G_32)# FP32 更新 │
└─────────────────────────────────────────────┘
为什么必须是 FP32 累加?
FP16 只有 10 位尾数,多次加法后舍入误差会累积。假设 4096 个 token 的梯度累加:
FP16: sum = 0.0
sum += 1e-7 # 4096 次
最终 sum ≈ 0.0 ← 每次加法都被吞掉了
FP32: sum = 1e-7 * 4096 = 4.096e-4 ← 正确
因此,所有累加操作(梯度 accumulation、softmax 内部求和、LayerNorm 内部求和)都必须在 FP32 下进行。
3. FP16 vs BF16 深入对比
Bit Layout 对比
FP32: [S][ E (8-bit) ][ M (23-bit) ]
FP16: [S][ E (5-bit) ][ M (10-bit) ]
BF16: [S][ E (8-bit) ][ M (7-bit) ]
| 属性 | FP16 | BF16 | FP32 |
|---|---|---|---|
| 总位数 | 16 | 16 | 32 |
| 符号位 | 1 | 1 | 1 |
| 指数位 | 5 | 8 ← 与 FP32 相同 | 8 |
| 尾数位 | 10 | 7 | 23 |
| 数值范围 | ~6.5×10⁻⁵ 至 6.5×10⁴ | ~1.2×10⁻³⁸ 至 3.4×10³⁸ | ≈1.2×10⁻³⁸ 至 3.4×10³⁸ |
| 最小正数 | 6.0×10⁻⁸ | 1.2×10⁻³⁸ | 1.2×10⁻³⁸ |
核心洞察:BF16 牺牲了尾数精度(7-bit vs 10-bit),换取了与 FP32 相同的指数范围(8 位指数)。这意味着 BF16 可以表示极小的梯度值而不会下溢——FP16 需要 Loss Scaling,BF16 不需要。
关键案例:小梯度下溢
梯度值 g = 2⁻³⁰ ≈ 9.31×10⁻¹⁰
FP16: 最小可表示 ≈ 6.0×10⁻⁸
→ g 变成 0(下溢,梯度信息丢失)
BF16: 最小可表示 ≈ 1.2×10⁻³⁸
→ g 正确保存为 ~2⁻³⁰(但尾数截断到 7-bit)
FP32: g 完全保存,精度无损
实际影响:训练大模型时,某些参数的梯度天然很小(如 embedding 层的低频 token、深层 decoder 的梯度残差)。FP16 下这些梯度直接消失 → 对应参数不再更新 → 训练进入 “blind spot” → loss 停滞甚至 NaN。
为什么 BF16 是训练的默认选择
- 无需 Loss Scaling:范围与 FP32 相同,梯度永不溢出/下溢
- 实现简单:直接
.to(torch.bfloat16)即可,少一个超参(scale factor) - 支撑 ICLR 论文:Google 在 TPU 上大量使用 BF16 训练 PaLM、Gemma 等模型
- A100 开始全支持:A100/H100/B200 的 Tensor Core 原生支持 BF16 运算
4. Loss Scaling 详解
Loss Scaling 是 FP16 时代的 “补丁”,解决梯度表示范围不足的问题。虽然 BF16 时代不需要它,但理解其原理对理解数值稳定性至关重要。
不用 Loss Scaling 会怎样?
小梯度 (|g| < 2⁻¹⁴) → FP16 无法表示 → 梯度归零
→ 参数停止更新 → 这些维度的 loss 不下降
→ 训练器误以为已收敛 → 实际在 "盲区" 里打转
严重情况:
某层出现巨大梯度 (|g| > 65504) → FP16 上溢
→ 梯度变成 +inf → 参数变成 NaN
→ optimizer.step() 把 NaN 播到所有层
→ 整个模型崩溃,必须从 checkpoint 恢复
Static Loss Scaling
手动选择固定放大系数(如 2¹² = 4096):
scale = 4096 # 固定值
for batch in dataloader:
with torch.autocast("cuda", dtype=torch.float16):
loss = model(batch)
# 手动 scale
scaled_loss = loss * scale
scaled_loss.backward()
# 手动 unscale + step
for param in model.parameters():
param.grad.data.div_(scale)
optimizer.step()
问题:scale 太小 → 下溢仍发生;scale 太大 → 上溢。手工调参痛苦。
Dynamic Loss Scaling(PyTorch GradScaler)
PyTorch 提供自动调整的 GradScaler:
from torch.cuda.amp import GradScaler, autocast
scaler = GradScaler(init_scale=2**16) # 初始 scale = 65536
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4)
for batch in dataloader:
optimizer.zero_grad()
with autocast(device_type="cuda", dtype=torch.float16):
output = model(batch["input"])
loss = criterion(output, batch["target"])
# GradScaler 自动处理 scale/backward/unscale
scaler.scale(loss).backward()
# 如果本次 step 没有 inf/NaN,scale 增大(乘 growth_factor)
# 如果检测到 inf/NaN,跳过本次 step 并减小 scale(除 backoff_factor)
scaler.step(optimizer)
scaler.update()
动态调整逻辑:
grads 无 inf/NaN → optimizer.step() + scale *= 2.0 (放大,加速尝试)
grads 有 inf/NaN → skip step + scale /= 2.0 (缩小,避免溢出)
scale 在 [min_scale, max_scale] 之间动态漂移
默认参数:init_scale=2^16, growth_factor=2.0, backoff_factor=0.5, growth_interval=2000。
5. FP8 on H100 — 新一代精度
NVIDIA H100 引入 FP8 硬件支持(Transformer Engine)。与 FP16/BF16 相比,显存再省一半。
FP8 的两种格式
E4M3 (Forward): [S][E(4)][M(3)] 范围 ≈ ±448, 精度 2^-3
E5M2 (Backward): [S][E(5)][M(2)] 范围 ≈ ±57344, 精度 2^-2
- E4M3:更窄范围但更多尾数位 → 适合前向传播(值域范围可控)
- E5M2:更宽范围但更少尾数位 → 适合反向传播(梯度范围跨度大,需防溢出)
Transformer Engine 集成
FP8 需要特殊的量化和反量化逻辑:
import transformer_engine.pytorch as te
from transformer_engine.common.recipe import Format, DelayedScaling
# 替换标准 Linear 层
class FP8TransformerLayer(torch.nn.Module):
def __init__(self, hidden_size, ffn_size, num_heads):
super().__init__()
self.attention = te.Linear(hidden_size, 3 * hidden_size)
self.proj = te.Linear(hidden_size, hidden_size)
self.ffn_1 = te.Linear(hidden_size, ffn_size)
self.ffn_2 = te.Linear(ffn_size, hidden_size)
def forward(self, x):
with te.fp8_autocast(enabled=True):
# 内部自动: FP8 quant → compute → FP8 dequant
# 每层独立计算 scaling factor
attn_out = self.attention(x) # FP8 matmul
out = self.proj(attn_out)
ffn_out = self.ffn_1(out)
out = self.ffn_2(ffn_out)
return out
# 多卡训练配置
fp8_format = Format.HYBRID # E4M3 forward, E5M2 backward
关键机制:Transformer Engine 按 tensor 粒度动态计算 scale factor,把每个 tensor 量化到 FP8 的表示范围内,避免截断误差。Scale factor 本身就是训练的一部分,通过 delayed scaling 策略反向传播。
FP8 收益
| 指标 | BF16 | FP8 |
|---|---|---|
| 参数显存 | 140 GB | 70 GB |
| 梯度显存 | 140 GB | 70 GB |
| 计算吞吐 | 1000 TFLOPS | 2000 TFLOPS |
| 代码改动 | 零 | 需 TE 替换 Linear 层 |
6. TF32 on A100 — 零代码改动的”免费午餐”
TF32(TensorFloat-32)是 A100 Tensor Core 内部使用的 19-bit 格式:
TF32: [S][E(8)][M(10)] — 与 FP32 相同的指��范围,FP16 级别的尾数
自动启用
import torch
# A100 上默认已启用
torch.backends.cuda.matmul.allow_tf32 = True # 控制 matmul
torch.backends.cudnn.allow_tf32 = True # 控制卷积
零代码改动:只要在 A100 上跑 PyTorch >= 1.7,矩阵乘法和卷积自动使用 TF32。
TF32 精度分析
普通 matmul:
FP16 A × FP16 B → FP32 累加 → FP16 输出
(输入和输出是 FP16,中间累加用 FP32 精度)
TF32 matmul:
FP32 A → truncate M to 10-bit → TF32 A
FP32 B → truncate M to 10-bit → TF32 B
TF32 A × TF32 B → FP32 累加 → FP32 输出
(输入尾数截断到 10-bit,但输出维持 FP32)
效果:相比 FP32 matmul,TF32 快 ≈ 8×;精度仅损失 13 位尾数(23 → 10),在大部分训练任务中几乎无精度损失。
局限
- 仅覆盖 matmul 和卷积,element-wise 操作仍是 FP32
- 需要 Ampere 及以上架构(A100/A6000/3090+/H100)
- 对于小型训练任务可能不明显,但在大 batch 训练中加速显著
7. Practical Guide — 什么场景用什么精度
按训练场景
| 场景 | 推荐精度 | 原因 |
|---|---|---|
| 从零预训练(>1B) | BF16 + TF32 | 范围安全、无 Loss Scaling、TF32 自动加速 matmul |
| 从零预训练(>70B, H100) | FP8 + TE | 显存省一半、吞吐翻倍,FP8 的精度损失可控 |
| 微调(LoRA/Full) | BF16 / FP16 | 微调在小 batch 下通常无梯度溢出风险 |
| 微调(H100 + 长上下文) | FP8 | 长序列 KV Cache 吃显存,FP8 是关键 |
| 推理 | INT8 / INT4 | 推理不需要梯度,量化后精度损失 < 0.5% |
| Embedding 模型训练 | TF32 / BF16 | Embedding 小梯度多,TF32 尾数优势明显 |
| RLHF(Reward Model) | BF16 | 小模型 + 稳定训练,BF16 省心 |
按硬件
| GPU | 最佳精度 | 备注 |
|---|---|---|
| V100 | FP16 + Loss Scaling | 无 BF16 TF Core 支持 |
| A100 | BF16 + TF32 | 原生 BF16 Tensor Core |
| A100-SXM 80GB | BF16 + TF32 | 80GB 给混合精度更多余量 |
| H100 | FP8 + Transformer Engine | FP8 吞吐是 BF16 的 2× |
| B200 | FP4 (推理) / FP8 (训练) | Blackwell FP4 原生支持 |
常见坑
- FP16 下掉点 NaN → 检查 Loss Scale 是否初始化太小。调大
init_scale或改用 BF16 - BF16 下精度下降 → 检查 LayerNorm/RMSNorm 是否开 FP32。Norm 层天然需要高精度
- FP8 训练不收敛 → 切换 E5M2 做前向。部分模型对 FP8 精度敏感
- Mixed Precision + Gradient Accumulation → 记得在
optimizer.step()前 unscale。否则 accumulation 的梯度被错误缩放
关联知识
学习时间
| 阶段 | 时间 | 备注 |
|---|---|---|
| 骨架创建 | 2026-06-30 | 框架搭建 |
| 内容完善 | 2026-06-30 | 七板块完整覆盖 |
状态标记
📖 已掌握 — FP16/BF16/FP8 格式对比与选型策略、Loss Scaling 原理与 PyTorch 实现、TF32 “免费加速”机制
📝 待补充 — FP8 在 Megatron-LM / DeepSpeed 的端到端实战配置、各精度在 A100/H100 的实测吞吐 Benchmark、FP4 训练(Blackwell)前瞻分析