混合精度分布式训练实践
混合精度训练使用低精度(FP16/BF16)进行前向和反向传播,同时保持高精度(FP32)的优化器状态,实现训练加速和显存节省。
FP16 vs BF16
| 特性 | FP16 | BF16 |
|---|---|---|
| 指数位 | 5 位 | 8 位 |
| 尾数位 | 10 位 | 7 位 |
| 动态范围 | 2^(-14) ~ 2^15 | 2^(-126) ~ 2^127 |
| 精度 | 高 | 较低 |
| Loss Scaling | 需要 | 不需要 |
| 硬件支持 | V100+ | A100+ |
选择建议
A100 及更新硬件推荐使用 BF16:不需要 Loss Scaling,训练更稳定。V100 不支持 BF16,只能使用 FP16。
PyTorch 混合精度
python
from torch.cuda.amp import autocast, GradScaler
model = MyModel().cuda()
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4)
scaler = GradScaler() # FP16 需要,BF16 不需要
for data, target in dataloader:
optimizer.zero_grad()
# FP16 自动混合精度
with autocast(dtype=torch.float16):
loss = model(data.cuda(), target.cuda())
# 缩放梯度防止下溢
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()BF16 训练
python
# BF16 不需要 GradScaler
model = MyModel().to(torch.bfloat16)
for data, target in dataloader:
optimizer.zero_grad()
with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
loss = model(data.cuda(), target.cuda())
loss.backward()
optimizer.step()DeepSpeed 混合精度
python
# DeepSpeed BF16 配置
ds_config = {
"bf16": {
"enabled": True,
},
"zero_optimization": {
"stage": 2,
},
}分布式训练中的精度问题
梯度同步精度
python
# FSDP 混合精度配置
from torch.distributed.fsdp import MixedPrecision
mp_policy = MixedPrecision(
param_dtype=torch.bfloat16, # 参数存储精度
reduce_dtype=torch.bfloat16, # 梯度归约精度
buffer_dtype=torch.bfloat16, # Buffer 精度
)
model = FSDP(model, mixed_precision=mp_policy)精度稳定性
| 问题 | FP16 解决方案 | BF16 解决方案 |
|---|---|---|
| 梯度下溢 | Loss Scaling | 不需要 |
| 溢出 | 降低 Loss Scale | 不需要 |
| 精度损失 | FP32 优化器状态 | FP32 优化器状态 |
| 归约精度 | FP32 AllReduce | BF16 AllReduce |
归约精度
梯度 AllReduce 的精度影响训练稳定性。对于 FP16,建议使用 FP32 归约后再转回 FP16。BF16 的动态范围足够,直接使用 BF16 归约通常安全。