AI 工程基础体系 · 第 66/100 篇。内容覆盖机器学习、深度学习与生成式 AI;模型、数据、评测、权限和成本会作为同一生产系统处理。

混合精度训练:FP16、BF16、Loss Scaling、溢出与精度验证

混合精度训练(mixed-precision training)不是“把模型全部改成半精度”,而是让不同计算使用不同数值格式:矩阵乘法、卷积等吞吐敏感的算子使用 FP16 或 BF16,参数更新、部分归约和优化器状态通常保留 FP32。这样做的目标是同时获得更低的显存占用、更高的硬件吞吐,以及足够接近 FP32 训练的数值结果。

但混合精度引入了新的故障模式:

  • FP16 的表示范围较小,容易发生上溢和下溢;
  • 较小的梯度可能在转换为 FP16 时变成零;
  • Loss Scaling 只能缓解梯度下溢,不能修复前向激活上溢;
  • BF16 的指数范围接近 FP32,但有效精度低于 FP16;
  • “没有出现 NaN”不等于训练数值正确;
  • AMP(Automatic Mixed Precision,自动混合精度)会根据算子类型、设备和 PyTorch 版本决定具体的数据类型,不能简单理解为全图转换。

因此,混合精度训练必须同时理解数值格式、自动类型转换、梯度缩放、溢出检测和精度验证。

一、先明确“精度”和“范围”不是同一个概念

浮点数通常可以写成:

x=(1)s×m×2ex = (-1)^s \times m \times 2^e

其中:

  • ss 是符号位;
  • mm 是有效数部分;
  • ee 是指数;
  • 指数位数主要决定可表示的数值范围;
  • 尾数或 fraction 位数主要决定相邻可表示数之间的间隔,也就是有效精度。

FP16 和 BF16 都使用 16 位,但位分配不同。

格式 符号位 指数位 fraction 位 最大有限值 最小正规格
FP32 1 8 23 3.4×10383.4\times10^{38} 1.18×10381.18\times10^{-38}
FP16 1 5 10 65504 6.10×1056.10\times10^{-5}
BF16 1 8 7 3.39×10383.39\times10^{38} 1.18×10381.18\times10^{-38}

如果硬件和实现支持 subnormal(非正规数),FP16 能表示到约 5.96×1085.96\times10^{-8},但精度会显著降低。许多高性能设备或算子可能对 subnormal 采用刷新为零(flush-to-zero)等处理,因此不能把理论最小值当作稳定可用范围。

FP16 有更多 fraction 位,因此在相同数量级下通常比 BF16 更精细;BF16 的指数位和 FP32 一样,因此可以表示非常大的数。可以把两者概括为:

  • FP16:范围小、精度相对高
  • BF16:范围大、精度相对低

这解释了一个常见现象:BF16 通常更不容易因为梯度或激活过大而溢出,但它不一定比 FP16 更接近 FP32。它只是更不容易因为“范围不够”失败。

1.1 一个表示精度的例子

在数量级接近 1 时:

  • FP16 的相邻数间隔大约是 2102^{-10}
  • BF16 的相邻数间隔大约是 272^{-7}

因此,假设一个数接近 1:

1.000000

转换为 BF16 后,能够保留的有效二进制位更少,舍入误差通常比 FP16 大。对于需要保持微小差异的计算,例如某些归约、归一化或概率计算,这种误差可能影响结果。

但如果数值为:

1e20

FP16 无法表示它,会变成无穷大;BF16 仍然可以表示。此时 BF16 的大范围比 FP16 的更高尾数精度重要。

1.2 FP32 主权重并不是多余的

常见的混合精度训练结构是:

FP32 master parameters
        │
        ├── 转换或以低精度参与前向
        │
        ▼
FP16/BF16 activations and temporary tensors
        │
        ▼
FP32 or mixed gradients
        │
        ▼
FP32 optimizer update

优化器状态通常也保留 FP32。例如 Adam 需要一阶矩和二阶矩:

mt=β1mt1+(1β1)gtm_t = \beta_1m_{t-1} + (1-\beta_1)g_t

vt=β2vt1+(1β2)gt2v_t = \beta_2v_{t-1} + (1-\beta_2)g_t^2

如果 mtm_tvtv_t 和参数更新全都使用 FP16,长期累积误差和小更新丢失的风险会明显增加。混合精度的常见做法不是牺牲这些状态,而是只让适合低精度的计算使用 FP16 或 BF16。

