Skip to content

Move ShardTensor ring ops to functional collectives; compile-safe ring attention and ball query - #1980

Open
coreyjadams wants to merge 5 commits into
dp-01-standalone-fixesfrom
dp-02-ring-funcol-compile
Open

coreyjadams wants to merge 5 commits into
dp-01-standalone-fixesfrom
dp-02-ring-funcol-compile

Conversation

@coreyjadams

Copy link
Copy Markdown
Collaborator

PhysicsNeMo Pull Request

Moves the ring collectives used by ShardTensor
attention, kNN and ball query onto functional collectives, and makes the
sharded attention and radius search usable under torch.compile and bf16.

  • ring.py: perform_ring_iteration_funcol / finish_ring_iteration
    (all-to-all with two nonzero splits; overlap by issuing before compute and
    waiting after). The stream-based perform_ring_iteration_async and
    get_comm_stream are removed.
  • Ring SDPA forward/backward now use funcol collectives, without comm streams,
    events or double buffers. ring_sdpa runs as an eager graph-break region
    under torch.compile. This can probably be addressed but hasn't been
    tested yet.
  • Sharded SDPA in bf16/fp16 (ring and replicated-query): inputs are cast to
    the active autocast dtype before the autograd Function, and the output
    saved for backward is returned in the input dtype and in the attention
    kernel's memory layout. Care is taken to use the right memory format
    and precision for these kenrels, too, or it leads to NaN in the backward
    pass.
  • Sharded radius_search traces end to end under torch.compile: the ring
    ball query's per-block merge was a raw warp launch and is now an opaque
    custom op with a fake impl.
  • torch.unbind on a Partial ShardTensor resolves the pending reduction
    first; aten.detach_ is handled.
  • Tests: ring iteration parametrized over the three transports; SDPA tests
    run in fp32 and bf16 autocast with O(1)-gradient losses; sharded ball-query
    compile test; unbind and detach tests.

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 4cf4ac3b6b3a. An approval covers every file listed for that owner; one owner is sufficient for shared files.

⏳ @coreyjadams — 11 file(s)
  • physicsnemo/domain_parallel/custom_ops/_tensor_ops.py
  • physicsnemo/domain_parallel/shard_utils/attention_patches.py
  • physicsnemo/domain_parallel/shard_utils/knn.py
  • physicsnemo/domain_parallel/shard_utils/point_cloud_ops.py
  • physicsnemo/domain_parallel/shard_utils/ring.py
  • test/domain_parallel/ops/test_compile_ops.py
  • test/domain_parallel/ops/test_detach.py
  • test/domain_parallel/ops/test_ring.py
  • test/domain_parallel/ops/test_sdpa.py
  • test/domain_parallel/ops/test_unbind.py
  • test/domain_parallel/ops/utils.py
⏳ @negin513 — 11 file(s)
  • physicsnemo/domain_parallel/custom_ops/_tensor_ops.py
  • physicsnemo/domain_parallel/shard_utils/attention_patches.py
  • physicsnemo/domain_parallel/shard_utils/knn.py
  • physicsnemo/domain_parallel/shard_utils/point_cloud_ops.py
  • physicsnemo/domain_parallel/shard_utils/ring.py
  • test/domain_parallel/ops/test_compile_ops.py
  • test/domain_parallel/ops/test_detach.py
  • test/domain_parallel/ops/test_ring.py
  • test/domain_parallel/ops/test_sdpa.py
  • test/domain_parallel/ops/test_unbind.py
  • test/domain_parallel/ops/utils.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 should not merge until the new detach_ handler preserves the in-place operation's semantics; the ring transport configuration should also be clarified.

Findings

  1. P1 Detach Is Not In-Place ▶
  2. P2 Transport Selection Is Ignored ▶

Summary

  • Adds functional all-to-all ring primitives and uses them in attention, kNN, and radius search.
  • Corrects attention output dtype and memory layout for half-precision backward kernels.
  • Wraps the ball-query merge as an opaque custom operation with fake-tensor support.
  • Expands distributed tests for compilation, autocasting, unbind, detach, and ring behavior.

