Repository navigation
Apply LAMB debias to the moments inside the trust ratio - #557
Open
shaneraphel wants to merge 1 commit into
Open
shaneraphel wants to merge 1 commit into
shaneraphel wants to merge 1 commit into
Conversation
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
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Summary
Fixes #443. With
debias=True, the Adam correction was multiplied into the learning rate andadam_normstayed 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 correctingmandv: 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 citedpytorch-lambreference also ships, with debiasing commented out). When the flag is on, the step uses copiesm̂ = 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.pydebias=Falsematches the hand-computed uncorrected direction, including the trust ratio.debias=Trueleavesexp_avgat 0.1 andexp_avg_sqat 0.001 (the raw EMA) and moves the parameter to the corrected value.