[Diffusion] Clean up obsolete worker plumbing, dead helpers, and tests (#39293)
This commit is contained in:
@@ -15,7 +15,6 @@ from typing import Any, Dict, List, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.distributed
|
||||
from torch.cuda import synchronize
|
||||
from torch.distributed import Backend, ProcessGroup
|
||||
|
||||
from sglang.multimodal_gen.runtime.distributed.device_communicators.base_device_communicator import (
|
||||
@@ -32,12 +31,6 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import (
|
||||
)
|
||||
from sglang.srt.utils import is_shm_available
|
||||
|
||||
try:
|
||||
import torch_musa # noqa: F401
|
||||
from torch_musa.core.device import synchronize
|
||||
except ModuleNotFoundError:
|
||||
pass
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
TensorMetadata = namedtuple("TensorMetadata", ["device", "dtype", "size"])
|
||||
@@ -321,20 +314,6 @@ class GroupCoordinator:
|
||||
world_size = self.world_size
|
||||
return (rank_in_group - 1) % world_size
|
||||
|
||||
@property
|
||||
def skip_rank(self):
|
||||
"""Return the global rank of the process that skip connects with the caller"""
|
||||
rank_in_group = self.rank_in_group
|
||||
world_size = self.world_size
|
||||
return self.ranks[(world_size - rank_in_group - 1) % world_size]
|
||||
|
||||
@property
|
||||
def group_skip_rank(self):
|
||||
"""Return the group rank of the process that skip connects with the caller"""
|
||||
rank_in_group = self.rank_in_group
|
||||
world_size = self.world_size
|
||||
return (world_size - rank_in_group - 1) % world_size
|
||||
|
||||
@contextmanager
|
||||
def graph_capture(self, graph_capture_context: GraphCaptureContext | None = None):
|
||||
if current_platform.is_cuda_alike():
|
||||
@@ -809,7 +788,7 @@ class GroupCoordinator:
|
||||
torch.distributed.barrier(group=self.cpu_group)
|
||||
|
||||
def send(self, tensor: torch.Tensor, dst: Optional[int] = None) -> None:
|
||||
"""Sends a tensor to the destination rank in a non-blocking way"""
|
||||
"""Send a tensor to the destination rank."""
|
||||
"""NOTE: `dst` is the rank_in_group of the destination rank."""
|
||||
if dst is None:
|
||||
dst = self.group_next_rank
|
||||
@@ -817,11 +796,7 @@ class GroupCoordinator:
|
||||
torch.distributed.send(
|
||||
tensor,
|
||||
self.ranks[dst],
|
||||
group=(
|
||||
self.device_groups[self.rank_in_group % 2]
|
||||
if self.world_size == 2
|
||||
else self.device_group
|
||||
),
|
||||
group=self.device_group,
|
||||
)
|
||||
|
||||
def recv(
|
||||
@@ -836,11 +811,7 @@ class GroupCoordinator:
|
||||
torch.distributed.recv(
|
||||
tensor,
|
||||
self.ranks[src],
|
||||
(
|
||||
self.device_groups[(self.rank_in_group + 1) % 2]
|
||||
if self.world_size == 2
|
||||
else self.device_group
|
||||
),
|
||||
group=self.device_group,
|
||||
)
|
||||
return tensor
|
||||
|
||||
@@ -860,438 +831,6 @@ class GroupCoordinator:
|
||||
self.mq_broadcaster = None
|
||||
|
||||
|
||||
class PipelineGroupCoordinator(GroupCoordinator):
|
||||
"""
|
||||
available attributes:
|
||||
rank: int # global rank
|
||||
ranks: List[int] # global ranks in the group
|
||||
world_size: int # size of the group
|
||||
difference between `local_rank` and `rank_in_group`:
|
||||
if we have a group of size 4 across two nodes:
|
||||
Process | Node | Rank | Local Rank | Rank in Group
|
||||
0 | 0 | 0 | 0 | 0
|
||||
1 | 0 | 1 | 1 | 1
|
||||
2 | 1 | 2 | 0 | 2
|
||||
3 | 1 | 3 | 1 | 3
|
||||
local_rank: int # local rank used to assign devices
|
||||
rank_in_group: int # rank inside the group
|
||||
cpu_group: ProcessGroup # group for CPU communication
|
||||
device_group: ProcessGroup # group for device communication
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
group_ranks: List[List[int]],
|
||||
local_rank: int,
|
||||
torch_distributed_backend: Union[str, Backend],
|
||||
group_name: str | None = None,
|
||||
):
|
||||
super().__init__(
|
||||
group_ranks=group_ranks,
|
||||
local_rank=local_rank,
|
||||
torch_distributed_backend=torch_distributed_backend,
|
||||
group_name=group_name,
|
||||
)
|
||||
self.rank = torch.distributed.get_rank()
|
||||
self.local_rank = local_rank
|
||||
self.device_group = None
|
||||
self.cpu_group = None
|
||||
self.cpu_groups = []
|
||||
self.device_groups = []
|
||||
if len(group_ranks[0]) > 2 or len(group_ranks[0]) == 1:
|
||||
for ranks in group_ranks:
|
||||
device_group = new_device_group(ranks, torch_distributed_backend)
|
||||
# a group with `gloo` backend, to allow direct coordination between
|
||||
# processes through the CPU.
|
||||
with suppress_stdout():
|
||||
cpu_group = torch.distributed.new_group(ranks, backend="gloo")
|
||||
if self.rank in ranks:
|
||||
self.ranks = ranks
|
||||
self.world_size = len(ranks)
|
||||
self.rank_in_group = ranks.index(self.rank)
|
||||
self.device_group = device_group
|
||||
self.cpu_group = cpu_group
|
||||
# when pipeline parallelism is 2, we need to create two groups to avoid
|
||||
# communication stall.
|
||||
# *_group_0_1 represents the group for communication from device 0 to
|
||||
# device 1.
|
||||
# *_group_1_0 represents the group for communication from device 1 to
|
||||
# device 0.
|
||||
elif len(group_ranks[0]) == 2:
|
||||
for ranks in group_ranks:
|
||||
device_group_0_1 = torch.distributed.new_group(
|
||||
ranks, backend=torch_distributed_backend
|
||||
)
|
||||
device_group_1_0 = torch.distributed.new_group(
|
||||
ranks, backend=torch_distributed_backend
|
||||
)
|
||||
# a group with `gloo` backend, to allow direct coordination between
|
||||
# processes through the CPU.
|
||||
with suppress_stdout():
|
||||
cpu_group_0_1 = torch.distributed.new_group(ranks, backend="gloo")
|
||||
cpu_group_1_0 = torch.distributed.new_group(ranks, backend="gloo")
|
||||
if self.rank in ranks:
|
||||
self.ranks = ranks
|
||||
self.world_size = len(ranks)
|
||||
self.rank_in_group = ranks.index(self.rank)
|
||||
self.device_groups = [device_group_0_1, device_group_1_0]
|
||||
self.cpu_groups = [cpu_group_0_1, cpu_group_1_0]
|
||||
self.device_group = device_group_0_1
|
||||
self.cpu_group = cpu_group_0_1
|
||||
|
||||
assert self.cpu_group is not None
|
||||
assert self.device_group is not None
|
||||
|
||||
self.device = current_platform.get_device(local_rank)
|
||||
|
||||
self.recv_buffer_set: bool = False
|
||||
self.recv_tasks_queue: List[Tuple[str, int]] = []
|
||||
self.receiving_tasks: List[Tuple[torch.distributed.Work, str, int]] = []
|
||||
self.dtype: Optional[torch.dtype] = None
|
||||
self.num_pipefusion_patches: Optional[int] = None
|
||||
|
||||
self.recv_shape: Dict[str, Dict[int, torch.Size]] = {}
|
||||
self.send_shape: Dict[str, Dict[int, torch.Size]] = {}
|
||||
self.recv_buffer: Dict[str, Dict[int, torch.Size]] = {}
|
||||
|
||||
self.skip_tensor_recv_buffer_set: bool = False
|
||||
self.recv_skip_tasks_queue: List[Union[int, Tuple[str, int]]] = []
|
||||
self.receiving_skip_tasks: List[Tuple[torch.distributed.Work, str, int]] = []
|
||||
self.skip_tensor_recv_buffer: Optional[
|
||||
Union[List[torch.Tensor], torch.Tensor]
|
||||
] = None
|
||||
self.skip_device_group = None
|
||||
for ranks in group_ranks:
|
||||
skip_device_group = new_device_group(ranks, torch_distributed_backend)
|
||||
if self.rank in ranks:
|
||||
self.skip_device_group = skip_device_group
|
||||
assert self.skip_device_group is not None
|
||||
|
||||
def reset_buffer(self):
|
||||
self.recv_tasks_queue = []
|
||||
self.receiving_tasks = []
|
||||
self.recv_shape = {}
|
||||
self.send_shape = {}
|
||||
self.recv_buffer = {}
|
||||
|
||||
self.recv_skip_tasks_queue = []
|
||||
self.receiving_skip_tasks = []
|
||||
self.skip_tensor_recv_buffer = {}
|
||||
|
||||
def set_config(self, dtype: torch.dtype):
|
||||
self.dtype = dtype
|
||||
|
||||
def set_recv_buffer(
|
||||
self,
|
||||
num_pipefusion_patches: int,
|
||||
patches_shape_list: List[List[int]],
|
||||
feature_map_shape: List[int],
|
||||
dtype: torch.dtype,
|
||||
):
|
||||
assert isinstance(dtype, torch.dtype), "dtype must be a torch.dtype object"
|
||||
assert (
|
||||
isinstance(num_pipefusion_patches, int) and num_pipefusion_patches >= 1
|
||||
), "num_pipefusion_patches must be greater than or equal to 1"
|
||||
self.dtype = dtype
|
||||
self.num_pipefusion_patches = num_pipefusion_patches
|
||||
self.recv_buffer = [
|
||||
torch.zeros(*shape, dtype=self.dtype, device=self.device)
|
||||
for shape in patches_shape_list
|
||||
]
|
||||
self.recv_buffer.append(
|
||||
torch.zeros(*feature_map_shape, dtype=self.dtype, device=self.device)
|
||||
)
|
||||
self.recv_buffer_set = True
|
||||
|
||||
def set_extra_tensors_recv_buffer(
|
||||
self,
|
||||
name: str,
|
||||
shape: List[int],
|
||||
num_buffers: int = 1,
|
||||
dtype: torch.dtype = torch.float16,
|
||||
):
|
||||
self.extra_tensors_recv_buffer[name] = [
|
||||
torch.zeros(*shape, dtype=dtype, device=self.device)
|
||||
for _ in range(num_buffers)
|
||||
]
|
||||
|
||||
def _check_shape_and_buffer(
|
||||
self,
|
||||
tensor_send_to_next=None,
|
||||
recv_prev=False,
|
||||
name: Optional[str] = None,
|
||||
segment_idx: int = 0,
|
||||
):
|
||||
send_flag = False
|
||||
name = name or "latent"
|
||||
if tensor_send_to_next is not None:
|
||||
shape_list = self.send_shape.get(name, None)
|
||||
if shape_list is None:
|
||||
self.send_shape[name] = {segment_idx: tensor_send_to_next.shape}
|
||||
send_flag = True
|
||||
elif shape_list.get(segment_idx, None) is None:
|
||||
self.send_shape[name][segment_idx] = tensor_send_to_next.shape
|
||||
send_flag = True
|
||||
|
||||
recv_flag = False
|
||||
if recv_prev:
|
||||
shape_list = self.recv_shape.get(name, None)
|
||||
if shape_list is None:
|
||||
recv_flag = True
|
||||
elif shape_list.get(segment_idx, None) is None:
|
||||
recv_flag = True
|
||||
|
||||
recv_prev_shape = self._communicate_shapes(
|
||||
tensor_send_to_next=tensor_send_to_next if send_flag else None,
|
||||
recv_prev=recv_flag,
|
||||
)
|
||||
|
||||
if recv_flag:
|
||||
if self.recv_shape.get(name, None) is None:
|
||||
self.recv_shape[name] = {segment_idx: recv_prev_shape}
|
||||
else:
|
||||
self.recv_shape[name][segment_idx] = recv_prev_shape
|
||||
|
||||
if self.recv_buffer.get(name, None) is None:
|
||||
self.recv_buffer[name] = {
|
||||
segment_idx: torch.zeros(
|
||||
recv_prev_shape, device=self.device, dtype=self.dtype
|
||||
)
|
||||
}
|
||||
else:
|
||||
if self.recv_buffer[name].get(segment_idx, None) is not None:
|
||||
logger.warning(
|
||||
f"Recv buffer [name: {name}, segment_idx: {segment_idx}] already exist. updating..."
|
||||
)
|
||||
self.recv_buffer[name][segment_idx] = torch.zeros(
|
||||
recv_prev_shape, device=self.device, dtype=self.dtype
|
||||
)
|
||||
|
||||
def _communicate_shapes(self, tensor_send_to_next=None, recv_prev=False):
|
||||
"""Communicate tensor shapes between stages. Used to communicate
|
||||
tensor shapes before the actual tensor communication happens.
|
||||
|
||||
Args:
|
||||
tensor_send_next: tensor to send to next rank (no tensor sent if
|
||||
set to None).
|
||||
recv_prev: boolean for whether tensor should be received from
|
||||
previous rank.
|
||||
"""
|
||||
|
||||
ops = []
|
||||
if recv_prev:
|
||||
recv_prev_dim_tensor = torch.empty(
|
||||
(1), device=self.device, dtype=torch.int64
|
||||
)
|
||||
recv_prev_dim_op = torch.distributed.P2POp(
|
||||
torch.distributed.irecv,
|
||||
recv_prev_dim_tensor,
|
||||
self.prev_rank,
|
||||
self.device_group,
|
||||
)
|
||||
ops.append(recv_prev_dim_op)
|
||||
|
||||
if tensor_send_to_next is not None:
|
||||
send_next_dim_tensor = torch.tensor(
|
||||
tensor_send_to_next.dim(), device=self.device, dtype=torch.int64
|
||||
)
|
||||
send_next_dim_op = torch.distributed.P2POp(
|
||||
torch.distributed.isend,
|
||||
send_next_dim_tensor,
|
||||
self.next_rank,
|
||||
self.device_group,
|
||||
)
|
||||
ops.append(send_next_dim_op)
|
||||
|
||||
if len(ops) > 0:
|
||||
reqs = torch.distributed.batch_isend_irecv(ops)
|
||||
for req in reqs:
|
||||
req.wait()
|
||||
|
||||
# To protect against race condition when using batch_isend_irecv().
|
||||
# should take this out once the bug with batch_isend_irecv is resolved.
|
||||
synchronize()
|
||||
|
||||
ops = []
|
||||
recv_prev_shape_tensor = None
|
||||
if recv_prev:
|
||||
recv_prev_shape_tensor = torch.empty(
|
||||
torch.Size(recv_prev_dim_tensor),
|
||||
device=self.device,
|
||||
dtype=torch.int64,
|
||||
)
|
||||
recv_prev_shape_op = torch.distributed.P2POp(
|
||||
torch.distributed.irecv,
|
||||
recv_prev_shape_tensor,
|
||||
self.prev_rank,
|
||||
self.device_group,
|
||||
)
|
||||
ops.append(recv_prev_shape_op)
|
||||
|
||||
if tensor_send_to_next is not None:
|
||||
send_next_shape_tensor = torch.tensor(
|
||||
tensor_send_to_next.size(),
|
||||
device=self.device,
|
||||
dtype=torch.int64,
|
||||
)
|
||||
send_next_shape_op = torch.distributed.P2POp(
|
||||
torch.distributed.isend,
|
||||
send_next_shape_tensor,
|
||||
self.next_rank,
|
||||
self.device_group,
|
||||
)
|
||||
ops.append(send_next_shape_op)
|
||||
|
||||
if len(ops) > 0:
|
||||
reqs = torch.distributed.batch_isend_irecv(ops)
|
||||
for req in reqs:
|
||||
req.wait()
|
||||
|
||||
synchronize()
|
||||
|
||||
recv_prev_shape = [0, 0, 0]
|
||||
if recv_prev_shape_tensor is not None:
|
||||
recv_prev_shape = recv_prev_shape_tensor
|
||||
return torch.Size(recv_prev_shape)
|
||||
|
||||
def pipeline_send(
|
||||
self, tensor: torch.Tensor, name: str = "latent", segment_idx: int = -1
|
||||
) -> None:
|
||||
tensor = tensor.contiguous()
|
||||
self._check_shape_and_buffer(
|
||||
tensor_send_to_next=tensor, name=name, segment_idx=segment_idx
|
||||
)
|
||||
self._pipeline_isend(tensor).wait()
|
||||
|
||||
def pipeline_isend(
|
||||
self, tensor: torch.Tensor, name: str = "latent", segment_idx: int = -1
|
||||
) -> None:
|
||||
tensor = tensor.contiguous()
|
||||
self._check_shape_and_buffer(
|
||||
tensor_send_to_next=tensor, name=name, segment_idx=segment_idx
|
||||
)
|
||||
self._pipeline_isend(tensor)
|
||||
|
||||
def pipeline_recv(self, idx: int = -1, name: str = "latent") -> torch.Tensor:
|
||||
name = name or "latent"
|
||||
self._check_shape_and_buffer(recv_prev=True, name=name, segment_idx=idx)
|
||||
self._pipeline_irecv(self.recv_buffer[name][idx]).wait()
|
||||
return self.recv_buffer[name][idx]
|
||||
|
||||
def add_pipeline_recv_task(self, idx: int = -1, name: str = "latent"):
|
||||
name = name or "latent"
|
||||
self.recv_tasks_queue.append((name, idx))
|
||||
|
||||
def recv_next(self):
|
||||
if len(self.recv_tasks_queue) == 0:
|
||||
raise ValueError("No more tasks to receive")
|
||||
elif len(self.recv_tasks_queue) > 0:
|
||||
name, idx = self.recv_tasks_queue.pop(0)
|
||||
self._check_shape_and_buffer(recv_prev=True, name=name, segment_idx=idx)
|
||||
self.receiving_tasks.append(
|
||||
(self._pipeline_irecv(self.recv_buffer[name][idx]), name, idx)
|
||||
)
|
||||
|
||||
def get_pipeline_recv_data(
|
||||
self, idx: int = -1, name: str = "latent"
|
||||
) -> torch.Tensor:
|
||||
assert len(self.receiving_tasks) > 0, (
|
||||
"No tasks to receive, call add_pipeline_recv_task first"
|
||||
)
|
||||
receiving_task = self.receiving_tasks.pop(0)
|
||||
receiving_task[0].wait()
|
||||
assert receiving_task[1] == name and receiving_task[2] == idx, (
|
||||
"Received tensor does not match the requested"
|
||||
)
|
||||
return self.recv_buffer[name][idx]
|
||||
|
||||
def _pipeline_irecv(self, tensor: torch.tensor):
|
||||
return torch.distributed.irecv(
|
||||
tensor,
|
||||
src=self.prev_rank,
|
||||
group=(
|
||||
self.device_groups[(self.rank_in_group + 1) % 2]
|
||||
if self.world_size == 2
|
||||
else self.device_group
|
||||
),
|
||||
)
|
||||
|
||||
def _pipeline_isend(self, tensor: torch.tensor):
|
||||
return torch.distributed.isend(
|
||||
tensor,
|
||||
dst=self.next_rank,
|
||||
group=(
|
||||
self.device_groups[self.rank_in_group % 2]
|
||||
if self.world_size == 2
|
||||
else self.device_group
|
||||
),
|
||||
)
|
||||
|
||||
def set_skip_tensor_recv_buffer(
|
||||
self,
|
||||
patches_shape_list: List[List[int]],
|
||||
feature_map_shape: List[int],
|
||||
):
|
||||
self.skip_tensor_recv_buffer = [
|
||||
torch.zeros(*shape, dtype=self.dtype, device=self.device)
|
||||
for shape in patches_shape_list
|
||||
]
|
||||
self.skip_tensor_recv_buffer.append(
|
||||
torch.zeros(*feature_map_shape, dtype=self.dtype, device=self.device)
|
||||
)
|
||||
self.skip_tensor_recv_buffer_set = True
|
||||
|
||||
def pipeline_send_skip(self, tensor: torch.Tensor) -> None:
|
||||
tensor = tensor.contiguous()
|
||||
self._pipeline_isend_skip(tensor).wait()
|
||||
|
||||
def pipeline_isend_skip(self, tensor: torch.Tensor) -> None:
|
||||
tensor = tensor.contiguous()
|
||||
self._pipeline_isend_skip(tensor)
|
||||
|
||||
def pipeline_recv_skip(self, idx: int = -1) -> torch.Tensor:
|
||||
self._pipeline_irecv_skip(self.skip_tensor_recv_buffer[idx]).wait()
|
||||
return self.skip_tensor_recv_buffer[idx]
|
||||
|
||||
def add_pipeline_recv_skip_task(self, idx: int = -1):
|
||||
self.recv_skip_tasks_queue.append(idx)
|
||||
|
||||
def get_pipeline_recv_skip_data(self, idx: int = -1) -> torch.Tensor:
|
||||
assert len(self.receiving_skip_tasks) > 0, (
|
||||
"No tasks to receive, call add_pipeline_recv_skip_task first"
|
||||
)
|
||||
receiving_skip_task = self.receiving_skip_tasks.pop(0)
|
||||
receiving_skip_task[0].wait()
|
||||
assert receiving_skip_task[2] == idx, (
|
||||
"Received tensor does not match the requested"
|
||||
)
|
||||
return self.skip_tensor_recv_buffer[idx]
|
||||
|
||||
def recv_skip_next(self):
|
||||
if len(self.recv_skip_tasks_queue) == 0:
|
||||
raise ValueError("No more tasks to receive")
|
||||
elif len(self.recv_skip_tasks_queue) > 0:
|
||||
task = self.recv_skip_tasks_queue.pop(0)
|
||||
idx = task
|
||||
self.receiving_skip_tasks.append(
|
||||
(
|
||||
self._pipeline_irecv_skip(self.skip_tensor_recv_buffer[idx]),
|
||||
None,
|
||||
idx,
|
||||
)
|
||||
)
|
||||
|
||||
def _pipeline_irecv_skip(self, tensor: torch.tensor):
|
||||
return torch.distributed.irecv(
|
||||
tensor, src=self.skip_rank, group=self.skip_device_group
|
||||
)
|
||||
|
||||
def _pipeline_isend_skip(self, tensor: torch.tensor):
|
||||
return torch.distributed.isend(
|
||||
tensor, dst=self.skip_rank, group=self.skip_device_group
|
||||
)
|
||||
|
||||
|
||||
class SequenceParallelGroupCoordinator(GroupCoordinator):
|
||||
def __init__(
|
||||
self,
|
||||
|
||||
@@ -34,25 +34,18 @@ If you only need to use the distributed environment without model parallelism,
|
||||
import contextlib
|
||||
import datetime
|
||||
import os
|
||||
import weakref
|
||||
from collections import namedtuple
|
||||
from collections.abc import Callable
|
||||
from contextlib import contextmanager
|
||||
from multiprocessing import shared_memory
|
||||
from typing import Any, List, Optional
|
||||
from unittest.mock import patch
|
||||
from typing import List, Optional
|
||||
|
||||
import torch
|
||||
import torch.distributed
|
||||
from torch.distributed import ProcessGroup
|
||||
|
||||
import sglang.multimodal_gen.envs as envs
|
||||
from sglang.multimodal_gen.runtime.distributed.utils import StatelessProcessGroup
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
|
||||
from .group_coordinator import (
|
||||
GroupCoordinator,
|
||||
PipelineGroupCoordinator,
|
||||
SequenceParallelGroupCoordinator,
|
||||
get_local_torch_device,
|
||||
new_device_group,
|
||||
@@ -64,7 +57,7 @@ logger = init_logger(__name__)
|
||||
_WORLD: GroupCoordinator | None = None
|
||||
_TP: GroupCoordinator | None = None
|
||||
_SP: SequenceParallelGroupCoordinator | None = None
|
||||
_PP: PipelineGroupCoordinator | None = None
|
||||
_PP: GroupCoordinator | None = None
|
||||
_CFG: GroupCoordinator | None = None
|
||||
_DP: GroupCoordinator | None = None
|
||||
# all ranks serving one pipeline replica (every dim except dp); with
|
||||
@@ -80,53 +73,6 @@ _VAE_DECODE_PARALLEL_AXES = "tp-sp-pp-cfg"
|
||||
_REPLICA_PARALLEL_AXES = "tp-sp-pp-cfg"
|
||||
_ENCODER_DP_PARALLEL_AXES = "sp-pp-cfg"
|
||||
|
||||
TensorMetadata = namedtuple("TensorMetadata", ["device", "dtype", "size"])
|
||||
|
||||
|
||||
def _split_tensor_dict(
|
||||
tensor_dict: dict[str, torch.Tensor | Any],
|
||||
) -> tuple[list[tuple[str, Any]], list[torch.Tensor]]:
|
||||
"""Split the tensor dictionary into two parts:
|
||||
1. A list of (key, value) pairs. If the value is a tensor, it is replaced
|
||||
by its metadata.
|
||||
2. A list of tensors.
|
||||
"""
|
||||
metadata_list: list[tuple[str, Any]] = []
|
||||
tensor_list: list[torch.Tensor] = []
|
||||
for key, value in tensor_dict.items():
|
||||
if isinstance(value, torch.Tensor):
|
||||
# Note: we cannot use `value.device` here,
|
||||
# because it contains not only the device type but also the device
|
||||
# index (e.g. "cuda:0"). We only need the device type.
|
||||
# receiving side will set the device index.
|
||||
device = value.device.type
|
||||
metadata_list.append(
|
||||
(key, TensorMetadata(device, value.dtype, value.size()))
|
||||
)
|
||||
tensor_list.append(value)
|
||||
else:
|
||||
metadata_list.append((key, value))
|
||||
return metadata_list, tensor_list
|
||||
|
||||
|
||||
_groups: dict[str, Callable[[], Optional["GroupCoordinator"]]] = {}
|
||||
|
||||
|
||||
def _register_group(group: "GroupCoordinator") -> None:
|
||||
_groups[group.unique_name] = weakref.ref(group)
|
||||
|
||||
|
||||
def all_reduce(tensor: torch.Tensor, group_name: str) -> torch.Tensor:
|
||||
assert group_name in _groups, f"Group {group_name} is not found."
|
||||
group = _groups[group_name]()
|
||||
if group is None:
|
||||
raise ValueError(f"Group {group_name} is destroyed.")
|
||||
return group._all_reduce_out_place(tensor)
|
||||
|
||||
|
||||
def all_reduce_fake(tensor: torch.Tensor, group_name: str) -> torch.Tensor:
|
||||
return torch.empty_like(tensor)
|
||||
|
||||
|
||||
def get_world_group() -> GroupCoordinator:
|
||||
assert _WORLD is not None, "world group is not initialized"
|
||||
@@ -227,14 +173,7 @@ def init_parallel_group_coordinator(
|
||||
"replica",
|
||||
"encoder_data",
|
||||
], f"parallel_mode {parallel_mode} is not supported"
|
||||
if parallel_mode == "pipeline":
|
||||
return PipelineGroupCoordinator(
|
||||
group_ranks=group_ranks,
|
||||
local_rank=local_rank,
|
||||
torch_distributed_backend=backend,
|
||||
group_name="pp_group",
|
||||
)
|
||||
elif parallel_mode == "sequence":
|
||||
if parallel_mode == "sequence":
|
||||
return SequenceParallelGroupCoordinator(
|
||||
group_ranks=group_ranks,
|
||||
local_rank=local_rank,
|
||||
@@ -244,6 +183,7 @@ def init_parallel_group_coordinator(
|
||||
)
|
||||
else:
|
||||
group_name = {
|
||||
"pipeline": "pp_group",
|
||||
"tensor": "tp_group",
|
||||
"vae_decode": "vae_decode_group",
|
||||
"replica": "replica_group",
|
||||
@@ -808,96 +748,6 @@ def cleanup_dist_env_and_memory(shutdown_ray: bool = False):
|
||||
ray.shutdown()
|
||||
|
||||
|
||||
def is_the_same_node_as(
|
||||
pg: ProcessGroup | StatelessProcessGroup, source_rank: int = 0
|
||||
) -> list[int]:
|
||||
"""
|
||||
This is a collective operation that returns if each rank is in the same node
|
||||
as the source rank. It tests if processes are attached to the same
|
||||
memory system (shared access to shared memory).
|
||||
"""
|
||||
if isinstance(pg, ProcessGroup):
|
||||
assert torch.distributed.get_backend(pg) != torch.distributed.Backend.NCCL, (
|
||||
"in_the_same_node_as should be tested with a non-NCCL group."
|
||||
)
|
||||
# local rank inside the group
|
||||
rank = torch.distributed.get_rank(group=pg)
|
||||
world_size = torch.distributed.get_world_size(group=pg)
|
||||
|
||||
# global ranks of the processes in the group
|
||||
ranks = torch.distributed.get_process_group_ranks(pg)
|
||||
else:
|
||||
rank = pg.rank
|
||||
world_size = pg.world_size
|
||||
ranks = list(range(world_size))
|
||||
|
||||
# local tensor in each process to store the result
|
||||
is_in_the_same_node = torch.tensor([0] * world_size, dtype=torch.int32)
|
||||
|
||||
magic_message = b"magic_message"
|
||||
shm = None
|
||||
|
||||
try:
|
||||
with contextlib.suppress(OSError):
|
||||
if rank == source_rank:
|
||||
# create a shared memory segment
|
||||
shm = shared_memory.SharedMemory(create=True, size=128)
|
||||
shm.buf[: len(magic_message)] = magic_message
|
||||
if isinstance(pg, ProcessGroup):
|
||||
torch.distributed.broadcast_object_list(
|
||||
[shm.name], src=ranks[source_rank], group=pg
|
||||
)
|
||||
else:
|
||||
pg.broadcast_obj(shm.name, src=source_rank)
|
||||
is_in_the_same_node[rank] = 1
|
||||
else:
|
||||
# try to open the shared memory segment
|
||||
if isinstance(pg, ProcessGroup):
|
||||
recv = [None]
|
||||
torch.distributed.broadcast_object_list(
|
||||
recv, src=ranks[source_rank], group=pg
|
||||
)
|
||||
name = recv[0]
|
||||
else:
|
||||
name = pg.broadcast_obj(None, src=source_rank)
|
||||
# fix to https://stackoverflow.com/q/62748654/9191338
|
||||
# Python incorrectly tracks shared memory even if it is not
|
||||
# created by the process. The following patch is a workaround.
|
||||
with patch(
|
||||
"multiprocessing.resource_tracker.register",
|
||||
lambda *args, **kwargs: None,
|
||||
):
|
||||
shm = shared_memory.SharedMemory(name=name)
|
||||
if shm.buf[: len(magic_message)] == magic_message:
|
||||
is_in_the_same_node[rank] = 1
|
||||
except Exception as e:
|
||||
logger.error("Error ignored in is_in_the_same_node: %s", e)
|
||||
finally:
|
||||
if shm:
|
||||
shm.close()
|
||||
|
||||
if isinstance(pg, ProcessGroup):
|
||||
torch.distributed.barrier(group=pg)
|
||||
else:
|
||||
pg.barrier()
|
||||
|
||||
# clean up the shared memory segment
|
||||
with contextlib.suppress(OSError):
|
||||
if rank == source_rank and shm:
|
||||
shm.unlink()
|
||||
|
||||
if isinstance(pg, ProcessGroup):
|
||||
torch.distributed.all_reduce(is_in_the_same_node, group=pg)
|
||||
aggregated_data = is_in_the_same_node
|
||||
else:
|
||||
aggregated_data = torch.zeros_like(is_in_the_same_node)
|
||||
for i in range(world_size):
|
||||
rank_data = pg.broadcast_obj(is_in_the_same_node, src=i)
|
||||
aggregated_data += rank_data
|
||||
|
||||
return [x == 1 for x in aggregated_data.tolist()]
|
||||
|
||||
|
||||
def get_tensor_model_parallel_world_size() -> int:
|
||||
"""Return world size for the tensor model parallel group."""
|
||||
return get_tp_world_size()
|
||||
@@ -950,7 +800,7 @@ def get_ring_ctx() -> tuple[int, int]:
|
||||
|
||||
|
||||
# PP
|
||||
def get_pp_group() -> PipelineGroupCoordinator:
|
||||
def get_pp_group() -> GroupCoordinator:
|
||||
assert _PP is not None, "pipeline model parallel group is not initialized"
|
||||
return _PP
|
||||
|
||||
|
||||
@@ -3,72 +3,8 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import argparse
|
||||
import os
|
||||
import shlex
|
||||
import subprocess
|
||||
import sys
|
||||
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class RaiseNotImplementedAction(argparse.Action):
|
||||
def __call__(self, parser, namespace, values, option_string=None):
|
||||
raise NotImplementedError(f"The {option_string} option is not yet implemented")
|
||||
|
||||
|
||||
def launch_distributed(
|
||||
num_gpus: int, args: list[str], master_port: int | None = None
|
||||
) -> int:
|
||||
"""
|
||||
Launch a distributed job with the given arguments
|
||||
|
||||
Args:
|
||||
num_gpus: Number of GPUs to use
|
||||
args: Arguments to pass to v1_sgl_diffusion_inference.py (defaults to sys.argv[1:])
|
||||
master_port: Port for the master process (default: random)
|
||||
"""
|
||||
|
||||
current_env = os.environ.copy()
|
||||
python_executable = sys.executable
|
||||
project_root = os.path.abspath(
|
||||
os.path.join(os.path.dirname(__file__), "../../../..")
|
||||
)
|
||||
main_script = os.path.join(
|
||||
project_root, "sgl_diffusion/sample/v1_sgl_diffusion_inference.py"
|
||||
)
|
||||
|
||||
cmd = [
|
||||
python_executable,
|
||||
"-m",
|
||||
"torch.distributed.run",
|
||||
f"--nproc_per_node={num_gpus}",
|
||||
]
|
||||
|
||||
if master_port is not None:
|
||||
cmd.append(f"--master_port={master_port}")
|
||||
|
||||
cmd.append(main_script)
|
||||
cmd.extend(args)
|
||||
|
||||
logger.info("Running inference with %d GPU(s)", num_gpus)
|
||||
logger.info("Launching command: %s", shlex.join(cmd))
|
||||
|
||||
current_env["PYTHONIOENCODING"] = "utf-8"
|
||||
process = subprocess.Popen(
|
||||
cmd,
|
||||
env=current_env,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.STDOUT,
|
||||
universal_newlines=True,
|
||||
bufsize=1,
|
||||
encoding="utf-8",
|
||||
errors="replace",
|
||||
)
|
||||
|
||||
if process.stdout:
|
||||
for line in iter(process.stdout.readline, ""):
|
||||
print(line.strip())
|
||||
|
||||
return process.wait()
|
||||
|
||||
@@ -150,24 +150,7 @@ def launch_server(server_args: ServerArgs, launch_http_server: bool = True):
|
||||
rank_offset = node_rank * local_num_gpus
|
||||
processes = []
|
||||
|
||||
# Pipes for master to talk to slaves (local to this node)
|
||||
task_pipes_to_slaves_w = []
|
||||
task_pipes_to_slaves_r = []
|
||||
for _ in range(local_num_gpus - 1):
|
||||
r, w = mp.Pipe(duplex=False)
|
||||
task_pipes_to_slaves_r.append(r)
|
||||
task_pipes_to_slaves_w.append(w)
|
||||
|
||||
# Pipes for slaves to talk to master (local to this node)
|
||||
result_pipes_from_slaves_w = []
|
||||
result_pipes_from_slaves_r = []
|
||||
for _ in range(local_num_gpus - 1):
|
||||
r, w = mp.Pipe(duplex=False)
|
||||
result_pipes_from_slaves_r.append(r)
|
||||
result_pipes_from_slaves_w.append(w)
|
||||
|
||||
# Launch this node's local worker processes
|
||||
master_port = server_args.master_port
|
||||
# Launch this node's local worker processes.
|
||||
scheduler_pipe_readers = []
|
||||
scheduler_pipe_writers = []
|
||||
|
||||
@@ -175,40 +158,12 @@ def launch_server(server_args: ServerArgs, launch_http_server: bool = True):
|
||||
rank = rank_offset + i
|
||||
reader, writer = mp.Pipe(duplex=False)
|
||||
scheduler_pipe_writers.append(writer)
|
||||
if i == 0: # This node's local pipe master
|
||||
process = mp.Process(
|
||||
target=run_scheduler_process,
|
||||
args=(
|
||||
i, # local_rank
|
||||
rank,
|
||||
master_port,
|
||||
server_args,
|
||||
writer,
|
||||
None, # No task pipe to read from master
|
||||
None, # No result pipe to write to master
|
||||
task_pipes_to_slaves_w,
|
||||
result_pipes_from_slaves_r,
|
||||
),
|
||||
name=f"sglang-diffusionWorker-{rank}",
|
||||
daemon=True,
|
||||
)
|
||||
else: # Slave workers
|
||||
process = mp.Process(
|
||||
target=run_scheduler_process,
|
||||
args=(
|
||||
i, # local_rank
|
||||
rank,
|
||||
master_port,
|
||||
server_args,
|
||||
writer,
|
||||
None, # No task pipe to read from master
|
||||
None, # No result pipe to write to master
|
||||
task_pipes_to_slaves_r[i - 1],
|
||||
result_pipes_from_slaves_w[i - 1],
|
||||
),
|
||||
name=f"sglang-diffusionWorker-{rank}",
|
||||
daemon=True,
|
||||
)
|
||||
process = mp.Process(
|
||||
target=run_scheduler_process,
|
||||
args=(i, rank, server_args, writer),
|
||||
name=f"sglang-diffusionWorker-{rank}",
|
||||
daemon=True,
|
||||
)
|
||||
scheduler_pipe_readers.append(reader)
|
||||
process.start()
|
||||
processes.append(process)
|
||||
@@ -218,16 +173,6 @@ def launch_server(server_args: ServerArgs, launch_http_server: bool = True):
|
||||
for writer in scheduler_pipe_writers:
|
||||
writer.close()
|
||||
|
||||
# Close unused pipe ends in parent process
|
||||
for p in task_pipes_to_slaves_w:
|
||||
p.close()
|
||||
for p in task_pipes_to_slaves_r:
|
||||
p.close()
|
||||
for p in result_pipes_from_slaves_w:
|
||||
p.close()
|
||||
for p in result_pipes_from_slaves_r:
|
||||
p.close()
|
||||
|
||||
for i, reader in enumerate(scheduler_pipe_readers):
|
||||
try:
|
||||
data = reader.recv()
|
||||
@@ -428,7 +373,7 @@ def launch_pool_disagg_server(
|
||||
|
||||
process = pool_ctx.Process(
|
||||
target=_run_disagg_role_process,
|
||||
args=(gpu_id, rank_idx, rank_idx, role_args, writer, [], []),
|
||||
args=(gpu_id, rank_idx, role_args, writer),
|
||||
name=f"sglang-pool-{role_type.value}-{inst_idx}-r{rank_idx}",
|
||||
daemon=True,
|
||||
)
|
||||
@@ -502,12 +447,9 @@ def launch_pool_disagg_server(
|
||||
|
||||
def _run_disagg_role_process(
|
||||
gpu_id: int,
|
||||
_local_rank: int,
|
||||
rank: int,
|
||||
server_args: ServerArgs,
|
||||
pipe_writer: mp.connection.Connection,
|
||||
task_pipes: list,
|
||||
result_pipes: list,
|
||||
):
|
||||
"""Entry point for a disagg role process.
|
||||
|
||||
@@ -519,13 +461,8 @@ def _run_disagg_role_process(
|
||||
run_scheduler_process(
|
||||
local_rank=gpu_id,
|
||||
rank=rank,
|
||||
master_port=server_args.master_port,
|
||||
server_args=server_args,
|
||||
pipe_writer=pipe_writer,
|
||||
task_pipe_r=None,
|
||||
result_pipe_w=None,
|
||||
task_pipes_to_slaves=task_pipes,
|
||||
result_pipes_from_slaves=result_pipes,
|
||||
)
|
||||
|
||||
|
||||
@@ -767,7 +704,7 @@ def launch_disagg_role(server_args: ServerArgs):
|
||||
|
||||
process = pool_ctx.Process(
|
||||
target=_run_disagg_role_process,
|
||||
args=(gpu_id, rank_idx, rank_idx, role_args, writer, [], []),
|
||||
args=(gpu_id, rank_idx, role_args, writer),
|
||||
name=f"sglang-{role_type.value}-r{rank_idx}",
|
||||
daemon=True,
|
||||
)
|
||||
|
||||
+2
-42
@@ -4,11 +4,8 @@
|
||||
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/attention/backends/abstract.py
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass, fields
|
||||
from typing import TYPE_CHECKING, Any, Generic, Protocol, TypeVar
|
||||
|
||||
if TYPE_CHECKING:
|
||||
pass
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Generic, TypeVar
|
||||
|
||||
import torch
|
||||
|
||||
@@ -78,15 +75,6 @@ class AttentionBackend(ABC):
|
||||
def get_metadata_cls() -> type["AttentionMetadata"]:
|
||||
raise NotImplementedError
|
||||
|
||||
# @staticmethod
|
||||
# @abstractmethod
|
||||
# def get_state_cls() -> Type["AttentionState"]:
|
||||
# raise NotImplementedError
|
||||
|
||||
# @classmethod
|
||||
# def make_metadata(cls, *args, **kwargs) -> "AttentionMetadata":
|
||||
# return cls.get_metadata_cls()(*args, **kwargs)
|
||||
|
||||
@staticmethod
|
||||
@abstractmethod
|
||||
def get_builder_cls() -> type["AttentionMetadataBuilder"]:
|
||||
@@ -100,18 +88,6 @@ class AttentionMetadata:
|
||||
# Current step of diffusion process
|
||||
current_timestep: int
|
||||
|
||||
def asdict_zerocopy(self, skip_fields: set[str] | None = None) -> dict[str, Any]:
|
||||
"""Similar to dataclasses.asdict, but avoids deepcopying."""
|
||||
if skip_fields is None:
|
||||
skip_fields = set()
|
||||
# Note that if we add dataclasses as fields, they will need
|
||||
# similar handling.
|
||||
return {
|
||||
field.name: getattr(self, field.name)
|
||||
for field in fields(self)
|
||||
if field.name not in skip_fields
|
||||
}
|
||||
|
||||
|
||||
T = TypeVar("T", bound=AttentionMetadata)
|
||||
|
||||
@@ -138,22 +114,6 @@ class AttentionMetadataBuilder(ABC, Generic[T]):
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class AttentionLayer(Protocol):
|
||||
_k_scale: torch.Tensor
|
||||
_v_scale: torch.Tensor
|
||||
_k_scale_float: float
|
||||
_v_scale_float: float
|
||||
|
||||
def forward(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
kv_cache: torch.Tensor,
|
||||
attn_metadata: AttentionMetadata,
|
||||
) -> torch.Tensor: ...
|
||||
|
||||
|
||||
class AttentionImpl(ABC, Generic[T]):
|
||||
@abstractmethod
|
||||
def __init__(
|
||||
|
||||
@@ -3,7 +3,6 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/attention/selector.py
|
||||
|
||||
import os
|
||||
from collections.abc import Generator
|
||||
from contextlib import contextmanager
|
||||
from contextvars import ContextVar
|
||||
@@ -23,47 +22,10 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
STR_BACKEND_ENV_VAR = "SGLANG_DIFFUSION_ATTENTION_BACKEND"
|
||||
|
||||
|
||||
def backend_name_to_enum(backend_name: str) -> AttentionBackendEnum | None:
|
||||
"""
|
||||
Convert a string backend name to a _Backend enum value.
|
||||
|
||||
Returns:
|
||||
* _Backend: enum value if backend_name is a valid in-tree type
|
||||
* None: otherwise it's an invalid in-tree type or an out-of-tree platform is
|
||||
loaded.
|
||||
"""
|
||||
assert backend_name is not None
|
||||
return (
|
||||
AttentionBackendEnum[backend_name]
|
||||
if backend_name in AttentionBackendEnum.__members__
|
||||
else None
|
||||
)
|
||||
|
||||
|
||||
def get_env_variable_attn_backend() -> AttentionBackendEnum | None:
|
||||
"""
|
||||
Get the backend override specified by the sglang-diffusion attention
|
||||
backend environment variable, if one is specified.
|
||||
|
||||
Returns:
|
||||
|
||||
* _Backend enum value if an override is specified
|
||||
* None otherwise
|
||||
"""
|
||||
backend_name = os.environ.get(STR_BACKEND_ENV_VAR)
|
||||
return None if backend_name is None else backend_name_to_enum(backend_name)
|
||||
|
||||
|
||||
# Global state allows a particular choice of backend
|
||||
# to be forced, overriding the logic which auto-selects
|
||||
# a backend based on system & workload configuration
|
||||
# (default behavior if this variable is None)
|
||||
#
|
||||
# THIS SELECTION TAKES PRECEDENCE OVER THE
|
||||
# FASTVIDEO ATTENTION BACKEND ENVIRONMENT VARIABLE
|
||||
forced_attn_backend: AttentionBackendEnum | None = None
|
||||
|
||||
|
||||
@@ -150,13 +112,6 @@ def _record_component_attn_backend(backend_name: str, reason: str | None) -> boo
|
||||
return True
|
||||
|
||||
|
||||
def record_component_attn_backend(
|
||||
backend: AttentionBackendEnum, reason: str | None = None
|
||||
) -> bool:
|
||||
"""Record a component backend selected outside layer construction."""
|
||||
return _record_component_attn_backend(backend.name.lower(), reason)
|
||||
|
||||
|
||||
def _log_component_attn_backend_summary(
|
||||
context: ComponentAttnBackendContext | None,
|
||||
) -> None:
|
||||
|
||||
@@ -380,20 +380,6 @@ class LayerNorm(CustomOp):
|
||||
else:
|
||||
self.register_parameter("weight", None)
|
||||
self.register_parameter("bias", None)
|
||||
# Lazy cache for ones vector (not a registered buffer to avoid FSDP/meta issues)
|
||||
self._weight_fallback_cache = None
|
||||
|
||||
def _get_weight_fallback(self, x: torch.Tensor) -> torch.Tensor:
|
||||
wf = getattr(self, "_weight_fallback_cache", None)
|
||||
if (
|
||||
wf is None
|
||||
or wf.device != x.device
|
||||
or wf.dtype != x.dtype
|
||||
or wf.numel() != self.hidden_size
|
||||
):
|
||||
wf = torch.ones(self.hidden_size, device=x.device, dtype=x.dtype)
|
||||
self._weight_fallback_cache = wf
|
||||
return wf
|
||||
|
||||
def forward_triton(self, x: torch.Tensor):
|
||||
# Fast inference kernel without residual/dropout branches
|
||||
|
||||
@@ -32,23 +32,6 @@ def get_group_rank(group) -> int:
|
||||
raise ValueError(f"Unsupported group type: {type(group)}")
|
||||
|
||||
|
||||
def get_token_bin_counts_and_mask(
|
||||
tokens: torch.Tensor,
|
||||
vocab_size: int,
|
||||
num_seqs: int,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
# Compute the bin counts for the tokens.
|
||||
# vocab_size + 1 for padding.
|
||||
bin_counts = torch.zeros(
|
||||
(num_seqs, vocab_size + 1), dtype=torch.long, device=tokens.device
|
||||
)
|
||||
bin_counts.scatter_add_(1, tokens, torch.ones_like(tokens))
|
||||
bin_counts = bin_counts[:, :vocab_size]
|
||||
mask = bin_counts > 0
|
||||
|
||||
return bin_counts, mask
|
||||
|
||||
|
||||
sglang_lib = Library("sglang", "FRAGMENT") # noqa
|
||||
|
||||
|
||||
|
||||
@@ -53,10 +53,8 @@ from sglang.multimodal_gen.runtime.entrypoints.utils import (
|
||||
save_outputs,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.managers.memory_managers.auto_residency import (
|
||||
DefaultWorkload,
|
||||
WarmupMemoryRecord,
|
||||
estimate_default_workload_peak_bytes,
|
||||
resolve_default_workload,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager import (
|
||||
get_global_component_residency_manager,
|
||||
@@ -244,23 +242,6 @@ class GPUWorker(GPUWorkerPostTrainingMixin):
|
||||
# per-rank memory measurements of server warmup forwards; consumed by
|
||||
# the auto-residency placement decision before the server turns ready
|
||||
self._auto_residency_warmup_records: list[WarmupMemoryRecord] = []
|
||||
# default workload resolved once for the per-request residency hint
|
||||
self._cached_default_workload: DefaultWorkload | None = None
|
||||
self._cached_default_workload_failed = False
|
||||
|
||||
def _default_workload_for_hint(self) -> DefaultWorkload | None:
|
||||
if (
|
||||
self._cached_default_workload is None
|
||||
and not self._cached_default_workload_failed
|
||||
):
|
||||
try:
|
||||
self._cached_default_workload = resolve_default_workload(
|
||||
self.server_args
|
||||
)
|
||||
except Exception:
|
||||
logger.debug("Default workload unresolvable", exc_info=True)
|
||||
self._cached_default_workload_failed = True
|
||||
return self._cached_default_workload
|
||||
|
||||
def release_realtime_session(self, session_id: str) -> OutputBatch:
|
||||
"""release the session of a realtime connection"""
|
||||
@@ -1553,23 +1534,10 @@ def _oom_exceptions():
|
||||
def run_scheduler_process(
|
||||
local_rank: int,
|
||||
rank: int,
|
||||
master_port: int,
|
||||
server_args: ServerArgs,
|
||||
pipe_writer: mp.connection.Connection,
|
||||
# For all workers: pipe to receive tasks from rank 0
|
||||
task_pipe_r: mp.connection.Connection,
|
||||
# For slave workers: pipe to send results back to rank 0
|
||||
result_pipe_w: mp.connection.Connection | None,
|
||||
# For rank 0 worker only: pipes to send tasks to slaves
|
||||
task_pipes_to_slaves: list[mp.connection.Connection] | None = None,
|
||||
# For rank 0 worker only: pipes to receive results from slaves
|
||||
result_pipes_from_slaves: list[mp.connection.Connection] | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
The entry point for the worker process.
|
||||
Rank 0 acts as the master, handling ZMQ requests and coordinating slaves.
|
||||
Ranks > 0 act as slaves, waiting for tasks from the master.
|
||||
"""
|
||||
"""Run a rank's scheduler and report readiness to the launching process."""
|
||||
kill_itself_when_parent_died()
|
||||
configure_logger(server_args)
|
||||
globally_suppress_loggers()
|
||||
@@ -1583,8 +1551,6 @@ def run_scheduler_process(
|
||||
port_args = PortArgs.from_server_args(server_args)
|
||||
|
||||
# start the scheduler event loop
|
||||
assert task_pipes_to_slaves is not None
|
||||
assert result_pipes_from_slaves is not None
|
||||
from sglang.multimodal_gen.runtime.managers.scheduler import Scheduler
|
||||
|
||||
try:
|
||||
@@ -1592,8 +1558,6 @@ def run_scheduler_process(
|
||||
server_args,
|
||||
gpu_id=rank,
|
||||
port_args=port_args,
|
||||
task_pipes_to_slaves=task_pipes_to_slaves,
|
||||
result_pipes_from_slaves=result_pipes_from_slaves,
|
||||
local_rank=local_rank,
|
||||
)
|
||||
logger.info(f"Worker {rank}: Scheduler loop started.")
|
||||
|
||||
@@ -90,8 +90,6 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag
|
||||
server_args: ServerArgs,
|
||||
gpu_id: int,
|
||||
port_args: PortArgs,
|
||||
task_pipes_to_slaves: list = None,
|
||||
result_pipes_from_slaves: list = None,
|
||||
local_rank: int | None = None,
|
||||
):
|
||||
self.server_args = server_args
|
||||
@@ -134,8 +132,6 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag
|
||||
server_args=server_args,
|
||||
)
|
||||
self.worker = worker
|
||||
self.task_pipes_to_slaves = task_pipes_to_slaves
|
||||
self.result_pipes_from_slaves = result_pipes_from_slaves
|
||||
self.gpu_id = gpu_id
|
||||
self._show_warmup_progress = gpu_id == 0
|
||||
self._running = True
|
||||
@@ -1300,21 +1296,6 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag
|
||||
self._cleanup_disagg()
|
||||
self.context.destroy(linger=0)
|
||||
|
||||
def _broadcast_task(self, payload: dict[str, Any]) -> None:
|
||||
"""Broadcast a task to all slave worker processes."""
|
||||
method = payload["method"]
|
||||
kwargs = {k: v for k, v in payload.items() if k != "method"}
|
||||
task = {"method": method, "kwargs": kwargs}
|
||||
for pipe in self.task_pipes_to_slaves:
|
||||
pipe.send(task)
|
||||
|
||||
def _collect_slave_results(self) -> List[dict[str, Any]]:
|
||||
"""Collect results from all slave worker processes."""
|
||||
results = []
|
||||
for pipe in self.result_pipes_from_slaves:
|
||||
results.append(pipe.recv())
|
||||
return results
|
||||
|
||||
def _handle_release_memory_occupation(self, _reqs: List[Any]) -> OutputBatch:
|
||||
logger.info(f"[SLEEP] handle_release_memory_occupation on rank={self.gpu_id}")
|
||||
return OutputBatch(output=self.worker.release_memory_occupation())
|
||||
|
||||
@@ -9,7 +9,6 @@ import torch
|
||||
import torch.nn as nn
|
||||
from torch.nn.attention.flex_attention import (
|
||||
BlockMask,
|
||||
create_block_mask,
|
||||
flex_attention,
|
||||
)
|
||||
|
||||
@@ -23,7 +22,6 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload im
|
||||
flex_attention = torch.compile(
|
||||
flex_attention, dynamic=False, mode="max-autotune-no-cudagraphs"
|
||||
)
|
||||
import torch.distributed as dist
|
||||
|
||||
from sglang.multimodal_gen.configs.models.dits import WanVideoConfig
|
||||
from sglang.multimodal_gen.configs.models.fsdp import is_block
|
||||
@@ -534,79 +532,6 @@ class CausalWanTransformer3DModel(BaseDiT, LayerwiseOffloadableModuleMixin):
|
||||
"blocks",
|
||||
]
|
||||
|
||||
@staticmethod
|
||||
def _prepare_blockwise_causal_attn_mask(
|
||||
device: torch.device | str,
|
||||
num_frames: int = 21,
|
||||
frame_seqlen: int = 1560,
|
||||
num_frame_per_block=1,
|
||||
local_attn_size=-1,
|
||||
) -> BlockMask:
|
||||
"""
|
||||
we will divide the token sequence into the following format
|
||||
[1 latent frame] [1 latent frame] ... [1 latent frame]
|
||||
We use flexattention to construct the attention mask
|
||||
"""
|
||||
total_length = num_frames * frame_seqlen
|
||||
|
||||
# we do right padding to get to a multiple of 128
|
||||
padded_length = math.ceil(total_length / 128) * 128 - total_length
|
||||
|
||||
ends = torch.zeros(
|
||||
total_length + padded_length, device=device, dtype=torch.long
|
||||
)
|
||||
|
||||
# Block-wise causal mask will attend to all elements that are before the end of the current chunk
|
||||
frame_indices = torch.arange(
|
||||
start=0,
|
||||
end=total_length,
|
||||
step=frame_seqlen * num_frame_per_block,
|
||||
device=device,
|
||||
)
|
||||
|
||||
for tmp in frame_indices:
|
||||
ends[tmp : tmp + frame_seqlen * num_frame_per_block] = (
|
||||
tmp + frame_seqlen * num_frame_per_block
|
||||
)
|
||||
|
||||
def attention_mask(b, h, q_idx, kv_idx):
|
||||
if local_attn_size == -1:
|
||||
return (kv_idx < ends[q_idx]) | (q_idx == kv_idx)
|
||||
else:
|
||||
return (
|
||||
(kv_idx < ends[q_idx])
|
||||
& (kv_idx >= (ends[q_idx] - local_attn_size * frame_seqlen))
|
||||
) | (q_idx == kv_idx)
|
||||
# return ((kv_idx < total_length) & (q_idx < total_length)) | (q_idx == kv_idx) # bidirectional mask
|
||||
|
||||
block_mask = create_block_mask(
|
||||
attention_mask,
|
||||
B=None,
|
||||
H=None,
|
||||
Q_LEN=total_length + padded_length,
|
||||
KV_LEN=total_length + padded_length,
|
||||
_compile=False,
|
||||
device=device,
|
||||
)
|
||||
|
||||
if not dist.is_initialized() or dist.get_rank() == 0:
|
||||
print(
|
||||
f" cache a block wise causal mask with block size of {num_frame_per_block} frames"
|
||||
)
|
||||
print(block_mask)
|
||||
|
||||
# import imageio
|
||||
# import numpy as np
|
||||
# from torch.nn.attention.flex_attention import create_mask
|
||||
|
||||
# mask = create_mask(attention_mask, B=None, H=None, Q_LEN=total_length +
|
||||
# padded_length, KV_LEN=total_length + padded_length, device=device)
|
||||
# import cv2
|
||||
# mask = cv2.resize(mask[0, 0].cpu().float().numpy(), (1024, 1024))
|
||||
# imageio.imwrite("mask_%d.jpg" % (0), np.uint8(255. * mask))
|
||||
|
||||
return block_mask
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
|
||||
@@ -235,9 +235,6 @@ class ComfyUIQwenImagePipelineBase(LoRAPipeline, ComposedPipelineBase):
|
||||
model = model_cls(**{"config": dit_config, "hf_config": hf_config})
|
||||
|
||||
use_fsdp = server_args.should_use_fsdp_for_component("transformer")
|
||||
component_starts_on_cpu = server_args.should_start_component_on_cpu(
|
||||
"transformer"
|
||||
)
|
||||
if current_platform.is_mps():
|
||||
use_fsdp = False
|
||||
logger.info("Disabling FSDP for MPS platform as it's not compatible")
|
||||
|
||||
-17
@@ -50,23 +50,6 @@ def _causal_block_count(batch: Req, server_args: ServerArgs) -> int:
|
||||
return latent_frames // block_size
|
||||
|
||||
|
||||
def expand_longlive2_shot_prompts(
|
||||
shot_prompts: list[str],
|
||||
*,
|
||||
num_blocks: int,
|
||||
shot_durations: list[int] | None = None,
|
||||
chunks_per_shot: int = 0,
|
||||
scene_cut_prefix: str = LONG_LIVE2_DEFAULT_SCENE_CUT_PREFIX,
|
||||
) -> list[str]:
|
||||
return expand_causal_block_prompts(
|
||||
shot_prompts,
|
||||
num_blocks=num_blocks,
|
||||
shot_durations=shot_durations,
|
||||
chunks_per_shot=chunks_per_shot,
|
||||
scene_cut_prefix=scene_cut_prefix,
|
||||
)[0]
|
||||
|
||||
|
||||
class LongLive2TextEncodingStage(TextEncodingStage):
|
||||
def build_dedup_fingerprint(self, batch: Req, server_args: ServerArgs):
|
||||
base = super().build_dedup_fingerprint(batch, server_args)
|
||||
|
||||
-142
@@ -12,11 +12,7 @@ from sglang.multimodal_gen.runtime.distributed import (
|
||||
get_local_torch_device,
|
||||
get_sp_world_size,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.distributed.cfg_parallel_utils import (
|
||||
dispatch_branches,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.distributed.communication_op import (
|
||||
cfg_model_parallel_all_gather,
|
||||
cfg_model_parallel_all_reduce,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.distributed.parallel_state import (
|
||||
@@ -256,144 +252,6 @@ class LTX2DenoisingStage(DenoisingStage):
|
||||
cfg_model_parallel_all_reduce(audio_partial),
|
||||
)
|
||||
|
||||
def _run_legacy_one_stage_multi_branch_cfg_parallel(
|
||||
self,
|
||||
*,
|
||||
base_model_kwargs: dict[str, object],
|
||||
ctx: "LTX2DenoisingContext",
|
||||
step: "DenoisingStepState",
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
audio_encoder_hidden_states: torch.Tensor,
|
||||
encoder_attention_mask: torch.Tensor | None,
|
||||
negative_encoder_hidden_states: torch.Tensor,
|
||||
negative_audio_encoder_hidden_states: torch.Tensor,
|
||||
negative_encoder_attention_mask: torch.Tensor | None,
|
||||
need_perturbed: bool,
|
||||
need_modality: bool,
|
||||
stage1_guider_params: dict[str, object],
|
||||
) -> dict[str, tuple[torch.Tensor, torch.Tensor]]:
|
||||
"""Multi-branch CFG parallel for the legacy LTX-2.3 one-stage path.
|
||||
|
||||
Distributes up to 4 forward passes (cond, neg, perturbed, modality)
|
||||
across CFG ranks via round-robin. Each rank runs only its assigned
|
||||
passes, then an all-gather collects every output so all ranks can
|
||||
compute the guidance combination locally.
|
||||
"""
|
||||
cfg_rank = get_classifier_free_guidance_rank()
|
||||
cfg_world_size = get_classifier_free_guidance_world_size()
|
||||
|
||||
# Build kwargs for every pass in canonical order.
|
||||
all_passes: list[tuple[str, dict[str, object]]] = [
|
||||
(
|
||||
"cond",
|
||||
self._build_ltx2_model_kwargs(
|
||||
ctx,
|
||||
base_model_kwargs,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
audio_encoder_hidden_states=audio_encoder_hidden_states,
|
||||
encoder_attention_mask=encoder_attention_mask,
|
||||
),
|
||||
),
|
||||
(
|
||||
"neg",
|
||||
self._build_ltx2_model_kwargs(
|
||||
ctx,
|
||||
base_model_kwargs,
|
||||
encoder_hidden_states=negative_encoder_hidden_states,
|
||||
audio_encoder_hidden_states=negative_audio_encoder_hidden_states,
|
||||
encoder_attention_mask=negative_encoder_attention_mask,
|
||||
),
|
||||
),
|
||||
]
|
||||
if need_perturbed:
|
||||
all_passes.append(
|
||||
(
|
||||
"perturbed",
|
||||
self._build_ltx2_model_kwargs(
|
||||
ctx,
|
||||
base_model_kwargs,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
audio_encoder_hidden_states=audio_encoder_hidden_states,
|
||||
encoder_attention_mask=encoder_attention_mask,
|
||||
skip_video_self_attn_blocks=tuple(
|
||||
stage1_guider_params["video_stg_blocks"]
|
||||
),
|
||||
skip_audio_self_attn_blocks=tuple(
|
||||
stage1_guider_params["audio_stg_blocks"]
|
||||
),
|
||||
),
|
||||
)
|
||||
)
|
||||
if need_modality:
|
||||
all_passes.append(
|
||||
(
|
||||
"modality",
|
||||
self._build_ltx2_model_kwargs(
|
||||
ctx,
|
||||
base_model_kwargs,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
audio_encoder_hidden_states=audio_encoder_hidden_states,
|
||||
encoder_attention_mask=encoder_attention_mask,
|
||||
disable_a2v_cross_attn=True,
|
||||
disable_v2a_cross_attn=True,
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
pass_names = [name for name, _ in all_passes]
|
||||
n_passes = len(pass_names)
|
||||
assignments = dispatch_branches(n_passes, cfg_world_size)
|
||||
my_indices = assignments[cfg_rank]
|
||||
max_local = max(len(a) for a in assignments)
|
||||
|
||||
local_videos: list[torch.Tensor] = []
|
||||
local_audios: list[torch.Tensor] = []
|
||||
|
||||
indices_to_run = my_indices if my_indices else [0]
|
||||
with set_forward_context(
|
||||
current_timestep=step.step_index, attn_metadata=step.attn_metadata
|
||||
):
|
||||
for idx in indices_to_run:
|
||||
_, kwargs = all_passes[idx]
|
||||
v, a = step.current_model(**kwargs)
|
||||
local_videos.append(v.float())
|
||||
local_audios.append(a.float())
|
||||
|
||||
if not my_indices:
|
||||
# This rank has no real branch, but it still needs tensor shapes for all-gather.
|
||||
# The dummy branch above provides the shapes; zeros keep this rank from contributing.
|
||||
local_videos = [torch.zeros_like(local_videos[0])]
|
||||
local_audios = [torch.zeros_like(local_audios[0])]
|
||||
|
||||
# Pad to max_local for unbalanced cases (n_passes not divisible by n_ranks).
|
||||
while len(local_videos) < max_local:
|
||||
local_videos.append(torch.zeros_like(local_videos[0]))
|
||||
local_audios.append(torch.zeros_like(local_audios[0]))
|
||||
|
||||
# Stack -> [max_local, B, ...], flatten to [max_local*B, ...] for all-gather.
|
||||
local_v = torch.stack(local_videos, dim=0)
|
||||
local_a = torch.stack(local_audios, dim=0)
|
||||
B = local_v.shape[1]
|
||||
local_v_flat = local_v.reshape(max_local * B, *local_v.shape[2:])
|
||||
local_a_flat = local_a.reshape(max_local * B, *local_a.shape[2:])
|
||||
|
||||
# All-gather along batch dim -> [cfg_world_size * max_local * B, ...].
|
||||
all_v_flat = cfg_model_parallel_all_gather(local_v_flat, dim=0)
|
||||
all_a_flat = cfg_model_parallel_all_gather(local_a_flat, dim=0)
|
||||
|
||||
# Reshape to [cfg_world_size, max_local, B, ...].
|
||||
all_v = all_v_flat.reshape(cfg_world_size, max_local, B, *all_v_flat.shape[1:])
|
||||
all_a = all_a_flat.reshape(cfg_world_size, max_local, B, *all_a_flat.shape[1:])
|
||||
|
||||
# Branch i was run by rank (i % cfg_world_size) at slot (i // cfg_world_size).
|
||||
return {
|
||||
name: (
|
||||
all_v[i % cfg_world_size, i // cfg_world_size],
|
||||
all_a[i % cfg_world_size, i // cfg_world_size],
|
||||
)
|
||||
for i, name in enumerate(pass_names)
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _get_video_latent_num_frames_for_model(
|
||||
batch: Req, server_args: ServerArgs, latents: torch.Tensor
|
||||
|
||||
@@ -114,38 +114,6 @@ def format_tcp_endpoint(host: str, port: int, field_name: str) -> str:
|
||||
return f"tcp://{host}:{port}"
|
||||
|
||||
|
||||
def configure_ipv6(dist_init_addr):
|
||||
addr = dist_init_addr
|
||||
end = addr.find("]")
|
||||
if end == -1:
|
||||
raise ValueError("invalid IPv6 address format: missing ']'")
|
||||
|
||||
host = addr[: end + 1]
|
||||
|
||||
# this only validates the address without brackets: we still need the below checks.
|
||||
# if it's invalid, immediately raise an error so we know it's not formatting issues.
|
||||
if not is_valid_ipv6_address(host[1:end]):
|
||||
raise ValueError(f"invalid IPv6 address: {host}")
|
||||
|
||||
port_str = None
|
||||
if len(addr) > end + 1:
|
||||
if addr[end + 1] == ":":
|
||||
port_str = addr[end + 2 :]
|
||||
else:
|
||||
raise ValueError("received IPv6 address format: expected ':' after ']'")
|
||||
|
||||
if not port_str:
|
||||
raise ValueError(
|
||||
"a port must be specified in IPv6 address (format: [ipv6]:port)"
|
||||
)
|
||||
|
||||
try:
|
||||
port = int(port_str)
|
||||
except ValueError:
|
||||
raise ValueError(f"invalid port in IPv6 address: '{port_str}'")
|
||||
return port, host
|
||||
|
||||
|
||||
def is_port_available(port):
|
||||
"""Return whether a port is available."""
|
||||
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
|
||||
|
||||
@@ -1,55 +0,0 @@
|
||||
import ast
|
||||
import re
|
||||
from pathlib import Path
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[5]
|
||||
REGISTRY_PATH = REPO_ROOT / "python/sglang/multimodal_gen/registry.py"
|
||||
CATALOG_PATH = REPO_ROOT / "docs/src/snippets/diffusion/model-catalog.jsx"
|
||||
|
||||
|
||||
def _registered_model_ids() -> set[str]:
|
||||
tree = ast.parse(REGISTRY_PATH.read_text())
|
||||
model_ids = set()
|
||||
|
||||
for node in ast.walk(tree):
|
||||
if not isinstance(node, ast.Call):
|
||||
continue
|
||||
if not isinstance(node.func, ast.Name) or node.func.id != "register_configs":
|
||||
continue
|
||||
|
||||
paths = next(
|
||||
(
|
||||
keyword.value
|
||||
for keyword in node.keywords
|
||||
if keyword.arg == "hf_model_paths"
|
||||
),
|
||||
None,
|
||||
)
|
||||
if not isinstance(paths, (ast.List, ast.Tuple)):
|
||||
continue
|
||||
|
||||
model_ids.update(
|
||||
item.value
|
||||
for item in paths.elts
|
||||
if isinstance(item, ast.Constant) and isinstance(item.value, str)
|
||||
)
|
||||
|
||||
return model_ids
|
||||
|
||||
|
||||
def _catalog_model_ids() -> set[str]:
|
||||
source = CATALOG_PATH.read_text()
|
||||
model_id_arrays = re.findall(r"modelIds:\s*\[(.*?)\]", source, re.DOTALL)
|
||||
return {
|
||||
model_id
|
||||
for model_id_array in model_id_arrays
|
||||
for model_id in re.findall(r'"([^"]+)"', model_id_array)
|
||||
}
|
||||
|
||||
|
||||
def test_explicit_registry_model_ids_are_documented():
|
||||
missing_model_ids = _registered_model_ids() - _catalog_model_ids()
|
||||
assert not missing_model_ids, (
|
||||
"Add newly registered model IDs to the Supported Models catalog: "
|
||||
f"{sorted(missing_model_ids)}"
|
||||
)
|
||||
@@ -47,35 +47,6 @@ def test_mlp_reuses_srt_activation_without_server_context(monkeypatch):
|
||||
assert isinstance(mlp.act_fn, SiluAndMul)
|
||||
|
||||
|
||||
def test_attention_keeps_diffusion_one_pass_qk_norm(monkeypatch):
|
||||
monkeypatch.setattr(qwen3, "get_tp_world_size", lambda: 1)
|
||||
monkeypatch.setattr(
|
||||
qwen3, "QKVParallelLinear", lambda **kwargs: torch.nn.Identity()
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
qwen3, "RowParallelLinear", lambda **kwargs: torch.nn.Identity()
|
||||
)
|
||||
monkeypatch.setattr(qwen3, "get_rope", lambda *args, **kwargs: torch.nn.Identity())
|
||||
monkeypatch.setattr(
|
||||
qwen3, "LocalAttention", lambda *args, **kwargs: torch.nn.Identity()
|
||||
)
|
||||
config = SimpleNamespace(
|
||||
head_dim=128,
|
||||
rms_norm_eps=1e-6,
|
||||
_supported_attention_backends=(),
|
||||
)
|
||||
|
||||
attention = qwen3.Qwen3Attention(
|
||||
config,
|
||||
hidden_size=256,
|
||||
num_heads=2,
|
||||
num_kv_heads=1,
|
||||
)
|
||||
|
||||
assert isinstance(attention.q_norm, qwen3.MMGenRMSNorm)
|
||||
assert isinstance(attention.k_norm, qwen3.MMGenRMSNorm)
|
||||
|
||||
|
||||
def test_default_position_ids_batch_shape():
|
||||
model = Qwen3ForCausalLM.__new__(Qwen3ForCausalLM)
|
||||
torch.nn.Module.__init__(model)
|
||||
|
||||
@@ -1,16 +1,63 @@
|
||||
"""Single-rank groups get gloo, so NCCL does not reserve device buffers for them."""
|
||||
|
||||
import unittest
|
||||
from datetime import timedelta
|
||||
from pathlib import Path
|
||||
from tempfile import TemporaryDirectory
|
||||
from unittest.mock import patch
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import torch.multiprocessing as mp
|
||||
|
||||
from sglang.multimodal_gen.runtime.distributed.group_coordinator import (
|
||||
new_device_group,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.distributed.parallel_state import (
|
||||
init_parallel_group_coordinator,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.platforms.cpu import CpuPlatform
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
NEW_GROUP_PATH = "torch.distributed.new_group"
|
||||
|
||||
|
||||
class TestSingleRankDeviceGroup(unittest.TestCase):
|
||||
def _check_pipeline_group_lifecycle(_rank):
|
||||
with TemporaryDirectory() as directory:
|
||||
dist.init_process_group(
|
||||
"gloo",
|
||||
init_method=Path(directory, "rendezvous").as_uri(),
|
||||
rank=0,
|
||||
world_size=1,
|
||||
timeout=timedelta(seconds=30),
|
||||
)
|
||||
try:
|
||||
process_groups = dist.distributed_c10d._world.pg_map
|
||||
initial_groups = len(process_groups)
|
||||
with patch(
|
||||
"sglang.multimodal_gen.runtime.distributed."
|
||||
"group_coordinator.current_platform",
|
||||
CpuPlatform(),
|
||||
):
|
||||
for _ in range(2):
|
||||
group = init_parallel_group_coordinator(
|
||||
[[0]],
|
||||
local_rank=0,
|
||||
backend="gloo",
|
||||
parallel_mode="pipeline",
|
||||
)
|
||||
try:
|
||||
assert (group.world_size, group.rank_in_group) == (1, 0)
|
||||
tensor = torch.tensor([3.0])
|
||||
torch.testing.assert_close(group.all_reduce(tensor), tensor)
|
||||
finally:
|
||||
group.destroy()
|
||||
assert len(process_groups) == initial_groups
|
||||
finally:
|
||||
dist.destroy_process_group()
|
||||
|
||||
|
||||
class TestSingleRankDeviceGroup(CustomTestCase):
|
||||
def test_single_rank_group_avoids_the_device_backend(self):
|
||||
for ranks, requested in [([0], "nccl"), ([3], "hccl"), ([0], None)]:
|
||||
with self.subTest(ranks=ranks, requested=requested):
|
||||
@@ -25,6 +72,10 @@ class TestSingleRankDeviceGroup(unittest.TestCase):
|
||||
new_device_group(ranks, requested)
|
||||
new_group.assert_called_once_with(ranks, backend=requested)
|
||||
|
||||
def test_pipeline_group_releases_process_groups(self):
|
||||
# Other unit tests can leave a default distributed group initialized.
|
||||
mp.spawn(_check_pipeline_group_lifecycle, nprocs=1)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
@@ -112,23 +112,6 @@ class TestTransformerLoaderFallbackAdmission(unittest.TestCase):
|
||||
native.assert_called_once()
|
||||
server_args.should_use_fsdp_for_component.assert_called_with("transformer_2")
|
||||
|
||||
def test_parallel_execution_rejects_native_fallback(self):
|
||||
cases = (
|
||||
({"tp_size": 2}, "tp_size=2"),
|
||||
({"sp_degree": 2}, "sp_degree=2"),
|
||||
({"ulysses_degree": 2}, "ulysses_degree=2"),
|
||||
({"ring_degree": 2}, "ring_degree=2"),
|
||||
({"kv_gather_degree": 2}, "kv_gather_degree=2"),
|
||||
({"fsdp_requested": True}, "FSDP"),
|
||||
)
|
||||
|
||||
for overrides, expected_error in cases:
|
||||
with self.subTest(overrides=overrides):
|
||||
with self.assertRaisesRegex(RuntimeError, expected_error):
|
||||
TransformerLoader().validate_native_fallback(
|
||||
self._server_args(**overrides), "transformer_2"
|
||||
)
|
||||
|
||||
def test_unreadable_checkpoint_is_not_a_missing_implementation(self):
|
||||
# the native fallback answers "no customized implementation for this
|
||||
# architecture"; a checkpoint that cannot be read is a different failure,
|
||||
@@ -162,13 +145,6 @@ class TestTransformerLoaderFallbackAdmission(unittest.TestCase):
|
||||
self.assertIn("shard-00002.safetensors", str(caught.exception))
|
||||
native.assert_not_called()
|
||||
|
||||
def test_replicated_execution_keeps_native_fallback_available(self):
|
||||
self.assertIsNone(
|
||||
TransformerLoader().validate_native_fallback(
|
||||
self._server_args(), "transformer_2"
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
@@ -1,11 +1,8 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import ast
|
||||
import os
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from dataclasses import dataclass, field
|
||||
from importlib.util import resolve_name
|
||||
from pathlib import Path
|
||||
from threading import Barrier
|
||||
|
||||
import pytest
|
||||
@@ -24,32 +21,6 @@ from sglang.multimodal_gen.runtime.utils.precision import (
|
||||
)
|
||||
|
||||
|
||||
def test_models_do_not_import_pipeline_stages():
|
||||
root = Path(__file__).resolve().parents[2]
|
||||
violations = []
|
||||
for path in sorted((root / "runtime/models").rglob("*.py")):
|
||||
package = "sglang.multimodal_gen." + str(path.parent.relative_to(root)).replace(
|
||||
"/", "."
|
||||
)
|
||||
for node in ast.walk(ast.parse(path.read_text())):
|
||||
names = []
|
||||
if isinstance(node, ast.Import):
|
||||
names = [alias.name for alias in node.names]
|
||||
elif isinstance(node, ast.ImportFrom):
|
||||
prefix = node.module or ""
|
||||
if node.level:
|
||||
prefix = resolve_name("." * node.level + prefix, package)
|
||||
names = [prefix] + [f"{prefix}.{alias.name}" for alias in node.names]
|
||||
if any(
|
||||
name.startswith("sglang.multimodal_gen.runtime.pipelines_core.stages")
|
||||
for name in names
|
||||
):
|
||||
violations.append(f"{path.relative_to(root)}:{node.lineno}")
|
||||
assert not violations, "Models must not depend on pipeline stages: " + ", ".join(
|
||||
violations
|
||||
)
|
||||
|
||||
|
||||
def test_argument_parser_preserves_config_and_explicit_values(tmp_path):
|
||||
config = tmp_path / "config.yaml"
|
||||
config.write_text("num_gpus: 2\nuse_cache: true\n")
|
||||
|
||||
Reference in New Issue
Block a user