Skip to content

混合精度分布式训练实践

混合精度训练使用低精度(FP16/BF16)进行前向和反向传播,同时保持高精度(FP32)的优化器状态,实现训练加速和显存节省。

混合精度分布式训练

FP16 vs BF16

特性FP16BF16
指数位5 位8 位
尾数位10 位7 位
动态范围2^(-14) ~ 2^152^(-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 AllReduceBF16 AllReduce

归约精度

梯度 AllReduce 的精度影响训练稳定性。对于 FP16,建议使用 FP32 归约后再转回 FP16。BF16 的动态范围足够,直接使用 BF16 归约通常安全。

相关资源

最近更新