Skip to content

Apply LAMB debias to the moments inside the trust ratio - #557

Open
shaneraphel wants to merge 1 commit into
jettify:masterfrom
shaneraphel:lamb-debias
Open

shaneraphel wants to merge 1 commit into
jettify:masterfrom
shaneraphel:lamb-debias

Conversation

@shaneraphel

@shaneraphel shaneraphel commented Sep 25, 2026 •

Copy link
Copy Markdown

Summary

Fixes #443. With debias=True, the Adam correction was multiplied into the learning rate and adam_norm stayed on the raw moments. Weight decay is added to the update before that norm, so a scalar on the learning rate is not the same thing as correcting m and v: it scales the decay term and the trust ratio never sees the corrected direction.

The flag still defaults to False. That path is the one in the paper (arXiv 1904.00962, the version the cited pytorch-lamb reference also ships, with debiasing commented out). When the flag is on, the step uses copies

m̂ = m / (1 − β1^t)
v̂ = v / (1 − β2^t)

and the trust ratio is taken on m̂ / (sqrt(v̂) + ε) + λθ. The stored averages are not divided. Dividing them in place would destroy the running average on the next step, which is not what the correction is.

On one step with θ = 2, gradient 1, lr 0.1, weight decay 0.1 and the default betas, the corrected update lands at 1.8. The old learning-rate scaling lands near 1.937.

Test plan

  • pytest tests/test_lamb_debias.py
  • debias=False matches the hand-computed uncorrected direction, including the trust ratio.
  • debias=True leaves exp_avg at 0.1 and exp_avg_sq at 0.001 (the raw EMA) and moves the parameter to the corrected value.

This change was produced with assistance from an AI coding tool. I compared the debiased direction with a learning-rate scale, kept the correction off the default path, and ran the tests locally.

debias=True multiplied the learning rate by the Adam correction and
left adam_norm on the raw moments. Weight decay, which is added
before the norm, was scaled by that factor too, so the trust ratio
did not match a debiased step (issue 443).

The correction is now m / (1 - beta1**t) and v / (1 - beta2**t) on
copies of the moments. The stored averages are unchanged, and the
default debias=False path is the paper's uncorrected step.

This branch has not been deployed

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

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

lamb optimizer mistake

1 participant