torch.cuda.amp不能直接套在model.forward()外面,因其autocast仅动态插入精度转换逻辑,不修改模型结构;若仅包裹forward而损失计算和backward在默认精度下运行,将导致梯度类型错配报错。

为什么 torch.cuda.amp 不能直接套在 model.forward() 外面?
因为 torch.cuda.amp 的核心是动态缩放(GradScaler)和上下文管理(autocast),它们必须精准包裹前向计算与反向传播的边界。如果只对 model.forward() 加 autocast,而损失计算、loss.backward() 仍在默认精度下运行,梯度更新会出错——尤其是当 loss 是 float16 张量时,backward() 可能触发 RuntimeError: expected scalar type Float but found Half。
正确做法是把整个前向 + loss 构建放进 autocast() 上下文,且 backward() 必须在该上下文外(否则梯度缩放失效):
scaler = torch.cuda.amp.GradScaler()
for data, target in dataloader:
optimizer.zero_grad()
with torch.cuda.amp.autocast(): # ← 这里开始 FP16 前向
output = model(data)
loss = criterion(output, target) # loss 也是 float16
scaler.scale(loss).backward() # ← backward 在 autocast 外,但用 scale 包裹
scaler.step(optimizer)
scaler.update()
GradScaler 的 scale() 和 step() 为什么必须配对使用?
GradScaler 不是简单地把梯度乘个数,它在每次 scale() 时记录当前缩放因子,并在 step() 中检查梯度是否溢出(inf/nan)。如果跳过 scale() 直接 backward(),或者 scale() 后没调 step(),就会导致:
- 梯度值过小,权重几乎不更新(漏掉
scale()) - optimizer 更新时用的是未缩放梯度,
step()内部无法做溢出检测(漏掉scale()) - 连续多次
scale()但只调一次step(),缩放因子累积错乱,后续update()失效
关键规则:scaler.scale(loss).backward() → scaler.step(optimizer) → scaler.update() 必须成组出现,且中间不能插入其他 backward() 或 step()。
哪些层或操作容易在 autocast 下出问题?
autocast 默认按算子类型决定输入输出精度,但部分操作没有 FP16 实现,或数值不稳定,需手动切回 FP32:
-
torch.nn.functional.softmax:FP16 下易溢出,建议用torch.nn.Softmax(已内置稳定实现)或显式转.float() - 自定义 loss(如带 log / exp 的):先
input.float()再计算 -
torch.where、torch.scatter_等索引类操作:某些版本不支持Half,报NotImplementedError - BatchNorm 层:PyTorch ≥ 1.7 默认在
autocast中保持 FP32 统计,无需干预;但若自己写 BN,需确保running_mean/var是float32
调试技巧:临时加 print(output.dtype, loss.dtype),确认关键张量是否意外变成 torch.float16。
混合精度训练时 model.eval() 需要关掉 autocast 吗?
不需要,但也不能无脑沿用训练逻辑。验证阶段不用 GradScaler,但 autocast 仍可加速前向(尤其大模型):
model.eval()
with torch.no_grad():
with torch.cuda.amp.autocast(): # ← 依然可以开,省显存、提速度
output = model(data)
loss = criterion(output, target)
注意两点:
- 去掉
scaler相关所有调用(scale/step/update) -
torch.no_grad()必须在外层,否则autocast可能尝试记录计算图,报RuntimeError: cudnn_batch_norm_backward not supported for tensors with less than 4 dimensions in autocast mode
实际中,验证时开 autocast 能降显存 30%~50%,且精度损失通常
真正容易被忽略的是:scaler.step() 失败后不会抛异常,而是静默跳过更新——务必检查 scaler.step(optimizer) 的返回值(None 表示成功,False 表示因梯度溢出被跳过),并在日志里记录跳过次数。
Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!











