混合精度训练(Mixed Precision):FP16、BF16、FP8——省显存还不降精度

有一个问题困扰了我很久:为什么显卡宣传的 FP16 算力是 FP32 的两倍,但我用 FP16 训练模型时,Loss 不是直接变成 NaN,就是精度反而比 FP32 差一大截?后来我才发现,把 FP32 的模型直接换成 FP16 训练,就像把一辆赛车的仪表盘刻度全部抹掉只留数字——不是不能跑,而是稍不留神就翻车。 混合精度训练,就是给这辆赛车装上一套额外的安全系统,让你既能享受高速,又不会失控。这篇文章,我想把 FP16、BF16、FP8 这三种主流半精度格式的门道一次讲清楚。

AI technology illustration

一张表看懂 FP32、FP16、BF16、FP8

先别管什么 loss scaling,咱们从起点开始:为什么不同的浮点格式会直接影响训练? 因为神经网络里的数字——权重、激活值、梯度——都靠浮点格式存在显存里。格式不同,能表示的数字范围和精度就不同。

格式 总位数 符号位 指数位 尾数位 动态范围 最小精度
FP32 32 1 8 23 ≈ 3.4×10³⁸ ≈ 1.2×10⁻⁷
FP16 16 1 5 10 ≈ 65504 ≈ 6.0×10⁻⁸
BF16 16 1 8 7 ≈ 3.4×10³⁸ ≈ 3.9×10⁻³
FP8 E4M3 8 1 4 3 ≈ 448 ≈ 0.125
FP8 E5M2 8 1 5 2 ≈ 57344 ≈ 0.25

这张表的关键信息就两条:FP16 的动态范围太小,BF16 的精度太粗。 FP16 的指数位只有 5 位,能表示的最大正数也就 65504,而 FP32 和 BF16 的 8 位指数能到 10³⁸ 级别。训练中梯度经常出现极端值,一旦超过 65504 就溢出变成 Inf,然后 Loss 就 NaN 了。反过来,BF16 虽然范围和 FP32 一样大,但它的尾数只有 7 位,精度比 FP16 还低,理论上收敛会更慢。TensorFlow 混合精度指南 里明确指出,FP16 训练的核心挑战就是数值稳定性。

我踩过的坑:直接上 FP16,Loss 当场爆炸

我最早尝试混合精度是在 2019 年,用一张 1080Ti 训一个中等规模的 Transformer。当时觉得 FP16 能省一半显存、速度还更快,为啥不直接用?结果训练了不到 100 步,Loss 从 3.5 直接跳到 NaN。查了半天才发现,不是模型写错了,而是梯度的数值在 FP16 下直接下溢变成了 0。

后来我读了 NVIDIA 的混合精度训练论文(2017) 才彻底想通:混合精度训练不是简单的“把参数和计算切成 FP16”,而是要三种技巧同时用

  1. 权重主副本保持 FP32:前向和反向用 FP16 计算,但权重更新时在 FP32 精度下进行,避免小梯度累加被截断。
  2. Loss Scaling:把 Loss 乘以一个大系数(比如 1024),反向传播时梯度也被放大,这样那些接近零的小梯度就能落进 FP16 的可表示范围,在更新前再除以同样的系数。
  3. 累积高精度操作:某些对精度敏感的操作(如 BatchNorm 的统计量)仍用 FP32。

用一句话总结:FP16 负责干粗活累活,FP32 负责在关键时刻兜底。 PyTorch 从 1.6 版本起提供了 torch.cuda.amp 自动混合精度,一行代码 with autocast(): 就能自动完成上述操作,告别手动 loss scaling。

BF16:为什么 Google 搞出这个“半精度”格式?

BF16 全称 Brain Floating Point,最早由 Google Brain 在 2018 年提出,用于 TPU 训练。它的设计思路非常粗暴:直接截断 FP32 的后 16 位尾数,保留完整的 8 位指数。 这样做的代价是精度只有 7 位尾数(相当于十进制约 2 位有效数字),但好处是动态范围和 FP32 完全一致,不需要 Loss Scaling

这意味着什么?你不再需要担心梯度溢出或下溢,训练流程大幅简化。而且 BF16 与 FP32 之间的转换只需简单截断或补零,硬件开销极低。NVIDIA 从 A100 开始支持 BF16,PyTorch 中只需设置 dtype=torch.bfloat16 即可。Google 的 BF16 论文 中实验表明,在 ResNet、Transformer 等主流模型上,BF16 训练的最终精度与 FP32 几乎无差异,有时甚至因为噪声的隐式正则化效果而略好。

现实中,BF16 已经成为大模型训练的事实标准。 因为它不需要复杂的 loss scaling 调参,且能同样节省一半显存和近似 2 倍的计算吞吐。但有一个前提:你的硬件得支持 BF16,比如 A100、H100、RTX 4090 等。

FP8:更激进,也更讲究

FP16 和 BF16 都是 16 位,FP8 直接砍到 8 位,显存再减半,计算吞吐再翻倍。但 FP8 的表示能力急剧下降,直接用就是找死。 NVIDIA 在 H100 上推出的 Transformer Engine 专门解决这个问题:它采用两种 FP8 格式按需切换——E4M3(偏重精度,适合前向传播)和 E5M2(偏重范围,适合反向传播)。同时,FP8 必须在每个张量级别进行缩放(per-tensor scaling),即计算张量的最大值,然后缩放所有值到 FP8 的可表示区间,反向时再缩放回来。这套机制被封装在 NVIDIA 的 Transformer Engine 库中,与 PyTorch 集成,对用户几乎透明。

NVIDIA 的 FP8 训练论文(2022) 在 GPT-3 175B 模型上验证,FP8 训练与 BF16 的最终损失曲线几乎重合,且训练吞吐提升 1.5-2 倍。但这也意味着,FP8 目前高度依赖特定硬件(H100 及以上)和软件生态,并非所有任务都能无痛迁移。

显存和速度到底能省多少?

理论上,FP16/BF16 相比 FP32 显存减半,FP8 再减半。但实际训练中,你还会保留 FP32 的主权重副本,所以参数显存节省不是 50%,而是约为 37.5%(假设模型大小 1x,FP16 权重+FP32 主权重共 1.5x,比 2x 节省 25%?这里需要仔细算)。更准确地说:假设模型参数量为 P,使用混合精度训练时,需存储 FP16 参数(2P 字节,因为 FP16 占 2 字节)、FP32 主权重(4P 字节)、FP16 梯度(2P 字节)、FP32 优化器状态(如 Adam 的 m 和 v,各 4P 字节,共 8P 字节)。总计 2+4+2+8=16P 字节。而纯 FP32 训练则需要 4+4+4+12=24P 字节。确实节省了 1/3 的显存。Hugging Face 性能指南 中有详细计算。速度方面,如果硬件支持 FP16/BF16 Tensor Core,矩阵运算吞吐能翻倍,端到端训练速度通常提升 1.5-2 倍。

什么时候混合精度会翻车?

尽管混合精度已经很成熟,但仍有几个坑:

  • FP16 + Loss Scaling 的调参:默认的 dynamic loss scaling 在大多数任务上没问题,但遇到极端模型(如某些 GAN)时,loss scale 可能不断下降导致性能退化。
  • 数值敏感操作:如 Softmax 中的大数值溢出、LayerNorm 的方差计算,需要临时回退到 FP32。
  • BF16 的精度瓶颈:在需要高精度累加的场景(如强化学习中的价值函数),BF16 的 7 位尾数可能不够,需要手动提升精度。
  • FP8 的生态壁垒:目前只有 NVIDIA H100 及以上 GPU 支持,且需要 Transformer Engine 库,迁移成本高。

FAQ:回答你最纠结的几个问题

Q:FP16 和 BF16 到底选哪个?
如果你的 GPU 支持 BF16(A100/RTX 4090 等),无脑选 BF16,省去 loss scaling 烦恼。如果是老卡(如 V100、2080Ti),只能选 FP16,那就用 PyTorch 的 autocast,它帮你自动处理 loss scaling。
Q:混合精度会不会影响模型收敛?
大量实验表明,在视觉和语言模型上,混合精度训练的最终精度与 FP32 几乎一致,甚至因为引入的噪声而略微提升泛化能力。但损失曲线可能略有不同,这是正常现象。
Q:FP8 现在能用吗?
如果你有 H100 或更高级的 GPU,并且使用 PyTorch 2.0+ 和 Transformer Engine,可以尝试。对 GPT 类大模型加速明显,但小模型或非 Transformer 架构仍需验证。
Q:混合精度训练能节省多少显存?
粗略估算,FP16/BF16 混合精度相比 FP32 能节省 30-40% 的显存,FP8 在 FP16 基础上再省约 30%。具体取决于模型中的其他内存占用(如激活值、临时缓冲区)。

最后说一个我自己的体会:当初我花了很多时间死磕 FP16 的 loss scaling 调参,后来 BF16 普及后,我整个人都轻松了。技术的演进有时候就是这样——最好的解决方案不是让你更聪明地解决问题,而是让问题本身消失。

原创文章,作者:guanweilu,如若转载,请注明出处:https://guanweilu.cn/article/136.html

(0)
上一篇 2026年8月16日 上午12:10
下一篇 2026年8月16日 上午12:13

相关推荐