PyTorch GradScaler 何时会跳过参数更新
来源:17golang原创
时间:2026-10-04 16:43:05 268浏览 收藏
在启用梯度缩放时,GradScaler 会先把该优化器负责参数的梯度反缩放,再检查这些梯度里是否出现 inf 或 NaN。只要发现非有限值,scaler.step(optimizer) 就不会调用这一次 optimizer.step(),因此参数保持不变。它并不是因为 loss 偶尔变大、scale 没增长,或者学习率很小就自动跳过。
判断核心只有一句:跳过是针对“这个优化器本轮反缩放后的梯度是否包含 inf/NaN”,而 scale 的下降是随后 scaler.update() 对溢出的响应。
真正触发跳过的是非有限梯度
标准 AMP 训练里,scaler.scale(loss).backward() 生成的是放大后的梯度。调用 scaler.step(optimizer) 时,如果之前没有显式执行 unscale_,GradScaler 会在内部完成反缩放,并记录该优化器的非有限值检查结果。梯度全部有限时,它才会调用底层优化器的 step;任意相关梯度出现 inf 或 NaN 时,本轮更新被跳过。

这里有三个容易混淆的边界:
- 检查对象是梯度,不是 loss 值本身。loss 为有限值不保证反向传播链路上的每个梯度都有限;反过来,loss 的数值看起来较大,也不等于一定跳过。
- 跳过发生在
scaler.step内。scaler.update不负责补做参数更新,它只根据本轮收集到的检查结果调整下一轮使用的 scale。 - 判断按优化器划分。一个优化器只检查它所管理参数的梯度;参数没有交给这个优化器,就不属于这次 step 的判断范围。
PyTorch 的 AMP 文档明确说明:梯度含 inf/NaN 时会跳过 optimizer.step(),随后 update() 使用 backoff_factor 降低 scale。官方仓库中的 GradScaler 实现也将每个优化器的检查状态分开保存。
用标准训练循环把职责分开
最小训练循环应保持 scale → backward → step → update 的调用关系。下面的写法同时记录更新前后的 scale,用于判断本轮是否发生了使 scale 回退的溢出:
import torch
device = "cuda"
model = MyModel().to(device)
optimizer = torch.optim.AdamW(model.parameters(), lr=3e-4)
scaler = torch.amp.GradScaler(device)
for inputs, targets in train_loader:
inputs = inputs.to(device)
targets = targets.to(device)
optimizer.zero_grad(set_to_none=True)
# 前向计算使用自动混合精度。
with torch.amp.autocast(device_type=device, dtype=torch.float16):
outputs = model(inputs)
loss = loss_fn(outputs, targets)
# 放大 loss 后反向传播,减少 float16 梯度下溢风险。
scaler.scale(loss).backward()
# 保存本轮 step 前的 scale,供 update 后判断是否回退。
scale_before = scaler.get_scale()
# 内部反缩放并检查梯度;发现 inf/NaN 时不会调用 optimizer.step()。
scaler.step(optimizer)
# 根据本轮检查结果更新下一轮的 scale。
scaler.update()
scale_after = scaler.get_scale()
# 在默认动态缩放且没有手动指定 new_scale 时,下降表示发生过溢出跳过。
skipped_for_overflow = scale_after
比较 scale 是单优化器训练里很实用的外部信号,但要正确理解它的含义。scale 在普通成功迭代里通常保持不变,达到 growth_interval 后才乘以 growth_factor;发现非有限梯度时才乘以 backoff_factor。因此“scale 没变”不代表 step 被跳过,而“scale 下降”才说明本轮检查发现了溢出。
这个判断还依赖两个前提:没有向 update(new_scale=...) 手动传入新值,并且你接受它表达的是“本轮至少有一次溢出检查失败”。如果要定位究竟哪个张量先出现非有限值,应在诊断阶段显式 unscale_ 后检查梯度,而不是读取 GradScaler 的私有字段。
梯度裁剪必须放在反缩放之后
如果训练中需要裁剪梯度,应先调用 scaler.unscale_(optimizer),再对普通尺度的梯度执行裁剪,最后调用 scaler.step(optimizer)。此时 step 不会重复反缩放,但仍会根据已经记录的检查结果决定是否跳过。
# 先反缩放,让裁剪阈值对应真实梯度尺度。 scaler.unscale_(optimizer) # 对反缩放后的梯度执行裁剪。 torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) # 即使已显式反缩放,step 仍会在发现 inf/NaN 时跳过底层优化器更新。 scaler.step(optimizer) # 每个完整训练迭代只在 step 之后更新一次 scale。 scaler.update()
官方 AMP 示例强调,unscale_ 对同一个优化器每个 step 只能调用一次,并且要等该优化器的梯度累积完成后再调用。若在梯度累积尚未结束时改变 scale 或把一部分梯度提前反缩放,后续 backward 会把不同尺度的梯度混在一起,无法再恢复正确值。
学习率调度器不要在跳过时盲目前进
参数更新被跳过时,如果仍然执行按迭代计数的 scheduler.step(),学习率计划就会比真实参数更新次数多走一步。对于单优化器、默认动态缩放的常见写法,可以在 update 后根据 scale 是否回退决定是否推进调度器:
scale_before = scaler.get_scale()
# 可能调用,也可能因非有限梯度而跳过 optimizer.step()。
scaler.step(optimizer)
scaler.update()
scale_after = scaler.get_scale()
# 只有本轮没有因溢出回退 scale 时,才推进按更新次数设计的调度器。
if scale_after >= scale_before:
scheduler.step()
这个条件适用于调度器语义确实是“每次参数更新推进一次”的场景。如果调度器按 epoch、验证指标或其他业务事件推进,就应保持它自己的调用约定,不要机械套用。另一个常见误区是利用 scaler.step(optimizer) 的返回值判断:该返回值会转发底层 optimizer.step() 的返回值,而大多数内置优化器成功执行时也返回 None,所以 None 不能证明发生了跳过。
多优化器时,每个优化器独立做决定
一个训练迭代可以包含多个优化器。PyTorch 会让每个优化器根据自己管理参数的梯度独立决定是否跳过:A 的梯度有限,A 可以更新;B 的梯度出现非有限值,B 可以在同一迭代里跳过。所有本轮使用的优化器都调用完 scaler.step 后,再统一调用一次 scaler.update()。

