Skip to content

Routed gather for sharded index_select and integer-tensor indexing - #1981

Open
coreyjadams wants to merge 1 commit into
dp-02-ring-funcol-compilefrom
dp-03-routed-gather
Open

coreyjadams wants to merge 1 commit into
dp-02-ring-funcol-compilefrom
dp-03-routed-gather

Conversation

@coreyjadams

Copy link
Copy Markdown
Collaborator

PhysicsNeMo Pull Request

Replaces 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. 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 the
full 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 (like ring.py
and halo.py is 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 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. All exchange sizes come from the specs, so there is
    no device-to-host sync and the ops trace under torch.compile. The
    backward 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.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.

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 💬

@github-actions

Copy link
Copy Markdown
Contributor

CODEOWNERS review map

Current for commit 66db6f26b20d. An approval covers every file listed for that owner; one owner is sufficient for shared files.

⏳ @coreyjadams — 10 file(s)
  • physicsnemo/domain_parallel/shard_tensor.py
  • physicsnemo/domain_parallel/shard_utils/init.py
  • physicsnemo/domain_parallel/shard_utils/exchange.py
  • physicsnemo/domain_parallel/shard_utils/halo_scatter.py
  • physicsnemo/domain_parallel/shard_utils/index_ops.py
  • test/ci_tests/interrogate_baseline.txt
  • test/domain_parallel/ops/test_compile_ops.py
  • test/domain_parallel/ops/test_exchange.py
  • test/domain_parallel/ops/test_index.py
  • test/domain_parallel/ops/test_select.py
⏳ @ktangsali — 1 file(s)
  • test/ci_tests/interrogate_baseline.txt
⏳ @negin513 — 9 file(s)
  • physicsnemo/domain_parallel/shard_tensor.py
  • physicsnemo/domain_parallel/shard_utils/init.py
  • physicsnemo/domain_parallel/shard_utils/exchange.py
  • physicsnemo/domain_parallel/shard_utils/halo_scatter.py
  • physicsnemo/domain_parallel/shard_utils/index_ops.py
  • test/domain_parallel/ops/test_compile_ops.py
  • test/domain_parallel/ops/test_exchange.py
  • test/domain_parallel/ops/test_index.py
  • test/domain_parallel/ops/test_select.py

No CODEOWNER

  • CHANGELOG.md

Comment /codeowners-info to refresh.

@coreyjadams
coreyjadams added this pull request to stack #1984 September 10, 2026 19:07
@greptile-apps

greptile-apps Bot commented Sep 10, 2026 •

Copy link
Copy Markdown
Contributor

Retrigger

The PR is not safe to merge until routed backward row selection, index bounds semantics, and DTensor placement handling are corrected.

Findings

  1. P1 Backward clamps to wrong size ▶
  2. P1 Invalid indices return data ▶
  3. P1 Sharded DTensor treated as replicated ▶

Summary

  • Adds compile-compatible routed gather and scatter-add custom operations.
  • Adds integer-tensor __getitem__ and method-form index_select handlers.
  • Adds distributed forward, backward, compile, exchange, negative-index, and placement tests.
  • The routed backward currently computes incorrect destination rows for common source/index shapes, and index validation and DTensor placement handling need correction.

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)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P1 Backward clamps to wrong size

_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.

Comment on lines +596 to +597
# Global positions, wrapped once; exchanged once and reused by the backward.
idx_local = _wrap_negative(idx_local.to(torch.int64), offsets[-1])

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P1 Invalid indices return data

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.

Comment on lines +500 to +501
if not isinstance(index, ShardTensor):
return False

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P1 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

codecov Bot commented Sep 10, 2026 •

Copy link
Copy Markdown

Codecov Report

❌ Patch coverage is 30.53097% with 157 lines in your changes missing coverage. Please review.
✅ Project coverage is 71.88%. Comparing base (46eaa16) to head (9831c24).

Files with missing lines Patch % Lines
...ysicsnemo/domain_parallel/shard_utils/index_ops.py 24.58% 135 Missing ⚠️
...hysicsnemo/domain_parallel/shard_utils/exchange.py 50.00% 22 Missing ⚠️
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     
Files with missing lines Coverage Δ
physicsnemo/domain_parallel/shard_tensor.py 59.40% <ø> (-2.20%) ⬇️
...hysicsnemo/domain_parallel/shard_utils/__init__.py 95.00% <ø> (ø)
...csnemo/domain_parallel/shard_utils/halo_scatter.py 63.07% <100.00%> (-24.29%) ⬇️
...hysicsnemo/domain_parallel/shard_utils/exchange.py 50.00% <50.00%> (ø)
...ysicsnemo/domain_parallel/shard_utils/index_ops.py 22.92% <24.58%> (-27.49%) ⬇️
🚀 New features to boost your workflow:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.

@copy-pr-bot

copy-pr-bot Bot commented Sep 14, 2026

Copy link
Copy Markdown

This pull request requires additional validation before any workflows can run on NVIDIA's runners.

Pull request vetters can view their responsibilities here.

Contributors can view more details about this message here.

@coreyjadams coreyjadams added the ci:multi-gpu Run this PR on multiGPU ci label Sep 14, 2026
@coreyjadams coreyjadams self-assigned this Sep 14, 2026
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.
Comment on lines +34 to +39
__all__ = [
"funcol_all_to_all_v_rows",
"get_fp32_scatter_accumulator",
"resolve_group_name",
"set_fp32_scatter_accumulator",
]

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

Comment on lines -49 to -53
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:

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The changes to this file should be exclusively moving things to exchange.py without functionality change.

Comment on lines +56 to +69
# ``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``).
#

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

Comment on lines +90 to +99
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).

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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()

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This was bad news previously.

spec = ctx.spec
dim = ctx.dim

local_index = index.full_tensor()

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

More bad news.

dtype=grad_output._local_tensor.dtype,
)
# local_grad_output = grad_output.to_local()
local_grad_output = grad_output.full_tensor()

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

still more bad news.

Comment on lines +63 to +68
# 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``).

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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()

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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,

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

shard_order is an upstream attribute we're simply adding in here.

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

ci:multi-gpu Run this PR on multiGPU ci

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant