torch.cuda.amp是最稳妥的混合精度方案,因其是pytorch 1.6+官方内置、cuda兼容性好、无需改模型结构,通过autocast自动切换算子精度、gradscaler防止梯度下溢,比手动half()或apex更轻量可靠。

为什么 torch.cuda.amp 是当前最稳妥的混合精度方案
PyTorch 1.6+ 内置的 torch.cuda.amp(Automatic Mixed Precision)是官方推荐、CUDA 驱动兼容性好、且无需修改模型结构的方案。它用 autocast + GradScaler 两层封装,自动处理 FP16 前向/反向中的数值下溢和梯度缩放,比手动 cast 张量或用 apex 更轻量、更少出错。
常见错误现象:直接把模型 .half() 后训练崩溃,报 RuntimeError: expected scalar type Half but found Float——这是因为部分算子(如 BatchNorm、Loss)仍需 FP32 输入,不能简单整体转半精度。
-
autocast只在前向传播中动态切换 dtype,不影响模型参数存储类型(仍是 FP32) -
GradScaler必须在loss.backward()前调用scale(loss),否则梯度会因 FP16 下溢而变成 0 - Optimizer 更新前必须调用
scaler.step(optimizer)和scaler.update(),漏掉任一都会导致训练停滞
如何正确插入 autocast 和 GradScaler
关键不是“加在哪”,而是“范围是否精准”。autocast 应仅包裹前向计算(含 loss 计算),不包含数据加载、metric 统计或日志打印;GradScaler 的生命周期必须与每次迭代对齐。
典型错误:把 with autocast(): 扩展到整个 epoch 循环外,导致 dataloader 返回的 float32 数据被强制转 float16,后续 transform(如 Normalize)出错。
scaler = torch.cuda.amp.GradScaler()
for data, target in dataloader:
optimizer.zero_grad()
with torch.cuda.amp.autocast():
output = model(data) # 自动用 FP16 运算
loss = criterion(output, target) # loss 也进 autocast
scaler.scale(loss).backward() # 缩放 loss 后反向
scaler.step(optimizer) # 更新前解缩放梯度
scaler.update() # 更新 scale 值(防 overflow)
哪些层/操作会退出 autocast 自动转换
autocast 不是“全有或全无”,它内置白名单机制:已知不稳定的算子(如 torch.nn.functional.batch_norm)会自动回落到 FP32,但某些自定义操作或第三方库函数不会被识别。
常见问题场景:用了自定义的 LayerNorm 或从 einops 来的 rearrange,结果梯度 nan —— 因为这些没在 autocast 白名单里,输入是 FP16,内部计算溢出。
- 显式指定 dtype:对关键不稳定层,手动
.float()输入,例如bn_layer(x.float()) - 检查白名单:PyTorch 源码中
torch/cuda/amp/autocast_mode.py有完整列表,v1.12+ 新增了更多支持(如F.interpolate) - 避免在 autocast 块内做非张量运算,如 Python
sum()、len(),它们不参与 dtype 推断
显存节省效果与实际瓶颈点
理论显存下降约 40%–50%,但实测常只省 25%–35%,因为模型参数仍以 FP32 存储(保留精度),optimizer state(如 Adam 的 momentums)也默认 FP32。真正吃显存的是 activation,而 autocast 对 activation 的压缩是有效的。
容易被忽略的细节:如果 batch size 已经卡在显存极限,单纯开 AMP 可能触发 CUDA OOM——因为 GradScaler 需额外缓存缩放前的原始梯度,反而多占几 MB。此时应配合 torch.utils.checkpoint(梯度检查点)进一步释放 activation 显存。
验证是否生效:运行时加 torch.cuda.memory_summary(),对比开启前后 “allocated memory by tensor type” 中 cuda:0 float16 行是否显著增长;若全是 float32,说明 autocast 没生效或被意外中断。
Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!











