[diffusion] feat: capture-safe pynccl all-to-all (#33775)

Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
Mick
2026-08-06 15:41:05 +08:00
committed by GitHub
co-authored by Claude Fable 5
parent 31c1e5943f
commit bfce378e5f
4 changed files with 281 additions and 0 deletions
@@ -5,6 +5,8 @@
# Adapted from https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/distributed/device_communicators/pynccl.py
# ===================== import region =====================
from contextlib import contextmanager
import torch
import torch.distributed as dist
from torch.distributed import ProcessGroup, ReduceOp
@@ -232,6 +234,106 @@ class PyNcclCommunicator:
cudaStream_t(stream.cuda_stream),
)
def group_start(self):
self.nccl.ncclGroupStart()
def group_end(self):
self.nccl.ncclGroupEnd()
def all_to_all_single(
self,
output: torch.Tensor,
input_: torch.Tensor,
output_split_sizes: list[int] | None = None,
input_split_sizes: list[int] | None = None,
stream=None,
) -> None:
"""Equal-split all-to-all, the dist.all_to_all_single equivalent.
Exists because ProcessGroupNCCL's collectives cannot be recorded into a
CUDA graph: their host-side per-op bookkeeping advances once at capture,
so replays leave the ranks disagreeing about which collective is in
flight and they hang. Raw ncclSend/ncclRecv inside a group carries no
such state, so a captured region can hold the exchange.
"""
if self.disabled:
raise RuntimeError(
"pynccl all_to_all_single called while the communicator is "
"disabled; wrap it in change_state(enable=True)"
)
assert output.dtype == input_.dtype, (output.dtype, input_.dtype)
assert input_.is_contiguous() and output.is_contiguous()
if input_split_sizes is None and output_split_sizes is None:
assert output.numel() == input_.numel(), (
output.numel(),
input_.numel(),
)
assert input_.numel() % self.world_size == 0, (
f"all_to_all_single without split sizes needs an equal split, "
f"got {input_.numel()} elements over {self.world_size} ranks"
)
if stream is None:
stream = current_stream()
# dist.all_to_all_single defines split sizes along dim 0; convert rows
# to element counts so n-D tensors split identically to torch
in_row = input_.numel() // input_.size(0) if input_.dim() else 1
out_row = output.numel() // output.size(0) if output.dim() else 1
chunk = input_.numel() // self.world_size
if input_split_sizes is None:
send_counts = [chunk] * self.world_size
else:
assert sum(input_split_sizes) == input_.size(0)
send_counts = [n * in_row for n in input_split_sizes]
if output_split_sizes is None:
recv_counts = [chunk] * self.world_size
else:
assert sum(output_split_sizes) == output.size(0)
recv_counts = [n * out_row for n in output_split_sizes]
assert len(send_counts) == len(recv_counts) == self.world_size
send = input_.view(-1)
recv = output.view(-1)
dtype = ncclDataTypeEnum.from_torch(input_.dtype)
itemsize = input_.element_size()
send_off = recv_off = 0
self.nccl.ncclGroupStart()
for peer in range(self.world_size):
self.nccl.ncclSend(
buffer_type(send.data_ptr() + send_off * itemsize),
send_counts[peer],
dtype,
peer,
self.comm,
cudaStream_t(stream.cuda_stream),
)
self.nccl.ncclRecv(
buffer_type(recv.data_ptr() + recv_off * itemsize),
recv_counts[peer],
dtype,
peer,
self.comm,
cudaStream_t(stream.cuda_stream),
)
send_off += send_counts[peer]
recv_off += recv_counts[peer]
self.nccl.ncclGroupEnd()
@contextmanager
def change_state(self, enable: bool | None = None):
"""Enable the communicator for the duration of the block.
The graph-capture path is the only caller that needs it, so ordinary
traffic keeps going through the process group and this stays a scoped
override rather than a mode switch.
"""
if enable is None:
enable = self.available
old_disabled = self.disabled
self.disabled = not enable
try:
yield
finally:
self.disabled = old_disabled
def broadcast(self, tensor: torch.Tensor, src: int, stream=None):
if self.disabled:
return
@@ -256,6 +256,12 @@ class NCCLLibrary:
# it is better not to call it at all.
# ncclResult_t ncclCommDestroy(ncclComm_t comm);
Function("ncclCommDestroy", ncclResult_t, [ncclComm_t]),
# Batches the enclosed send/recv into one collective, which is what
# makes an all-to-all out of the pairwise primitives.
# ncclResult_t ncclGroupStart();
Function("ncclGroupStart", ncclResult_t, []),
# ncclResult_t ncclGroupEnd();
Function("ncclGroupEnd", ncclResult_t, []),
]
# class attribute to store the mapping from the path to the library
@@ -436,6 +442,12 @@ class NCCLLibrary:
)
)
def ncclGroupStart(self) -> None:
self.NCCL_CHECK(self._funcs["ncclGroupStart"]())
def ncclGroupEnd(self) -> None:
self.NCCL_CHECK(self._funcs["ncclGroupEnd"]())
def ncclCommDestroy(self, comm: ncclComm_t) -> None:
self.NCCL_CHECK(self._funcs["ncclCommDestroy"](comm))
@@ -1145,6 +1145,7 @@ STANDALONE_FILES = {
"../single_test_file/test_ar_models.py",
"../single_test_file/test_ipc_a2a_2_gpu.py",
"../single_test_file/test_dp_serving_2_gpu.py",
"../single_test_file/test_pynccl_a2a_capture_2_gpu.py",
],
}
@@ -1181,6 +1182,8 @@ STANDALONE_FILE_EST_TIMES = {
"../single_test_file/test_ipc_a2a_2_gpu.py": 240.0,
# zimage server startup dominates; six short requests after warmup
"../single_test_file/test_dp_serving_2_gpu.py": 900.0,
# one capture plus three replays on a 32K-element exchange
"../single_test_file/test_pynccl_a2a_capture_2_gpu.py": 180.0,
},
}
@@ -0,0 +1,164 @@
"""A pynccl all-to-all must be replay-safe inside a CUDA graph.
ProcessGroupNCCL's collectives are not: their host-side per-op bookkeeping
advances once at capture, so replays leave the ranks disagreeing about which
collective is in flight and both hang. That is why the full-forward DiT graph
cannot capture `dist.all_to_all_single`. Raw ncclSend/ncclRecv inside a group
carries no such state, so this exchange can live in a captured region -- this
test is what keeps that true.
pytest -v python/sglang/multimodal_gen/test/single_test_file/test_pynccl_a2a_capture_2_gpu.py
"""
from __future__ import annotations
import os
import subprocess
import sys
import unittest
import torch
from sglang.multimodal_gen.runtime.platforms import current_platform
from sglang.test.test_utils import CustomTestCase
_WORLD = 2
def _worker() -> int:
import torch.distributed as dist
from sglang.multimodal_gen.runtime.distributed.device_communicators.pynccl import (
PyNcclCommunicator,
)
rank = int(os.environ["RANK"])
world = int(os.environ["WORLD_SIZE"])
torch.cuda.set_device(rank)
dist.init_process_group("nccl", rank=rank, world_size=world)
cpu_group = dist.new_group(ranks=list(range(world)), backend="gloo")
comm = PyNcclCommunicator(group=cpu_group, device=torch.device(f"cuda:{rank}"))
if not comm.available:
print("SKIP pynccl unavailable", flush=True)
return 0
failures = []
def make(seed: int) -> torch.Tensor:
g = torch.Generator(device="cuda").manual_seed(seed)
return (
torch.randn(world * 4096, dtype=torch.bfloat16, device="cuda", generator=g)
+ rank
)
def reference(x: torch.Tensor) -> torch.Tensor:
out = torch.empty_like(x)
dist.all_to_all_single(out, x.contiguous())
return out
x = make(11)
got = torch.empty_like(x)
with comm.change_state(enable=True):
comm.all_to_all_single(got, x)
torch.cuda.synchronize()
if not torch.equal(reference(x), got):
failures.append("eager result differs from dist.all_to_all_single")
# the point of the test: capture once, then replay against fresh inputs
static_in = make(22)
static_out = torch.empty_like(static_in)
with comm.change_state(enable=True):
comm.all_to_all_single(static_out, static_in) # NCCL wants a warm path
torch.cuda.synchronize()
graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(graph, capture_error_mode="thread_local"):
comm.all_to_all_single(static_out, static_in)
torch.cuda.synchronize()
for i, seed in enumerate((33, 44, 55)):
fresh = make(seed)
static_in.copy_(fresh)
expected = reference(fresh)
graph.replay()
torch.cuda.synchronize()
if not torch.equal(expected, static_out):
failures.append(f"replay {i} (seed {seed}) diverged")
# uneven dim-0 splits must match dist semantics (rows, not elements)
rows = (
torch.arange(4 * 3 * world, dtype=torch.bfloat16, device="cuda").reshape(
4 * world, 3
)
+ 100 * rank
)
in_splits = (
[1, 4 * world - 1] if world == 2 else [1] * (world - 1) + [3 * world + 1]
)
out_splits = [
in_splits[rank] for _ in range(world)
] # every rank sends in_splits[j] rows to rank j
ref_out = torch.empty(sum(out_splits), 3, dtype=torch.bfloat16, device="cuda")
dist.all_to_all_single(
ref_out, rows, output_split_sizes=out_splits, input_split_sizes=in_splits
)
got_out = torch.empty_like(ref_out)
with comm.change_state(enable=True):
comm.all_to_all_single(got_out, rows, out_splits, in_splits)
torch.cuda.synchronize()
if not torch.equal(ref_out, got_out):
failures.append("uneven dim-0 split diverged from dist.all_to_all_single")
# the raise-when-disabled contract, with the state set explicitly: this
# communicator initializes enabled, so exiting change_state restores that
with comm.change_state(enable=False):
try:
comm.all_to_all_single(static_out, static_in)
failures.append("disabled communicator did not refuse the exchange")
except RuntimeError:
pass
verdict = torch.tensor([len(failures)], device="cuda")
dist.all_reduce(verdict)
if failures:
print(f"rank{rank} FAIL {failures}", flush=True)
if rank == 0:
print(
f"PYNCCL_A2A_CAPTURE {'FAIL' if verdict.item() else 'PASS'}",
flush=True,
)
dist.barrier()
dist.destroy_process_group()
return 1 if verdict.item() else 0
class TestPyncclA2ACapture(CustomTestCase):
def test_all_to_all_survives_graph_replay(self):
if not current_platform.is_cuda():
self.skipTest("pynccl graph capture is exercised on CUDA only")
if torch.cuda.device_count() < _WORLD:
self.skipTest(f"needs {_WORLD} GPUs")
proc = subprocess.run(
[
sys.executable,
"-m",
"torch.distributed.run",
f"--nproc-per-node={_WORLD}",
"--master-port=29519",
__file__,
"--worker",
],
capture_output=True,
text=True,
timeout=1200,
)
print(proc.stdout[-4000:])
if proc.returncode != 0:
print(proc.stderr[-4000:], file=sys.stderr)
self.assertEqual(proc.returncode, 0, "pynccl all-to-all is not replay-safe")
self.assertIn("PYNCCL_A2A_CAPTURE PASS", proc.stdout)
if __name__ == "__main__":
if "--worker" in sys.argv:
raise SystemExit(_worker())
unittest.main()