Skip to content

[pat] (4/5) Accelerate coupled pruning for sharded DTensors - #4986

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

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

Conversation

@lisjin

@lisjin lisjin commented Oct 8, 2026 •

Copy link
Copy Markdown
Contributor

Stack from ghstack (oldest at bottom):

Avoid materializing full coupled clusters for eligible two-dimensional Dim0/Dim1 Shard(0) layouts. Aggregate reader score partials, gather writer blocks, mask shards locally, and batch literal-zero counts using at most three mesh-scoped collectives; retain the materializing reference for unsupported layouts. This PR depends on the preceding coupled-pruning core PR and is adapted from qpat ebc76a1 (#57).

Verify real multi-rank fast/reference parity, reader-only and writer-only clusters, uneven and reshaped fallbacks, both replicated/sharded mesh-axis orders, replica accounting, subset meshes, and tied-score rank agreement. This completes the original feature set without changing the final tree; CPU Gloo correctness is not a NCCL/FSDP2 performance benchmark.

Test Plan:

source /home/lvj/.ao/bin/activate
OMP_NUM_THREADS=1 python -m pytest /home/lvj/ao/test/prototype/pat -q

181 passed, 160 subtests passed, including all eight real multi-rank tests. Ruff F/I and formatting checks, git diff --check, and read-only split-boundary review passed. The final tree exactly matches the saved original four-PR tip.

@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/4986

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

✅ No Failures

As of commit 46248df 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 2m 10s —— View job


Review: Accelerate coupled pruning for sharded DTensors

  • Read the PR diff against origin/gh/lisjin/5/base
  • Compare the fast path (_coupled_prox_sharded) with the materializing reference
  • Check eligibility, collective safety, and replica accounting
  • Review the tests and README
  • Post feedback

Overall this is a clean, well-scoped optimization. The eligibility gate is conservative. It only uses global metadata (shapes, placements, mesh), so every participant reaches the same decision and no rank can hang waiting on a collective the others skipped. Non-participants return early in _validate_coupled_specs before reaching the fast path. The writer gather math (view(n_shard, n_writers, L).permute(1,0,2).reshape(n_writers, -1)) is correct, and so is the per-writer offset. Counting zeros only over the shard group avoids double-counting replicas. The tests are strong: they mock out full_tensor and grouped_view, cap the collective count, check that every collective uses the shard group, and cover the even and uneven cases on both sides of the gate. Comments below, most important first.

1. Replica groups could pick different channels on NCCL (medium)

prox_executor.py _coupled_prox_sharded: dist.all_reduce(reader_sq, group=shard_group)

The docstring says "selection is consistent within a mesh." Within one shard group that holds, because ring and tree all-reduce give every member the same bits. On a 2-D mesh, though, each replica group (each row of (Replicate, Shard(0)), for example) runs its own all-reduce. NCCL can choose a different algorithm, channel count or reduction order per communicator depending on topology, for example when one group spans an NVLink island and another crosses nodes. If that happens, reader_sq can differ by an ULP between replica groups. With a near-tie at the top-k boundary, the groups then zero different channels. Weights that are supposed to be Replicate end up different, and nothing reports it. The reference path cannot hit this: full_tensor() gathers exact bits and every rank sums them locally in the same order. CPU Gloo tests won't show the problem.

Suggested fix: gather the reader partials instead of all-reducing them, then sum them locally in a fixed order. The partials are only n_channels floats per reader-sum, so this costs little. They can also go into the same all_gather_into_tensor as the writer blocks: stack reader_sq (length n_channels) next to the writer blocks, or use two gathers. That drops you to 2 collectives and makes the result bitwise identical on every rank of the mesh, which matches the reference. If you keep the all-reduce, please make the docstring and README say that agreement across replica groups depends on the backend reducing in a deterministic order.

Fix this →

2. The fast path skips prox_map.zero_groups_ (low)

The reference path zeroes through prox_map.zero_groups_(view, zero_idx). The fast path writes local[:, zero_idx] = 0.0 and local[local_idx] = 0.0 directly. The two match today, but if CoupledMinSparsityConstraint or a subclass ever changes zero_groups_ (soft shrinkage, a custom fill), the paths will quietly disagree. Consider calling prox_map.zero_groups_(local.t(), zero_idx) for readers and prox_map.zero_groups_(local, local_idx) for writers. The local.t() view is 2-D and indexing writes through it.

3. Smaller points

  • Dead check: _coupled_fast_path returns None when mesh.get_coordinate() is None, but _validate_coupled_specs has already returned False for non-participants by then. You could drop the check, or keep it as an assert with a comment.
  • Kwargs gate is stricter than it needs to be: {"end_dim": 1} on a 2-D tensor is the same as the default -1, but it is sent to the fallback. That costs only speed, so it's fine as is. Normalizing end_dim % p.dim() would cover it if you want.
  • Precision note: readers are summed into one fp32 vector locally before the reduction, while the reference sums one norm per tensor. The docstring and README already mention this, and the tests only assert exact parity for well-separated scores, which is right.
  • Optional simplification: local_idx = zero_idx[(zero_idx >= offset) & (zero_idx < offset + n)] - offset could be a boolean mask: build a full n_channels bool mask once, then use mask[offset:offset+n] for writers and mask for readers. This avoids advanced-index scatter with duplicate-free indices and reads a little more simply.

Tests

  • test_pat_coupled_sharded_parity.py is thorough. Two things you might add:
    • a mixed-dtype cluster (bf16 reader with an fp32 writer) to exercise the .float() casts and torch.stack of fp32 blocks;
    • a reader whose size(0) < n_shard, so some ranks hold empty local shards and sum(dim=0) returns zeros.
  • In the subset-mesh test refactor, route.call_count == 3 (one direct call plus two optimizer steps) is a little implicit. A short comment would help.

README

The new paragraph is accurate. If item 1 is fixed with a gather, update it to say "at most two/three collectives" and that selection is bitwise identical across the mesh.

I didn't run the test suite in this job. The review is based on reading the code.

@lisjin lisjin changed the title [pat] Accelerate coupled pruning for sharded DTensors [pat] (4/5) Accelerate coupled pruning for sharded DTensors 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