神经网络里的公式通常是在实数上推导的,但 GPU 实际操作的是有限精度的浮点数。于是,同一个公式在纸面上是正确的,放进 FP16、BF16 或 FP8 的计算流程后,却可能遇到舍入误差、溢出、下溢和 NaN。
混合精度训练要解决的不是“低精度一定不准确”这个简单问题,而是要把不同工作分给合适的数值格式:大规模矩阵乘法用低精度换取吞吐,累加、归约和优化器更新在关键位置保留更高精度。
1
2
3
4
5
6
7
浮点数只能表示有限集合
↓
精度、动态范围和舍入误差
↓
溢出、下溢、消减与条件数
↓
混合精度、Loss Scaling、稳定且可融合的算法
1. 浮点数是有限集合
1.1 计算机保存的不是任意实数
数学中的实数有无穷多个,但一个数据类型只有有限个 bit,因此只能表示有限个数。
先用十进制做一个玩具例子。假设某种格式只保留 3 位有效数字:
1
2
3
1234 → 1.23 × 10^3
0.001234 → 1.23 × 10^-3
1.234 → 1.23
真实值 1.234 不在这套有限的数字表里,计算机只能把它舍入到附近的可表示值 1.23。
浮点数的想法和科学计数法类似:用一部分信息表示符号,一部分表示有效数字,再用指数表示数量级。典型二进制浮点数可以抽象成:
$(-1)^s\times\text{significand}\times2^{\text{exponent}}$
例如:
$(-1)^0\times1.5\times2^3=1.5\times8=12$
这里:
1
2
3
s = 0 符号为正
significand=1.5 有效数字
exponent=3 乘以 2^3
真实 IEEE 754 编码还包含偏置指数、隐含的前导位,以及对 0、inf、NaN 和 subnormal 的特殊编码。现在先抓住最重要的直觉:bit 是有限的,所以可表示的数也是有限的。
1.2 不同格式怎样分配 bit?
浮点格式通常把 bit 分成:
1
1 个符号位 + 指数位 + 尾数字段位
教程中的对比表可以这样读:
| 格式 | 总位数 | 指数位 | 尾数字段位 | 主要特点 |
|---|---|---|---|---|
| FP32 | 32 | 8 | 23 | 范围和精度较均衡 |
| TF32 | 19(乘法输入语义) | 8 | 10 | 保留 FP32 范围,降低乘法精度,常用 FP32 累加 |
| FP16 | 16 | 5 | 10 | 精度尚可但动态范围较小 |
| BF16 | 16 | 8 | 7 | 接近 FP32 动态范围,精度低于 FP16 |
| FP8 E4M3 | 8 | 4 | 3 | 精度优先、范围较小,具体编码依实现规范 |
| FP8 E5M2 | 8 | 5 | 2 | 范围优先、精度更低,具体编码依实现规范 |
表中的“尾数字段位”没有把规格化二进制数默认存在的前导 1 算进去。例如 FP16 的尾数字段是 10 bit,但有效数字通常可以粗略理解为有 11 bit 的信息。
E4M3 和 E5M2 的名字也直接描述了 bit 的分配:
1
2
E4M3 = 4 个 exponent bit + 3 个 mantissa bit
E5M2 = 5 个 exponent bit + 2 个 mantissa bit
TF32 需要特别注意:它通常不是一个像 FP16 那样单独存放的常规 dtype,而是 NVIDIA Tensor Core 使用的一种乘法输入精度语义。实际内部累加精度要看硬件和算子实现。
2. 精度和动态范围是两件事
这两个词很容易混在一起,但它们回答的是不同问题:
1
2
动态范围:最小能表示多小的数,最大能表示多大的数?
精度:同一个数量级附近,能区分多接近的两个数?
可以把浮点格式想成一张地图。动态范围是地图覆盖的地域,精度是地图在局部区域画得有多细。
FP16 和 BF16 都占 16 bit,但分配方式不同:
1
2
3
FP16 = 1 + 5 + 10
BF16 = 1 + 8 + 7
符号 指数 尾数
2.1 FP16:局部更细,但视野较窄
FP16 的指数位只有 5 位,最大有限值约为:
$65504$
在 1 附近,FP16 的相邻规格化数间隔约为:
$2^{-10}=0.0009765625$
因此 FP16 在 1 附近能区分相对细小的差异,但它表示不了 100000:
1
100000 > 65504
结果可能变成 inf。
2.2 BF16:视野更广,但局部更粗
BF16 使用 8 个指数位,和 FP32 相同,因此它的动态范围接近 FP32。在 1 附近,BF16 的相邻数间隔约为:
$2^{-7}=0.0078125$
所以 1.000、1.001、1.002 这样的细小差异,可能在转换为 BF16 后落到同一个可表示值;但 100000 通常仍能表示。
可以这样记:
FP16 像“近处看得更细、但视野窄”;BF16 像“视野很广、但局部细节粗”。
大模型训练里的激活、梯度和 logits 可能处于完全不同的数量级:
1
2
3
某些梯度:1e-8
某些激活:1e5
某些 logits:几十到几百
因此选择 dtype 不是单纯追求“位数越多越好”,而是要看这个张量更怕溢出,还是更怕细小差异被舍掉。
3. 舍入与机器精度
3.1 为什么 0.1 + 0.2 不一定等于 0.3?
很多十进制小数不能被二进制有限位精确表示。在 Python 中:
1
2
3
4
>>> 0.1 + 0.2 == 0.3
False
>>> 0.1 + 0.2
0.30000000000000004
这不是加法规则错了,而是:
1
2
3
真实的 0.1 → 存储为附近的二进制浮点数
真实的 0.2 → 存储为附近的二进制浮点数
两者相加 → 再次舍入
打印时通常会隐藏很小的误差,但误差会参与后续计算。
3.2 浮点数间隔不是固定的
在 1 附近,FP16 的相邻间隔大约是 0.001;到了 1000 附近,间隔会变成更大的量级。因此,同一个 0.01 更新,在不同大小的权重附近可能有完全不同的结果。
假设当前权重为:
1
w = 1000.00
数学上想做一个很小的更新:
1
w + delta = 1000.00 - 0.01 = 999.99
如果 1000 附近的 FP16 间隔大约是 0.5,那么 999.99 离 1000 太近,舍入后仍可能保存成:
1
1000.00
这个现象可写成:
$\operatorname{fl}(x+\delta)=x$
其中 fl 表示“把结果舍入到当前浮点格式可表示的数”。更准确地说,当 delta 小到不足以让结果跨过附近的舍入边界时,存储值就不会变化。
3.3 为什么需要 FP32 master weights?
经典 FP16 混合精度训练通常维护两份参数:
1
2
FP16/BF16 compute weights:用于前向和大部分反向计算
FP32 master weights: 用于优化器累积更新
一次更新的流程是:
1
2
3
4
5
6
7
8
FP32 master weight
↓ 转换
FP16/BF16 compute weight
↓
前向、反向,得到梯度
↓ 梯度转 FP32
FP32 master weight ← 优化器更新
↓ 下一步再转换
如果直接更新 FP16 的 1000:
1
1000.00 - 0.01 → 仍然是 1000.00
而 FP32 master weight 可以记录:
1
1000.00 → 999.99 → 999.98 → 999.97
所以 master weight 更像优化器手中的“权威账本”。计算用的低精度副本暂时看不出每一个小变化,但这些更新会先在 FP32 中积累,积累到足够大后再反映到低精度副本。
现代 BF16 训练是否额外保留一份 FP32 master weights,则取决于框架和优化器实现;但优化器状态通常仍会使用 FP32,以保留更新信息。
4. 非结合性和归约顺序
数学中的加法满足结合律:
$(a+b)+c=a+(b+c)$
浮点计算中,每一步都可能舍入,因此实际结果可能不同。
用一个有限精度的直觉例子。设:
1
2
3
a = 1e8
b = -1e8
c = 1
如果先让两个大数抵消:
1
(a + b) + c = 0 + 1 = 1
如果先计算 b + c,由于 1 相对于 1e8 太小,低精度下可能被舍掉:
1
2
b + c ≈ -1e8
a + (b + c) ≈ 1e8 - 1e8 = 0
于是两个本应相同的加法顺序可能得到 1 和 0。
4.1 GPU 为什么特别在意加法顺序?
假设要计算:
$x_1+x_2+x_3+x_4$
CPU 可以按顺序计算:
1
(((x1 + x2) + x3) + x4)
GPU 则可能让不同线程先计算局部和:
1
2
3
线程 1:x1 + x2
线程 2:x3 + x4
最后: (x1 + x2) + (x3 + x4)
这种把一组数压缩成一个结果的过程叫归约。它常出现在:
1
2
3
4
Softmax 的最大值和指数和
LayerNorm 的均值和方差
一个 batch 的 loss 平均值
GEMM 中沿 K 维累加部分积
因此不同的 block 划分、线程数或原子操作顺序,可能让结果的最后几位不同。结果不完全 bitwise 一致,不一定意味着 GPU 算错;工程上要确认误差处于可接受范围,并按需求选择确定性实现。
4.2 减少归约误差的方法
第一种:用更高精度累加。
典型 GEMM 采用:
1
FP16/BF16 输入 × FP16/BF16 输入 → FP32 累加
低精度输入提供更高吞吐,FP32 中间和减少大量部分积相加时的误差。
第二种:成对或树形求和。
1
2
3
4
5
x1 x2 x3 x4
\ / \ /
x1+x2 x3+x4
\ /
总和
它通常先合并规模接近的数,减少“很大的数直接加很小的数”造成的小数被吞掉的问题。
第三种:先局部归约,再合并部分和。
1
2
3
4
5
6
7
大量元素
↓
每个线程 / warp / block 算局部和
↓
少量 partial sums
↓
使用更高精度合并
这样可以减少全局同步和显存访问,也让最终合并更容易使用 FP32。
第四种:Kahan 补偿求和。
Kahan 求和额外保存一个补偿量,记录之前因为舍入而没有真正加进去的小误差:
1
2
3
4
5
6
7
8
total = 0.0
correction = 0.0
for x in values:
y = x - correction
t = total + y
correction = (t - total) - y
total = t
它通常比普通求和更准确,但需要更多指令和状态,并行实现也更复杂。因此 GPU 大规模 GEMM 不会对每个累加都使用 Kahan,而是在对精度特别敏感的归约中考虑它。
5. 溢出、下溢和非有限值
5.1 溢出:数太大了
FP16 的最大有限值约为 65504。如果计算:
1
60000 × 2 = 120000
这个结果超出 FP16 范围,可能变成:
1
inf
后续计算还可能继续传播:
1
2
3
inf + 1 = inf
inf - inf = NaN
0 × inf = NaN
大模型中可能溢出的来源包括过大的激活值、梯度爆炸、过大的 logits,以及未做稳定化的指数运算,例如 exp(100)。
5.2 下溢:数太小了
如果梯度是:
1
1e-8
而当前低精度格式无法表示这么小的数,存储时可能变成:
1
0
之后再把它转为 FP32,也只是:
1
0(FP16)→ 0(FP32)
转换无法找回已经丢失的信息。Loss Scaling 就是后面专门用来缓解这类小梯度下溢的机制。
5.3 NaN:非法或未定义的结果
常见来源包括:
$0/0,\qquad \infty-\infty,\qquad \sqrt{-1}$
在 Attention 中,如果某一行被 mask 后变成:
1
[-inf, -inf, -inf]
稳定 Softmax 会先求最大值:
1
m = -inf
接着计算 z - m,就出现:
1
-inf - (-inf) = NaN
所以排查 loss 变成 NaN 时,应该寻找第一个非有限张量,而不是只盯着最后的 loss:
1
输入 → embedding → attention score → softmax → norm → loss
还要检查 mask 是否制造了全屏蔽行、是否存在除零、负数开方或自定义 Kernel 的边界错误。
6. 灾难性消减
灾难性消减发生在两个非常接近的大数相减时:高位数字彼此抵消,最后剩下的小差异却带着前面累积的舍入误差。
先用简单例子说明。真实计算是:
1
2
3
a = 1000000.1
b = 1000000.0
a - b = 0.1
如果使用的格式只能把这两个数都近似保存为 1000000,结果就变成:
1
1000000 - 1000000 = 0
真正有用的 0.1 被消掉了。
6.1 方差公式中的消减
方差可以写成:
$\operatorname{Var}(X)=\mathbb{E}[X^2]-\mathbb{E}[X]^2$
考虑数据:
1
[1000000, 1000001]
平均值为 1000000.5,真实方差为:
$\frac{(1000000-1000000.5)^2+(1000001-1000000.5)^2}{2}=0.25$
但 E[X²] 和 E[X]² 都约为 10^12,最后要做两个巨大且接近的数相减:
1
2
3
4
约 1,000,001,000,000.5
-约 1,000,001,000,000.25
--------------------------------
0.25
只要上面两个大数各自有一点舍入误差,下面的 0.25 就可能被严重污染。
更稳定的方式是先求均值,再直接计算偏差:
1
2
3
偏差:-0.5, +0.5
平方: 0.25, 0.25
平均: 0.25
常见稳定算法包括 two-pass variance 和 Welford 在线算法。它们的共同点是:尽量围绕均值处理小偏差,避免直接相减两个巨大的近似值。
这说明一个重要事实:
数学上等价的公式,放进有限精度计算后,数值稳定性可能完全不同。
7. 条件数:输入误差会被放大多少
前面讲的是算法如何产生误差。这一节讨论另一个问题:问题本身是否对输入误差敏感。
设:
$y=f(x)$
输入有一个小扰动 delta x,输出发生 delta y。如果近似满足:
$\frac{|\delta y|}{|y|}\approx\kappa\frac{|\delta x|}{|x|}$
那么 kappa 就描述了相对误差被放大的程度,它叫条件数。
7.1 一个条件数很大的问题
考虑:
$y=x_1-x_2$
如果:
1
2
3
x1 = 1000000
x2 = 999999
y = 1
两个输入都在约 10^6 的量级,但输出只有 1。只要输入各自有大约 0.1 的误差,输出就可能从 1 变成 0.8、1.2,甚至更糟。
对于这个减法问题,一个直观的相对敏感程度约为:
$\frac{|x_1|+|x_2|}{|x_1-x_2|}\approx2\times10^6$
分母很小,意味着条件数很大:输入的微小相对误差可能被放大很多倍。
7.2 矩阵条件数
对于可逆矩阵,在某个一致范数下:
$\kappa(\mathbf A)=\lVert\mathbf A\rVert\,\lVert\mathbf A^{-1}\rVert$
条件数大,表示矩阵接近不可逆,或者某些方向会对误差特别敏感。矩阵求逆、解线性方程、SVD 和低精度矩阵运算都可能受到影响。
要区分两件事:
1
2
条件数:问题本身有多敏感
算法稳定性:实现是否额外制造了不必要的误差
增加 bit 可以缓解舍入误差,但不能从根本上消除一个病态问题的敏感性。数值稳定算法能减少额外误差,却不能让病态问题突然变成良态问题。
8. 混合精度训练的基本模式
混合精度不是“把所有张量都换成 FP16”,而是根据算子的工作特点分配精度。
一次训练 step 可以概括为:
1
2
3
4
5
6
7
8
9
10
11
FP32 参数 / master weights
↓ 转换为 FP16 或 BF16
低精度前向:Linear、GEMM、Attention 中的大型矩阵乘
↓
loss(必要时进行 Loss Scaling)
↓
反向传播:低精度输入,关键累加或敏感操作使用 FP32
↓
梯度反缩放并转为 FP32
↓
FP32 优化器状态和参数更新
常见的选择是:
| 计算部分 | 常见精度策略 | 原因 |
|---|---|---|
| 大型 GEMM 输入 | FP16/BF16 | Tensor Core 吞吐高、显存占用低 |
| GEMM 累加 | FP32 或硬件指定的更高精度 | 减少大量部分积累加的误差 |
| Softmax、Norm、归约 | 保留或内部提升到 FP32 | 包含 exp、除法、均值、方差和大量求和 |
| 优化器状态 | 通常 FP32 | 保留动量、二阶矩和小更新 |
| 参数更新 | 通常 FP32 master weights | 避免低精度舍入吞掉小步更新 |
这里的“通常”不能替代实际检查。不同 GPU、框架和算子可能使用不同的内部累加路径,不能只根据输入 dtype 猜测整个计算过程;需要查算子文档,必要时用 profiler 验证。
9. Loss Scaling 为什么有用
9.1 它解决的是小梯度下溢
假设真实梯度为:
1
grad = 1e-8
它在 FP16 中可能下溢为 0。于是先把 loss 放大 S 倍:
$L'=S L$
根据链式法则:
$\nabla L'=\nabla(SL)=S\nabla L$
如果取:
1
S = 1024
那么反向传播中看到的梯度是:
1
1e-8 × 1024 = 1.024e-5
这个数更有机会在 FP16 中被表示。反向结束后再除以 S:
$\nabla L=\frac{\nabla L'}{S}$
理论上就回到了原始梯度。Loss Scaling 改变的是中间表示的尺度,不是最终想要的梯度。
9.2 为什么 scale 不能无限增大?
如果 scale 太小,小梯度仍然可能下溢;如果 scale 太大,正常梯度也可能溢出:
1
2
3
原始梯度:100
scale: 1024
缩放后: 102400
因此动态 Loss Scaling 通常这样工作:
1
2
3
4
1. 使用当前 scale 做前向和反向
2. 检查梯度中是否存在 inf 或 NaN
3. 如果溢出:跳过本次更新,并减小 scale
4. 如果连续多步稳定:尝试增大 scale
在实际训练中,梯度裁剪应该放在反缩放之后。以 PyTorch 的典型写法为例:
1
2
3
4
5
scaler.scale(loss).backward()
scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
scaler.step(optimizer)
scaler.update()
如果先对“被放大了 S 倍的梯度”做裁剪,裁剪阈值就被错误地缩放了。unscale_ 的作用就是先恢复真实梯度尺度,再进行梯度裁剪和优化器更新。
BF16 的动态范围接近 FP32,通常不依赖 Loss Scaling;但它仍可能因为学习率过大、激活爆炸、非法 mask 或模型本身不稳定而出现 inf 和 NaN。
10. 稳定算法通常也更适合融合
GPU Kernel 的性能很大一部分取决于显存访问。若一个操作被拆成多个 Kernel,就可能不断发生:
1
2
读显存 → 计算 → 写显存
读显存 → 计算 → 写显存
融合就是把连续的操作合并到一个 Kernel 中,尽量让中间结果停留在寄存器或共享内存中。
有意思的是,数值稳定的算法经常也更容易融合,因为它们会把计算改写成“分块后只需维护少量状态”的形式。
10.1 Stable Softmax 与 Online Softmax
朴素 Softmax 是:
$p_i=\frac{e^{z_i}}{\sum_j e^{z_j}}$
当 z 很大时,exp(z) 可能溢出。稳定写法先求最大值:
$m=\max_j z_j$
$p_i=\frac{e^{z_i-m}}{\sum_j e^{z_j-m}}$
因为 z_i - m <= 0,指数项不会因为正指数太大而溢出。
处理一个 block 时,可以维护两个状态:
1
2
m:当前看到的最大值
l:以 m 为基准的指数和
如果已有一块状态 (m_A, l_A),又读到一块状态 (m_B, l_B),令:
$m=\max(m_A,m_B)$
$l=l_Ae^{m_A-m}+l_Be^{m_B-m}$
就能把两块合并到同一个数值尺度。这种状态可合并的性质,使 Softmax 能够分块扫描,避免先物化完整的指数矩阵。
10.2 log_softmax + NLL 融合
分类训练常见的逻辑链路是:
1
logits → softmax → 概率 → log → NLL loss
但交叉熵可以直接使用 log-softmax:
$\log\operatorname{softmax}(z_i)=z_i-\operatorname{LSE}(z)$
其中:
$\operatorname{LSE}(z)=\log\sum_j e^{z_j}$
这样可以避免显式保存完整概率矩阵,再读取它并取对数。融合之后通常能够:
1
2
3
4
少一次中间张量写回
少一次中间张量读取
减少显存流量
同时保留 log-sum-exp 的数值稳定性
10.3 Welford 与 LayerNorm
LayerNorm 需要计算均值和方差。直接使用:
$\operatorname{Var}(X)=\mathbb E[X^2]-\mathbb E[X]^2$
容易遇到灾难性消减;而 Welford 算法只维护:
1
2
3
n:已经处理的样本数
mean:当前均值
M2:当前平方离差的累计量
每个 block 可以先计算自己的局部状态,之后再合并这些状态。它既减少了不稳定的大数相减,也符合 GPU 的局部归约与最终合并模式。
10.4 低精度输入、高精度累加
GEMM 也是类似的折中:
1
2
低精度输入 → Tensor Core 快速乘法
高精度累加 → 减少部分积归约误差
因此,稳定性和性能并不天然冲突。关键是寻找一种数学表达,使它同时满足:
1
2
3
数值上不容易溢出或消减
状态数量少、可以分块维护
中间结果不必频繁写回显存
11. 把第 8 节串起来
1
2
3
4
5
6
7
8
9
10
8.1 浮点数只能表示有限集合
8.2 精度决定局部细节,动态范围决定能覆盖的数量级
8.3 舍入会让小更新消失
8.4 浮点加法不满足结合律,GPU 归约顺序会影响结果
8.5 太大、太小或非法运算会产生 inf、0、NaN
8.6 接近大数相减会损失真正关心的小差异
8.7 条件数描述输入误差会被问题放大多少
8.8 不同算子选择不同精度,而不是全模型统一低精度
8.9 Loss Scaling 缓解 FP16 小梯度下溢
8.10 稳定算法往往也更容易分块和融合
从 AI Infra 的角度,混合精度可以浓缩成一句话:
用低精度承担适合吞吐的计算,用高精度保护容易丢失信息的地方,再用稳定且可融合的算法把两者组织起来。