Skip to content

Preserve the row-normalization axis in batched Adafactor - #561

Open
Afloat16 wants to merge 1 commit into
jettify:masterfrom
Afloat16:fix/adafactor-batched-row-normalization-20261009
Open

Afloat16 wants to merge 1 commit into
jettify:masterfrom
Afloat16:fix/adafactor-batched-row-normalization-20261009

Conversation

@Afloat16

@Afloat16 Afloat16 commented Oct 8, 2026

Copy link
Copy Markdown

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 approximation Vhat[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:

  • All 20 new parameter cases fail before the fix and pass after it.
  • All 12 original Adafactor cases retain their outcomes: 5 pass and 7 fail at the existing checkpoint-state comparison. The retained failures have matching complete traces, numerical values, assertion locations and exception messages before and after; they remain failures. Original 800-step benchmark and 1,500-step neural-network loops are unchanged.
  • The full-project reStructuredText, Pyroma, Bandit and isort checks pass. Real sdist/wheel builds and Twine checks of both artifacts pass. These tools ran with CPython 3.12.13.
  • Full-project Flake8 reports the same two F401 diagnostics for unchanged typing.Dict / typing.List imports in torch_optimizer/__init__.py before and after. Full-project Black likewise reports the same three unchanged files: torch_optimizer/__init__.py, torch_optimizer/aggmo.py and torch_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.

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

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.

Adafactor fails to run on a custom (rfs) resnet12 (with MAML)

1 participant