pytorch amp中autocast必须与gradscaler配对使用,仅用autocast会导致梯度下溢崩溃;gradscaler需正确调用scale/step/update,且不能在torch.compile或fsdp包装后再初始化,否则参数跟踪失效。

PyTorch AMP的torch.cuda.amp.autocast必须和GradScaler配对用
单独加autocast不仅不省显存,还可能因梯度下溢直接训练崩溃。它只负责前向计算时自动切fp16,但反向传播中fp16梯度容易变成inf或nan,必须靠GradScaler动态调整loss scale来避免。
典型错误写法是只包with torch.cuda.amp.autocast():,漏掉scaler.step()和scaler.update()——这样loss scale不会更新,几轮后梯度全归零。
-
autocast作用范围仅限代码块内,模型forward、损失计算都要包进去 -
scaler.scale(loss).backward()替代原始loss.backward() -
scaler.step(optimizer)内部会检查unscale后的梯度是否有效,无效则跳过更新 -
scaler.update()必须每步都调,否则scale值不会自适应调整
optimizer传给GradScaler前不能已用torch.compile或FSDP
如果用了torch.compile(model)再初始化GradScaler,scaler会无法正确跟踪参数——因为compile后参数地址已变,scaler内部的param_groups引用失效,导致scaler.step()报RuntimeError: unscale_() called on optimizer with no params。
FSDP同理:FSDP包装后的模型参数被分片管理,GradScaler默认不兼容。必须用fsdp.shard_module配合torch.cuda.amp.GradScaler的enabled=False模式,或改用FSDP内置的mixed_precision配置。
Python 3.14.2是Python编程语言在2025年12月5日发布的稳定版本,属于3.14系列的第二个维护更新。该版本包含了18项修复,重点解决了多进程、数据类及正则表达式等模块的回归问题,并修复了CVE-2025-12084等安全漏洞。此版本标志着自由线程模式(移除GIL)正式获得官方支持,是Python发展的重要里程碑。
- 顺序必须是:定义optimizer → 初始化
GradScaler→ 再做torch.compile或FSDP包装 - 验证方式:打印
scaler._per_optimizer_states[id(optimizer)],确认"found_inf_per_device"字段存在 - 若用
torch.compile,建议搭配mode="reduce-overhead",避免JIT干扰autocast上下文
autocast对不同层的精度选择不是完全透明的
PyTorch AMP不是“所有算子都转fp16”,而是按白名单逐个判断——比如nn.Linear、nn.Conv2d默认支持fp16,但nn.LayerNorm、nn.Softmax在某些CUDA版本里仍走fp32,尤其当输入含inf/nan时会自动fallback。
这意味着显存节省幅度取决于模型结构:卷积多的CV模型通常能压到原显存的55%~65%,而Transformer类模型因大量LayerNorm和Softmax,可能只降到70%~75%。
- 可通过
torch.backends.cuda.matmul.allow_tf32 = False强制禁用TF32,让AMP更激进地使用fp16(但可能降低吞吐) - 自定义算子若没注册AMP支持,会静默回退到fp32,可用
torch.cuda.amp.custom_fwd/custom_bwd手动标注 - 验证是否生效:在
autocast块内打印layer.weight.dtype,确认关键层权重参与了fp16计算
batch size调大后显存不降反升?检查GradScaler的growth_factor
默认growth_factor=2.0,意味着每次成功step后scale翻倍。大batch下梯度norm天然更大,scale涨得太快会导致后续step频繁触发found_inf,scaler反复缩放——实际显存占用反而比小batch+稳定scale更高。
这不是AMP bug,而是动态scale机制在大batch下的震荡表现。实测中把growth_factor设为1.01、backoff_factor设为0.8,能显著改善收敛稳定性与显存波动。
- 修改方式:
scaler = GradScaler(growth_factor=1.01, backoff_factor=0.8) - 监控指标:每100步打印
scaler.get_scale(),若发现持续在65536和1之间跳变,说明需要调参 - 极端情况可固定scale:
scaler = GradScaler(init_scale=2048, growth_interval=1000),适合数据/模型高度确定的场景
nvidia-smi对比baseline,再调。Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!










