[Diffusion] Delete useless _ulysses_input_split func (#19786)
This commit is contained in:
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user