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/BF16/FP32 精度对比

这张图很关键: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 画了一个完整的流程图:

flowchart TD A[开始训练批次] --> B{显卡支持 BF16?} B -->|是| C[使用 BF16 模式] B -->|否| D[使用 FP16 模式] C --> E[前向传播<br>自动转换为 BF16] D --> F[前向传播<br>自动转换为 FP16] E --> G[计算 Loss] F --> G G --> H{FP16 模式?} H -->|是| I[Loss Scaling<br>放大梯度] H -->|否| J[直接反向传播] I --> K[反向传播] J --> K K --> L[检查梯度状态<br>NaN/Inf?] L -->|正常| M[优化器更新] L -->|异常| N[跳过更新<br>调整 Scaling] N --> O[继续下一批次] M --> O

这张图想说明一件事:混合精度训练不是简单地把 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 参数):

配置训练速度显存占用训练稳定性最终精度
FP321.0x16.2GB稳定基准
FP161.6x9.8GB需调参基本持平
BF161.5x10.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/) 转载或引用必须申明原指尖魔法屋来源及源地址!