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 版本决定具体的数据类型,不能简单理解为全图转换。
因此,混合精度训练必须同时理解数值格式、自动类型转换、梯度缩放、溢出检测和精度验证。
一、先明确“精度”和“范围”不是同一个概念
浮点数通常可以写成:
其中:
- 是符号位;
- 是有效数部分;
- 是指数;
- 指数位数主要决定可表示的数值范围;
- 尾数或 fraction 位数主要决定相邻可表示数之间的间隔,也就是有效精度。
FP16 和 BF16 都使用 16 位,但位分配不同。
| 格式 | 符号位 | 指数位 | fraction 位 | 最大有限值 | 最小正规格 |
|---|---|---|---|---|---|
| FP32 | 1 | 8 | 23 | 约 | 约 |
| FP16 | 1 | 5 | 10 | 65504 | 约 |
| BF16 | 1 | 8 | 7 | 约 | 约 |
如果硬件和实现支持 subnormal(非正规数),FP16 能表示到约 ,但精度会显著降低。许多高性能设备或算子可能对 subnormal 采用刷新为零(flush-to-zero)等处理,因此不能把理论最小值当作稳定可用范围。
FP16 有更多 fraction 位,因此在相同数量级下通常比 BF16 更精细;BF16 的指数位和 FP32 一样,因此可以表示非常大的数。可以把两者概括为:
- FP16:范围小、精度相对高;
- BF16:范围大、精度相对低。
这解释了一个常见现象:BF16 通常更不容易因为梯度或激活过大而溢出,但它不一定比 FP16 更接近 FP32。它只是更不容易因为“范围不够”失败。
1.1 一个表示精度的例子
在数量级接近 1 时:
- FP16 的相邻数间隔大约是 ;
- BF16 的相邻数间隔大约是 。
因此,假设一个数接近 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 需要一阶矩和二阶矩:
如果 、 和参数更新全都使用 FP16,长期累积误差和小更新丢失的风险会明显增加。混合精度的常见做法不是牺牲这些状态,而是只让适合低精度的计算使用 FP16 或 BF16。
二、什么是 AMP:它不是简单的 .half()
直接调用:
model.half()
会把模型参数和缓冲区整体转换成 FP16。这种做法有几个问题:
- 不同算子对低精度的容忍度不同;
- 某些归约、指数、除法或归一化操作需要更高精度;
- 梯度和优化器状态的处理不会自动变得安全;
- 输入、标签、参数和中间结果可能出现不一致的 dtype;
- 前向溢出后,后面的 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 下溢:小数值被舍入为零
假设某个真实梯度为:
它小于 FP16 常用的可稳定表示范围。在 FP16 计算或存储中,它可能被舍入为:
如果一个参数的梯度变成零,该参数本次更新不会发生。对于深层网络、长序列 Transformer 或经过多次链式求导的路径,梯度可能自然地变得很小,因此下溢会累积为训练质量问题。
3.2 上溢:超过最大有限值
如果某个值超过 FP16 的最大有限值 65504,则可能变成:
之后的计算可能产生:
inf - inf = nan
0 * inf = nan
一旦 NaN 进入梯度,优化器更新就可能把参数污染为 NaN,后续所有输出都失效。
上溢可能发生在:
- 激活值;
- logits;
- loss;
- 梯度;
- 梯度缩放后的梯度;
- 中间归约结果。
必须区分这些位置,因为不同位置需要不同处理方法。Loss Scaling 主要针对梯度下溢,并不能修复前向激活已经变成 inf 的情况。
四、Loss Scaling:为什么放大 loss 可以减少梯度下溢
Loss Scaling(损失缩放)是在反向传播前把损失乘以一个缩放因子 :
根据链式法则,对参数 的梯度为:
反向传播时,梯度从 变为 。如果原始梯度太小,放大后就可能落入 FP16 的可表示范围。
更新前再除以 :
这样理想情况下,优化器看到的仍然是原始梯度。
4.1 完整数值算例
设真实梯度为:
如果 FP16 存储这个值时下溢为零:
直接转换:
1e-8 → 0
取缩放因子:
则反向传播中使用:
这个数通常可以被 FP16 表示。更新前除以 1024:
于是梯度信息得以保留。
4.2 Loss Scaling 的反例:它不能修复前向溢出
设前向激活为:
FP16 无法表示 100000,因此在前向阶段已经可能得到:
a → inf
即使之后使用:
也只是放大一个已经包含 inf 或 nan 的损失,无法恢复正确的激活。此时应从前向数值入手,例如:
- 使用 BF16;
- 将敏感算子保留为 FP32;
- 检查 logits、归一化和指数运算;
- 调整初始化、输入范围或模型结构;
- 检查是否存在异常数据。
这一区分很重要:
梯度太小 → Loss Scaling 可能有帮助
前向值太大 → Loss Scaling 无法修复
缩放后梯度太大 → 动态缩放需要回退
五、动态 Loss Scaling 与溢出处理
固定缩放因子可以工作,但不同训练阶段的梯度范围会变化。训练初期、学习率变化、序列长度变化或异常 batch 都可能导致梯度范围发生改变,因此实践中常使用动态 Loss Scaling。
动态缩放维护一个状态 :
- 用当前 放大 loss;
- 执行反向传播;
- 检查缩放后的梯度是否包含
inf或nan; - 如果溢出,则跳过本次参数更新,并减小 ;
- 如果连续多个 step 没有溢出,则增大 。
典型状态转移如下:
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,仍然调用优化器更新,参数就可能被写成 inf 或 nan。缩小 scale 后重新计算下一批数据,才有机会恢复。
5.1 PyTorch 中的标准训练顺序
下面是一个可运行的 CUDA 示例。它使用现代 PyTorch 的 torch.autocast 和 torch.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 相同,因此它可以表示非常小和非常大的数量级。对于前面的小梯度例子:
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,而不是接近 。
因此,某些归约的累积类型会使用 FP32,即使输入来自低精度。对于自定义 kernel,应明确指定 accumulator dtype,而不能只看输入输出 dtype。
8.2 指数、对数与 softmax
softmax 的形式是:
直接计算 容易上溢,因此稳定实现通常先减去最大值:
即使使用稳定公式,低精度仍可能带来舍入误差。交叉熵也通常应使用框架提供的融合实现,而不是手动执行 softmax 后再 log。
8.3 归一化
LayerNorm、RMSNorm、BatchNorm 等包含平方、均值、方差和倒数平方根。常见实现会对统计量使用 FP32 或更高精度路径,但自定义版本必须检查:
- 分母是否过小;
- epsilon 是否合理;
- 平方是否上溢;
- 统计量累积是否使用 FP32。
8.4 注意力计算
缩放点积注意力包含:
其中 可能产生大幅度值,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 梯度,再比较:
应该关注:
- 梯度是否有限;
- 梯度的最大绝对误差;
- 梯度范数;
- 梯度方向余弦相似度;
- 是否有大量梯度变为零。
梯度方向余弦相似度为:
它比单纯比较每个元素的相对误差更适合判断优化方向是否一致。
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)
此时梯度仍然乘以 ,裁剪阈值针对的是缩放后的梯度。结果可能是每一步都被过度裁剪,实际更新远小于预期。
正确顺序是:
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 成本
成本收益至少包括:
如果 FP16 理论吞吐更高,但频繁溢出、跳步和重跑,实际成本可能高于更稳定的 BF16。正确比较应同时测量:
- 每秒有效样本数;
- 峰值显存;
- 训练总时长;
- 跳过的更新数;
- 失败重启次数;
- 最终达到目标指标所需的计算量。
十六、一套可执行的排查顺序
当混合精度训练出现 NaN、loss 突然变大或指标退化时,可以按以下因果链定位:
-
确认输入是否已有非有限值
在进入模型前检查输入、标签和 mask。 -
定位第一次出现非有限值的算子
对模块输出、loss、梯度和参数逐层检查,而不是只看最终指标。 -
区分前向溢出与反向溢出
如果 autocast 前向输出已经是inf,Loss Scaling 不是修复手段。 -
检查 scale 和 skipped step
如果 scale 持续下降并且大量跳步,说明缩放后的梯度范围仍然不适合当前路径。 -
检查梯度裁剪顺序
裁剪必须发生在unscale_之后。 -
切换 BF16 或局部 FP32
如果 BF16 稳定而 FP16 不稳定,通常说明主要问题是 FP16 动态范围,而不是模型逻辑错误。 -
检查自定义算子和归约累积类型
框架自动策略覆盖不到自定义 kernel 时,需要显式指定 accumulator 和输出类型。 -
与 FP32 单步结果对照
固定参数、输入和随机状态,比较输出、loss、梯度范数和梯度方向。 -
最后再调整学习率或模型结构
不能把所有混合精度问题都归因于学习率。先确定故障发生在格式转换、前向算子、梯度缩放还是优化器更新。
结语
FP16 和 BF16 的核心差异是指数范围与有效精度的权衡:FP16 更精细但范围更小,BF16 范围接近 FP32 但尾数更短。AMP 通过算子级别的自动类型转换,让适合低精度的计算获得性能收益,同时保留敏感路径和优化器状态的较高精度。
Loss Scaling 的数学作用是把过小的梯度移动到 FP16 可表示的范围,再在优化器更新前恢复原始尺度。动态 Loss Scaling 进一步通过检测 inf 和 nan,在“梯度下溢”和“缩放后上溢”之间寻找可用范围。但它不能修复已经在前向传播中发生的溢出。
可靠的混合精度训练不能只依赖默认配置或最终指标,而应验证完整数值链路:
输入
→ 前向激活
→ loss
→ 缩放后的梯度
→ unscale 后的梯度
→ 优化器状态
→ 参数
→ 验证指标
只有当有限性、FP32 对照、训练轨迹、任务指标、checkpoint 状态和实际成本都在可接受范围内时,FP16 或 BF16 才适合作为生产训练配置。
系列导航与关联阅读
- 系列入口:AI 工程完整学习路线:从机器学习与 Transformer 到 RAG、Agent 和生产治理
- 上一篇:AI GPU 与 CUDA 基础:核函数、显存、带宽、算子和性能证据
- 下一篇:模型检查点与训练恢复:状态、随机数、分片、容错和一致性
官方资料
本文依据研究论文、标准组织与主流框架官方文档重新梳理;正文、示例与工程清单由 WR BLOG 编写。

评论
0 条讨论