# 两个优化器分别清空自己管理参数的梯度。
optimizer_a.zero_grad(set_to_none=True)
optimizer_b.zero_grad(set_to_none=True)
with torch.amp.autocast(device_type="cuda", dtype=torch.float16):
loss_a, loss_b = model.compute_losses(batch)
# 两个 loss 都通过同一个 scaler 建立缩放后的反向传播。
scaler.scale(loss_a).backward(retain_graph=True)
scaler.scale(loss_b).backward()
# 每个优化器只依据自己参数的梯度做独立 step 决策。
scaler.step(optimizer_a)
scaler.step(optimizer_b)
# 所有优化器的 step 都结束后,本轮只更新一次 scale。
scaler.update()
这也是为什么多优化器场景下仅比较全局 scale 不能告诉你“哪个优化器跳过了”。scale 下降只能说明本轮至少有一个优化器记录到了非有限梯度。若必须做优化器级别的告警,可在诊断代码里分别 unscale_,再遍历各自参数的 .grad 检查有限性;完成诊断后仍让 GradScaler 负责最终 step 决策。
连续跳过时按这个顺序排查
训练初期偶尔跳过并不一定是故障。动态梯度缩放会尝试找到适合当前模型和数据的 scale,初始 scale 偏高时可能先回退几次。真正需要关注的是持续多轮下降、loss 很快变成非有限值,或者模型长期没有任何有效更新。
- 先确认调用顺序。每轮是否确实执行了
scale(loss).backward()、对应优化器的step,以及末尾一次update。 - 确认梯度累积边界。只有在完整有效 batch 累积结束时才
step与update;累积过程中不要改变 scale。 - 反缩放后检查首个非有限梯度。记录参数名和梯度是否有限,而不是把所有问题都归因于 GradScaler。
- 再看数值来源。检查除零、对负数取对数、过大的指数、归一化分母、异常输入和过高学习率等产生非有限值的上游操作。
- 最后调整缩放参数。降低
init_scale可以减少训练初期连续回退,但它不能修复模型本身持续产生 NaN 的计算。
如果设置 enabled=False,GradScaler 的缩放相关方法会退化为无操作,step 会直接调用底层 optimizer.step();此时不会由 GradScaler 替你检查并跳过。这个开关适合让同一套代码在启用或关闭 AMP 时共用控制流,但不能把它当作溢出保护。
几个常见判断问题
loss 是 NaN 时一定会跳过吗?
最终仍以该优化器相关梯度的有限性检查为准。NaN loss 通常会产生 NaN 梯度并触发跳过,但判断对象不是 loss 标量本身。排障时应同时检查 loss 和反缩放后的梯度。
scale 没有增长是不是说明更新失败?
不是。成功迭代会累计增长计数,只有连续达到 growth_interval 后 scale 才增长。多数成功迭代里 scale 保持不变是正常现象。
梯度裁剪能阻止所有跳过吗?
不能。裁剪适合限制有限梯度的范数,但已经出现 inf/NaN 时,GradScaler 仍会跳过 step。正确顺序是先反缩放,再裁剪,再让 scaler.step 完成检查与决策。
可以直接读取 found_inf 之类的内部字段吗?
不建议。私有字段和内部状态可能随 PyTorch 版本变化。常规监控使用公开的 get_scale(),精确诊断则在显式 unscale_ 后检查公开的 param.grad。
归根结底,GradScaler 跳过参数更新是一项数值安全动作:它保护优化器不把非有限梯度写入参数,同时通过降低 scale 给下一轮留下恢复空间。把 step 的局部判断、update 的全局缩放调整和调度器的推进条件分开,训练循环就不会再把“没有更新”误判成程序卡住。
-
284 收藏
-
387 收藏
-
328 收藏
-
426 收藏
-
147 收藏
-
316 收藏
-
150 收藏
-
148 收藏
-
251 收藏
-
333 收藏
-
145 收藏
-
479 收藏
-
236 收藏
-
314 收藏
-
134 收藏
-
357 收藏
-
206 收藏
-
- 前端进阶之JavaScript设计模式
- 设计模式是开发人员在软件开发过程中面临一般问题时的解决方案,代表了最佳的实践。本课程的主打内容包括JS常见设计模式以及具体应用场景,打造一站式知识长龙服务,适合有JS基础的同学学习。
- 立即学习 543次学习
-
- GO语言核心编程课程
- 本课程采用真实案例,全面具体可落地,从理论到实践,一步一步将GO核心编程技术、编程思想、底层实现融会贯通,使学习者贴近时代脉搏,做IT互联网时代的弄潮儿。
- 立即学习 516次学习
-
- 简单聊聊mysql8与网络通信
- 如有问题加微信:Le-studyg;在课程中,我们将首先介绍MySQL8的新特性,包括性能优化、安全增强、新数据类型等,帮助学生快速熟悉MySQL8的最新功能。接着,我们将深入解析MySQL的网络通信机制,包括协议、连接管理、数据传输等,让
- 立即学习 500次学习
-
- JavaScript正则表达式基础与实战
- 在任何一门编程语言中,正则表达式,都是一项重要的知识,它提供了高效的字符串匹配与捕获机制,可以极大的简化程序设计。
- 立即学习 487次学习
-
- 从零制作响应式网站—Grid布局
- 本系列教程将展示从零制作一个假想的网络科技公司官网,分为导航,轮播,关于我们,成功案例,服务流程,团队介绍,数据部分,公司动态,底部信息等内容区块。网站整体采用CSSGrid布局,支持响应式,有流畅过渡和展现动画。
- 立即学习 485次学习