二、什么是 AMP:它不是简单的 .half()

直接调用:

model.half()

会把模型参数和缓冲区整体转换成 FP16。这种做法有几个问题:

  1. 不同算子对低精度的容忍度不同;
  2. 某些归约、指数、除法或归一化操作需要更高精度;
  3. 梯度和优化器状态的处理不会自动变得安全;
  4. 输入、标签、参数和中间结果可能出现不一致的 dtype;
  5. 前向溢出后,后面的 Loss Scaling 无法挽救。

AMP 的思路是由 autocast 根据算子策略选择计算类型。以 PyTorch 为例:

with torch.autocast(device_type="cuda", dtype=torch.float16):
    output = model(input)
    loss = criterion(output, target)

在这个上下文中,某些矩阵乘法和卷积通常会使用 FP16,某些对数值范围敏感的操作可能保留 FP32,具体行为由 PyTorch、CUDA、设备架构和算子实现共同决定。不能假设 autocast 会把上下文中的每个 Tensor 都转换成指定 dtype。

可以用下面的例子观察结果 dtype:

import torch

if not torch.cuda.is_available():
    raise RuntimeError("此示例需要 CUDA 设备")

x = torch.randn(1024, 1024, device="cuda")
w = torch.randn(1024, 1024, device="cuda")

with torch.autocast(device_type="cuda", dtype=torch.float16):
    y = x @ w
    z = torch.softmax(y, dim=-1)

print("matmul dtype:", y.dtype)
print("softmax dtype:", z.dtype)

具体输出可能随 PyTorch 和设备变化,重点是:autocast 的 dtype 是算子级别策略,不是整个程序的全局 dtype 开关

2.1 autocast 的生命周期

autocast 只影响上下文内部的前向计算:

with torch.autocast(device_type="cuda", dtype=torch.float16):
    output = model(input)
    loss = loss_fn(output, target)

# 退出上下文后,后续操作不再自动使用 autocast 规则

通常应将模型前向和损失计算放入 autocast。反向传播会使用前向计算保存的类型信息,但不应手动再包一层独立的 autocast backward 上下文。

验证或推理阶段可以这样写:

model.eval()

with torch.inference_mode():
    with torch.autocast(device_type="cuda", dtype=torch.float16):
        output = model(input)

如果需要严格的 FP32 参考结果,则不要启用 autocast:

model.eval()

with torch.inference_mode():
    output_fp32 = model(input)

三、FP16 的两个核心问题:下溢与上溢

3.1 下溢:小数值被舍入为零

假设某个真实梯度为:

g=108g = 10^{-8}

它小于 FP16 常用的可稳定表示范围。在 FP16 计算或存储中,它可能被舍入为:

roundFP16(g)=0\operatorname{round}_{FP16}(g)=0

如果一个参数的梯度变成零,该参数本次更新不会发生。对于深层网络、长序列 Transformer 或经过多次链式求导的路径,梯度可能自然地变得很小,因此下溢会累积为训练质量问题。

3.2 上溢:超过最大有限值

如果某个值超过 FP16 的最大有限值 65504,则可能变成:

roundFP16(x)=+\operatorname{round}_{FP16}(x)=+\infty

之后的计算可能产生:

inf - inf = nan
0 * inf = nan

一旦 NaN 进入梯度,优化器更新就可能把参数污染为 NaN,后续所有输出都失效。

上溢可能发生在:

  • 激活值;
  • logits;
  • loss;
  • 梯度;
  • 梯度缩放后的梯度;
  • 中间归约结果。

必须区分这些位置,因为不同位置需要不同处理方法。Loss Scaling 主要针对梯度下溢,并不能修复前向激活已经变成 inf 的情况。

四、Loss Scaling:为什么放大 loss 可以减少梯度下溢

Loss Scaling(损失缩放)是在反向传播前把损失乘以一个缩放因子 SS

L=SLL' = S L

根据链式法则,对参数 θ\theta 的梯度为:

Lθ=(SL)θ=SLθ=Sg\frac{\partial L'}{\partial \theta} = \frac{\partial (SL)}{\partial \theta} = S\frac{\partial L}{\partial \theta} = Sg

反向传播时,梯度从 gg 变为 SgSg。如果原始梯度太小,放大后就可能落入 FP16 的可表示范围。

更新前再除以 SS

g^=SgS=g\hat{g} = \frac{Sg}{S}=g

这样理想情况下,优化器看到的仍然是原始梯度。

4.1 完整数值算例

设真实梯度为:

g=108g=10^{-8}

如果 FP16 存储这个值时下溢为零:

直接转换:
1e-8 → 0

取缩放因子:

S=1024S=1024

则反向传播中使用:

g=Sg=1024×108=1.024×105g'=Sg=1024\times10^{-8}=1.024\times10^{-5}

这个数通常可以被 FP16 表示。更新前除以 1024:

g^=1.024×1051024108\hat{g}=\frac{1.024\times10^{-5}}{1024}\approx10^{-8}

于是梯度信息得以保留。

4.2 Loss Scaling 的反例:它不能修复前向溢出

设前向激活为:

a=100000a=100000

FP16 无法表示 100000,因此在前向阶段已经可能得到:

a → inf

即使之后使用:

L=SLL'=SL

也只是放大一个已经包含 infnan 的损失,无法恢复正确的激活。此时应从前向数值入手,例如:

  • 使用 BF16;
  • 将敏感算子保留为 FP32;
  • 检查 logits、归一化和指数运算;
  • 调整初始化、输入范围或模型结构;
  • 检查是否存在异常数据。

这一区分很重要:

梯度太小       → Loss Scaling 可能有帮助
前向值太大     → Loss Scaling 无法修复
缩放后梯度太大 → 动态缩放需要回退

五、动态 Loss Scaling 与溢出处理

固定缩放因子可以工作,但不同训练阶段的梯度范围会变化。训练初期、学习率变化、序列长度变化或异常 batch 都可能导致梯度范围发生改变,因此实践中常使用动态 Loss Scaling。

动态缩放维护一个状态 SS

  1. 用当前 SS 放大 loss;
  2. 执行反向传播;
  3. 检查缩放后的梯度是否包含 infnan
  4. 如果溢出,则跳过本次参数更新,并减小 SS
  5. 如果连续多个 step 没有溢出,则增大 SS

典型状态转移如下:

flowchart TD
    A[读取当前 scale S] --> B[计算 loss]
    B --> C[计算 S * loss]
    C --> D[反向传播]
    D --> E{梯度是否包含 inf/nan}
    E -- 是 --> F[跳过 optimizer.step]
    F --> G[减小 scale]
    G --> A
    E -- 否 --> H[unscale 梯度]
    H --> I[梯度裁剪或其他检查]
    I --> J[optimizer.step]
    J --> K[更新 scale 状态]
    K --> A

“跳过 optimizer.step”是必要的。如果梯度已经为 inf,仍然调用优化器更新,参数就可能被写成 infnan。缩小 scale 后重新计算下一批数据,才有机会恢复。

5.1 PyTorch 中的标准训练顺序

下面是一个可运行的 CUDA 示例。它使用现代 PyTorch 的 torch.autocasttorch.amp.GradScaler 接口:

import torch
from torch import nn
from torch.utils.data import DataLoader, TensorDataset

if not torch.cuda.is_available():
    raise RuntimeError("此示例需要 CUDA")

device = torch.device("cuda")

# 构造一个可复现实验数据集
torch.manual_seed(0)
x = torch.randn(4096, 128)
y = torch.randint(0, 10, (4096,))

loader = DataLoader(
    TensorDataset(x, y),
    batch_size=128,
    shuffle=True,
    pin_memory=True,
)

model = nn.Sequential(
    nn.Linear(128, 512),
    nn.GELU(),
    nn.Linear(512, 10),
).to(device)

criterion = nn.CrossEntropyLoss()
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3)

# 现代 PyTorch 写法;具体构造形式需以所安装版本文档为准
scaler = torch.amp.GradScaler("cuda")

for epoch in range(3):
    model.train()

    for inputs, targets in loader:
        inputs = inputs.to(device, non_blocking=True)
        targets = targets.to(device, non_blocking=True)

        optimizer.zero_grad(set_to_none=True)

        with torch.autocast(
            device_type="cuda",
            dtype=torch.float16,
        ):
            logits = model(inputs)
            loss = criterion(logits, targets)

        # 先放大 loss,再 backward
        scaler.scale(loss).backward()

        # unscale 后,grad 才回到优化器实际使用的尺度
        scaler.unscale_(optimizer)

        # 如果需要梯度裁剪,必须放在 unscale_ 之后
        grad_norm = torch.nn.utils.clip_grad_norm_(
            model.parameters(),
            max_norm=1.0,
        )

        # 如果检测到 inf/nan,内部会跳过这次 step
        scaler.step(optimizer)

        # 根据是否溢出调整 scale
        scaler.update()

    print(
        f"epoch={epoch}, "
        f"loss={loss.item():.6f}, "
        f"grad_norm={float(grad_norm):.6f}, "
        f"scale={scaler.get_scale():.1f}"
    )

