登录
推荐 文章 Go 技术 课程 下载 专题 AI
首页 >  科技周边 >  人工智能

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 时,本轮更新被跳过。

GradScaler 对缩放梯度进行反缩放和有限性检查后决定是否调用 optimizer.step 的结构
图1:GradScaler 的缩放梯度、反缩放检查与参数更新边界静态关系,不是运行截图。

这里有三个容易混淆的边界:

  • 检查对象是梯度,不是 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()。

共享 GradScaler 下两个优化器分别检查局部梯度并在最后统一更新 scale 的结构
图2:共享 GradScaler 下两个优化器的局部梯度检查和一次 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 很快变成非有限值,或者模型长期没有任何有效更新。

  1. 先确认调用顺序。每轮是否确实执行了 scale(loss).backward()、对应优化器的 step,以及末尾一次 update。
  2. 确认梯度累积边界。只有在完整有效 batch 累积结束时才 step 与 update;累积过程中不要改变 scale。
  3. 反缩放后检查首个非有限梯度。记录参数名和梯度是否有限,而不是把所有问题都归因于 GradScaler。
  4. 再看数值来源。检查除零、对负数取对数、过大的指数、归一化分母、异常输入和过高学习率等产生非有限值的上游操作。
  5. 最后调整缩放参数。降低 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 的全局缩放调整和调度器的推进条件分开,训练循环就不会再把“没有更新”误判成程序卡住。

声明:本文转载于:17golang原创 如有侵犯,请联系study_golang@163.com删除
相关阅读
更多>
最新阅读
更多>
课程推荐
更多>