AI混合精度折腾手记
混合精度不是新概念,但真正把 FP16/BF16 落地到生产环境,还是踩了不少坑。
背景:为什么需要混合精度
深度学习训练默认使用 FP32(32位单精度浮点数),这几乎是所有框架的标准配置。FP32 用 32 个比特位表示一个数字,其中 1 位符号位、8 位指数位、23 位尾数位。能表示的数值范围大约是 ±3.4×10³⁸,精度达到约 7 位有效数字。
听起来很完美,但有两个问题:
显存占用太大:训练一个模型,显存主要消耗在三个地方:模型参数、梯度、优化器状态。以 Adam 优化器为例,每个参数需要保存 2 倍参数量的动量信息,再加上参数本身,总共就是 4 倍的 FP32 存储开销。模型一大,这个数字很快就破 16G。
计算效率不高:现代 GPU 的 Tensor Core 对 FP16/BF16 有专门的硬件加速,FP32 反而跑不出满性能。比如 NVIDIA A100 的 FP16 理论算力是 FP32 的 2 倍,实际训练中配合合适的混合精度策略,性能提升可以接近这个数字。
混合精度的思路很直接:把 FP16/BF16 作为主力计算格式,FP32 仅作为备份和精度保障。
精度选择:FP16 还是 BF16
真正动手前,先得搞清楚 FP16 和 BF16 的区别。这两个都是半精度,但性格完全不同。
FP16(Half Precision)用 1 位符号位、5 位指数位、10 位尾数位表示数字,数值范围大约是 ±6.5×10⁴,精度约 3 位有效数字。它的特点是精度高,但动态范围小。小到一定程度会变成零,大到一定程度会溢出。
BF16(BFloat16)是 Google 推出来的,专为 AI 场景设计。它用 1 位符号位、8 位指数位、7 位尾数位,数值范围和 FP32 一样大(±3.4×10³⁸),但精度只有约 2 位有效数字。它的特点是动态范围大,不容易溢出,但精度牺牲明显。
我用 Python 画了一张对比图,直观感受一下差异:
import matplotlib.pyplot as plt
import numpy as np
# 设置中文显示
plt.rcParams['font.sans-serif'] = ['DejaVu Sans']
plt.rcParams['axes.unicode_minus'] = False
# 创建对比数据
formats = {
'FP32': {'exp': 8, 'mant': 23, 'range': 38},
'BF16': {'exp': 8, 'mant': 7, 'range': 38},
'FP16': {'exp': 5, 'mant': 10, 'range': 4}
}
# 创建图形
fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(12, 4))
# 左图:指数位对比
ax1.bar(formats.keys(), [v['exp'] for v in formats.values()], color=['#1f77b4', '#ff7f0e', '#2ca02c'])
ax1.set_ylabel('指数位 (Bits)')
ax1.set_title('动态范围能力(指数位)')
ax1.set_ylim(0, 10)
# 右图:尾数位对比
ax2.bar(formats.keys(), [v['mant'] for v in formats.values()], color=['#1f77b4', '#ff7f0e', '#2ca02c'])
ax2.set_ylabel('尾数位 (Bits)')
ax2.set_title('数值精度能力(尾数位)')
ax2.set_ylim(0, 25)
plt.tight_layout()
plt.savefig('/home/liqinsi/Documents/project/thinkblog/static/images/posts/274-ai-mixed-precision-fp32-fp16-bf16-practice/precision-comparison.webp', dpi=150, bbox_inches='tight')

