Skip to content

Adam 对延后获得梯度的参数使用了错误的步数 #235

Description

@superme1one

问题描述

当前 Adam 实现使用优化器级别的全局步数进行偏差修正,没有区分不同参数实际获得梯度并参与更新的次数。

如果一个参数在前面的优化器调用中没有梯度,之后才第一次获得梯度,其首次更新仍会使用已经累加的全局步数,导致更新结果与逐参数计步的 Adam 实现不一致。

在本次最小复现中,该参数的首次更新幅度比预期小约 25.59%。

测试环境

  • InfiniTrain 提交:a8b2d3776f04a9eee5c0dbff74b76fdf6c4d5c9a
  • 数据类型:FP32。
  • 已在原生 CPU 后端复现。
  • 同时在天数 Iluvatar TG-V200 的本地 CoreX CUDA 兼容适配版本上复现。
  • 参考实现:PyTorch 2.4.1,CPU。

CPU 复现没有修改优化器实现或 CPU 数值内核,因此该问题不依赖天数 GPU 的本地适配。

复现步骤

创建两个标量参数 A 和 B,初始值均为 1.0,使用以下 Adam 配置:

learning_rate = 0.01
beta1 = 0.9
beta2 = 0.999
eps = 1e-8

执行以下操作:

  1. 第一次调用 Step():设置 A.grad = 0.5,B.grad 保持为空。
  2. 第二次调用 Step():设置 A.grad = 0.5、B.grad = 0.5。
  3. 检查参数 B 首次实际更新后的数值。

对应的 PyTorch 参考代码:

import torch

a = torch.tensor([1.0], requires_grad=True)
b = torch.tensor([1.0], requires_grad=True)

optimizer = torch.optim.Adam(
[a, b],
lr=0.01,
betas=(0.9, 0.999),
eps=1e-8,
foreach=False,
)

第一次调用:只有 A 有梯度。

a.grad = torch.tensor([0.5])
optimizer.step()

第二次调用:B 第一次获得梯度,A 的梯度仍为 0.5。

b.grad = torch.tensor([0.5])
optimizer.step()

print(b.item())
print(optimizer.state[b]["step"].item())

预期结果

没有梯度的参数应跳过更新,其独立的 Adam 步数也不应增加。

虽然优化器已经调用两次,但参数 B 只实际更新了一次,因此应使用 t = 1 进行偏差修正。

预期 FP32 结果:

B = 0.9900000095
B 的独立步数 = 1

实际结果

实现 B 首次实际更新后的值
原版 InfiniTrain CPU 0.9925586581
原版 InfiniTrain,本地天数 GPU 适配版 0.9925586581
PyTorch CPU 0.9900000095
本地修复后的 InfiniTrain GPU 0.9900000095

原版 InfiniTrain 的实际更新幅度约为 0.00744134,而预期约为 0.01,更新幅度偏小约 25.59%。

原因分析

在 infini_train/src/optimizer.cc 中,Adam::Step() 在遍历参数之前增加全局计数器 t_:

++t_;

for (...) {
// 跳过梯度为空的参数。
...
// 其余参数均使用相同的全局 t_ 进行偏差修正。
}

对于参数 B,一阶矩和二阶矩只累计了一次梯度,但偏差修正使用的是 t = 2,两者对应的更新历史不一致。

另外,测试版本的优化器状态仅保存全局 adam.t,无法表达不同参数具有不同实际更新次数的情况。

当所有参数每一步都有梯度时,不会触发这一计步问题。但对于条件分支、逐步解冻参数,以及间歇性获得梯度的参数,该问题可能影响训练结果。

建议修复方法

为每个参数分别维护实际更新步数:

  • 仅在该参数梯度非空、参与更新时增加其步数。
  • 使用该参数自己的步数计算 Adam 偏差修正。
  • 梯度为空时,参数、一阶矩、二阶矩和该参数步数均保持不变。
  • 区分“梯度为空”和“梯度数值为零”:零值梯度仍应参与更新并增加步数。
  • 在优化器状态中保存和恢复逐参数步数,优先使用稳定的参数名称进行对应。

不能只把全局 ++t_ 移入参数循环,因为这仍不是独立计步,还会使同一次优化器调用中的不同参数使用不同的全局步数。

对于仅包含 adam.t 的旧检查点,可以提供带警告的兼容回退策略。但旧检查点没有记录各参数缺失梯度的历史,无法精确还原真实的逐参数更新次数。

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions