训练一个大模型,时间和显存是两笔硬开销。让矩阵乘法跑在更短的数字上(FP16 或 BF16),两笔账能同时省——这就是混合精度训练。它与之前讲推理侧量化的《模型量化的精度账:从 FP16 到 INT4 该怎么选》分工明确:那篇说的是训练完成后怎么把模型压小拿去部署,本文说的是训练过程中怎么在低精度上把梯度算得又快又稳。
它解决什么问题
训练的开销大头是矩阵乘法。半精度数字只占 FP32 一半的存储,显存里能塞下约两倍的权重与激活;在现代 GPU 的 Tensor Core 上,半精度乘法的吞吐还能成倍提升。看起来只是把 float 换成 half 的举手之劳,实际却撞上浮点数表示的根本约束。
浮点数在计算机里是三段式的:符号位、指数位、尾数位。指数位决定「能表示多大、多小」(动态范围),尾数位决定「相邻两个数之间有多密」(精度)。固定总位宽下,位分给谁多一点,另一个就得少一点。常用格式的分配如下:
| 格式 | 符号位 | 指数位 | 尾数位 | 最大值 | 最小正规数 |
|---|---|---|---|---|---|
| FP32 | 1 | 8 | 23 | 约 3.4×10^38 | 约 1.2×10^-38 |
| FP16 | 1 | 5 | 10 | 65504 | 约 6.1×10^-5 |
| BF16 | 1 | 8 | 7 | 约 3.4×10^38 | 约 1.2×10^-38 |
| FP8 E4M3 | 1 | 4 | 3 | 448 | 约 1.6×10^-2 |
| FP8 E5M2 | 1 | 5 | 2 | 57344 | 约 6.1×10^-5 |
(FP8 的两种编码来自 NVIDIA 等的论文 FP8 Formats for Deep Learning:E4M3 精度略高,多用于前向;E5M2 范围略大,多用于反向。表中上下限可由位分配直接推算。)
表格里最扎眼的是 FP16 那一行:最大值只有 65504,最小正规数 6.1×10^-5——从上限到下限(含次正规数)只覆盖约 40 个 2 的次方。而神经网络里的梯度偏偏又小又散,麻烦就来了。
它怎么工作
为什么直接用 FP16 训练会崩
两个死法:上溢和下溢。
上溢好理解:某个激活或梯度超过 65504,FP16 只能记成无穷大(inf),随后污染整条反向链路变成 NaN,训练直接报废。
下溢更隐蔽也更常见。网络里大量梯度的量级在 10^-8 甚至更小——比 FP16 的最小次正规数 5.96×10^-8 还小,转成 FP16 后直接归零。梯度是零,这一步这个参数就不更新:小贡献被系统性抹掉,模型「悄悄变笨」。NVIDIA 的混合精度指南给过实测:SSD 目标检测网络的梯度值直接转成 FP16,约 31% 会下溢成零。想靠扩大动态范围自救?FP16 的指数位只有 5 位,没门。
混合精度的三件事
混合精度不是「全用 FP16」,而是各取所长。Mixed Precision Training(NVIDIA,ICLR 2018)给出的方案是三件事配合:
第一,前向与反向用半精度算。 权重、激活、梯度的存储与矩阵乘法用 FP16,这是速度与显存收益的来源——论文报告显存消耗接近减半。
第二,FP32 主权重副本(master weights)。 优化器维护一份 FP32 的权威权重,每步迭代从它拷出一份 FP16 副本去跑前向反向。为什么要多一份?参数更新量常常远小于权重本身的量级,若直接在 FP16 上做 w ← w + Δw,小更新可能低于当前权重最后一位尾数所对应的台阶,加法等于没加(被舍入吞掉)。FP32 主权重保证小更新能持续累积。
第三,动态损失缩放(loss scaling)。 算出 loss 后乘一个大系数 S 再反向,链式法则让所有梯度同步放大 S 倍,把 10^-8 量级的小梯度抬进 FP16 可表示区间;优化器更新前在 FP32 里除以 S 还原。S 不是拍死的常数,而是动态调整的,工作循环是:初始给一个大值(NVIDIA 文档示例用 2^24,PyTorch GradScaler 默认 2^16);每连续 N 步(默认 2000)无溢出,就把 S 乘 2 试着抬高;一旦在梯度里检出 inf 或 NaN,本步更新作废跳过、S 减半重来——只要跳步不频繁,对收敛没有影响。
动手验证:不到 60 行 numpy 看清三件事
下面的演示代码不需要 GPU,本地即可运行(本文写作时已实测跑通,输出附后):
# 混合精度训练演示:FP16 下溢与损失缩放 / FP16 vs BF16 舍入
import numpy as np
f16, f32 = np.finfo(np.float16), np.finfo(np.float32)
print(f"FP16 最大 {f16.max} 最小正规数 {f16.tiny:.3e} 最小次正规数 {f16.smallest_subnormal:.3e}")
print(f"FP32 最大 {f32.max:.3e} 最小正规数 {f32.tiny:.3e}")
def to_bf16(x):
"""模拟 BF16:保留 FP32 的符号位+8 位指数位,尾数舍入到 7 位(round-to-nearest-even)。"""
x = np.asarray(x, dtype=np.float32)
u = x.view(np.uint32)
return (((u + 0x7FFF + ((u >> 16) & 1)) >> 16) << 16).view(np.float32)
print("\n== 演示一:小梯度在 FP16 下下溢,损失缩放把它救回来 ==")
grads = np.array([2.0e-8, 5.0e-8, 3.0e-7, 1.0e-5, 2.0e-2], dtype=np.float32)
S = 2.0**16 # 缩放系数,GradScaler 默认初始值
direct = grads.astype(np.float16)
scaled = (grads * np.float32(S)).astype(np.float16) # 先在 FP32 里放大,再转 FP16
restored = scaled.astype(np.float32) / np.float32(S) # 反缩放留在 FP32 里做
for g, d, s, r in zip(grads, direct, scaled, restored):
print(f"真值 {g:10.2e} | 直接转FP16 {d:12.4e} | x2^16后 {s:12.4e} | 反缩放恢复 {r:12.3e}")
print("\n== 演示二:动态损失缩放状态机(对照 GradScaler 默认参数,间隔 2000 缩短为 2 便于演示)==")
scale, growth_factor, backoff_factor, growth_interval = 2.0**16, 2.0, 0.5, 2
steps = [ # 每步的 FP32 真实梯度(模拟)
np.array([3.0e-8, 1.0e-3], dtype=np.float32),
np.array([5.0e-8, 2.0e-3], dtype=np.float32),
np.array([0.8, 4.0e-8], dtype=np.float32), # 乘缩放系数后超 65504,上溢
np.array([2.0e-8, 1.0e-3], dtype=np.float32),
]
tracker = 0
with np.errstate(over="ignore"):
for i, g32 in enumerate(steps, 1):
s = scale # 本步实际使用的缩放系数
g16 = (g32 * np.float32(s)).astype(np.float16)
overflow = bool(np.isinf(g16.astype(np.float64)).any() or np.isnan(g16).any())
if overflow:
scale = s * backoff_factor
note = f"溢出 -> 跳过本步更新,缩放系数 x{backoff_factor} -> {scale:g}(下一步生效)"
tracker = 0
else:
note = "正常 -> 执行更新"
tracker += 1
if tracker == growth_interval:
scale = s * growth_factor
note += f";连续 {growth_interval} 步无溢出,缩放系数 x{growth_factor} -> {scale:g}(下一步生效)"
tracker = 0
print(f"step {i}: 缩放系数={s:g} 缩放后梯度={g16} {note}")
print("\n== 演示三:FP16 舍入 vs BF16 舍入 ==")
for v in [0.1, 12.345, 70000.0, 1.0e-30]:
with np.errstate(over="ignore"):
a = np.float16(v)
b = to_bf16([v])[0]
ea = abs(float(a) - v) / abs(v) if np.isfinite(a) and v != 0 else float("inf")
eb = abs(float(b) - v) / abs(v) if np.isfinite(b) and v != 0 else float("inf")
print(f"真值 {v:10.3e} | FP16 {a:14.8g} 误差 {ea:9.1e} | BF16 {b:14.8g} 误差 {eb:9.1e}")
三段输出对应三个结论:
- 演示一:2×10^-8 的梯度直接转 FP16 变成 0;乘上 2^16 后是 1.31×10^-3,稳稳落在可表示区间,反缩放后恢复出 1.999×10^-8。损失缩放没有创造信息,只是把信息挪进 FP16 看得见的窗口。
- 演示二:完整的状态机——连续无溢出就试着抬系数,一旦出现 inf 立即跳步减半,权重始终不被坏梯度污染。
- 演示三:12.345 在 FP16 下相对误差 1.0×10^-4,BF16 下 2.4×10^-3——FP16 尾数多 3 位,范围内确实更准;但 70000 与 10^-30 在 FP16 分别上溢为 inf、下溢为 0,BF16 泰然处之。一个让精度换范围,一个让范围换精度。
BF16 为什么省心
BF16 的思路与 FP16 反着来:尾数砍到 7 位,指数位保住完整的 8 位——与 FP32 相同。Google 团队在博客里说得直白:BF16 的动态范围与 FP32 完全相同("the dynamic range of bfloat16 is identical to that of FP32"),下溢、上溢、NaN 的行为都和 FP32 一致,因此几乎不需要损失缩放,代码也基本不用改,接近 FP32 的「drop-in 替换」。
代价是精度变粗:有效位数只有 8 位(FP16 是 11 位)。演示三里已经看到,0.1 的舍入误差 BF16 约为 FP16 的 4 倍。深度学习恰好对此不敏感——Google 团队的判断是,网络「对指数位宽的敏感度远高于尾数位宽」,丢一点精度基本无伤,甚至有研究把某些模型的精度略升归因于正则化效应。
硬件支持上(截至 2026-10):NVIDIA 从 Ampere 架构(2020 年的 A100 起)为 BF16 提供 Tensor Core 原生支持,此后 Hopper、Blackwell 一脉相承;Google TPU 则更早就在矩阵单元 MXU 里用 BF16 相乘、FP32 累加。今天的新卡上 BF16 已是事实上的训练默认精度。怎么选?一句话:Ampere 及之后,能 BF16 就 BF16,省掉整套缩放机制;老卡或 BF16 不可用时,FP16 + 动态损失缩放是同样成熟的路线。
边界在哪
- 哪些算子留在 FP32。 softmax 与归一化类(LayerNorm、BatchNorm 的统计量)、所有跨大量元素的累加,官方都建议留在 FP32:NVIDIA 指南明确「大归约算出的值应留在 FP32」,Google 的 MXU 也是 BF16 相乘、FP32 累加。乘法的误差是局部的,求和的误差会随项数累积——这是低精度的第一条红线。
- 梯度累加的精度。 用损失缩放做梯度累加时,多个微批的梯度要保持带缩放状态累加,直到最后一个微批才反缩放、判溢出、更新——中途反缩放会破坏下溢保护。PyTorch 官方文档对这套流程有专门说明,并规定 unscale_ 每个 optimizer 每步只能调用一次。
- FP8 训练能跑,但还不是开箱即用。 E4M3/E5M2 的范围(最大 448 / 57344)比 BF16 窄了几个数量级,必须按张量动态缩放。现状是 Hopper 起的 H 系列卡配合 NVIDIA Transformer Engine,前向用 E4M3、反向用 E5M2,属于面向专家的优化,本文点到为止。
- 训练与推理的精度需求不同。 推理没有反向传播,权重与激活的分布训练完就固定了,压到 INT8/INT4 也常能接受;训练要反复求梯度、跨步累加更新,对动态范围与舍入误差敏感得多。所以训练侧的主流是「低精度浮点算 + FP32 主权重」的混合精度,而不是直接压到定点整数。
参考资料
- Mixed Precision Training(NVIDIA,ICLR 2018)
- NVIDIA Docs: Training With Mixed Precision
- PyTorch: Automated Mixed Precision examples(autocast 与 GradScaler)
- PyTorch GradScaler 源码(默认参数 init_scale=2^16 等)
- Google Cloud Blog: BFloat16, the secret to high performance on Cloud TPUs
- FP8 Formats for Deep Learning(arXiv 2209.05433)
读者留言
COMMENTS 暂无还没有留言,来说第一句?