Skip to content

[pat] (2/5) Add opt-in global group-count score correction - #4983

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

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

Conversation

@lisjin

@lisjin lisjin commented Oct 8, 2026 •

Copy link
Copy Markdown
Contributor

Stack from ghstack (oldest at bottom):

Compensate for compatible width-dependent initialization by optionally weighting global scores by the square root of group count. Preserve default selection and batch per-parameter count synchronization. Tests distinguish selection from the unchanged sparsity budget. Adapted from qpat ebc76a1 (#57).

Test Plan: OMP_NUM_THREADS=1 python -m pytest /home/lvj/ao/test/prototype/pat -q: 137 passed, 27 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/4983

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

✅ No Failures

As of commit c83e862 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 44s —— View job


PR Review: [pat] Add opt-in global group-count score correction

  • Gather context (diff, prox_executor.py, pruneopt.py, min_sparsity.py)
  • Review correctness of score correction
  • Review count-sync refactor
  • Review tests and README
  • Post review

Summary

This change is small and opt-in. The default path doesn't change: when score_group_count_ref is None the scores are untouched. The budget (ceil(min_sparsity * total_groups)) is computed after scoring, so it doesn't depend on the correction, and the tests check exactly that. I found no correctness bugs. The comments below are about design and validation.

Correctness

  • Score scaling (prox_executor.py:351-352): scores.numel() is the number of groups in the 2-D view, so the factor is sqrt(n_groups / ref) as documented. For DTensors the scores come from the fully materialized view, so n_groups is the global count, not the local shard count. That is the right choice.
  • Batched sync (prox_executor.py:404-411): GlobalMinSparsityConstraint.zero_groups_ already returns a 0-d long tensor on p.device. The device check at line 360 guarantees the views share a device, so torch.stack is safe. Going from one .item() per parameter to a single .tolist() is a real improvement. The torch.tensor(zeros, device=view.device) fallback for a plain Python int only adds a host-to-device-to-host round trip on an unused path, which is fine.

Design suggestions

  1. score_group_count_ref is effectively a boolean. The README says so itself: every positive ref gives the same ranking, because it scales all candidates by the same factor. Making it an integer means users have to choose a value that has no effect, and they may expect it to change something. Options:

    • Make it a bool flag such as score_group_count_correction: true, or
    • State the reason for keeping the integer (for example, keeping corrected scores at a similar magnitude to uncorrected ones for logging or parity with qpat), and say plainly in the README that any value is fine (for example, "use 1 unless you need parity with …").

    Minor, but it would avoid confusion.

  2. Validation happens at step(), not at construction (pruneopt.py:406). Other options are validated in PruneOptimizer.__init__, for example _validate_prox_through_heal (pruneopt.py:63). With this change:

    • An invalid score_group_count_ref only raises on the first prox step after warmup, which could be thousands of steps into a run.
    • On a non-global group (for example MinSparsityConstraint), the key is silently ignored.

    I'd move the positive-int check into the __init__ loop and reject the key when prox_type != "GlobalMinSparsityConstraint". apply_global_prox can keep its own check for direct callers.

  3. Side effect before the error in the optimizer path. In PruneOptimizer.step, self.state[p]["latent"].copy_(p) (pruneopt.py:398-399) runs before apply_global_prox raises. That's harmless, since latent already equals p, but moving validation to construction (suggestion 2) removes the question entirely.

Tests

  • The tests are good: they check that selection changes while zero_elts and numel stay the same, that the result doesn't depend on ref, that equal group counts leave the result unchanged, that invalid input fails before any mutation (including True and 1.5), and that the optimizer forwards the setting.
  • test_reference_magnitude_preserves_selection: the chosen values avoid ties (1.5·√2 vs 1·√8 differ clearly), so rounding can't flip the ranking. Good choice.
  • Gap: there's no test for the key on a non-global group. If you take suggestion 2, add an assertRaises at construction.
  • I couldn't run the test suite here, because the command needed approval in this CI sandbox. The PR description reports 137 passed.

Nits

  • The PR combines two separate changes: the score correction and the batched device-to-host sync. Both are fine, but the sync refactor would be easier to review or bisect as its own commit. The description does mention it.
  • README: "is not automatically enabled by head_dim" is useful, but readers may not know why head_dim is relevant. Half a sentence would help, e.g. "attention-head groupers with head_dim set do not imply this correction".

CI

The only failure is PR Label Check. The module: not user facing label now appears to be applied, so a re-run should fix it.

Overall the change looks good. The main thing I'd change is validating at construction time and rejecting the key on non-global groups.

@lisjin lisjin changed the title [pat] Add opt-in global group-count score correction [pat] (2/5) Add opt-in global group-count score correction 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