每一步的因果关系是:

  • autocast 控制前向中适合低精度的算子;
  • scaler.scale(loss) 只改变反向传播的数值尺度;
  • backward() 得到的是缩放后的梯度;
  • unscale_(optimizer) 把梯度除回原尺度;
  • 梯度裁剪必须在 unscale_ 之后,否则裁剪阈值会被缩放因子放大;
  • scaler.step(optimizer) 检查梯度是否有限,并在正常时执行更新;
  • scaler.update() 根据本次是否溢出调整 scale。

较旧的 PyTorch 版本常见写法是:

from torch.cuda.amp import autocast, GradScaler

scaler = GradScaler()

with autocast(dtype=torch.float16):
    ...

这属于版本相关接口,应以当前安装版本的 PyTorch 文档为准。核心生命周期不变。

5.2 溢出时发生了什么

假设当前:

scale = 65536
真实梯度 = 2
缩放后梯度 = 131072

如果某个保存梯度的路径使用 FP16,131072 超过 65504,可能变成 inf。GradScaler 发现非有限梯度后通常会:

本次不更新参数
scale 从 65536 降低
下一次重新尝试

如果连续很多次没有溢出,则可能逐步增大 scale,以减少梯度下溢。

这里有一个重要边界:动态缩放因子可能下降到小于 1。不能假设它永远大于等于 1,也不能在监控或恢复逻辑中硬编码“scale 只会增长”。

六、BF16 为什么通常不需要 Loss Scaling

BF16 的指数位与 FP32 相同,因此它可以表示非常小和非常大的数量级。对于前面的小梯度例子:

g=108g=10^{-8}

BF16 通常不会像 FP16 那样因为范围不足而直接下溢为零。它的主要问题是尾数较短,即:

  • 它能表示这个数量级;
  • 但表示值可能有较大的舍入误差;
  • 连续计算和归约仍然可能积累误差。

所以在支持 BF16 的硬件上,常见策略是:

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

并且通常不使用 FP16 专用的动态 Loss Scaling。原因不是 BF16 没有误差,而是它通常不需要通过放大梯度来扩大指数范围。

“BF16 永远不需要缩放”仍然是过度绝对化的说法。具体是否使用缩放取决于:

  • PyTorch 版本;
  • 设备和内核实现;
  • 自定义算子;
  • 梯度是否在某些阶段被强制转换成 FP16;
  • 模型是否包含非常规数值范围。

如果模型的某条路径仍然把 BF16 梯度写入 FP16 缓冲区,BF16 本身的大范围也无法保护那条路径。

七、FP16 与 BF16 的选择

可以从三个维度判断。

7.1 硬件支持

首先确认设备是否原生支持目标低精度。某些 GPU、TPU 或专用加速器对 BF16 有高吞吐支持,另一些设备对 FP16 支持更成熟。没有硬件支持时,软件模拟可能失去性能收益,甚至增加转换开销。

可以先检查:

import torch

print("PyTorch:", torch.__version__)
print("CUDA available:", torch.cuda.is_available())

if torch.cuda.is_available():
    print("GPU:", torch.cuda.get_device_name())
    print("CUDA capability:", torch.cuda.get_device_capability())
    print("BF16 supported:", torch.cuda.is_bf16_supported())

torch.cuda.is_bf16_supported() 的具体行为和版本有关;它是设备能力检查,不等于每一个算子、每一种布局或每一条自定义 kernel 路径都支持 BF16。

7.2 数值范围

以下类型更容易暴露 FP16 范围问题:

  • 很长的序列;
  • 大幅度 logits;
  • 训练初期激活变化剧烈;
  • 梯度裁剪前梯度跨度很大;
  • 自定义归一化或指数运算;
  • 需要较大动态范围的生成式模型组件。

这类模型通常可以优先测试 BF16。

7.3 有效精度

