Routed gather for sharded index_select and integer-tensor indexing - #1981
coreyjadams wants to merge 1 commit into
Conversation
CODEOWNERS review mapCurrent for commit ⏳ @coreyjadams — 10 file(s)
⏳ @ktangsali — 1 file(s)
⏳ @negin513 — 9 file(s)
No CODEOWNER
Comment |
|
The PR is not safe to merge until routed backward row selection, index bounds semantics, and DTensor placement handling are corrected. Findings
Summary
Reviews (1) · Last reviewed commit: "Routed gather for sharded index_select a..." |
|
|
||
| # Keep a request iff this rank owns its row and it falls in the sender's | ||
| # responsible share. | ||
| local_rows, owned = _owned_rows(grad, all_requests, offsets, rank) |
There was a problem hiding this comment.
_owned_rows derives its clamp bound from grad.shape[0], but that is the first index dimension, not the number of local source rows. When the local source shard has more rows than this index dimension, valid requests for higher rows are clamped to grad.shape[0] - 1. The following index_add_ therefore assigns those gradients to the wrong source row. The local row IDs need to be bounded using n_local_rows instead.
| # Global positions, wrapped once; exchanged once and reused by the backward. | ||
| idx_local = _wrap_negative(idx_local.to(torch.int64), offsets[-1]) |
There was a problem hiding this comment.
The routed path wraps every negative index and performs no bounds validation. As a result, torch.index_select accepts negative indices that eager PyTorch rejects. For both indexing APIs, indices at least as large as the global extent or below its negative extent are converted into clamped filler-row selections instead of raising IndexError, silently returning incorrect data. The routed path must apply the bounds and negative-index rules of the specific indexing operation.
| if not isinstance(index, ShardTensor): | ||
| return False |
There was a problem hiding this comment.
Sharded DTensor treated as replicated
DTensor indices are accepted and converted to local tensors, but _index_is_sharded classifies every DTensor as replicated without inspecting its placements. A sharded DTensor can therefore provide different local requests on each rank while the result and backward path assume that all requests are identical. This produces incorrect output or gradients, and uneven shards can also give ranks inconsistent collective split sizes. Either require a replicated DTensor index or handle its sharding explicitly.
Codecov Report❌ Patch coverage is
Additional details and impacted files@@ Coverage Diff @@
## dp-02-ring-funcol-compile #1981 +/- ##
=============================================================
- Coverage 72.17% 71.88% -0.30%
=============================================================
Files 928 929 +1
Lines 69968 70111 +143
Branches 10540 10560 +20
=============================================================
- Hits 50502 50400 -102
- Misses 16009 16250 +241
- Partials 3457 3461 +4
🚀 New features to boost your workflow:
|
66db6f2 to
f106f1e
Compare
f106f1e to
e2720f7
Compare
Replace the all-gathering ShardedIndexSelect with a routed exchange: a Shard(d) source and an index of global positions (sharded or replicated) move only the requested rows between ranks. All exchange sizes come from the specs, so the custom ops have no device-to-host sync and trace under torch.compile; the backward reuses the forward's request exchange and is a single collective. - shard_utils/exchange.py (new): row all-to-all, group-name resolution and the fp32 scatter accumulator policy (float64 default, configurable), shared by halo_scatter and index_ops. - index_ops.py: routed_gather / sharded_index_select; handlers for torch.index_select, Tensor.index_select and Tensor.__getitem__ with a single int-tensor key; replicated-index backward splits responsibility across ranks; negative indices wrap; 1-D mesh and int dtype guards. - shard_tensor.py: _ToTorchTensor.backward adopts the gradient's shard order with its placements. - Tests: test_index.py (getitem), index_select along the sharded dim in test_select.py, test_exchange.py, routed-gather compile test.
e2720f7 to
9831c24
Compare
| __all__ = [ | ||
| "funcol_all_to_all_v_rows", | ||
| "get_fp32_scatter_accumulator", | ||
| "resolve_group_name", | ||
| "set_fp32_scatter_accumulator", | ||
| ] |
There was a problem hiding this comment.
This file, exchange.py, is a brand new file that consolidates the halo exchange work supporting alchemi merged in v2.2. It was useful and necessary to reproduce almost all of that for index and index_select handlers properly, so I moved core functionality here to mirror halo.py and ring.py as a slightly-lower-level utility than ShardTensor. We could do optimizations on these functions and it would now hit all paths, etc.
| def _funcol_group_arg(group: object) -> object: | ||
| r"""Return *group* in the form functional collectives accept (funcol rejects ``None``).""" | ||
| if isinstance(group, DeviceMesh): | ||
| return (group, 0) | ||
| if group is None: |
There was a problem hiding this comment.
The changes to this file should be exclusively moving things to exchange.py without functionality change.
| # ``values`` is ``Shard(d)`` with known per-rank sizes; ``index`` holds GLOBAL | ||
| # positions along ``d`` that may live on any rank. Everything about the exchange | ||
| # is known from the specs before any data moves: the shard boundaries of | ||
| # ``values`` (``offsets``) and how many requests every rank holds | ||
| # (``request_sizes``, from the index's sharding shapes). So the exchange uses | ||
| # STATIC split sizes and no host sync: | ||
| # | ||
| # 1. Every rank sends its full request list to every owner (sizes: | ||
| # ``request_sizes``). | ||
| # 2. Each owner gathers, for each requester, a full-length row buffer in | ||
| # request order -- rows it does not own are filler -- and sends it back. | ||
| # 3. The requester holds ``W`` buffers of its own length and picks, per | ||
| # request, the row from the owner it computed locally (``bucketize``). | ||
| # |
There was a problem hiding this comment.
Major changes here; we never properly had an algorithm here, it was always a hack. Now, it's sort of an algorithm. RDMA could help but it's not really part of physics nemo or anything yet.
| def _wrap_negative(index_flat: torch.Tensor, n_global: int) -> torch.Tensor: | ||
| r"""Wrap negative global positions like eager indexing does. | ||
|
|
||
| Parameters | ||
| ---------- | ||
| index_flat : torch.Tensor | ||
| 1-D global positions; negatives count from the end. | ||
| n_global : int | ||
| Global extent along the indexed dim. Positions ``>= n_global`` are not | ||
| checked (they select filler rows on the owner side). |
There was a problem hiding this comment.
This is a handler for negative indexes, which is good practice for sanitizing our arguments here.
|
|
||
| # First - Make sure we have the full input tensor | ||
| # Triggers an all_gather(_v) for (uneven) tensors. | ||
| local_tensor = tensor.full_tensor() |
There was a problem hiding this comment.
This was bad news previously.
| spec = ctx.spec | ||
| dim = ctx.dim | ||
|
|
||
| local_index = index.full_tensor() |
There was a problem hiding this comment.
More bad news.
| dtype=grad_output._local_tensor.dtype, | ||
| ) | ||
| # local_grad_output = grad_output.to_local() | ||
| local_grad_output = grad_output.full_tensor() |
There was a problem hiding this comment.
still more bad news.
| # 1. Every rank sends its full request list to every owner (sizes: | ||
| # ``request_sizes``). | ||
| # 2. Each owner gathers, for each requester, a full-length row buffer in | ||
| # request order -- rows it does not own are filler -- and sends it back. | ||
| # 3. The requester holds ``W`` buffers of its own length and picks, per | ||
| # request, the row from the owner it computed locally (``bucketize``). |
There was a problem hiding this comment.
The algorithm here could be a two-stage pass of
- ask everyone, "what do you need, on all ranks?" in an all-gather
- then, it's an all-to-allv to collect that.
The reason I didn't do that, even though it's less bandwidth, is that you would necessarily have a cuda.sync to allocate ragged memory buffers based on the size of arrays from each other rank. Instead, I take a path where we over-allocate and over-spend on comms, since it's still relatively cheap, and have 100% non-blocking ops.
| "index_select off the sharded dim with a sharded index needs an " | ||
| "all-gather of the index; not supported under torch.compile" | ||
| ) | ||
| idx = index.full_tensor() |
There was a problem hiding this comment.
In this case, there is not a nice way I've come up with to avoid resolving the full index tensor here.
| cached_spec = dataclasses.replace( | ||
| cached_spec, | ||
| placements=grad_placements, | ||
| shard_order=grad_output._spec.shard_order, |
There was a problem hiding this comment.
shard_order is an upstream attribute we're simply adding in here.
PhysicsNeMo Pull Request
Replaces the all-gathering
ShardedIndexSelectwith a routed exchange: a
Shard(d)source and an index of global positions(sharded or replicated) move only the requested rows between ranks. This is
what lets the domain-parallel mesh readers (next PR) keep global vertex ids
in the cell arrays and resolve
points[cells]without materializing thefull point array on any rank.
There is a lot of intellectual overlap here with the halo work that alchemi
supports. Therefore, In this PR a new primitive
exchange.py(likering.pyand
halo.pyis introduced to support both, and used as a shared backend.The Alchemi pieces shouldn't be affected by this...
shard_utils/exchange.py(new): row all-to-all, group-name resolution andthe fp32 scatter-accumulator policy (float64 default, configurable),
shared by
halo_scatterandindex_ops.index_ops.py:routed_gather/sharded_index_select; handlers fortorch.index_select,Tensor.index_selectandTensor.__getitem__with asingle int-tensor key. All exchange sizes come from the specs, so there is
no device-to-host sync and the ops trace under
torch.compile. Thebackward reuses the forward's request exchange and is a single collective;
a replicated index splits backward responsibility across ranks. Negative
indices wrap; 1-D mesh and integer dtype are enforced.
_ToTorchTensor.backwardadopts the gradient's shard order with itsplacements.
test_index.py(getitem),index_selectalong the sharded dim intest_select.py,test_exchange.py, routed-gather compile test.Description
Checklist
Dependencies
Review Process
All PRs are reviewed by the PhysicsNeMo team before merging.
Depending on which files are changed, GitHub may automatically assign a maintainer for review.
We are also testing AI-based code review tools (e.g., Greptile), which may add automated comments with a confidence score.
This score reflects the AI’s assessment of merge readiness and is not a qualitative judgment of your work, nor is
it an indication that the PR will be accepted / rejected.
AI-generated feedback should be reviewed critically for usefulness.
You are not required to respond to every AI comment, but they are intended to help both authors and reviewers.
Please react to Greptile comments with 👍 or 👎 to provide feedback on their accuracy.
Stack created with GitHub Stacks CLI • Give Feedback 💬