📑 本页目录(点开跳转)
06 · 混合精度
⏱ 26 分钟 | ⭐ 性价比最高的一个优化
🎯 一句话
用一半的位数存数据,速度快一倍、显存省一半。 听起来像作弊,但它能work是因为 —— 深度学习需要的是"动态范围",而不是"小数点后多少位"。
🔢 一、五种数字格式
对照
FP32 [符号 1][指数 8][尾数 23] ← 传统标准
TF32 [符号 1][指数 8][尾数 10] ← A100 的默认,硬件自动用
FP16 [符号 1][指数 5][尾数 10] ⚠️ 范围窄
BF16 [符号 1][指数 8][尾数 7] ⭐ 范围同 FP32,精度低
FP8 [符号 1][指数 4-5][尾数 2-3] ← H100 起
↑ 决定【范围】 ↑ 决定【精度】
| 格式 | 最大值 | 最小正规值 | 十进制有效位 |
|---|---|---|---|
| FP32 | ~3.4e38 | ~1.2e-38 | ~7 位 |
| FP16 | 65504 ⚠️ | ~6e-5 ⚠️ | ~3 位 |
| BF16 | ~3.4e38 ⭐ | ~1.2e-38 ⭐ | ~2 位 |
🔑 BF16 vs FP16 是这一章最重要的对比: BF16 牺牲精度换范围,FP16 牺牲范围换精度。 而深度学习更怕"溢出"(范围不够)而不是"不精确" —— 梯度本来就是噪声估计,差几个百分点无所谓; 但一旦变成
inf或0,训练直接崩。所以现在的默认是 BF16,不是 FP16。
💰 二、为什么快:三个来源
流程图
💡 注意第 ② 点:即使某个操作没用 Tensor Core,减半的数据量也让它变快。 这就是为什么混合精度对整个模型都有收益,而不只是矩阵乘。
⚙️ 三、"混合"精度:哪些用低精度,哪些不能
关键信息
- ✅ 用 BF16/FP16:
- 矩阵乘、卷积(占 90%+ 的计算量)
- 激活值(占大部分显存)
- 前向和反向的中间结果
- ❌ 必须保持 FP32:
- ⭐ 主权重副本(master weights)
- 优化器状态(Adam 的 m、v)
- Softmax / LayerNorm 的【累加部分】
- 损失的求和归约
⭐ 为什么必须有 FP32 主权重副本
因果链
🔑 这就是第 3 章那个 "12 字节/参数" 的来源: FP32 主权重 4 + Adam 的 m 4 + v 4 = 12。 省显存的是激活,不是优化器状态。
🩹 四、Loss Scaling:FP16 专属的补丁
流程图
⭐ BF16 不需要 loss scaling —— 它的指数位和 FP32 一样宽, 梯度不会下溢。这是 BF16 在工程上最大的省心之处。
# 🧩 骨架:`model` 来自你自己的代码,这一段只看写法
# FP16:需要 GradScaler
import torch
scaler = torch.amp.GradScaler()
with torch.autocast("cuda", dtype=torch.float16):
loss = model(x)
scaler.scale(loss).backward()
scaler.unscale_(opt) # ⭐ 裁剪前必须先还原
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
scaler.step(opt); scaler.update()
# BF16:干净得多 ⭐
with torch.autocast("cuda", dtype=torch.bfloat16):
loss = model(x)
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
opt.step()
⚠️
scaler.unscale_()这一步经常被漏掉: 不还原就做梯度裁剪,等于按放大后的值裁剪,裁剪阈值完全失效。
⚠️ 五、五个真实的坑
| 坑 | 症状 | 解法 |
|---|---|---|
| 维度没对齐 | 精度改了但速度没变 | 维度补到 8 的倍数(FP16)/ 16 的倍数(INT8)⭐ |
在 autocast 里手动 .half() |
类型混乱、报错或静默降速 | 让 autocast 自己管,别手动转 |
| LayerNorm/Softmax 用低精度累加 | loss 抖动、NaN | PyTorch 的 autocast 默认已把它们留在 FP32 ✅ |
| 梯度裁剪前没 unscale | 裁剪失效,训练发散 | 见上面代码 |
| FP16 训练中期突然 NaN | 某个激活溢出 65504 | 换 BF16(首选)或降低 loss scale ⭐ |
💥 一个高频的困惑:"我开了 AMP,为什么没变快?" 排查顺序:
① 矩阵维度是 8 的倍数吗?(4095 → 4096) ② 模型是不是太小?(小 kernel 的启动开销占主导) ③ 瓶颈是不是根本不在 GPU?(第 4 章的 dataloader 实验) ④ 用 profiler 确认 Tensor Core kernel 真的被调用了
🚀 六、FP8:H100 之后的下一步
信息关系
💡 FP8 目前的实际状态: 推理上已经相当成熟(尤其是权重和 KV Cache); 训练上仍需谨慎 —— 大模型预训练用 FP8 需要仔细的逐层策略, 通常敏感层(第一层、最后一层、LayerNorm)仍保持更高精度。
🔗 和站内其他章的关系
| 相关的地方 | 这里的位置 |
|---|---|
| 第 2 章 Tensor Core | 混合精度快的真正原因 ⭐ |
| 第 3 章 带宽瓶颈 | 数据减半 → 带宽压力减半 |
| 第 3 章 12 字节/参数 | FP32 主权重的来源 |
| 第 18 章 量化 | 更激进的同一思路(INT8/INT4) |
| ML 基础第 11 章 NaN 排查 | FP16 溢出是常见原因 |
| 《机器学习与深度学习基础》09 优化器与学习率 lr 常在 1e-4 量级 | lr × 梯度 比权重小几个数量级 —— 这才是必须留 FP32 主权重的原因 ⭐ |
✅ 检查点
- BF16 和 FP16 的区别是什么?为什么现在默认用 BF16?
- 混合精度快的三个来源是什么?哪个最主要?
- 为什么即使不用 Tensor Core 的操作也会变快?
- 为什么必须保留 FP32 主权重副本?举个具体数字。
- 混合精度省的是哪部分显存?不省哪部分?
- Loss scaling 解决什么问题?动态版本怎么工作?
- 为什么 BF16 不需要 loss scaling?
- "开了 AMP 但没变快"的四步排查顺序?
- FP8 的两种格式各用在哪?它的主要代价是什么?
👀 答案
- BF16 指数 8 位(范围同 FP32)尾数 7 位;FP16 指数 5 位(最大 65504)尾数 10 位。BF16 牺牲精度换范围。默认用 BF16 是因为深度学习更怕溢出而不是不精确——梯度本来就是噪声估计,但变成 inf/0 训练直接崩。
- ①Tensor Core(最主要,16 倍)②带宽减半 ③显存减半能开更大 batch。
- 因为数据量减半,同样的带宽一次能搬两倍元素——对带宽瓶颈的操作(逐元素、LayerNorm、Softmax 等)直接快一倍。
- 因为低精度下小更新会丢失。例:权重 1.0,更新量 0.0001,FP16 在 1.0 附近的最小间隔约 0.001,所以 1.0 + 0.0001 = 1.0,更新完全丢失。
- 省的是激活值(占大部分显存)。不省优化器状态——FP32 主权重 4 + m 4 + v 4 = 12 字节/参数照样要。
- 解决 FP16 下梯度普遍太小(1e-7)会下溢成 0 的问题。动态版:没溢出就每隔 N 步把 S 翻倍试探;出现 inf/nan 就 S 减半并跳过这一步更新。
- 因为 BF16 的指数位和 FP32 一样宽(8 位),梯度不会下溢。这是它工程上最省心之处。
- ①矩阵维度是 8 的倍数吗(4095→4096)②模型是不是太小(kernel 启动开销占主导)③瓶颈是不是根本不在 GPU(dataloader 实验)④profiler 确认 Tensor Core kernel 真被调用。
- E4M3 用于前向的激活和权重(精度稍好),E5M2 用于反向的梯度(范围稍大)。代价:动态范围极窄,必须 per-tensor scaling。
🛑 可以停在这里
⚡ 走神救援
用一半位数存数据:快一倍、省一半显存。 ⭐ 它能 work 是因为深度学习需要的是动态范围,而不是小数点后多少位。
⭐⭐ BF16 和 FP16 的对比是全章的核心:一个把位数花在指数上(范围和单精度一样),一个花在尾数上(精度更高但范围小得多)。⭐ 而深度学习更怕溢出、不怕不精确——梯度本来就是噪声估计,可一旦变成无穷或零就直接崩。所以默认选前者。
快的三个来源:只有半精度才用得上专用的矩阵乘单元、带宽减半(所以连不吃那个单元的操作也变快)、显存减半能开大 batch。
⭐ 「混合」的含义:矩阵乘和激活用半精度,⭐ 但主权重副本、优化器的动量、归一化和 softmax 的累加必须留全精度——⚠️ 因为权重加上一个很小的更新时,半精度的最小间隔比更新本身还大,更新会完全丢失。 ⭐ 这也解释了为什么混合精度省的是激活,不是优化器状态。
⭐ 只有小范围那种格式才需要 loss scaling(先把损失放大再反向,更新前还原),⭐ 而大范围那种根本不需要——这是它工程上最省心的地方。
⚠️ 五个坑里最典型的是「改了精度速度没变」:多半是维度没对齐到硬件要求的倍数。排查顺序:先查对齐 → 模型是不是太小 → 瓶颈是不是根本不在 GPU → 最后用 profiler 确认那个单元真的被调用了。 ⚠️ 另一个隐蔽的:梯度裁剪之前忘了还原缩放,裁剪阈值就完全失效了。
🗓️ 更低精度的格式推理已经成熟、训练仍需谨慎——⭐ 它的动态范围极窄,必须逐张量地做缩放,敏感层还要保持高精度。
下一节 👉 07-算子融合与编译.md