Reviews (1) · Last reviewed commit: "Fix numerical precision error in sdpa sh..."

Comment on lines +873 to +877
return ShardTensor(
tensor._local_tensor.detach(),
tensor._spec,
requires_grad=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 Detach Is Not In-Place

The new aten.detach_ handler returns a separate wrapper and leaves the original ShardTensor attached to its autograd graph. If a caller invokes torch.ops.aten.detach_(tensor) without replacing every reference, later operations on the original tensor can still backpropagate through the graph that was supposed to be detached. The handler should preserve the in-place contract or use an explicitly functional operation.

Comment on lines +257 to +259
``ring_config.communication_method`` is not consulted: the functional
collective is always an all-to-all with two nonzero splits. (There is
no P2P equivalent)

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.

P2 Transport Selection Is Ignored

perform_ring_iteration_funcol accepts a RingPassingConfig whose communication_method can be "p2p", but it silently ignores that setting and always runs all_to_all_single. Existing attention, kNN, and ball-query callers still request "p2p", so configuration and diagnostics describe a transport that is not used. Removing or overriding the field at those call sites, or rejecting unsupported selections, would prevent future backend decisions from silently taking the wrong collective path.

Note: If this suggestion doesn't match your team's coding style, reply to this and let me know. I'll remember it for next time!

@coreyjadams
coreyjadams force-pushed the dp-02-ring-funcol-compile branch from 4cf4ac3 to 7072815 Compare September 14, 2026 14:26
@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
@coreyjadams
coreyjadams force-pushed the dp-02-ring-funcol-compile branch from 7072815 to 3607b3d Compare September 14, 2026 16:02
@codecov

codecov Bot commented Sep 14, 2026 •

Copy link
Copy Markdown

Codecov Report

❌ Patch coverage is 33.00000% with 67 lines in your changes missing coverage. Please review.
✅ Project coverage is 72.17%. Comparing base (f69588c) to head (46eaa16).

Files with missing lines Patch % Lines
...o/domain_parallel/shard_utils/attention_patches.py 33.33% 30 Missing ⚠️
...emo/domain_parallel/shard_utils/point_cloud_ops.py 29.41% 24 Missing ⚠️
physicsnemo/domain_parallel/shard_utils/ring.py 31.25% 11 Missing ⚠️
...sicsnemo/domain_parallel/custom_ops/_tensor_ops.py 50.00% 2 Missing ⚠️
Additional details and impacted files
@@                    Coverage Diff                     @@
##           dp-01-standalone-fixes    #1980      +/-   ##
==========================================================
- Coverage                   72.35%   72.17%   -0.18%     
==========================================================
  Files                         928      928              
  Lines                       70005    69968      -37     
  Branches                    10549    10540       -9     
==========================================================
- Hits                        50650    50502     -148     
- Misses                      15868    16009     +141     
+ Partials                     3487     3457      -30     
Files with missing lines Coverage Δ
physicsnemo/domain_parallel/shard_utils/knn.py 60.00% <100.00%> (-31.43%) ⬇️
...sicsnemo/domain_parallel/custom_ops/_tensor_ops.py 58.73% <50.00%> (-8.61%) ⬇️
physicsnemo/domain_parallel/shard_utils/ring.py 54.28% <31.25%> (-26.49%) ⬇️
...emo/domain_parallel/shard_utils/point_cloud_ops.py 70.18% <29.41%> (-10.71%) ⬇️
...o/domain_parallel/shard_utils/attention_patches.py 43.70% <33.33%> (-16.05%) ⬇️
🚀 New features to boost your workflow:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.

…g SDPA

- ring.py: add perform_ring_iteration_funcol / finish_ring_iteration (a2a
  with two nonzero splits; overlap by issuing before compute and waiting
  after). Remove the stream-based perform_ring_iteration_async and
  get_comm_stream, which no longer have callers.
- attention_patches.py: ring SDPA forward/backward on the funcol step; drop
  the comm stream, events, double buffers and record_stream calls. ring_sdpa
  runs as an eager graph-break region under torch.compile.
- knn.py, point_cloud_ops.py: use the funcol step (overlap left as TODO).
- _tensor_ops.py: unbind resolves a Partial input before slicing; add an
  aten.detach_ dispatch handler.
- Tests: parametrize test_ring over dist-p2p / dist-a2a / funcol; compile
  helper passes dynamic=False and runs twice; ring SDPA compile test now
  asserts success and matches eager; new test_unbind.py and test_detach.py;
  numerical check asserts plain parameters receive plain gradients.
@coreyjadams
coreyjadams force-pushed the dp-02-ring-funcol-compile branch from 3607b3d to 46eaa16 Compare September 14, 2026 20:47

import torch
import torch.distributed as dist
import torch.distributed._functional_collectives as funcol

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 change is really not a change in algorithm at all. it's a move from cuda.streams to funcol collectives to ring.py so that we can compile it. Instead of explicit overlap, we use a launch-and-wait paradigm.

Note that totally disables the P2P route which is called out in the docstring.

_comm_streams: dict[int, torch.cuda.Stream] = {}


def get_comm_stream(device: torch.device | int) -> torch.cuda.Stream:

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.

No need to get streams anymore

tensor: torch.Tensor,
mesh: DeviceMesh,
ring_config: RingPassingConfig,
recv_tensor: torch.Tensor | None = 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.

THere is no need to send in the receive buffer explicitly when using funcol, we can allocate it internally here.

recv_shape: torch.Size | None = None,
) -> tuple[torch.Tensor, list[dist.Work]]:
r"""Non-blocking single step of ring collective communication.
wait: bool = True,

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.

Passing wait=True makes this a synchronous function

Comment on lines +308 to +318
def finish_ring_iteration(
flat_recv: torch.Tensor,
recv_shape: torch.Size,
) -> torch.Tensor:
r"""Complete an in-flight ring step started with ``wait=False``.

