混合精度训练(AMP)
本教程共 60 篇 · 第 54 篇 · 更新于 2026-08-17 · 约 4 分钟阅读
本节目标:理解 FP16、BF16 和 FP32 的区别,学会用
autocast+GradScaler改造训练循环,知道什么情况该用哪种精度。
为什么要把精度混着用
训练默认全程 FP32。问题在于:矩阵运算多,FP32 又慢又占显存。半精度(16 位)浮点体积减半,运算更快。那为什么不全用半精度?因为两种半精度各有短板:
- FP16:指数位只有 5 位,数值范围最多到 ±65504。大数容易溢出成 inf,小梯度容易下溢成 0。
- BF16:指数位和 FP32 一样多,范围完全够用,但尾数位少,精度只有约 2.4 位有效数字。
混合精度训练(Mixed Precision Training)的思路是:每个算子挑最合适的精度。线性层、卷积这类重计算吃 FP16/BF16 的红利;归一化、损失计算这类精度敏感的活,自动回退到 FP32。
三种浮点格式的结构值得看一眼。一个浮点数由符号位、指数位、尾数位组成:指数位决定能表示多大的数(范围),尾数位决定有多精细(精度)。FP32 是 8 位指数 + 23 位尾数;FP16 只有 5 位指数 + 10 位尾数,范围小;BF16 是 8 位指数 + 7 位尾数,范围和 FP32 一样,只是精度粗。理解了这张表,后面所有选型决策都顺理成章。
加速从哪来?主要靠 GPU 上的 Tensor Core 专用单元。它一个时钟周期能算完一次 4×4 矩阵乘加,普通 CUDA 核心要好几条指令。V100 以后的 NVIDIA GPU 基本都带,RTX 3060 以上消费卡也能享受。
Note收益大致是:训练速度 2-3 倍、显存占用减半、内存带宽减半。模型太小或 GPU 没吃饱时加速不明显,网络太小可能瓶颈在 CPU。
两个核心组件
PyTorch 的 AMP(Automatic Mixed Precision,自动混合精度)由两部分组成。
torch.autocast 是上下文管理器。它包住前向传播,自动判断每个算子用 FP16 还是 FP32。敏感算子(softmax、LayerNorm、交叉熵等)会自动留在 FP32,你不用管。
torch.amp.GradScaler 负责防下溢。FP16 里很小的梯度会直接变成 0,模型学不动。GradScaler 把损失先放大(默认初始 65536 倍)再反向传播,梯度跟着变大就不会丢了。优化器更新前再缩回去。
缩放系数是动态的:某一步梯度出现 inf/NaN,说明放大过头,系数减半并跳过这次更新;连续 2000 步平安无事,系数翻倍。这套反馈循环全自动。
动态损失缩放(dynamic loss scaling)的完整逻辑是这样:每一步先检查反缩放后的梯度里有没有 inf/NaN,有就说明当前 scale 太大,乘以 0.5 回退,同时跳过这一步的参数更新;没有就正常更新,攒够 2000 次连续成功,scale 乘 2 试探更大值。整个过程像开车自动换挡:路况好就提速,颠簸就降速。默认参数(初始 65536、增长 2 倍、回退 0.5 倍、间隔 2000 步)在绝大多数模型上不用调,真出问题再碰。
Warning
torch.cuda.amp.autocast和torch.cuda.amp.GradScaler是旧入口,2.4 起已弃用。新写法是torch.amp.autocast和torch.amp.GradScaler,第一个参数传设备类型"cuda",用法一致。
改造训练循环
普通训练改成 AMP,只动四处。以 FP16 为例:
import torch
import torch.nn as nn
model = nn.Sequential(nn.Linear(128, 256), nn.ReLU(), nn.Linear(256, 10)).cuda()
optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)
loss_fn = nn.CrossEntropyLoss()
scaler = torch.amp.GradScaler("cuda") # ① 创建 scaler,整个训练共用一个
for x, y in train_loader:
x, y = x.cuda(), y.cuda()
optimizer.zero_grad()
with torch.autocast(device_type="cuda", dtype=torch.float16): # ② 包住前向
outputs = model(x)
loss = loss_fn(outputs, y)
scaler.scale(loss).backward() # ③ 缩放后反向
scaler.step(optimizer) # 内部先反缩放,检查 inf/NaN 再更新
scaler.update() # ④ 更新缩放系数
就这么简单。反向传播不用包在 autocast 里,梯度会自动沿用前向时选的精度。
如果你的显存小到要玩梯度累积(gradient accumulation),AMP 照样兼容。注意每个子 batch 的 loss 要除以累积步数,不然累积梯度会被放大:
accum_steps = 4
for i, (x, y) in enumerate(train_loader):
with torch.autocast(device_type="cuda"):
loss = loss_fn(model(x), y) / accum_steps
scaler.scale(loss).backward()
if (i + 1) % accum_steps == 0:
scaler.step(optimizer)
scaler.update()
optimizer.zero_grad()
Note梯度缩放只对 FP16 有意义。BF16 的范围和 FP32 相同,天然不会下溢,所以全程不需要
GradScaler。这也是 BF16 训练代码更简洁的原因。
完整对比一下两种用法:FP16 四件套(autocast + scale + step + update)一个都不能少;BF16 只要 autocast 加普通 backward(),训练循环和不用 AMP 时几乎一模一样。刚入门建议直接上 BF16,踩坑最少。
梯度裁剪、检查点、推理
如果你要梯度裁剪,顺序有讲究。裁剪前必须先反缩放,不然裁的是放大后的值:
scaler.scale(loss).backward()
scaler.unscale_(optimizer) # 先还原真实梯度
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
scaler.step(optimizer)
scaler.update()
保存检查点时,scaler 的状态也要带上,否则恢复训练后缩放系数对不上:
torch.save({"model": model.state_dict(),
"optimizer": optimizer.state_dict(),
"scaler": scaler.state_dict()}, "ckpt.pt")
推理不需要 GradScaler,也没有反向传播。包个 autocast 就行,再配合 torch.inference_mode() 更快:
@torch.inference_mode()
def predict(model, x):
with torch.autocast(device_type="cuda"):
return model(x)
Tip想对比「开不开 AMP」的效果,给
autocast和GradScaler都传enabled=False,它们就变成空操作,不用写两套循环。
提速不明显时,按顺序排查三个嫌疑:GPU 没吃饱(模型太小,瓶颈在 CPU 数据加载,加大 batch 或网络宽度试试)、同步太频繁(训练循环里别调 .item() 打印数值,每步一次同步会拖垮流水线)、Tensor Core 吃不上(矩阵维度最好是 8 的倍数,NLP 模型的词表维度经常踩这个坑)。
FP16 还是 BF16
选型只看硬件。Ampere 架构(RTX 30 系、A100)往后的 GPU 都支持 BF16,优先用 BF16:数值范围和 FP32 完全一样,不会溢出,也不需要 GradScaler,训练最稳。
with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
outputs = model(x)
loss.backward() # BF16 直接反向,不用 scaler
optimizer.step()
V100、RTX 20 系这种老卡只支持 FP16,老老实实配 GradScaler。判断方法一行代码:torch.cuda.is_bf16_supported()。
再提醒一次训练中的玄学:同一套代码换卡后,把 dtype 从 float16 换成 bfloat16 往往能救回不收敛的模型,反过来也一样。硬件允许的情况下,BF16 永远优先,省心。
Warning损失变成 NaN/Inf 时别慌。先分别关掉
autocast和GradScaler定位是谁的问题;再检查学习率是否过高、损失函数是否对 FP16 敏感。少见的类型不匹配报错,多半是自定义算子没进 autocast 白名单。
还有一个高频坑:精度敏感层不要全交给 autocast 赌。LayerNorm、Softmax 这类官方已纳入白名单,会自动回退 FP32;但你自己写的归一化、缩放类算子不在白名单里,就得手动处理。在 autocast 区域里局部禁用,强制回 FP32:
with torch.autocast(device_type="cuda"):
outputs = model(x)
with torch.autocast(device_type="cuda", enabled=False):
loss = loss_fn(outputs.float(), y) # 强制 FP32
另外 Ampere 架构还有一招 TF32:矩阵乘法用 19 位精度,速度接近 FP16、精度接近 FP32,torch.backends.cuda.matmul.allow_tf32 = True 即可开启。它是 FP32 的替身,不是 AMP 的替代品,但可以叠加。
AMP 还能和 torch.compile、DDP 分布式训练叠加使用,代码结构完全不变。这三件套是现代深度学习训练的标配,学会这套循环,后面所有章节都用得上。