ByteDance-Seed / ByteDance-Seed/Triton-distributed

[BUG] Multi-node reduce-scatter crash due to unregistered output buffer in NVSHMEM P2P

Open
#155 2 comments 0 reactions 0 assignees View on GitHub
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

Open the contributing 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

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.