Repository navigation
Move ShardTensor ring ops to functional collectives; compile-safe ring attention and ball query - #1980
Move ShardTensor ring ops to functional collectives; compile-safe ring attention and ball query#1980coreyjadams wants to merge 5 commits into
Conversation
CODEOWNERS review mapCurrent for commit ⏳ @coreyjadams — 11 file(s)
⏳ @negin513 — 11 file(s)
No CODEOWNER
Comment |
|
The PR should not merge until the new FindingsSummary
Reviews (1) · Last reviewed commit: "Fix numerical precision error in sdpa sh..." |
| return ShardTensor( | ||
| tensor._local_tensor.detach(), | ||
| tensor._spec, | ||
| requires_grad=False, | ||
| ) |
There was a problem hiding this comment.
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.
| ``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) |
There was a problem hiding this comment.
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!
4cf4ac3 to
7072815
Compare
7072815 to
3607b3d
Compare
Codecov Report❌ Patch coverage is 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
🚀 New features to boost your workflow:
|
…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.
…m tensor memory layouts.
3607b3d to
46eaa16
Compare
|
|
||
| import torch | ||
| import torch.distributed as dist | ||
| import torch.distributed._functional_collectives as funcol |
There was a problem hiding this comment.
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: |
There was a problem hiding this comment.
No need to get streams anymore
| tensor: torch.Tensor, | ||
| mesh: DeviceMesh, | ||
| ring_config: RingPassingConfig, | ||
| recv_tensor: torch.Tensor | None = None, |
There was a problem hiding this comment.
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, |
There was a problem hiding this comment.
Passing wait=True makes this a synchronous function
| 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) |
There was a problem hiding this comment.
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], | ||
| ) |
There was a problem hiding this comment.
ball query here is a synchronous function still, FYI.
| mesh, | ||
| ring_config, | ||
| recv_shape=recv_shape, | ||
| ) |
There was a problem hiding this comment.
knn is a synchronous function still (wait=True by default)
| aten = torch.ops.aten | ||
|
|
||
|
|
||
| def _apply_autocast(*tensors: torch.Tensor) -> tuple[torch.Tensor, ...]: |
There was a problem hiding this comment.
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: |
There was a problem hiding this comment.
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.
| 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 | ||
| ) |
There was a problem hiding this comment.
Note here that the ring collectives are overlapped
| current_k = finish_ring_iteration(next_k_flat, k.shape) | ||
| current_v = finish_ring_iteration(next_v_flat, v.shape) |
There was a problem hiding this comment.
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) |
There was a problem hiding this comment.
Like the comment says: the last iteration in the ring does not launch async. it blocks immediately.
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.compileand 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_asyncandget_comm_streamare removed.funcolcollectives, without comm streams,events or double buffers.
ring_sdparuns as an eager graph-break regionunder
torch.compile. This can probably be addressed but hasn't beentested yet.
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.
radius_searchtraces end to end undertorch.compile: the ringball query's per-block merge was a raw warp launch and is now an opaque
custom op with a fake impl.
torch.unbindon aPartialShardTensor resolves the pending reductionfirst;
aten.detach_is handled.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 💬