From e0680ca39c7bc3b693ff2d62853ac8a8b7645152 Mon Sep 17 00:00:00 2001 From: Mick Date: Thu, 6 Aug 2026 11:24:20 +0800 Subject: [PATCH] Match dist's dim-0 split semantics, and pin the disabled-state check Review caught both: split sizes were treated as flattened element counts where dist.all_to_all_single defines them along dim 0 -- rows now convert to element counts so n-D tensors split identically to torch, with an uneven 2-D parity case in the test. And this communicator initializes enabled, so the raise-when-disabled check now sets its state explicitly instead of relying on change_state's exit value. Verified on 2xH200: the full test file passes (bitwise parity, capture+replay, uneven 2-D splits, disabled refusal). Co-Authored-By: Claude Fable 5 --- .../device_communicators/pynccl.py | 102 +++++++++++ .../device_communicators/pynccl_wrapper.py | 12 ++ .../multimodal_gen/test/server/gpu_cases.py | 3 + .../test_pynccl_a2a_capture_2_gpu.py | 164 ++++++++++++++++++ 4 files changed, 281 insertions(+) create mode 100644 python/sglang/multimodal_gen/test/single_test_file/test_pynccl_a2a_capture_2_gpu.py diff --git a/python/sglang/multimodal_gen/runtime/distributed/device_communicators/pynccl.py b/python/sglang/multimodal_gen/runtime/distributed/device_communicators/pynccl.py index 6eca1e35a2c9..45a9ff40fae6 100644 --- a/python/sglang/multimodal_gen/runtime/distributed/device_communicators/pynccl.py +++ b/python/sglang/multimodal_gen/runtime/distributed/device_communicators/pynccl.py @@ -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 @@ def recv(self, tensor: torch.Tensor, src: int, stream=None): 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 diff --git a/python/sglang/multimodal_gen/runtime/distributed/device_communicators/pynccl_wrapper.py b/python/sglang/multimodal_gen/runtime/distributed/device_communicators/pynccl_wrapper.py index d28618f4f51c..048f117dd679 100644 --- a/python/sglang/multimodal_gen/runtime/distributed/device_communicators/pynccl_wrapper.py +++ b/python/sglang/multimodal_gen/runtime/distributed/device_communicators/pynccl_wrapper.py @@ -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 @@ def ncclBroadcast( ) ) + 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)) diff --git a/python/sglang/multimodal_gen/test/server/gpu_cases.py b/python/sglang/multimodal_gen/test/server/gpu_cases.py index 558f1cab3ec0..40903b85ff58 100644 --- a/python/sglang/multimodal_gen/test/server/gpu_cases.py +++ b/python/sglang/multimodal_gen/test/server/gpu_cases.py @@ -1145,6 +1145,7 @@ def _discover_unit_tests() -> list[str]: "../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 @@ def _discover_unit_tests() -> list[str]: "../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, }, } diff --git a/python/sglang/multimodal_gen/test/single_test_file/test_pynccl_a2a_capture_2_gpu.py b/python/sglang/multimodal_gen/test/single_test_file/test_pynccl_a2a_capture_2_gpu.py new file mode 100644 index 000000000000..310e4d90cca7 --- /dev/null +++ b/python/sglang/multimodal_gen/test/single_test_file/test_pynccl_a2a_capture_2_gpu.py @@ -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()