[Diffusion] Clean up obsolete worker plumbing, dead helpers, and tests (#39293)

This commit is contained in:
Xiaoyu Zhang
2026-09-16 12:11:45 +08:00
committed by GitHub
parent 954bb6804a
commit 444b29c932
20 changed files with 72 additions and 1336 deletions
@@ -15,7 +15,6 @@ from typing import Any, Dict, List, Optional, Tuple, Union
import torch import torch
import torch.distributed import torch.distributed
from torch.cuda import synchronize
from torch.distributed import Backend, ProcessGroup from torch.distributed import Backend, ProcessGroup
from sglang.multimodal_gen.runtime.distributed.device_communicators.base_device_communicator import ( 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 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__) logger = init_logger(__name__)
TensorMetadata = namedtuple("TensorMetadata", ["device", "dtype", "size"]) TensorMetadata = namedtuple("TensorMetadata", ["device", "dtype", "size"])
@@ -321,20 +314,6 @@ class GroupCoordinator:
world_size = self.world_size world_size = self.world_size
return (rank_in_group - 1) % 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 @contextmanager
def graph_capture(self, graph_capture_context: GraphCaptureContext | None = None): def graph_capture(self, graph_capture_context: GraphCaptureContext | None = None):
if current_platform.is_cuda_alike(): if current_platform.is_cuda_alike():
@@ -809,7 +788,7 @@ class GroupCoordinator:
torch.distributed.barrier(group=self.cpu_group) torch.distributed.barrier(group=self.cpu_group)
def send(self, tensor: torch.Tensor, dst: Optional[int] = None) -> None: 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.""" """NOTE: `dst` is the rank_in_group of the destination rank."""
if dst is None: if dst is None:
dst = self.group_next_rank dst = self.group_next_rank
@@ -817,11 +796,7 @@ class GroupCoordinator:
torch.distributed.send( torch.distributed.send(
tensor, tensor,
self.ranks[dst], self.ranks[dst],
group=( group=self.device_group,
self.device_groups[self.rank_in_group % 2]
if self.world_size == 2
else self.device_group
),
) )
def recv( def recv(
@@ -836,11 +811,7 @@ class GroupCoordinator:
torch.distributed.recv( torch.distributed.recv(
tensor, tensor,
self.ranks[src], self.ranks[src],
( group=self.device_group,
self.device_groups[(self.rank_in_group + 1) % 2]
if self.world_size == 2
else self.device_group
),
) )
return tensor return tensor
@@ -860,438 +831,6 @@ class GroupCoordinator:
self.mq_broadcaster = None 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): class SequenceParallelGroupCoordinator(GroupCoordinator):
def __init__( def __init__(
self, self,
@@ -34,25 +34,18 @@ If you only need to use the distributed environment without model parallelism,
import contextlib import contextlib
import datetime import datetime
import os import os
import weakref
from collections import namedtuple
from collections.abc import Callable
from contextlib import contextmanager from contextlib import contextmanager
from multiprocessing import shared_memory from typing import List, Optional
from typing import Any, List, Optional
from unittest.mock import patch
import torch import torch
import torch.distributed import torch.distributed
from torch.distributed import ProcessGroup from torch.distributed import ProcessGroup
import sglang.multimodal_gen.envs as envs 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 sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
from .group_coordinator import ( from .group_coordinator import (
GroupCoordinator, GroupCoordinator,
PipelineGroupCoordinator,
SequenceParallelGroupCoordinator, SequenceParallelGroupCoordinator,
get_local_torch_device, get_local_torch_device,
new_device_group, new_device_group,
@@ -64,7 +57,7 @@ logger = init_logger(__name__)
_WORLD: GroupCoordinator | None = None _WORLD: GroupCoordinator | None = None
_TP: GroupCoordinator | None = None _TP: GroupCoordinator | None = None
_SP: SequenceParallelGroupCoordinator | None = None _SP: SequenceParallelGroupCoordinator | None = None
_PP: PipelineGroupCoordinator | None = None _PP: GroupCoordinator | None = None
_CFG: GroupCoordinator | None = None _CFG: GroupCoordinator | None = None
_DP: GroupCoordinator | None = None _DP: GroupCoordinator | None = None
# all ranks serving one pipeline replica (every dim except dp); with # 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" _REPLICA_PARALLEL_AXES = "tp-sp-pp-cfg"
_ENCODER_DP_PARALLEL_AXES = "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: def get_world_group() -> GroupCoordinator:
assert _WORLD is not None, "world group is not initialized" assert _WORLD is not None, "world group is not initialized"
@@ -227,14 +173,7 @@ def init_parallel_group_coordinator(
"replica", "replica",
"encoder_data", "encoder_data",
], f"parallel_mode {parallel_mode} is not supported" ], f"parallel_mode {parallel_mode} is not supported"
if parallel_mode == "pipeline": if parallel_mode == "sequence":
return PipelineGroupCoordinator(
group_ranks=group_ranks,
local_rank=local_rank,
torch_distributed_backend=backend,
group_name="pp_group",
)
elif parallel_mode == "sequence":
return SequenceParallelGroupCoordinator( return SequenceParallelGroupCoordinator(
group_ranks=group_ranks, group_ranks=group_ranks,
local_rank=local_rank, local_rank=local_rank,
@@ -244,6 +183,7 @@ def init_parallel_group_coordinator(
) )
else: else:
group_name = { group_name = {
"pipeline": "pp_group",
"tensor": "tp_group", "tensor": "tp_group",
"vae_decode": "vae_decode_group", "vae_decode": "vae_decode_group",
"replica": "replica_group", "replica": "replica_group",
@@ -808,96 +748,6 @@ def cleanup_dist_env_and_memory(shutdown_ray: bool = False):
ray.shutdown() 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: def get_tensor_model_parallel_world_size() -> int:
"""Return world size for the tensor model parallel group.""" """Return world size for the tensor model parallel group."""
return get_tp_world_size() return get_tp_world_size()
@@ -950,7 +800,7 @@ def get_ring_ctx() -> tuple[int, int]:
# PP # PP
def get_pp_group() -> PipelineGroupCoordinator: def get_pp_group() -> GroupCoordinator:
assert _PP is not None, "pipeline model parallel group is not initialized" assert _PP is not None, "pipeline model parallel group is not initialized"
return _PP return _PP
@@ -3,72 +3,8 @@
# SPDX-License-Identifier: Apache-2.0 # SPDX-License-Identifier: Apache-2.0
import argparse 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): class RaiseNotImplementedAction(argparse.Action):
def __call__(self, parser, namespace, values, option_string=None): def __call__(self, parser, namespace, values, option_string=None):
raise NotImplementedError(f"The {option_string} option is not yet implemented") 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 rank_offset = node_rank * local_num_gpus
processes = [] processes = []
# Pipes for master to talk to slaves (local to this node) # Launch this node's local worker processes.
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
scheduler_pipe_readers = [] scheduler_pipe_readers = []
scheduler_pipe_writers = [] scheduler_pipe_writers = []
@@ -175,40 +158,12 @@ def launch_server(server_args: ServerArgs, launch_http_server: bool = True):
rank = rank_offset + i rank = rank_offset + i
reader, writer = mp.Pipe(duplex=False) reader, writer = mp.Pipe(duplex=False)
scheduler_pipe_writers.append(writer) scheduler_pipe_writers.append(writer)
if i == 0: # This node's local pipe master process = mp.Process(
process = mp.Process( target=run_scheduler_process,
target=run_scheduler_process, args=(i, rank, server_args, writer),
args=( name=f"sglang-diffusionWorker-{rank}",
i, # local_rank daemon=True,
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,
)
scheduler_pipe_readers.append(reader) scheduler_pipe_readers.append(reader)
process.start() process.start()
processes.append(process) processes.append(process)
@@ -218,16 +173,6 @@ def launch_server(server_args: ServerArgs, launch_http_server: bool = True):
for writer in scheduler_pipe_writers: for writer in scheduler_pipe_writers:
writer.close() 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): for i, reader in enumerate(scheduler_pipe_readers):
try: try:
data = reader.recv() data = reader.recv()
@@ -428,7 +373,7 @@ def launch_pool_disagg_server(
process = pool_ctx.Process( process = pool_ctx.Process(
target=_run_disagg_role_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}", name=f"sglang-pool-{role_type.value}-{inst_idx}-r{rank_idx}",
daemon=True, daemon=True,
) )
@@ -502,12 +447,9 @@ def launch_pool_disagg_server(
def _run_disagg_role_process( def _run_disagg_role_process(
gpu_id: int, gpu_id: int,
_local_rank: int,
rank: int, rank: int,
server_args: ServerArgs, server_args: ServerArgs,
pipe_writer: mp.connection.Connection, pipe_writer: mp.connection.Connection,
task_pipes: list,
result_pipes: list,
): ):
"""Entry point for a disagg role process. """Entry point for a disagg role process.
@@ -519,13 +461,8 @@ def _run_disagg_role_process(
run_scheduler_process( run_scheduler_process(
local_rank=gpu_id, local_rank=gpu_id,
rank=rank, rank=rank,
master_port=server_args.master_port,
server_args=server_args, server_args=server_args,
pipe_writer=pipe_writer, 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( process = pool_ctx.Process(
target=_run_disagg_role_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}", name=f"sglang-{role_type.value}-r{rank_idx}",
daemon=True, daemon=True,
) )
@@ -4,11 +4,8 @@
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/attention/backends/abstract.py # Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/attention/backends/abstract.py
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from dataclasses import dataclass, fields from dataclasses import dataclass
from typing import TYPE_CHECKING, Any, Generic, Protocol, TypeVar from typing import Any, Generic, TypeVar
if TYPE_CHECKING:
pass
import torch import torch
@@ -78,15 +75,6 @@ class AttentionBackend(ABC):
def get_metadata_cls() -> type["AttentionMetadata"]: def get_metadata_cls() -> type["AttentionMetadata"]:
raise NotImplementedError 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 @staticmethod
@abstractmethod @abstractmethod
def get_builder_cls() -> type["AttentionMetadataBuilder"]: def get_builder_cls() -> type["AttentionMetadataBuilder"]:
@@ -100,18 +88,6 @@ class AttentionMetadata:
# Current step of diffusion process # Current step of diffusion process
current_timestep: int 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) T = TypeVar("T", bound=AttentionMetadata)
@@ -138,22 +114,6 @@ class AttentionMetadataBuilder(ABC, Generic[T]):
raise NotImplementedError 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]): class AttentionImpl(ABC, Generic[T]):
@abstractmethod @abstractmethod
def __init__( def __init__(
@@ -3,7 +3,6 @@
# SPDX-License-Identifier: Apache-2.0 # SPDX-License-Identifier: Apache-2.0
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/attention/selector.py # 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 collections.abc import Generator
from contextlib import contextmanager from contextlib import contextmanager
from contextvars import ContextVar from contextvars import ContextVar
@@ -23,47 +22,10 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__) 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 # Global state allows a particular choice of backend
# to be forced, overriding the logic which auto-selects # to be forced, overriding the logic which auto-selects
# a backend based on system & workload configuration # a backend based on system & workload configuration
# (default behavior if this variable is None) # (default behavior if this variable is None)
#
# THIS SELECTION TAKES PRECEDENCE OVER THE
# FASTVIDEO ATTENTION BACKEND ENVIRONMENT VARIABLE
forced_attn_backend: AttentionBackendEnum | None = None forced_attn_backend: AttentionBackendEnum | None = None
@@ -150,13 +112,6 @@ def _record_component_attn_backend(backend_name: str, reason: str | None) -> boo
return True 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( def _log_component_attn_backend_summary(
context: ComponentAttnBackendContext | None, context: ComponentAttnBackendContext | None,
) -> None: ) -> None:
@@ -380,20 +380,6 @@ class LayerNorm(CustomOp):
else: else:
self.register_parameter("weight", None) self.register_parameter("weight", None)
self.register_parameter("bias", 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): def forward_triton(self, x: torch.Tensor):
# Fast inference kernel without residual/dropout branches # 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)}") 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 sglang_lib = Library("sglang", "FRAGMENT") # noqa
@@ -53,10 +53,8 @@ from sglang.multimodal_gen.runtime.entrypoints.utils import (
save_outputs, save_outputs,
) )
from sglang.multimodal_gen.runtime.managers.memory_managers.auto_residency import ( from sglang.multimodal_gen.runtime.managers.memory_managers.auto_residency import (
DefaultWorkload,
WarmupMemoryRecord, WarmupMemoryRecord,
estimate_default_workload_peak_bytes, estimate_default_workload_peak_bytes,
resolve_default_workload,
) )
from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager import ( from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager import (
get_global_component_residency_manager, get_global_component_residency_manager,
@@ -244,23 +242,6 @@ class GPUWorker(GPUWorkerPostTrainingMixin):
# per-rank memory measurements of server warmup forwards; consumed by # per-rank memory measurements of server warmup forwards; consumed by
# the auto-residency placement decision before the server turns ready # the auto-residency placement decision before the server turns ready
self._auto_residency_warmup_records: list[WarmupMemoryRecord] = [] 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: def release_realtime_session(self, session_id: str) -> OutputBatch:
"""release the session of a realtime connection""" """release the session of a realtime connection"""
@@ -1553,23 +1534,10 @@ def _oom_exceptions():
def run_scheduler_process( def run_scheduler_process(
local_rank: int, local_rank: int,
rank: int, rank: int,
master_port: int,
server_args: ServerArgs, server_args: ServerArgs,
pipe_writer: mp.connection.Connection, 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: ) -> None:
""" """Run a rank's scheduler and report readiness to the launching process."""
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.
"""
kill_itself_when_parent_died() kill_itself_when_parent_died()
configure_logger(server_args) configure_logger(server_args)
globally_suppress_loggers() globally_suppress_loggers()
@@ -1583,8 +1551,6 @@ def run_scheduler_process(
port_args = PortArgs.from_server_args(server_args) port_args = PortArgs.from_server_args(server_args)
# start the scheduler event loop # 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 from sglang.multimodal_gen.runtime.managers.scheduler import Scheduler
try: try:
@@ -1592,8 +1558,6 @@ def run_scheduler_process(
server_args, server_args,
gpu_id=rank, gpu_id=rank,
port_args=port_args, port_args=port_args,
task_pipes_to_slaves=task_pipes_to_slaves,
result_pipes_from_slaves=result_pipes_from_slaves,
local_rank=local_rank, local_rank=local_rank,
) )
logger.info(f"Worker {rank}: Scheduler loop started.") logger.info(f"Worker {rank}: Scheduler loop started.")
@@ -90,8 +90,6 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag
server_args: ServerArgs, server_args: ServerArgs,
gpu_id: int, gpu_id: int,
port_args: PortArgs, port_args: PortArgs,
task_pipes_to_slaves: list = None,
result_pipes_from_slaves: list = None,
local_rank: int | None = None, local_rank: int | None = None,
): ):
self.server_args = server_args self.server_args = server_args
@@ -134,8 +132,6 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag
server_args=server_args, server_args=server_args,
) )
self.worker = worker 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.gpu_id = gpu_id
self._show_warmup_progress = gpu_id == 0 self._show_warmup_progress = gpu_id == 0
self._running = True self._running = True
@@ -1300,21 +1296,6 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag
self._cleanup_disagg() self._cleanup_disagg()
self.context.destroy(linger=0) 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: def _handle_release_memory_occupation(self, _reqs: List[Any]) -> OutputBatch:
logger.info(f"[SLEEP] handle_release_memory_occupation on rank={self.gpu_id}") logger.info(f"[SLEEP] handle_release_memory_occupation on rank={self.gpu_id}")
return OutputBatch(output=self.worker.release_memory_occupation()) return OutputBatch(output=self.worker.release_memory_occupation())
@@ -9,7 +9,6 @@ import torch
import torch.nn as nn import torch.nn as nn
from torch.nn.attention.flex_attention import ( from torch.nn.attention.flex_attention import (
BlockMask, BlockMask,
create_block_mask,
flex_attention, flex_attention,
) )
@@ -23,7 +22,6 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload im
flex_attention = torch.compile( flex_attention = torch.compile(
flex_attention, dynamic=False, mode="max-autotune-no-cudagraphs" 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.dits import WanVideoConfig
from sglang.multimodal_gen.configs.models.fsdp import is_block from sglang.multimodal_gen.configs.models.fsdp import is_block
@@ -534,79 +532,6 @@ class CausalWanTransformer3DModel(BaseDiT, LayerwiseOffloadableModuleMixin):
"blocks", "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( def forward(
self, self,
hidden_states: torch.Tensor, hidden_states: torch.Tensor,
@@ -235,9 +235,6 @@ class ComfyUIQwenImagePipelineBase(LoRAPipeline, ComposedPipelineBase):
model = model_cls(**{"config": dit_config, "hf_config": hf_config}) model = model_cls(**{"config": dit_config, "hf_config": hf_config})
use_fsdp = server_args.should_use_fsdp_for_component("transformer") 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(): if current_platform.is_mps():
use_fsdp = False use_fsdp = False
logger.info("Disabling FSDP for MPS platform as it's not compatible") logger.info("Disabling FSDP for MPS platform as it's not compatible")
@@ -50,23 +50,6 @@ def _causal_block_count(batch: Req, server_args: ServerArgs) -> int:
return latent_frames // block_size 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): class LongLive2TextEncodingStage(TextEncodingStage):
def build_dedup_fingerprint(self, batch: Req, server_args: ServerArgs): def build_dedup_fingerprint(self, batch: Req, server_args: ServerArgs):
base = super().build_dedup_fingerprint(batch, server_args) base = super().build_dedup_fingerprint(batch, server_args)
@@ -12,11 +12,7 @@ from sglang.multimodal_gen.runtime.distributed import (
get_local_torch_device, get_local_torch_device,
get_sp_world_size, 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 ( from sglang.multimodal_gen.runtime.distributed.communication_op import (
cfg_model_parallel_all_gather,
cfg_model_parallel_all_reduce, cfg_model_parallel_all_reduce,
) )
from sglang.multimodal_gen.runtime.distributed.parallel_state import ( from sglang.multimodal_gen.runtime.distributed.parallel_state import (
@@ -256,144 +252,6 @@ class LTX2DenoisingStage(DenoisingStage):
cfg_model_parallel_all_reduce(audio_partial), 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 @staticmethod
def _get_video_latent_num_frames_for_model( def _get_video_latent_num_frames_for_model(
batch: Req, server_args: ServerArgs, latents: torch.Tensor 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}" 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): def is_port_available(port):
"""Return whether a port is available.""" """Return whether a port is available."""
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s: 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) 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(): def test_default_position_ids_batch_shape():
model = Qwen3ForCausalLM.__new__(Qwen3ForCausalLM) model = Qwen3ForCausalLM.__new__(Qwen3ForCausalLM)
torch.nn.Module.__init__(model) torch.nn.Module.__init__(model)
@@ -1,16 +1,63 @@
"""Single-rank groups get gloo, so NCCL does not reserve device buffers for them.""" """Single-rank groups get gloo, so NCCL does not reserve device buffers for them."""
import unittest import unittest
from datetime import timedelta
from pathlib import Path
from tempfile import TemporaryDirectory
from unittest.mock import patch 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 ( from sglang.multimodal_gen.runtime.distributed.group_coordinator import (
new_device_group, 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" 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): def test_single_rank_group_avoids_the_device_backend(self):
for ranks, requested in [([0], "nccl"), ([3], "hccl"), ([0], None)]: for ranks, requested in [([0], "nccl"), ([3], "hccl"), ([0], None)]:
with self.subTest(ranks=ranks, requested=requested): with self.subTest(ranks=ranks, requested=requested):
@@ -25,6 +72,10 @@ class TestSingleRankDeviceGroup(unittest.TestCase):
new_device_group(ranks, requested) new_device_group(ranks, requested)
new_group.assert_called_once_with(ranks, backend=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__": if __name__ == "__main__":
unittest.main() unittest.main()
@@ -112,23 +112,6 @@ class TestTransformerLoaderFallbackAdmission(unittest.TestCase):
native.assert_called_once() native.assert_called_once()
server_args.should_use_fsdp_for_component.assert_called_with("transformer_2") 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): def test_unreadable_checkpoint_is_not_a_missing_implementation(self):
# the native fallback answers "no customized implementation for this # the native fallback answers "no customized implementation for this
# architecture"; a checkpoint that cannot be read is a different failure, # 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)) self.assertIn("shard-00002.safetensors", str(caught.exception))
native.assert_not_called() 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__": if __name__ == "__main__":
unittest.main() unittest.main()
@@ -1,11 +1,8 @@
# SPDX-License-Identifier: Apache-2.0 # SPDX-License-Identifier: Apache-2.0
import ast
import os import os
from concurrent.futures import ThreadPoolExecutor from concurrent.futures import ThreadPoolExecutor
from dataclasses import dataclass, field from dataclasses import dataclass, field
from importlib.util import resolve_name
from pathlib import Path
from threading import Barrier from threading import Barrier
import pytest 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): def test_argument_parser_preserves_config_and_explicit_values(tmp_path):
config = tmp_path / "config.yaml" config = tmp_path / "config.yaml"
config.write_text("num_gpus: 2\nuse_cache: true\n") config.write_text("num_gpus: 2\nuse_cache: true\n")