Repository navigation
Conversation
Retain the reduced row axis when normalizing factored second moments so tensor parameters with leading axes broadcast within each matrix. Add independent scalar-oracle regressions for batched and noncontiguous parameters, multi-step updates, axis permutation and optimizer checkpoint continuation.
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.
Fixes #405.
Preserve the reduced row axis when normalizing Adafactor's factored second moments. For parameters shaped
(..., rows, columns), the row accumulator has shape(..., rows)and its mean must have shape(..., 1). Dropping that axis aligns leading dimensions with the row dimension during broadcasting:(2, 3, 4)parameters raise a shape error, while(2, 2, 3)parameters can silently mix normalization statistics between matrices.mean(dim=-1, keepdim=True)preserves the factored approximationVhat[b,r,c] = R[b,r]*C[b,c]/mean_r(R[b,r]), where R and C are the row and column second-moment accumulators. This implements the correction proposed by @ionutmodo in the existing #405 discussion. The scope is factored second-moment normalization; tensor-wide RMS clipping and parameter scaling retain their current semantics.The public optimizer regressions use independent scalar moment recurrences and closed-form rank-one updates. They cover multiple leading-axis shapes, noncontiguous parameters, permutation of leading axes, multi-step momentum/weight-decay trajectories, and continuation from a real copied optimizer checkpoint.
Validation on Windows CPU with CPython 3.12.10, PyTorch 2.7.1+cpu and NumPy 1.26.4:
typing.Dict/typing.Listimports intorch_optimizer/__init__.pybefore and after. Full-project Black likewise reports the same three unchanged files:torch_optimizer/__init__.py,torch_optimizer/aggmo.pyandtorch_optimizer/shampoo.py. Neither modified file produces a lint or formatting diagnostic.The test scope includes every original Adafactor parameter case and all new cases; it does not claim execution of unrelated optimizer tests or the historical CI matrix. The existing checkpoint, lint and formatting failures are disclosed rather than presented as a green repository run.