混合精度训练(Mixed Precision Training)部署涉及训练时如何使用 FP16/BF16+FP32 来加速并省显存,以及部署/推理时如何落地混合精度模型。下面按「原理 → 训练 → 部署 → 常见坑」系统讲一遍,偏工程可落地。
一、核心原理(先搞清楚在做什么)
混合精度 ≠ 全用半精度,而是:
- 前向/反向主计算:FP16 或 BF16(快、省显存)
- 关键状态保持精度:
- 权重副本(Master Weight):FP32
- BatchNorm / LayerNorm 统计量:FP32
- Softmax / Loss 计算:FP32 或 FP32 累加
- 梯度:FP16 计算,FP32 更新
BF16:动态范围大,几乎不会溢出,训练更稳
FP16:更快但易溢出,需要 Loss Scaling
二、训练阶段怎么做(以 PyTorch 为例)
1️⃣ 使用 AMP(Automatic Mixed Precision)
from torch.cuda.amp import autocast, GradScaler
scaler = GradScaler()
for x, y in dataloader:
optimizer.zero_grad()
with autocast(dtype=torch.float16): # 或 torch.bfloat16
out = model(x)
loss = criterion(out, y)
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()
✅ 一行 autocast 自动决定哪层用半精度
2️⃣ BF16(推荐,尤其大模型)
with autocast(dtype=torch.bfloat16):
优点:
- 不需要 Loss Scaling
- 训练更稳定
- A100 / H100 / 新 GPU 原生支持
3️⃣ 分布式 + 混合精度
- DDP / FSDP 都支持 AMP
- FSDP 可设置:
mixed_precision=MixedPrecision(
param_dtype=torch.bfloat16,
reduce_dtype=torch.float32
)
三、部署 / 推理阶段怎么用混合精度
✅ 常见部署方案
1️⃣ 推理直接用 FP16 / BF16(最常见)
model = model.half().cuda() # FP16
# 或
model = model.to(torch.bfloat16).cuda()
注意:
- 输入也要对应 dtype
- 某些层(softmax、layernorm)仍内部用 FP32
2️⃣ TensorRT(极致加速)
流程:
- PyTorch → ONNX
- ONNX → TensorRT (FP16)
- 推理引擎部署
trtexec --onnx=model.onnx --fp16 --saveEngine=model.engine
✅ 延迟最低,适合生产
3️⃣ 量化 + 混合精度(进阶)
- Weight FP16 + Activation INT8
- 或 W8A8(需要校准)
工具:
- TensorRT
- PyTorch Quantization
- vLLM / TensorRT-LLM(LLM 专用)
四、训练 → 部署的完整流程建议
训练(BF16 + FP32 master)
↓
导出模型(state_dict / ONNX)
↓
转换部署格式(TorchScript / ONNX / Engine)
↓
推理使用 FP16/BF16
五、常见坑(非常重要)
❌ 1. Loss 变成 NaN
- FP16 + 无 Loss Scaling
- 解决:用 BF16 或开启 GradScaler
❌ 2. 部署精度下降明显
- 某些层强行 FP16
- 解决:Norm / Softmax 保持 FP32
❌ 3. CPU 不支持 BF16
❌ 4. ONNX 导出 dtype 不一致
六、不同场景推荐方案
| 场景 | 推荐 |
|---|
| 大模型训练 | BF16 + FSDP |
| 普通 CV/NLP | FP16 + AMP |
| 线上推理 | FP16 / BF16 |
| 极致延迟 | TensorRT FP16 |
| 边缘设备 | INT8 + FP16 |
如果你愿意,可以告诉我:
- 用的框架(PyTorch / TensorFlow / JAX)
- 是 训练加速 还是 推理部署
- 模型类型(LLM / CNN / Transformer)
我可以直接给你可复制的配置或部署脚本。