From 6a28e3b88c16c7312187e1c0fbc05d59ee94d9ba Mon Sep 17 00:00:00 2001 From: hnyls2002 Date: Fri, 25 Sep 2026 15:38:04 -0700 Subject: [PATCH] use torch *_single collectives --- .../base_device_communicator.py | 9 +++++---- .../runtime/distributed/utils.py | 3 +++ .../device_communicators/hpu_communicator.py | 3 ++- .../device_communicators/npu_communicator.py | 7 ++++--- .../device_communicators/xpu_communicator.py | 5 ++--- .../sglang/srt/distributed/parallel_state.py | 18 ++++++++---------- python/sglang/srt/distributed/utils.py | 10 ++++++++++ python/sglang/srt/layers/dp_attention.py | 5 +++-- python/sglang/srt/managers/prefill_delayer.py | 3 ++- .../managers/scheduler_components/dp_attn.py | 3 ++- .../cuda_graph_setup.py | 3 ++- .../dspark_components/dspark_planner.py | 5 ++--- 12 files changed, 45 insertions(+), 29 deletions(-) diff --git a/python/sglang/multimodal_gen/runtime/distributed/device_communicators/base_device_communicator.py b/python/sglang/multimodal_gen/runtime/distributed/device_communicators/base_device_communicator.py index be3dd2552199..c7a7dbdc0804 100644 --- a/python/sglang/multimodal_gen/runtime/distributed/device_communicators/base_device_communicator.py +++ b/python/sglang/multimodal_gen/runtime/distributed/device_communicators/base_device_communicator.py @@ -11,7 +11,10 @@ from torch import Tensor from torch.distributed import ProcessGroup, ReduceOp -from sglang.multimodal_gen.runtime.distributed.utils import all_gather_single +from sglang.multimodal_gen.runtime.distributed.utils import ( + all_gather_single, + reduce_scatter_single, +) def _ipc_all_to_all_4d(group, input_, scatter_dim): @@ -131,9 +134,7 @@ def backward(ctx: Any, grad_output: Tensor) -> tuple[None, Tensor, None, None]: grad_input = torch.empty( ctx.input_shape, dtype=grad_output.dtype, device=grad_output.device ) - dist.reduce_scatter_tensor( - grad_input, grad_chunks.contiguous(), group=ctx.group - ) + reduce_scatter_single(grad_input, grad_chunks.contiguous(), group=ctx.group) return None, grad_input, None, None diff --git a/python/sglang/multimodal_gen/runtime/distributed/utils.py b/python/sglang/multimodal_gen/runtime/distributed/utils.py index 0f185e6910b0..d21179f91575 100644 --- a/python/sglang/multimodal_gen/runtime/distributed/utils.py +++ b/python/sglang/multimodal_gen/runtime/distributed/utils.py @@ -24,12 +24,15 @@ try: from torch.distributed import all_gather_single as _all_gather_single + from torch.distributed import reduce_scatter_single as _reduce_scatter_single except ImportError: from torch.distributed import all_gather_into_tensor as _all_gather_single + from torch.distributed import reduce_scatter_tensor as _reduce_scatter_single from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger all_gather_single = _all_gather_single +reduce_scatter_single = _reduce_scatter_single logger = init_logger(__name__) diff --git a/python/sglang/srt/distributed/device_communicators/hpu_communicator.py b/python/sglang/srt/distributed/device_communicators/hpu_communicator.py index cfc0eff9ed11..0331b57f46c4 100644 --- a/python/sglang/srt/distributed/device_communicators/hpu_communicator.py +++ b/python/sglang/srt/distributed/device_communicators/hpu_communicator.py @@ -6,6 +6,7 @@ import torch.distributed as dist from torch.distributed import ProcessGroup +from sglang.srt.distributed.utils import all_gather_single from sglang.srt.utils import is_hpu if is_hpu(): @@ -41,7 +42,7 @@ def all_gather(self, x: torch.Tensor, dim: int = -1) -> torch.Tensor: ) # All-gather. htorch.core.mark_step() - dist.all_gather_into_tensor(output_tensor, x, group=self.group) + all_gather_single(output_tensor, x, group=self.group) # Reshape output_tensor = output_tensor.movedim(0, dim) output_tensor = output_tensor.reshape( diff --git a/python/sglang/srt/distributed/device_communicators/npu_communicator.py b/python/sglang/srt/distributed/device_communicators/npu_communicator.py index ec3f28fc66da..348e81a525d2 100644 --- a/python/sglang/srt/distributed/device_communicators/npu_communicator.py +++ b/python/sglang/srt/distributed/device_communicators/npu_communicator.py @@ -2,6 +2,7 @@ import torch.distributed as dist from torch.distributed import ProcessGroup +from sglang.srt.distributed.utils import all_gather_single from sglang.srt.utils import is_npu _is_npu = is_npu() @@ -39,8 +40,8 @@ def quant_all_reduce(self, x: torch.Tensor) -> torch.Tensor: output_size[:1], dtype=scale.dtype, device=scale.device ) # All-gather. - dist.all_gather_into_tensor(output_tensor, x_q, group=self.group) - dist.all_gather_into_tensor(output_scale, scale, group=self.group) + all_gather_single(output_tensor, x_q, group=self.group) + all_gather_single(output_scale, scale, group=self.group) output_tensor = output_tensor.to(x.dtype) * output_scale.unsqueeze(-1).to( x.dtype @@ -60,7 +61,7 @@ def all_gather(self, x: torch.Tensor, dim: int = -1) -> torch.Tensor: # Allocate output tensor. output_tensor = torch.empty(output_size, dtype=x.dtype, device=x.device) # All-gather. - dist.all_gather_into_tensor(output_tensor, x, group=self.group) + all_gather_single(output_tensor, x, group=self.group) # Reshape output_tensor = output_tensor.reshape((world_size,) + input_size) output_tensor = output_tensor.movedim(0, dim) diff --git a/python/sglang/srt/distributed/device_communicators/xpu_communicator.py b/python/sglang/srt/distributed/device_communicators/xpu_communicator.py index 2bca063bf2e1..4d84f5b05d70 100644 --- a/python/sglang/srt/distributed/device_communicators/xpu_communicator.py +++ b/python/sglang/srt/distributed/device_communicators/xpu_communicator.py @@ -6,6 +6,7 @@ import torch.distributed as dist from torch.distributed import ProcessGroup +from sglang.srt.distributed.utils import all_gather_single from sglang.srt.utils import is_xpu @@ -33,9 +34,7 @@ def gather( (self.world_size,) + input_size, dtype=input_.dtype, device=input_.device ) # All-gather. - torch.distributed.all_gather_into_tensor( - output_tensor, input_, group=self.group - ) + all_gather_single(output_tensor, input_, group=self.group) if rank_in_group == dst: # Reshape output_tensor = output_tensor.movedim(0, dim) diff --git a/python/sglang/srt/distributed/parallel_state.py b/python/sglang/srt/distributed/parallel_state.py index 7b078f2be733..7676ef3676c9 100644 --- a/python/sglang/srt/distributed/parallel_state.py +++ b/python/sglang/srt/distributed/parallel_state.py @@ -46,7 +46,11 @@ from sglang.srt import platforms from sglang.srt.compilation.compilation_config import register_split_op -from sglang.srt.distributed.utils import set_global_tcp_store +from sglang.srt.distributed.utils import ( + all_gather_single, + reduce_scatter_single, + set_global_tcp_store, +) from sglang.srt.environ import envs from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( is_in_tc_piecewise_cuda_graph, @@ -1093,9 +1097,7 @@ def _reduce_scatter_tensor( with pynccl_comm.change_state(enable=True): pynccl_comm.reduce_scatter(output, input) else: - torch.distributed.reduce_scatter_tensor( - output, input, group=self.device_group - ) + reduce_scatter_single(output, input, group=self.device_group) return output def reduce_scatter_tensor(self, output: torch.Tensor, input: torch.Tensor): @@ -1304,9 +1306,7 @@ def _all_gather_into_tensor(self, output: torch.Tensor, input: torch.Tensor): with pynccl_comm.change_state(enable=True): pynccl_comm.all_gather(output, input) else: - torch.distributed.all_gather_into_tensor( - output, input, group=self.device_group - ) + all_gather_single(output, input, group=self.device_group) def _has_aiter_custom_all_gather(self) -> bool: if self._deterministic_collectives_enabled(): @@ -1397,9 +1397,7 @@ def all_gather( if is_shm_available(input_.dtype, self.world_size, self.local_size): return torch.ops.sgl_kernel.shm_allgather(input_, dim) else: - torch.distributed.all_gather_into_tensor( - output_tensor, input_, group=self.device_group - ) + all_gather_single(output_tensor, input_, group=self.device_group) else: self.all_gather_into_tensor(output_tensor, input_) diff --git a/python/sglang/srt/distributed/utils.py b/python/sglang/srt/distributed/utils.py index 8862658fa44d..cbe3b65d40cc 100644 --- a/python/sglang/srt/distributed/utils.py +++ b/python/sglang/srt/distributed/utils.py @@ -17,8 +17,18 @@ import torch from torch.distributed import TCPStore +try: + from torch.distributed import all_gather_single as _all_gather_single + from torch.distributed import reduce_scatter_single as _reduce_scatter_single +except ImportError: # older torch builds only have the *_tensor names + from torch.distributed import all_gather_into_tensor as _all_gather_single + from torch.distributed import reduce_scatter_tensor as _reduce_scatter_single + from sglang.srt.runtime_context import get_resources +all_gather_single = _all_gather_single +reduce_scatter_single = _reduce_scatter_single + logger = logging.getLogger(__name__) diff --git a/python/sglang/srt/layers/dp_attention.py b/python/sglang/srt/layers/dp_attention.py index 44760c3d0b1e..a8c8d0bb0b38 100644 --- a/python/sglang/srt/layers/dp_attention.py +++ b/python/sglang/srt/layers/dp_attention.py @@ -518,6 +518,7 @@ def get_dp_local_slice_cpu( from sglang.kernels.ops.memory.memcpy_triton import memcpy_triton +from sglang.srt.distributed.utils import all_gather_single # TODO: write c++ kernel for cpu @@ -606,7 +607,7 @@ def _dp_gather_via_all_gather( if get_parallel().attn_tp_size == 1: if use_world: - torch.distributed.all_gather_into_tensor( + all_gather_single( global_tokens, local_tokens, group=torch.distributed.group.WORLD, @@ -625,7 +626,7 @@ def _dp_gather_via_all_gather( scattered_local_tokens, local_tokens ) if use_world: - torch.distributed.all_gather_into_tensor( + all_gather_single( global_tokens, scattered_local_tokens, group=torch.distributed.group.WORLD, diff --git a/python/sglang/srt/managers/prefill_delayer.py b/python/sglang/srt/managers/prefill_delayer.py index edcfa669696a..88e3d5fcb51f 100644 --- a/python/sglang/srt/managers/prefill_delayer.py +++ b/python/sglang/srt/managers/prefill_delayer.py @@ -7,6 +7,7 @@ import torch +from sglang.srt.distributed.utils import all_gather_single from sglang.srt.environ import envs from sglang.srt.runtime_context import ( get_parallel, @@ -369,7 +370,7 @@ def _gather_info( device=self._gather_device, dtype=torch.int64, ) - torch.distributed.all_gather_into_tensor( + all_gather_single( self._global_info_buffer.flatten(), local_info, group=self._gather_group, diff --git a/python/sglang/srt/managers/scheduler_components/dp_attn.py b/python/sglang/srt/managers/scheduler_components/dp_attn.py index 405f11f5d799..58483c2be190 100644 --- a/python/sglang/srt/managers/scheduler_components/dp_attn.py +++ b/python/sglang/srt/managers/scheduler_components/dp_attn.py @@ -7,6 +7,7 @@ from sglang.srt.batch_overlap.two_batch_overlap import TboDPAttentionPreparer from sglang.srt.configs.model_config import ModelConfig +from sglang.srt.distributed.utils import all_gather_single from sglang.srt.environ import envs from sglang.srt.layers.cp.utils import get_cp_strategy from sglang.srt.layers.dp_attention import dp_gather_width, world_dp_gather_enabled @@ -180,7 +181,7 @@ def all_gather( missing = flat_info.abs().sum(dim=1) == 0 flat_info[missing] = fallback_tensor else: - torch.distributed.all_gather_into_tensor( + all_gather_single( global_info_tensor.flatten(), local_info_tensor, group=group, diff --git a/python/sglang/srt/model_executor/model_runner_components/cuda_graph_setup.py b/python/sglang/srt/model_executor/model_runner_components/cuda_graph_setup.py index 50e1ad2ed8ef..5151de063c25 100644 --- a/python/sglang/srt/model_executor/model_runner_components/cuda_graph_setup.py +++ b/python/sglang/srt/model_executor/model_runner_components/cuda_graph_setup.py @@ -15,6 +15,7 @@ from sglang.srt.distributed.device_communicators.pynccl_allocator import ( prealloc_symmetric_memory_pool, ) +from sglang.srt.distributed.utils import all_gather_single from sglang.srt.environ import envs from sglang.srt.hardware_backend.npu.graph_runner.npu_graph_runner import NPUGraphRunner from sglang.srt.hardware_backend.xpu.graph_runner.xpu_graph_runner import XPUGraphRunner @@ -252,7 +253,7 @@ def sync_elastic_cuda_graph_config( gathered_hashes = torch.empty( dist.get_world_size(world_group), dtype=torch.int64, device=device ) - dist.all_gather_into_tensor(gathered_hashes, local_hash, group=world_group) + all_gather_single(gathered_hashes, local_hash, group=world_group) hashes = gathered_hashes.cpu().tolist() if any(value != hashes[0] for value in hashes[1:]): raise RuntimeError( diff --git a/python/sglang/srt/speculative/dspark_components/dspark_planner.py b/python/sglang/srt/speculative/dspark_components/dspark_planner.py index a807674826d6..b025f1000962 100644 --- a/python/sglang/srt/speculative/dspark_components/dspark_planner.py +++ b/python/sglang/srt/speculative/dspark_components/dspark_planner.py @@ -11,6 +11,7 @@ ScheduleVerifyLensTopk, compute_sort_survival, ) +from sglang.srt.distributed.utils import all_gather_single from sglang.srt.environ import InvariantCheckLevel, envs from sglang.srt.layers.dp_attention import is_dp_attention_enabled from sglang.srt.managers.overlap_utils import ( @@ -321,9 +322,7 @@ def _maybe_gather_dp_verify_tier( gathered = torch.empty( (torch.distributed.get_world_size(group=cpu_group),), dtype=torch.int64 ) - torch.distributed.all_gather_into_tensor( - gathered, local_tensor, group=cpu_group - ) + all_gather_single(gathered, local_tensor, group=cpu_group) batch.global_spec_verify_tier_num_tokens = gathered.tolist() def note_non_decode_step(self) -> None: