torch.autocast 混合精度加速:bfloat16 vs float16 怎么选

torch.autocast 是 PyTorch 混合精度(AMP)的核心 API——把部分算子从 FP32 切到 BF16 / FP16,推理速度和训练吞吐都能翻倍,显存也能降 30-50%。

最简用法

推理

import torch

model = model.cuda().eval()

with torch.no_grad(), torch.autocast("cuda", dtype=torch.bfloat16):
    output = model(input.cuda())

训练

model = model.cuda().train()
optimizer = torch.optim.AdamW(model.parameters())

for x, y in loader:
    x, y = x.cuda(), y.cuda()
    optimizer.zero_grad()

    with torch.autocast("cuda", dtype=torch.bfloat16):
        pred = model(x)
        loss = loss_fn(pred, y)

    loss.backward()
    optimizer.step()

关键autocast 只包在 forward + loss 计算里,backwardoptimizer.step() 在外面。

autocast 到底改了什么

进入 autocast 上下文后,PyTorch 会自动选择每个算子的精度

算子类别会被降精度
Linear / Matmul✅ 用 BF16/FP16
Conv2d / Conv3d✅ 用 BF16/FP16
GEMM 类操作✅ 用 BF16/FP16
Softmax❌ 保持 FP32(防溢出)
LayerNorm❌ 保持 FP32
Loss❌ 保持 FP32
Reduction❌ 保持 FP32

PyTorch 有个内置白名单/黑名单,你不用手动 .half()、不用改模型结构。

bfloat16 vs float16

BF16(bfloat16)

  • 动态范围和 FP32 一样(8 位 exponent),不会像 FP16 那样容易溢出/下溢
  • 不需要 GradScaler(训练时不用缩放梯度)
  • ✅ 稳定性接近 FP32,一般不影响 loss
  • ⚠️ 需要 Ampere 及以上(A100、RTX 30 系及以后)——旧卡不支持
  • ⚠️ 比 FP16 精度低(7 位 mantissa)

FP16(float16)

  • ✅ 所有现代 GPU 都支持(Volta 起)
  • ✅ 稍快一点(在某些卡上)
  • ⚠️ 容易 NaN(动态范围窄)
  • ⚠️ 必须配 torch.cuda.amp.GradScaler()
  • ⚠️ 大模型容易训崩

该选谁

硬件训练推理
A100 / H100BF16BF16
RTX 30 / 40 系BF16BF16
RTX 20 系 / V100FP16 + GradScalerFP16
GTX 10 系及以下不支持,用 FP32不支持

没有理由用 FP16 训练时选 BF16——除非你的卡不支持。

训练带 GradScaler(FP16 必须)

scaler = torch.cuda.amp.GradScaler()

for x, y in loader:
    optimizer.zero_grad()
    with torch.autocast("cuda", dtype=torch.float16):
        pred = model(x)
        loss = loss_fn(pred, y)

    scaler.scale(loss).backward()
    scaler.step(optimizer)
    scaler.update()

GradScaler 把 loss 放大再反传,避免小梯度被 FP16 归零。BF16 用了反而干扰训练,别加。

常见坑

1. 输入没上 GPU

with torch.autocast("cuda", ...):
    output = model(input)   # input 在 CPU,报错

务必:

output = model(input.cuda())

2. 和 .half() 混用

别这样

model.half()                                # 全模型转 FP16
with torch.autocast("cuda", dtype=torch.bfloat16):
    output = model(input)

冲突。要么全模型 .half() 不用 autocast,要么用 autocast 不 .half()。混合精度靠 autocast 就够了。

3. CPU 上 autocast 不等价

with torch.autocast("cpu", dtype=torch.bfloat16):

CPU 上 autocast 支持有限,别期望和 GPU 一样效果。CPU 推理想加速走 torch.compile 或 ONNX Runtime。

4. 自定义 op 不支持 BF16

第三方 op / native extension 如果没写 BF16 kernel,autocast 会 fallback 到 FP32——不是报错,只是没加速效果。

5. 推理不用 GradScaler

# 推理
with torch.no_grad(), torch.autocast("cuda", dtype=torch.bfloat16):
    output = model(x)

torch.no_grad()autocast 一起用最省显存。

加速效果参考

Transformer 类模型(LLaMA 7B、Qwen 8B 之类)BF16 vs FP32:

  • 推理速度:1.8 – 2.3x
  • 显存占用:约 50%
  • 精度损失:几乎无(bf16 动态范围和 fp32 一致)

CNN 训练(YOLO 之类):

  • 训练吞吐:1.3 – 1.7x
  • 显存节省:30 – 45%
  • mAP:几乎无差

一句话总结

Ampere 及以上 GPU 就用 torch.autocast("cuda", dtype=torch.bfloat16),无需 GradScaler、稳定性接近 FP32。老卡 / V100 才考虑 FP16 + GradScaler。别和 model.half() 混用。