ByteDance-Seed / ByteDance-Seed/Triton-distributed
[BUG] Multi-node reduce-scatter crash due to unregistered output buffer in NVSHMEM P2P
- Dominant language
- Python
- Stars
- 1.5k
- Forks
- 172
- PR merge metrics
- No merged PRs in 30d
Description
When I use `gemm_rs` op with `fuse_scatter=False` in multi node environment, it crashes with errors like:
```
File "/workspace/triton-distributed/python/triton_dist/kernels/nvidia/reduce_scatter.py", line 864, in reduce_scatter_2d_op
output = reduce_scatter_multi_node(input, ctx, output)
File "/workspace/triton-distributed/python/triton_dist/kernels/nvidia/reduce_scatter.py", line 843, in reduce_scatter_multi_node
rs_result_per_node = reduce_scatter_for_each_node_ring(input, ctx, out_each_node)
File "/workspace/triton-distributed/python/triton_dist/kernels/nvidia/reduce_scatter.py", line 574, in reduce_scatter_for_each_node_ring
scatter_buf.data_ptr(),
AttributeError: 'NoneType' object has no attribute 'data_ptr'
```
When I read code, I find that in the current multi-node `reduce_scatter_multi_node` implementation, the output tensor is set to None when `nnodes > 1`. But it is used by `pynvshmem.putmem_on_stream` in `reduce_scatter_for_each_node_ring` like:
```
pynvshmem.putmem_on_stream(
p2p_buf[M_start:M_end].data_ptr(),
scatter_buf.data_ptr(),
nbytes_per_rank,
peer_rank,
torch.cuda.current_stream().cuda_stream,
)
```
So ther error happens. However, after modifying the code to pass a valid output tensor for multi-node runs, the process crashes during pynvshmem.putmem_on_stream with errors like:
```
[rank0]:[rank0]: File "/workspace/triton-distributed/python/triton_dist/kernels/nvidia/gemm_reduce_scatter.py", line 580, in gemm_rs
[rank0]:[rank0]: c = gemm_rs_op(a, b, ctx, persistent, fuse_scatter)
[rank0]:[rank0]: File "/workspace/triton-distributed/python/triton_dist/kernels/nvidia/gemm_reduce_scatter.py", line 558, in gemm_rs_op
[rank0]:[rank0]: reduce_scatter_2d_op(gemm_out, ctx.rs_ctx, output)
[rank0]:[rank0]: File "/workspace/triton-distributed/python/triton_dist/kernels/nvidia/reduce_scatter.py", line 866, in reduce_scatter_2d_op
[rank0]:[rank0]: output = reduce_scatter_multi_node(input, ctx, output)
[rank0]:[rank0]: File "/workspace/triton-distributed/python/triton_dist/kernels/nvidia/reduce_scatter.py", line 844, in reduce_scatter_multi_node
[rank0]:[rank0]: rs_result_per_node = reduce_scatter_for_each_node_ring(input, ctx, output)
[rank0]:[rank0]: File "/workspace/triton-distributed/python/triton_dist/kernels/nvidia/reduce_scatter.py", line 572, in reduce_scatter_for_each_node_ring
[rank0]:[rank0]: pynvshmem.putmem_on_stream(
[rank0]:[rank0]: File "nvshmem/bindings/nvshmem.pyx", line 1617, in nvshmem.bindings.nvshmem.putmem_on_stream
[rank0]:[rank0]: TypeError: putmem_on_stream() takes exactly 5 positional arguments (4 given)
```
And I use `nvshmem4py-cu12==0.2.1`. I think this is a bug and I want you to help me resolve it.
Contributor guide
Research direction
Start in python/triton_dist/kernels/nvidia/reduce_scatter.py at reduce_scatter_multi_node and reduce_scatter_for_each_node_ring, then inspect gemm_reduce_scatter.py and the installed nvshmem4py-cu12==0.2.1 putmem_on_stream signature. Reproduce gemm_rs with fuse_scatter=False across multiple nodes; done means the output buffer is valid and the NVSHMEM call matches the installed API without either reported crash.
Written by the indexing model from the issue text.
Assessment
- Tech stack
- python, pytorch
- Domain
- distributed-systems
- Issue type
- Bug
- Difficulty
- 4/5
- Estimated time
- 3-5 days
- Activity status
- Stale
- Clarity
- Mostly clear
- Newbie friendliness
- 38/100