BF16 的尾数比 FP16 短,某些对微小差异敏感的计算可能更适合保留 FP32,或者使用 FP16 加 FP32 累积。最终选择不应只看单步吞吐,还要看:

  • 训练是否稳定;
  • 验证集损失和任务指标;
  • 收敛速度;
  • 重试和跳步次数;
  • 显存;
  • 单位有效样本的成本。

八、哪些计算应保留 FP32

autocast 会自动处理许多常见算子,但自定义代码仍需要理解数值边界。以下操作通常值得重点检查:

8.1 归约与平均

求和、均值、方差等操作可能积累大量舍入误差。一个简单反例是:

[1e4, 1e-3, -1e4]

如果低精度先把小数项舍入掉,求和结果可能是 0,而不是接近 10310^{-3}

因此,某些归约的累积类型会使用 FP32,即使输入来自低精度。对于自定义 kernel,应明确指定 accumulator dtype,而不能只看输入输出 dtype。

8.2 指数、对数与 softmax

softmax 的形式是:

softmax(xi)=exijexj\operatorname{softmax}(x_i) = \frac{e^{x_i}}{\sum_j e^{x_j}}

直接计算 exie^{x_i} 容易上溢,因此稳定实现通常先减去最大值:

softmax(xi)=eximax(x)jexjmax(x)\operatorname{softmax}(x_i) = \frac{e^{x_i-\max(x)}}{\sum_j e^{x_j-\max(x)}}

即使使用稳定公式,低精度仍可能带来舍入误差。交叉熵也通常应使用框架提供的融合实现,而不是手动执行 softmax 后再 log

8.3 归一化

LayerNorm、RMSNorm、BatchNorm 等包含平方、均值、方差和倒数平方根。常见实现会对统计量使用 FP32 或更高精度路径,但自定义版本必须检查:

  • 分母是否过小;
  • epsilon 是否合理;
  • 平方是否上溢;
  • 统计量累积是否使用 FP32。

8.4 注意力计算

缩放点积注意力包含:

A=softmax(QKdk)VA=\operatorname{softmax}\left(\frac{QK^\top}{\sqrt{d_k}}\right)V

其中 QKQK^\top 可能产生大幅度值,softmax 又对指数范围敏感。现代 fused attention kernel 可能拥有专门的数值稳定实现,但自定义注意力实现不能假设它自动安全。

九、训练状态:参数、梯度、优化器和 scale 必须一起管理

混合精度训练不仅改变 Tensor dtype,还增加了 scaler 状态。完整 checkpoint 至少应考虑:

checkpoint = {
    "model": model.state_dict(),
    "optimizer": optimizer.state_dict(),
    "scaler": scaler.state_dict(),
    "epoch": epoch,
}
torch.save(checkpoint, "checkpoint.pt")

恢复时:

checkpoint = torch.load("checkpoint.pt", map_location="cuda")

model.load_state_dict(checkpoint["model"])
optimizer.load_state_dict(checkpoint["optimizer"])
scaler.load_state_dict(checkpoint["scaler"])

start_epoch = checkpoint["epoch"]

如果只恢复模型参数而不恢复优化器状态,Adam 的动量和二阶矩会丢失;如果只恢复模型和优化器而不恢复 scaler,训练会从不同的缩放状态继续。对于严格复现实验,这些差异都可能影响后续轨迹。

如果训练使用梯度累积,缩放逻辑也要保持一致:

accumulation_steps = 4
optimizer.zero_grad(set_to_none=True)

for step, (inputs, targets) in enumerate(loader):
    inputs = inputs.to(device, non_blocking=True)
    targets = targets.to(device, non_blocking=True)

    with torch.autocast(device_type="cuda", dtype=torch.float16):
        loss = criterion(model(inputs), targets)
        loss = loss / accumulation_steps

    scaler.scale(loss).backward()

    if (step + 1) % accumulation_steps == 0:
        scaler.unscale_(optimizer)
        torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
        scaler.step(optimizer)
        scaler.update()
        optimizer.zero_grad(set_to_none=True)

除非明确设计了不同策略,否则不要在每个 micro-batch 都执行 optimizer step。梯度累积期间的 loss 缩放、除以累积步数和最终 unscale 必须保持一致。

十、精度验证不能只看最终准确率

混合精度验证应至少包含四层。

10.1 数值有限性

首先检查输出、损失、梯度和参数是否有限:

def check_finite(name, tensor):
    if not torch.isfinite(tensor).all():
        bad = (~torch.isfinite(tensor)).sum().item()
        raise FloatingPointError(
            f"{name} contains {bad} non-finite values"
        )

check_finite("loss", loss.detach())

for name, parameter in model.named_parameters():
    check_finite(f"parameter:{name}", parameter.data)
    if parameter.grad is not None:
        check_finite(f"gradient:{name}", parameter.grad)

这个检查应放在有意义的生命周期节点:

  • backward 后检查缩放梯度;
  • unscale_ 后检查真实尺度梯度;
  • optimizer step 后检查参数;
  • 验证前向后检查 logits 和 loss。

只检查最终 loss 可能太晚,因为参数早已被污染。

10.2 与 FP32 参考结果比较

给定相同模型参数和相同输入,可以比较 FP32 与 AMP 的输出:

import copy
import torch

reference = copy.deepcopy(model).float().eval()
candidate = copy.deepcopy(model).eval()

inputs = torch.randn(32, 128, device=device)

with torch.inference_mode():
    output_fp32 = reference(inputs.float())

    with torch.autocast(
        device_type="cuda",
        dtype=torch.float16,
    ):
        output_amp = candidate(inputs)

diff = (output_fp32 - output_amp.float()).abs()
max_abs_error = diff.max().item()
max_rel_error = (
    diff / output_fp32.abs().clamp_min(1e-8)
).max().item()

print("max absolute error:", max_abs_error)
print("max relative error:", max_rel_error)
print(
    "allclose:",
    torch.allclose(
        output_fp32,
        output_amp.float(),
        rtol=1e-2,
        atol=1e-3,
    ),
)

这里的阈值不是通用标准。它取决于:

  • 模型深度;
  • 输出尺度;
  • 是否包含随机算子;
  • 输入分布;
  • 低精度格式;
  • 任务对误差的敏感程度。

绝对误差适合观察接近零的输出,绝对误差可能很小但相对误差很大;相对误差在参考值接近零时又会失真,因此通常同时报告两者。

10.3 比较梯度而不是只比较输出

在固定模型参数、输入和 loss 的条件下,可以分别计算 FP32 梯度与 AMP 梯度,再比较:

Δg=gampgfp32\Delta_g = g_{\text{amp}} - g_{\text{fp32}}

应该关注:

  • 梯度是否有限;
  • 梯度的最大绝对误差;
  • 梯度范数;
  • 梯度方向余弦相似度;
  • 是否有大量梯度变为零。

梯度方向余弦相似度为:

cos(θ)=g1g2g1g2\cos(\theta) = \frac{g_1^\top g_2} {\|g_1\|\|g_2\|}

它比单纯比较每个元素的相对误差更适合判断优化方向是否一致。

10.4 比较训练轨迹和任务指标

最终验证集准确率接近,并不能证明两次训练过程等价。还应比较:

  • 每个 step 或每个 epoch 的 loss;
  • 学习率;
  • 梯度范数;
  • scale 的变化;
  • 跳过的 optimizer step 数量;
  • 验证集 loss;
  • 任务特定指标;
  • 吞吐、显存和实际成本。

如果 AMP 训练频繁跳步,即使最终指标暂时正常,也说明当前配置可能在稳定性上付出了代价。

十一、一个小型 FP16 下溢实验

下面的实验展示 Loss Scaling 的数学作用。它不依赖模型,只模拟低精度存储:

import torch

g = torch.tensor([1e-8], dtype=torch.float32)

g_fp16 = g.to(torch.float16)
g_scaled_fp16 = (g * 1024).to(torch.float16)
g_recovered = g_scaled_fp16.to(torch.float32) / 1024

print("original:", g.item())
print("direct fp16:", g_fp16.item())
print("scaled fp16:", g_scaled_fp16.item())
print("recovered:", g_recovered.item())

可能看到类似结果:

original: 9.99999993922529e-09
direct fp16: 0.0
scaled fp16: 1.0251998901367188e-05
recovered: 1.0011717677116394e-08

这里的关键不是某个具体打印值,而是:

直接转 FP16 丢失
先放大再转 FP16 可以保留数量级
再除回 scale 后得到近似原值

这不是“提高了 FP16 的有效精度”,而是把数值移动到了 FP16 更容易表示的范围内。

十二、常见错误及其失败表现

12.1 直接把整个模型 .half()

