PyTorch AMP 的核心是 torch.cuda.amp 模块,通过 autocast 自动插入 float16 前向计算、GradScaler 缩放并恢复梯度,二者必须配对使用,否则易出现 NaN 或 RuntimeError;autocast 仅对 CUDA ops 生效,模型和优化器仍需手动移至 GPU,训练循环中 loss 计算、scaler.scale(loss).backward()、scaler.step() 和 scaler.update() 必须严格按序执行。

PyTorch AMP 的核心是 torch.cuda.amp 模块
混合精度训练不是靠手动 cast 张量实现的,而是通过 torch.cuda.amp.autocast 和 torch.cuda.amp.GradScaler 协同工作:前者在前向中自动插入 float16 计算,后者在反向传播时缩放梯度、避免下溢,并在更新参数前恢复为 float32。不配对使用两者,大概率会遇到 NaN 梯度或 RuntimeError: expected dtype float32。
-
autocast默认只对 CUDA ops 生效,CPU 上调用会静默忽略——别指望它加速 CPU 训练 -
GradScaler必须在optimizer.step()前调用step()和update(),漏掉update()会导致后续迭代 scaler 失效 - 模型和优化器仍需显式放在 GPU 上(
.cuda()或.to(device)),autocast不负责设备搬运
典型训练循环里 AMP 的写法必须严格对齐
常见错误是把 scaler.step(optimizer) 放在 autocast 上下文外,或忘记 scaler.update()。正确结构如下:
for data, target in dataloader:
optimizer.zero_grad()
<pre class="brush:php;toolbar:false;">with torch.cuda.amp.autocast():
output = model(data)
loss = criterion(output, target)
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update() # ← 这行不能少,且必须在 step() 后
- 损失计算必须在
autocast内,否则 loss 是float32,scaler.scale(loss)会报错 -
scaler.scale(loss).backward()是关键——直接loss.backward()会绕过 scaler,导致梯度未缩放 - 如果
scaler.step()发现梯度含inf/NaN,它会跳过本次更新,但不会中断训练;可通过scaler.get_scale()监控是否频繁下降
哪些层/操作容易在 AMP 下出问题?
不是所有 CUDA op 都支持 float16 输入,尤其是一些自定义算子或较老版本 PyTorch 中的边缘操作。典型表现是报错 RuntimeError: "xxx" not implemented for 'Half'。
SkillSub Pro - Python 题解与代码注释双功能技能功能概述SkillSub Pro - Python 题解与代码注释双功能技能是一项面向实际任务的技能,主要用于SkillSub Pro 是一个 Python 题解生成与代码注释的 双功能合体技能 ,专为学生、算法学习者和开发者设计;✅ 一个技能,两种用途 :;核心要点📝 题解模式 :输入题目/题号,自动生成完整 Python 题解(含详细注释、解题思路、复杂度分析);💬 注释模式 :输入 Python 代码,自动添加详细中。它将相关步骤、
立即学习“Python免费学习笔记(深入)”;
-
nn.AdaptiveLogSoftmaxWithLoss和部分torch.nn.functional中的归一化函数(如F.layer_norm)在旧版 PyTorch(float16,需手动切回float32 - 使用
torchvision.models时,大部分 ResNet、ViT 等主干已适配,但若自己加了torch.fft或torch.linalg.svd,得确认对应版本是否支持 half - 验证阶段建议关闭
autocast(或明确用torch.float32),避免评估指标因精度抖动产生偏差
显存节省效果取决于模型结构和 batch size
AMP 本身不减少参数量,显存下降主要来自三方面:激活值(activation)存储从 float32 → float16、部分权重缓存用 half、以及更小的梯度张量。但收益不是线性的。
- 纯 Transformer 类模型(如 BERT、ViT)通常能省 25%–35% 显存;CNN 类模型(如 ResNet)因卷积核计算密集,节省常在 15%–20%
- batch size 越大,激活值占比越高,AMP 省显存越明显;但若原始 batch size 已卡在显存极限,开 AMP 后可能仍需配合梯度检查点(
torch.utils.checkpoint)才能进一步提升 -
torch.cuda.memory_allocated()可在 epoch 前后打印,比看 nvidia-smi 更准——因为后者包含缓存碎片
真正难调的是数值稳定性边界:有些模型在特定 learning rate 下,GradScaler 的初始 scale 设为 65536 会频繁 downscale,这时得手动设更低初始值(init_scale=2048)并观察 scaler.get_scale() 走势。

















