数值计算与混合精度

从浮点数、Loss Scaling 到 Kernel 融合

Posted by Liu Mengxuan on August 28, 2026

神经网络里的公式通常是在实数上推导的,但 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 编码还包含偏置指数、隐含的前导位,以及对 0infNaN 和 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 的信息。

E4M3E5M2 的名字也直接描述了 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.0001.0011.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.991000 太近,舍入后仍可能保存成:

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

于是两个本应相同的加法顺序可能得到 10

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.81.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 或模型本身不稳定而出现 infNaN

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 的角度,混合精度可以浓缩成一句话:

用低精度承担适合吞吐的计算,用高精度保护容易丢失信息的地方,再用稳定且可融合的算法把两者组织起来。