失败表现可能包括:

  • loss 很快变成 NaN;
  • LayerNorm、softmax 或自定义归约不稳定;
  • 优化器状态被低精度污染;
  • 某些算子报 dtype 或设备错误;
  • 训练指标明显劣于 FP32。

应优先使用 autocast,并让优化器状态和敏感计算保留在合适的精度。

12.2 忘记调用 unscale_ 就梯度裁剪

错误顺序:

scaler.scale(loss).backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
scaler.step(optimizer)

此时梯度仍然乘以 SS,裁剪阈值针对的是缩放后的梯度。结果可能是每一步都被过度裁剪,实际更新远小于预期。

正确顺序是:

scaler.scale(loss).backward()
scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
scaler.step(optimizer)
scaler.update()

12.3 手动把 loss 除掉或重复除 scale

错误做法通常是:

scaled_loss = scaler.scale(loss)
scaled_loss.backward()

for p in model.parameters():
    if p.grad is not None:
        p.grad.div_(scaler.get_scale())

scaler.step(optimizer)

这会绕过 GradScaler 的内部状态和溢出处理。应使用 scaler.unscale_(optimizer),不要自己读取 scale 并修改梯度。

12.4 认为 scale 越大越好

scale 太小,不能充分缓解梯度下溢;scale 太大,缩放后的梯度会更容易上溢。动态 scaler 的目的正是在两者之间寻找可用范围。

因此不应只监控:

scale 是否越来越大

而应同时监控:

是否发生 skipped step
梯度是否有限
验证指标是否异常
scale 是否频繁上下波动

12.5 只看 loss,不看参数和梯度

某些参数可能已经出现 NaN,但当前 batch 的损失还没有立即变成 NaN。下一次前向才暴露问题。生产训练应在关键位置增加有限性检查,或记录异常 batch 的输入范围、序列长度、标签分布和模型阶段。

12.6 用一次随机实验判断数值等价

低精度训练本来就可能产生不同的舍入路径。一次实验中指标接近,不能证明所有数据、所有随机种子和所有训练阶段都安全。至少应进行固定种子单步对比、短程训练对比,以及完整验证集对比。

十三、分布式训练中的额外边界

在 DDP、FSDP 或其他分布式训练中,梯度可能先在 rank 内产生,再参与通信或归约。Loss Scaling 和梯度裁剪的顺序仍然重要:

反向传播
→ 得到缩放梯度
→ 在正确时机 unscale
→ 再进行梯度裁剪
→ 再执行 optimizer step

具体的通信、参数分片和 scaler 集成会随并行方案和 PyTorch 版本变化。需要明确:

  • 溢出检测是在本地完成还是需要跨 rank 协调;
  • 梯度裁剪针对的是本地梯度还是全局范数;
  • optimizer state 是否分片;
  • checkpoint 是否包含每个 rank 的 scaler、优化器和随机状态;
  • 某个 rank 跳过更新时,其他 rank 是否保持一致。

如果各 rank 对是否执行 optimizer step 的判断不一致,参数会失去同步,后续通信可能失败或产生隐蔽的模型偏差。分布式框架提供的集成方案应优先于自行拼接逻辑。

十四、严格验证时还要控制 TF32 和随机性

在 CUDA 上,某些 FP32 矩阵乘法可能使用 TF32。TF32 不是 FP16 或 BF16,但它会改变 FP32 矩阵乘法的有效精度,因此可能干扰“FP32 参考结果”的比较。

做严格对照实验时,可以显式设置:

torch.set_float32_matmul_precision("highest")

这个 API 和具体后端行为具有版本与设备相关性,不能把它理解为对所有算子提供绝对 bitwise FP32 保证。还应控制:

  • 随机种子;
  • DataLoader 顺序;
  • dropout 状态;
  • CUDA 算子确定性设置;
  • 数据预处理;
  • 初始模型参数;
  • checkpoint 恢复位置。

即使这些都相同,GPU 并行归约和原子操作仍可能导致非 bitwise 一致。工程上更合理的目标通常是“误差在预先定义的容忍范围内”,而不是要求每个浮点位完全相同。

十五、面向生产系统的验证与取舍

混合精度的收益应按完整系统衡量,而不是只比较一次前向耗时。

15.1 模型与数据

