Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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

Expand Down
3 changes: 3 additions & 0 deletions python/sglang/multimodal_gen/runtime/distributed/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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__)

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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():
Expand Down Expand Up @@ -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(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down Expand Up @@ -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
Expand All @@ -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)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down Expand Up @@ -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)
Expand Down
18 changes: 8 additions & 10 deletions python/sglang/srt/distributed/parallel_state.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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):
Expand Down Expand Up @@ -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():
Expand Down Expand Up @@ -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_)

Expand Down
10 changes: 10 additions & 0 deletions python/sglang/srt/distributed/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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__)


Expand Down
5 changes: 3 additions & 2 deletions python/sglang/srt/layers/dp_attention.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand All @@ -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,
Expand Down
3 changes: 2 additions & 1 deletion python/sglang/srt/managers/prefill_delayer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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 (
Expand Down Expand Up @@ -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:
Expand Down
Loading