From 115e9a1acd3cbe57c2d17ad0d6eab8b4f90d5d21 Mon Sep 17 00:00:00 2001 From: Xiaoyu Zhang <35585791+BBuf@users.noreply.github.com> Date: Wed, 4 Mar 2026 10:45:11 +0800 Subject: [PATCH] [Diffusion] Delete useless _ulysses_input_split func (#19786) --- .../sglang/multimodal_gen/runtime/layers/usp.py | 17 ----------------- 1 file changed, 17 deletions(-) diff --git a/python/sglang/multimodal_gen/runtime/layers/usp.py b/python/sglang/multimodal_gen/runtime/layers/usp.py index f4de11797..e82235009 100644 --- a/python/sglang/multimodal_gen/runtime/layers/usp.py +++ b/python/sglang/multimodal_gen/runtime/layers/usp.py @@ -9,7 +9,6 @@ from torch.distributed.tensor.experimental._attention import _cp_options from sglang.multimodal_gen.runtime.distributed.parallel_state import ( get_sp_group, - get_ulysses_parallel_rank, get_ulysses_parallel_world_size, ) from sglang.srt.utils.common import torch_release @@ -159,22 +158,6 @@ def _usp_output_all_to_all(x: torch.Tensor, head_dim: int = 1) -> torch.Tensor: return x -def _ulysses_input_split(x: torch.Tensor, dim: int = 1) -> torch.Tensor: - world_size = get_ulysses_parallel_world_size() - if world_size <= 1: - return x - rank = get_ulysses_parallel_rank() - assert x.ndim == 4, f"x must have 4 dimensions, got {x.ndim}" - - dim_to_split_size = x.shape[dim] - - assert ( - dim_to_split_size % world_size == 0 - ), f"The size of dimension {dim} ({dim_to_split_size}) must be divisible by world_size ({world_size})" - - return torch.tensor_split(x, world_size, dim=dim)[rank].contiguous() - - def ring_attn( query: torch.Tensor, key: torch.Tensor,