这张图很关键:FP16 的尾数位有 10 个,比 BF16 的 7 个多,所以 FP16 在小数值上的精度更好。但 FP16 的指数位只有 5 个,远小于 BF16 和 FP32 的 8 个,这意味着 FP16 的动态范围小,容易溢出。
实际选择时,我的判断是:
如果你的显卡支持 BF16(比如 NVIDIA A100、RTX 40 系列),优先用 BF16。动态范围和 FP32 一样大,不容易出现梯度消失或爆炸的问题。精度虽然低一点,但对大部分深度学习任务来说够用。
如果只支持 FP16(比如 V100、RTX 30 系列),需要搭配 Loss Scaling。Loss Scaling 的作用是把梯度放大,避免小梯度被 FP16 的精度限制截断成零。放大后再缩回来,保持整体数值稳定。
实现:PyTorch 中的混合精度训练
PyTorch 从 1.6 版本开始就原生支持混合精度,核心是 torch.cuda.amp 模块。实现起来其实很简单,但有一些细节需要注意。
基础配置
先看一个最简单的混合精度训练循环:
import torch
from torch.cuda.amp import autocast, GradScaler
# 模型和优化器
model = MyModel().cuda()
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4)
# GradScaler 仅在 FP16 时需要
use_bf16 = torch.cuda.is_bf16_supported()
scaler = None if use_bf16 else GradScaler()
for batch in dataloader:
optimizer.zero_grad()
# 前向传播自动使用混合精度
with autocast(dtype=torch.bfloat16 if use_bf16 else torch.float16):
loss = model(batch)
# 反向传播
if use_bf16:
loss.backward()
optimizer.step()
else:
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()
这个代码片段已经能跑起来,但有几个坑点。
坑点一:GradScaler 的更新逻辑
scaler.update() 这一步很容易漏掉。它的作用是检查梯度是否正常(没有变成 NaN 或 Inf),如果正常就更新缩放因子,否则跳过这次优化器更新。
# 错误示范:不调用 update()
scaler.scale(loss).backward()
scaler.step(optimizer)
# 缺少 scaler.update() 会导致缩放因子一直不更新
# 正确写法
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()
坑点二:Loss Scaling 的初始值和增长策略
默认的 Loss Scaling 初始值是 65536,对大部分任务合适,但有些场景需要调整。比如如果你的损失值本身就很小,初始值可能太小;如果模型容易梯度爆炸,初始值可能太大。
# 自定义 Loss Scaling 策略
scaler = GradScaler(
init_scale=2.0, # 初始缩放因子,根据你的 loss 范围调整
growth_factor=2.0, # 每次成功后的增长倍数
backoff_factor=0.5, # 失败后的回退倍数
growth_interval=2000 # 连续成功多少次后才增长
)
坑点三:手动启用 BF16 的陷阱
PyTorch 的 torch.cuda.is_bf16_supported() 并不总是可靠的。有些显卡硬件上支持 BF16,但驱动或 CUDA 版本不支持。这种情况下,强制用 BF16 会报错或导致不可预测的行为。
# 更安全的 BF16 检测方式
def safe_bf16_supported():
if not torch.cuda.is_available():
return False
device = torch.cuda.current_device()
device_name = torch.cuda.get_device_name(device)
# 检查已知支持 BF16 的显卡
bf16_gpus = ['A100', 'RTX 40', 'H100', 'L40']
return any(gpu in device_name for gpu in bf16_gpus) and torch.cuda.is_bf16_supported()
use_bf16 = safe_bf16_supported()
混合精度训练的完整流程
为了更清楚地理解混合精度训练的工作原理,我用 Mermaid 画了一个完整的流程图:
这张图想说明一件事:混合精度训练不是简单地把 FP32 换成 FP16/BF16,而是一套完整的数值稳定性保障机制。FP16 需要额外的 Loss Scaling 和梯度检查,BF16 则可以直接用,省去这些步骤。
踩坑记录
真正上手后,还是遇到了几个意料之外的问题。
问题一:精度下降不明显,但训练不稳定
现象是 Loss 曲线波动变大,偶尔会出现突然飙升的情况。排查了很久,发现是某些自定义的算子没有正确支持 FP16。
# 错误示范:自定义算子没有考虑精度
def custom_attention(q, k, v):
# 默认用 FP32 计算注意力分数
scores = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(q.size(-1))
attn = torch.softmax(scores, dim=-1)
return torch.matmul(attn, v)
# 正确做法:显式处理精度
def custom_attention(q, k, v):
with autocast():
scores = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(q.size(-1))
attn = torch.softmax(scores, dim=-1)
return torch.matmul(attn, v)
把所有自定义算子都包在 autocast() 里后,训练稳定性明显提升。
问题二:显存节省不如预期
理论上混合精度能节省 50% 的显存,但实际只节省了 30% 左右。最后发现是 PyTorch 默认会保留一些中间结果用于反向传播,这些结果没有及时释放。
# 手动清理中间结果
for batch in dataloader:
optimizer.zero_grad()
with autocast():
loss = model(batch)
loss.backward()
# 清理不需要的中间结果
torch.cuda.empty_cache()
optimizer.step()
另外,显存节省的幅度也和模型结构有关。如果你的模型激活值占用很大,混合精度的显存优势会更明显;如果模型参数量特别大,优化器状态的显存占比高,混合精度的优势会被稀释。
问题三:某些层需要强制用 FP32
发现一些特殊的层在 FP16 下数值不稳定,比如 LayerNorm 的分母容易出现极小值导致除零。解决办法是在这些层强制使用 FP32。
class StableLayerNorm(nn.Module):
def __init__(self, dim, eps=1e-12):
super().__init__()
self.weight = nn.Parameter(torch.ones(dim))
self.bias = nn.Parameter(torch.zeros(dim))
self.eps = eps
def forward(self, x):
# LayerNorm 的均值和方差计算强制用 FP32
mean = x.mean(-1, keepdim=True, dtype=torch.float32)
var = x.var(-1, keepdim=True, unbiased=False, dtype=torch.float32)
x = (x - mean) / torch.sqrt(var + self.eps)
return (self.weight * x + self.bias).to(x.dtype)
实际效果测试
在一台 RTX 4090 上跑了对比测试,数据来自一个中等规模的 Transformer 模型(约 500M 参数):
| 配置 | 训练速度 | 显存占用 | 训练稳定性 | 最终精度 |
|---|---|---|---|---|
| FP32 | 1.0x | 16.2GB | 稳定 | 基准 |
| FP16 | 1.6x | 9.8GB | 需调参 | 基本持平 |
| BF16 | 1.5x | 10.5GB | 稳定 | 基本持平 |
训练速度提升 50%-60%,显存节省 35%-40%,这个收益是实实在在的。最终精度和 FP32 基本持平,有些任务甚至略好(推测是 FP16/BF16 的噪声起到了类似正则化的作用)。
结论
混合精度不是万能药,但确实是一个性价比很高的优化手段。如果你正在为显存不足或训练太慢而头疼,值得尝试一下。
总结几个关键点:
- 硬件支持的优先用 BF16,省去 Loss Scaling 的麻烦,稳定性更好。
- FP16 需要 Loss Scaling,初始值和增长策略要调一下,避免梯度消失或爆炸。
- 自定义算子要显式处理精度,不要默认假定所有操作都能正确支持混合精度。
- 显存节省幅度因模型而异,不是固定 50%,实际在 30%-40% 之间。
最后想说一句:混合精度是手段,不是目的。如果你的模型训练已经很稳定、显存也够用,没必要为了"优化"而强行上混合精度。技术选择要看实际需求,而不是看"别人都在用"。
折腾完这次混合精度,最大的感受是:很多看起来复杂的技术,本质上都是在解决限制条件。FP32 太占显存,FP16 动态范围太小,BF16 精度牺牲太多。混合精度就是把这些限制条件折中一下,找到一个既能跑得快又不容易炸的平衡点。
这就是工程实践,没有什么完美的方案,只有更适合当前场景的选择。
版权声明: 本文首发于 指尖魔法屋-AI混合精度折腾手记(https://blog.thinkmoon.cn/post/274-ai-mixed-precision-fp32-fp16-bf16-practice/) 转载或引用必须申明原指尖魔法屋来源及源地址!
评论
使用 GitHub 账号登录后即可留言,支持 Markdown。