文章

混合精度训练

混合精度训练

用 FP16/BF16 算前向反向(快、省显存),用 FP32 做累加和参数更新(准)。理解混合精度是理解为什么训练能吃这么多显存的关键。从 FP16 Loss Scaling 到黑科技 FP8 Transformer Engine,再到 “零成本加速” TF32,一文讲透。


1. Why Mixed Precision — 显存的故事

纯 FP32 训练:每个参数占 4 bytes,纯 FP16 占 2 bytes。对于 LLaMA-70B:

精度参数显存梯度显存优化器显存(Adam)总显存(估算)
纯 FP32280 GB280 GB560 GB(m + v)~1120 GB
纯 FP16140 GB140 GB280 GB~560 GB
混合精度(FP16 master + FP32 optimizer)140 GB140 GB560 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)  ]
属性FP16BF16FP32
总位数161632
符号位111
指数位58 ← 与 FP32 相同8
尾数位10723
数值范围~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 是训练的默认选择

  1. 无需 Loss Scaling:范围与 FP32 相同,梯度永不溢出/下溢
  2. 实现简单:直接 .to(torch.bfloat16) 即可,少一个超参(scale factor)
  3. 支撑 ICLR 论文:Google 在 TPU 上大量使用 BF16 训练 PaLM、Gemma 等模型
  4. 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 收益

指标BF16FP8
参数显存140 GB70 GB
梯度显存140 GB70 GB
计算吞吐1000 TFLOPS2000 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 / BF16Embedding 小梯度多,TF32 尾数优势明显
RLHF(Reward Model)BF16小模型 + 稳定训练,BF16 省心

按硬件

GPU最佳精度备注
V100FP16 + Loss Scaling无 BF16 TF Core 支持
A100BF16 + TF32原生 BF16 Tensor Core
A100-SXM 80GBBF16 + TF3280GB 给混合精度更多余量
H100FP8 + Transformer EngineFP8 吞吐是 BF16 的 2×
B200FP4 (推理) / FP8 (训练)Blackwell FP4 原生支持

常见坑

  1. FP16 下掉点 NaN → 检查 Loss Scale 是否初始化太小。调大 init_scale 或改用 BF16
  2. BF16 下精度下降 → 检查 LayerNorm/RMSNorm 是否开 FP32。Norm 层天然需要高精度
  3. FP8 训练不收敛 → 切换 E5M2 做前向。部分模型对 FP8 精度敏感
  4. 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)前瞻分析