数值计算与精度
03-线性代数与 GEMM 与 05-反向传播与梯度优化 里的公式都写在实数上。但公式与实现在硬件上是否可靠,取决于浮点数这套有限集合的行为 —— 它会舍入、溢出、下溢,而且加法不满足结合律。
这一篇讲的是「同一个公式,换个累加顺序结果就不一样」这件事的成因与对策。
浮点数是一个有限集合
位数被分给符号、指数、尾数三部分。指数位决定动态范围,尾数位决定相邻可表示数的精细程度 —— 这两个维度相互独立,是最容易混淆的一点。
| 格式 | 总位数 | 指数位 | 尾数字段位 | 主要特点 |
|---|---|---|---|---|
| FP32 | 32 | 8 | 23 | 范围与精度较均衡 |
| TF32 | 19 | 8 | 10 | 保留 FP32 范围,降低乘法精度;通常是 Tensor Core 的计算模式,不是内存存储 dtype |
| FP16 | 16 | 5 | 10 | 精度尚可,动态范围小 |
| BF16 | 16 | 8 | 7 | 接近 FP32 的动态范围,精度低于 FP16 |
| FP8 E4M3 | 8 | 4 | 3 | 精度优先,范围小 |
| FP8 E5M2 | 8 | 5 | 2 | 范围优先,精度更低 |
表中不计入隐含前导位。
位数被分给符号、指数、尾数三部分,两个维度相互独立。
(-1)^s × significand × 2^exponent
│ │ │
│ │ └─ 指数位 ──▶ 决定动态范围
│ └────────────── 尾数位 ──▶ 决定相邻可表示数的精细程度
└──────────────────────── 符号位
┌──────────┬──────┬────────┬────────┬────────────────────────────┐
│ 格式 │ 总位 │ 指数位 │ 尾数位 │ 主要特点 │
├──────────┼──────┼────────┼────────┼────────────────────────────┤
│ FP32 │ 32 │ 8 │ 23 │ 范围与精度较均衡 │
│ TF32 │ 19 │ 8 │ 10 │ 保留 FP32 范围、降乘法精度; │
│ │ │ │ │ 通常是 Tensor Core 的计算 │
│ │ │ │ │ 模式,不是内存存储 dtype │
│ FP16 │ 16 │ 5 │ 10 │ 精度尚可,动态范围小 │
│ BF16 │ 16 │ 8 │ 7 │ 接近 FP32 动态范围, │
│ │ │ │ │ 精度低于 FP16 │
│ FP8 E4M3 │ 8 │ 4 │ 3 │ 精度优先,范围小 │
│ FP8 E5M2 │ 8 │ 5 │ 2 │ 范围优先,精度更低 │
└──────────┴──────┴────────┴────────┴────────────────────────────┘
(表中不计入隐含前导位)
FP16 与 BF16 的取舍就是这两维的取舍
FP16 尾数多 3 位 ──▶ 同数量级下能表示更细的差别
指数少 3 位 ──▶ 最大有限值只有 65,504,非零数范围窄得多
BF16 牺牲精度换与 FP32 相同的指数位数
──▶ 训练中更不易溢出 / 下溢
──▶ 但仍有明显舍入误差,尤其在把很小的增量加到很大的数上时
└─ 这两条直接决定了「为什么 FP16 需要 Loss Scaling
而 BF16 通常不需要」以及「为什么大模型训练默认用 BF16」。精度与动态范围是两件事
FP16 尾数比 BF16 多 3 位,所以同一数量级下能表示更细的差别;但 FP16 指数少 3 位,最大有限值只有 65,504,非零数范围也窄得多。
BF16 牺牲精度换取与 FP32 相同的指数位数,因此训练中更不易溢出/下溢。它仍然有明显舍入误差 —— 尤其在把很小的增量加到很大的数上时。
这两条直接决定了「为什么 FP16 需要 Loss Scaling 而 BF16 通常不需要」以及「为什么大模型训练默认用 BF16」。
舍入与机器精度
>>> 0.1 + 0.2 == 0.3
False很多十进制小数不能被二进制浮点精确表示。更关键的一条:浮点数间距随数值绝对值增大而增大。若
这一步的更新等于没做。
这就是优化器必须保留 FP32 master weights 的原因,也是 05-反向传播与梯度优化 里 Adam 状态用 FP32 的数学依据。低精度权重上过小的更新会被直接舍掉 —— 即使梯度算得再准。
加法不满足结合律
由于每步都舍入:
并行归约用树形加法,不同线程数、block 划分、原子操作顺序都可能产生略不同的结果。这也解释了 03-线性代数与 GEMM 里「不同 tile 划分结果不逐位一致」以及 03-集合通信与 NCCL 里「AllReduce 的顺序影响结果」。
非确定性不等于实现错误
它必须在可接受误差内,并符合任务对复现性的要求。要求逐位可复现时,就得固定归约顺序与并行度 —— 这会牺牲性能,是个需要明确决策的取舍,不是默认行为。
提高归约精度的手段:
| 手段 | 说明 |
|---|---|
| 用更高精度累加 | 最直接,代价是带宽与算力 |
| 成对 / 树形求和 | 避免极端尺度长期相加 |
| Kahan 等补偿求和 | 精度高,但额外指令成本 |
| 先局部归约,再合并部分和 | 并行场景的标准做法 |
溢出、下溢与非有限值
| 现象 | 成因 |
|---|---|
| overflow | 绝对值超过最大有限值 → 常变成 |
| underflow | 太接近 0 → 进入 subnormal 或舍入为 0 |
| NaN | |
| 参与后续运算可能迅速变成 NaN |
排查 loss 变 NaN 时,要找「第一个」非有限张量,而不是只盯最终 loss。 最终 loss 是传播结果,源头可能在几十层之前。常见源头:过大的 logits、除零、负数开方、无效 mask、梯度溢出、自定义 kernel 的边界。
PyTorch 里可以用
torch.autograd.set_detect_anomaly(True)或逐层 hook 来定位第一个异常点。
非有限值有四类现象,但排查的关键在「找第一个」。
现象 成因
────────────────────────────────────────────────────
overflow 绝对值超过最大有限值 → 常变成 ±∞
underflow 太接近 0 → 进入 subnormal 或舍入为 0
NaN 0/0、∞−∞、非法运算
∞ 传播 参与后续运算可能迅速变成 NaN
排查顺序:找「第一个」非有限张量,而不是只盯最终 loss
最终 loss 是传播结果,源头可能在几十层之前。
│
▼
常见源头(按出现频率)
过大的 logits / 除零 / 负数开方 /
无效 mask / 梯度溢出 / 自定义 kernel 的边界
│
▼
定位工具
torch.autograd.set_detect_anomaly(True)
或逐层 hook
└─ 盯 loss 只能知道「已经坏了」,找第一个异常点才知道「哪里坏的」。灾难性消减
两个非常接近的大数相减时,高位有效数字抵消,剩余结果的相对误差可能巨大。
经典例子是用下面这个公式算方差:
两项可能都很大且接近,相减后精度损失严重。「大均值、小方差」的数据上尤其明显。
更稳定的做法是 two-pass variance,或 Welford 在线算法:
总体方差为
Welford 的关键附加价值:它能合并不同分块的统计量。 这使它天然适合并行的块级归约 —— 一个「数值稳定」的算法同时也是一个「容易分块」的算法,后面会看到这不是巧合。
稳定算法通常也容易分块,共同点是「可组合的状态」。
算法 双层收益
──────────────────────────────────────────────────────────────
Stable Softmax 先做 max 归约再做 exp/sum 归约;
Online Softmax 还能合并分块状态
log_softmax + NLL 融合 既避开「先 Softmax 再 log」的不稳定路径,
又减少中间张量写回
Welford 同时维护统计量,天然适合 LayerNorm 的块级归约
GEMM 低精度输入 + 高精度累加,兼顾吞吐与误差
Welford 的在线更新
n ← n + 1
δ = x_n − μ_{n−1}
μ_n = μ_{n−1} + δ / n
M_{2,n} = M_{2,n−1} + δ (x_n − μ_n)
└─ 总体方差为 M_{2,n} / n
└─ 关键附加价值:它能合并不同分块的统计量
⟹ 天然适合并行的块级归约
共同点:这些算法的状态是「少量统计量」
分块处理时只需传递:最大值 / 指数和 / 均值 / 平方离差 ——
而不必保留完整中间结果。
└─ 这正是它们既能数值稳定、又能写成分块 kernel 的原因。
一个「数值稳定」的算法同时也是一个「容易分块」的算法,
这不是巧合。条件数
描述输入微小扰动对输出的影响。对可逆矩阵,在某个一致范数下:
条件数大表示问题病态:即使算法实现正确、输入只含很小的舍入误差,输出也可能变化很大。
要区分两件事:
- 问题本身是否病态(
大) - 所选算法是否数值稳定
更多位数能缓解舍入误差,但不能从根本上消除病态问题的敏感性。这是「加精度解决不了所有数值问题」的原因。
混合精度训练
混合精度不是把所有东西都换成 FP16。常见模式是五条:
| # | 做法 |
|---|---|
| 1 | 大型矩阵乘用 FP16/BF16 输入,走 Tensor Core |
| 2 | 乘加在 FP32 或更高的指定累加精度中完成 |
| 3 | 归约、归一化、Softmax 等敏感操作保留或内部提升到 FP32 |
| 4 | 优化器状态和/或 master weights 保留 FP32 |
| 5 | 输出 dtype 按硬件与算子选择 |
第 5 条的实践含义:每个框架与硬件的具体策略不同,应查实际算子文档与 profiler,不要从输入 dtype 猜内部计算路径。「我输入了 BF16,所以内部就是 BF16」是一个常见的错误假设。
Loss Scaling 为什么有用
反向传播中的 FP16 小梯度可能下溢为 0。令缩放因子为
梯度先被放大到 FP16 可表示范围,反向完成后再除回来:
动态 Loss Scaling 的四步:
- 用当前 scale 计算反向
- 检查梯度是否含
inf/NaN - 若溢出 → 跳过更新并减小 scale
- 连续稳定若干步 → 尝试增大 scale
Loss Scaling 解决的是小梯度下溢;scale 过大会反过来造成溢出。BF16 因为动态范围更大,通常不依赖 Loss Scaling —— 但仍可能因模型本身不稳定产生非有限值。
Loss Scaling 解决的是「小梯度下溢」,动态调整分四步。
问题:反向传播中的 FP16 小梯度可能下溢为 0
│
▼
令缩放因子为 S,损失放大 S 倍
L' = S·L ⟹ ∇L' = S·∇L
└─ 梯度先被放大到 FP16 可表示范围
│
▼
反向完成后再除回来
∇L = ∇L' / S
动态 Loss Scaling 的四步
① 用当前 scale 计算反向
② 检查梯度是否含 inf / NaN
③ 若溢出 ──▶ 跳过更新,并减小 scale
④ 连续稳定若干步 ──▶ 尝试增大 scale
└─ scale 过大会反过来造成溢出 —— 所以是一个双向调节。
BF16 因为动态范围更大,通常不依赖 Loss Scaling,
但仍可能因模型本身不稳定产生非有限值。
混合精度不是「把所有东西都换成 FP16」,常见模式是五条
① 大型矩阵乘用 FP16 / BF16 输入,走 Tensor Core
② 乘加在 FP32 或更高的指定累加精度中完成
③ 归约、归一化、Softmax 等敏感操作保留或内部提升到 FP32
④ 优化器状态和 / 或 master weights 保留 FP32
⑤ 输出 dtype 按硬件与算子选择
└─ 第 ⑤ 条的实践含义:每个框架与硬件的具体策略不同,
应查实际算子文档与 profiler,不要从输入 dtype 猜内部计算路径。
「我输入了 BF16,所以内部就是 BF16」是一个常见的错误假设。稳定算法通常也更适合融合
数值稳定与性能优化不总是冲突,很多时候方向一致:
| 算法 | 双层收益 |
|---|---|
| Stable Softmax | 先做 max 归约再做 exp/sum 归约;Online Softmax 还能合并分块状态 |
log_softmax + NLL 融合 | 既避开「先 Softmax 再 log」的不稳定路径,又减少中间张量写回 |
| Welford | 同时维护统计量,天然适合 LayerNorm 的块级归约 |
| GEMM | 低精度输入 + 高精度累加,兼顾吞吐与误差 |
共同点是这些算法具有「可组合的状态」 —— 分块处理时只需传递少量统计量(最大值、指数和、均值、平方离差),而不必保留完整中间结果。这正是它们既能数值稳定、又能写成分块 kernel 的原因。
这条判据在 11-Self-Attention 机制 的 Online Softmax 与 FlashAttention 里体现得最清楚。
相关
- 03-线性代数与 GEMM —— 浮点非结合性对 GEMM 结果的影响
- 05-反向传播与梯度优化 —— FP32 主权重与 Adam 状态的显存账
- 04-概率、Softmax 与信息论 —— 稳定 Softmax 与方差的不稳定算法
- 01-GPU 硬件架构与存储层次 —— 各精度格式的硬件支持矩阵
- 02-分布式训练总论与显存账本 —— 16Ψ 与 18Ψ 的口径差异来自梯度精度
参考
- https://caomaolufei.github.io/AIInfraGuide/guides/模块一-前置知识/第2章-数学基础
- https://www.deeplearningbook.org/contents/numerical.html
- https://docs.nvidia.com/deeplearning/performance/mixed-precision-training/
- https://blogs.nvidia.com/blog/tensorfloat-32-precision-format/
- https://docs.pytorch.org/docs/stable/notes/numerical_accuracy.html
- https://pytorch.org/docs/stable/amp.html
YJ