Waits on the functional collective and reshapes the flat buffer to
``recv_shape``. This is the synchronization point: place it after the
compute that should overlap the communication.
"""
return funcol.wait_tensor(flat_recv).reshape(recv_shape)

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 helper function is needed to finish the ring later, if we are using the async version.

mesh,
ring_config,
recv_shape=shard_sizes[next_source_rank],
)

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.

ball query here is a synchronous function still, FYI.

mesh,
ring_config,
recv_shape=recv_shape,
)

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.

knn is a synchronous function still (wait=True by default)

aten = torch.ops.aten


def _apply_autocast(*tensors: torch.Tensor) -> tuple[torch.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 autocasting is necessary especially for FLARE. We end up with NaNs if we do not do this: we get funky precision errors without this.

return tuple(t.to(dtype) if t.is_floating_point() else t for t in tensors)


def _kernel_output_layout(t: torch.Tensor) -> torch.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 one was nasty to debug. sdpa has expectations for layout that we have to preserve, So we wrap it here to maintain it the right way the memory layout for the attention kernel. Yuck.

Comment on lines +265 to +270
next_k_flat = perform_ring_iteration_funcol(
current_k, mesh, ring_config, wait=False
)
next_v_flat = perform_ring_iteration_funcol(
current_v, mesh, ring_config, wait=False
)

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.

Note here that the ring collectives are overlapped

Comment on lines +306 to +307
current_k = finish_ring_iteration(next_k_flat, k.shape)
current_v = finish_ring_iteration(next_v_flat, v.shape)

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.

And here we explicitly wait on the next iteration's collectives to finish (but not on the last iteration, don't need it there).

grad_v = perform_ring_iteration(grad_v, mesh, ring_config)
# Last iteration: shift grads to the owning rank; nothing left
# to overlap with, so wait immediately.
grad_k = perform_ring_iteration_funcol(grad_k, mesh, ring_config)

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.

Like the comment says: the last iteration in the ring does not launch async. it blocks immediately.

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