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] (2/5) Add opt-in global group-count score correction - #4983
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).
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
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
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.
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.
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
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
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):
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.