Skip to content

[pat] (3/5) Add direct projected-weight optimization mode - #4984

Open
lisjin wants to merge 1 commit into
gh/lisjin/3/basefrom
gh/lisjin/3/head
Open

lisjin wants to merge 1 commit into
gh/lisjin/3/basefrom
gh/lisjin/3/head

Conversation

@lisjin

@lisjin lisjin commented Oct 8, 2026 •

Copy link
Copy Markdown
Contributor

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.

[ghstack-poisoned]
@pytorch-bot

pytorch-bot Bot commented Oct 8, 2026 •

Copy link
Copy Markdown

🔗 Helpful Links

🧪 See artifacts and rendered test results at hud.pytorch.org/pr/pytorch/ao/4984

Note: Links to docs will display an error until the docs builds have been completed.

✅ No Failures

As of commit c62775f with merge base cff77b4 (image):
💚 Looks good so far! There are no failures yet. 💚

This comment was automatically generated by Dr. CI and updates every 15 minutes.

@meta-cla meta-cla Bot 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
@lisjin lisjin added the module: not user facing Use this tag if you don't want this PR to show up in release notes label Oct 8, 2026
@lisjin

lisjin commented Oct 8, 2026

Copy link
Copy Markdown
Contributor Author

@claude review

@claude

claude Bot commented Oct 8, 2026 •

Copy link
Copy Markdown

Claude finished @lisjin's task in 1m 12s —— View job


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-shrunk p 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.

Fix this →

2. Cross-mode resume fails with an unclear error

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 lisjin changed the title [pat] Add direct projected-weight optimization mode [pat] (3/5) Add direct projected-weight optimization mode Oct 8, 2026
@lisjin
lisjin marked this pull request as ready for review October 9, 2026 14:17

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

CLA Signed This label is managed by the Facebook bot. Authors need to sign the CLA before a PR can be reviewed. module: not user facing Use this tag if you don't want this PR to show up in release notes

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant