You signed in with another tab or window. Reload to refresh your session.You signed out in another tab or window. Reload to refresh your session.You switched accounts on another tab or window. Reload to refresh your session.Dismiss alert
{{ message }}
Repository navigation
[pat] (3/5) Add direct projected-weight optimization mode - #4984
Add opt-in latent_weights=False to step from projected parameters without allocating PAT latent copies. Preserve the default trajectory, delegated checkpoint format, and base optimizer moments. Document same-mode checkpoint reconstruction and cover warmup, frequency skips, healing, and resume. Adapted from qpat ebc76a1 (#57).
meta-claBot
added
the
CLA Signed
This label is managed by the Facebook bot. Authors need to sign the CLA before a PR can be reviewed.
label
Oct 8, 2026
Review: latent_weights=False direct projection mode
Read diff (pruneopt.py, __init__.py, README, new tests)
Trace step() control flow in both modes
Check gamma / prox-map semantics in direct mode
Post review
The gating is minimal and correct. Every state["latent"] read and write in step() is now behind self.latent_weights, and _state is only created in save_latent_params, so in direct mode the temp-buffer branch can't run. Nothing else under torchao/prototype/pat reads state["latent"] (utils.latent_svd doesn't use it), and with the default True the old code path is unchanged. The tests are thorough for hard constraints: they compare against a manual base-step-plus-projection oracle and check moment equality and resume in each phase.
1. Soft prox maps combined with cumulative gamma probably over-shrink in direct mode (main concern)
_set_gamma adds lr every step, so gamma is the running total of step sizes (proxmap.py:48: "typically the cumulative sum over step sizes"). That is the right scaling when prox is applied to the latent iterate: it acts like dual averaging, with the latent holding the full gradient history and one cumulative threshold applied to it.
Direct mode applies prox to the already-shrunkp on every step. For the soft maps (Lasso, GroupLasso, NuclearNorm, where threshold = reg_lambda * tau * gamma), the total shrinkage after t prox steps then grows like Σ_s γ_s ≈ lr·t²/2, where you would expect lr·t (a proximal gradient step uses lr·λ per step). Weights would collapse to zero much faster than reg_lambda implies, and the effect gets worse as training runs longer.
Every test uses MinSparsityConstraint, which ignores gamma (min_sparsity.py:42), so this case isn't covered.
Options:
In direct mode, pass a per-step gamma (lr, or the lr accumulated since the last prox when prox_freq > 1) instead of the cumulative value, and keep group["gamma"] cumulative so checkpoints stay compatible.
Or restrict latent_weights=False to hard-constraint prox types for now, and raise in __init__ for the others.
Either way, please add a test with prox_type="Lasso" or "NuclearNorm" and reg_lambda > 0.
The README says cross-mode resume is unsupported, but nothing enforces it. Suppose a direct-mode checkpoint is loaded into a default-mode optimizer with num_steps > warmup_steps. On the next step(), restore_latent_params raises a bare KeyError: 'latent'. Two cheap fixes:
Raise a clear error there when "latent" is missing.
Or record the mode in the regularized param group (for example group.setdefault("latent_weights", ...), like num_steps/gamma) and assert it matches in load_state_dict / patch_state_dict.
The second option would change test_state_dict_is_delegated_and_does_not_serialize_mode, but it turns a documented rule into a check that actually runs. (The other direction, a default checkpoint loaded into direct mode, works silently and just carries an unused latent. That's harmless, but it is also undetected.)
3. Minor
pruneopt.py:395: if not p.requires_grad or not self.latent_weights: continue checks the mode inside the per-param loop. Hoisting it (if self.latent_weights: for p in ...) would match the other two call sites and read more clearly.
restore_latent_params and save_latent_params are public. In direct mode they fail or allocate latents. Consider a guard or a docstring note saying they're only valid when latent_weights=True.
The README and PR body say "Gradients or momentum may revive zeros during pruning". With prox_freq > 1, revived values persist until the next prox step, and the test at L238 confirms this. A short README note that sparsity_frac / relative_sparsity can be stale between prox steps in this mode would help users reading metrics.
CI: "Check PR Labels" is failing. The PR already has module: not user facing, so a re-run should clear it.
Summary
The mechanics are correct and the tests cover hard constraints (MinSparsity*, NMSparse, MinRank) well. Please fix issue 1 before soft prox maps are used with latent_weights=False. Issue 2 is a small UX improvement.
lisjin
changed the title
[pat] Add direct projected-weight optimization mode
[pat] (3/5) Add direct projected-weight optimization mode
Oct 8, 2026
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
CLA SignedThis label is managed by the Facebook bot. Authors need to sign the CLA before a PR can be reviewed.module: not user facingUse this tag if you don't want this PR to show up in release notes
1 participant
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.
Stack from ghstack (oldest at bottom):
Add opt-in latent_weights=False to step from projected parameters without allocating PAT latent copies. Preserve the default trajectory, delegated checkpoint format, and base optimizer moments. Document same-mode checkpoint reconstruction and cover warmup, frequency skips, healing, and resume. Adapted from qpat ebc76a1 (#57).
Test Plan: OMP_NUM_THREADS=1 python -m pytest /home/lvj/ao/test/prototype/pat -q: 147 passed, 99 subtests passed. Ruff F/I checks, formatting, and git diff --check passed.