应记录:

  • 模型结构和初始化方式;
  • 输入 dtype、范围和异常值比例;
  • 序列长度或图像尺寸分布;
  • loss 和 label 的尺度;
  • 自定义算子和 fused kernel 版本;
  • 训练硬件、驱动、CUDA 和 PyTorch 版本。

同一个模型在短序列和长序列上的溢出风险可能完全不同。只用平均长度数据验证,可能掩盖长尾输入导致的生产故障。

15.2 评测

评测集应固定并与训练数据隔离。除了最终准确率或生成质量,还应保存:

  • FP32 基线;
  • FP16 训练结果;
  • BF16 训练结果;
  • 关键子集上的结果;
  • 极端长度和极端数值输入结果;
  • NaN、Inf、跳步和 scale 变化日志。

对于生成式 AI,还应检查长序列生成、极端提示词、采样温度变化和输出截断行为。低精度差异有时不会影响短文本平均指标,却会在长上下文中放大。

15.3 权限与日志

AMP 不改变模型、数据集和 checkpoint 的访问权限。训练日志若记录异常 batch、输入样本或生成结果,仍应遵守数据权限和脱敏要求。数值诊断需要记录足够信息来复现问题,但不应为了定位 NaN 而无控制地保存原始敏感数据。

15.4 成本

成本收益至少包括:

单位有效样本成本=总训练成本未因失败或重试而丢弃的有效样本数\text{单位有效样本成本} = \frac{\text{总训练成本}} {\text{未因失败或重试而丢弃的有效样本数}}

如果 FP16 理论吞吐更高,但频繁溢出、跳步和重跑,实际成本可能高于更稳定的 BF16。正确比较应同时测量:

  • 每秒有效样本数;
  • 峰值显存;
  • 训练总时长;
  • 跳过的更新数;
  • 失败重启次数;
  • 最终达到目标指标所需的计算量。

十六、一套可执行的排查顺序

当混合精度训练出现 NaN、loss 突然变大或指标退化时,可以按以下因果链定位:

  1. 确认输入是否已有非有限值
    在进入模型前检查输入、标签和 mask。

  2. 定位第一次出现非有限值的算子
    对模块输出、loss、梯度和参数逐层检查,而不是只看最终指标。

  3. 区分前向溢出与反向溢出
    如果 autocast 前向输出已经是 inf,Loss Scaling 不是修复手段。

  4. 检查 scale 和 skipped step
    如果 scale 持续下降并且大量跳步,说明缩放后的梯度范围仍然不适合当前路径。

  5. 检查梯度裁剪顺序
    裁剪必须发生在 unscale_ 之后。

  6. 切换 BF16 或局部 FP32
    如果 BF16 稳定而 FP16 不稳定,通常说明主要问题是 FP16 动态范围,而不是模型逻辑错误。

  7. 检查自定义算子和归约累积类型
    框架自动策略覆盖不到自定义 kernel 时,需要显式指定 accumulator 和输出类型。

  8. 与 FP32 单步结果对照
    固定参数、输入和随机状态,比较输出、loss、梯度范数和梯度方向。

  9. 最后再调整学习率或模型结构
    不能把所有混合精度问题都归因于学习率。先确定故障发生在格式转换、前向算子、梯度缩放还是优化器更新。

结语

FP16 和 BF16 的核心差异是指数范围与有效精度的权衡:FP16 更精细但范围更小,BF16 范围接近 FP32 但尾数更短。AMP 通过算子级别的自动类型转换,让适合低精度的计算获得性能收益,同时保留敏感路径和优化器状态的较高精度。

Loss Scaling 的数学作用是把过小的梯度移动到 FP16 可表示的范围,再在优化器更新前恢复原始尺度。动态 Loss Scaling 进一步通过检测 infnan,在“梯度下溢”和“缩放后上溢”之间寻找可用范围。但它不能修复已经在前向传播中发生的溢出。

可靠的混合精度训练不能只依赖默认配置或最终指标,而应验证完整数值链路:

输入
→ 前向激活
→ loss
→ 缩放后的梯度
→ unscale 后的梯度
→ 优化器状态
→ 参数
→ 验证指标

只有当有限性、FP32 对照、训练轨迹、任务指标、checkpoint 状态和实际成本都在可接受范围内时,FP16 或 BF16 才适合作为生产训练配置。


系列导航与关联阅读

官方资料

本文依据研究论文、标准组织与主流框架官方文档重新梳理;正文、示例与工程清单由 WR BLOG 编写。