混合精度训练怎样选择参数

GPU
小华
2026-08-06

混合精度训练(Mixed Precision Training)通常结合使用 FP16(或 BF16)FP32,在保持数值稳定性的同时提升训练速度和节省显存。下面从 核心参数选择、常见框架配置、注意事项 三个层面来讲。

一、混合精度训练中的关键参数

1️⃣ 精度格式选择(最重要)

精度适用场景稳定性
FP16(float16)GPU(V100 / T4 / A100)中等,需要 loss scaling
BF16(bfloat16)TPU / A100 / H100更稳定,动态范围接近 FP32
FP32不稳定层、BatchNorm、EMA标准

推荐默认组合:

  • 主计算:FP16 / BF16
  • 主权重 & 优化器状态:FP32
  • BatchNorm:FP32

2️⃣ Loss Scaling(仅 FP16 必须)

目的:防止 FP16 下梯度下溢(变成 0)

关键参数:

参数含义
init_scale初始缩放因子
growth_factor稳定时放大
backoff_factor溢出时缩小
growth_interval多少次连续无溢出才增长

PyTorch 默认值(通常不用改):

GradScaler(
init_scale=2**16,
growth_factor=2.0,
backoff_factor=0.5,
growth_interval=2000
)

经验选择

  • 通常 不用手动调
  • 若频繁 NaN:
  • 降低 init_scale
  • 检查模型结构(尤其 softmax / log)

3️⃣ 哪些层必须保持 FP32?

建议保持 FP32:

  • BatchNorm / LayerNorm
  • Softmax + CrossEntropy(部分情况下)
  • EMA(指数移动平均)
  • 数值敏感的小模型

可放心 FP16:

  • Conv / Linear
  • Attention
  • Dropout

二、常见框架下的参数配置示例


✅ PyTorch(AMP)

基本用法

from torch.cuda.amp import autocast, GradScaler
scaler = GradScaler()
for x, y in dataloader:
optimizer.zero_grad()
with autocast():
y_hat = model(x)
loss = criterion(y_hat, y)
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()

自定义参数

GradScaler(
init_scale=2**10,        # 如有 NaN
growth_interval=500
)

✅ PyTorch + BF16(无 scaler)

with torch.autocast(device_type='cuda', dtype=torch.bfloat16):
y_hat = model(x)
loss = criterion(y_hat, y)
loss.backward()
optimizer.step()

更稳定,推荐新 GPU

✅ TensorFlow / Keras

from tensorflow.keras.mixed_precision import set_global_policy
set_global_policy('mixed_float16')

可选:

set_global_policy('mixed_bfloat16')

✅ Hugging Face Transformers

TrainingArguments(
fp16=True,              # 或 bf16=True
fp16_opt_level="O1",    # 默认即可
)

三、如何判断参数是否选对了?

✅ 正常表现

  • 收敛曲线与 FP32 基本一致
  • 无 NaN / Inf
  • 显存下降 30–50%
  • 训练速度提升 1.5x–3x

❌ 常见异常与对策

现象原因解决方案
Loss 突然 NaN梯度下溢降低 init_scale
收敛慢scale 过小增大 init_scale
精度明显下降BN / softmax 精度不足强制 FP32
完全不收敛模型不适合 FP16改用 BF16

四、经验总结(直接照着选)

无脑首选

  • A100 / H100 → BF16
  • V100 / T4 → FP16 + AMP

不调参数原则

  • 能用默认就一定用默认
  • 出问题先检查模型,再调 scaler

优先 FP32 的部分

  • 小模型
  • 数值敏感任务
  • 强化学习 / GAN

如果你愿意,可以告诉我:

  • 使用的 GPU 型号
  • 框架(PyTorch / TF / Huggingface)
  • 模型类型(CNN / Transformer / 生成模型)

我可以直接给你一套最优参数配置

亿速云提供售前/售后服务

售前业务咨询

售后技术保障

400-100-2938

7*24小时售后电话

官方微信小程序