[diffusion] refactor: refactor utility ownership and document helper placement (#38699)
Co-authored-by: Mick Qian <mickqian@users.noreply.github.com>
This commit is contained in:
@@ -10,7 +10,7 @@ from typing import Any
|
||||
import torch
|
||||
|
||||
from sglang.multimodal_gen.configs.models.base import ArchConfig, ModelConfig
|
||||
from sglang.multimodal_gen.utils import StoreBoolean
|
||||
from sglang.multimodal_gen.runtime.utils.argparse import StoreBoolean
|
||||
|
||||
AUTO_PARALLEL_DECODE_MODE = "auto"
|
||||
SPATIAL_SHARD_PARALLEL_DECODE_MODES = ("spatial_shard", "spatial")
|
||||
|
||||
@@ -7,6 +7,7 @@ import os
|
||||
from collections.abc import Callable
|
||||
from dataclasses import asdict, dataclass, field, fields
|
||||
from enum import Enum, auto
|
||||
from operator import attrgetter
|
||||
from typing import Any, ClassVar
|
||||
|
||||
import numpy as np
|
||||
@@ -35,13 +36,12 @@ from sglang.multimodal_gen.runtime.distributed.parallel_state import (
|
||||
get_sp_parallel_rank,
|
||||
get_sp_world_size,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.multimodal_gen.runtime.utils.vision import get_default_height_width
|
||||
from sglang.multimodal_gen.utils import (
|
||||
from sglang.multimodal_gen.runtime.utils.argparse import (
|
||||
FlexibleArgumentParser,
|
||||
StoreBoolean,
|
||||
shallow_asdict,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.multimodal_gen.runtime.utils.vision import get_default_height_width
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -1154,7 +1154,7 @@ class PipelineConfig:
|
||||
)
|
||||
|
||||
def dump_to_json(self, file_path: str):
|
||||
output_dict = shallow_asdict(self)
|
||||
output_dict = {f.name: attrgetter(f.name)(self) for f in fields(self)}
|
||||
del_keys = []
|
||||
for key, value in output_dict.items():
|
||||
if isinstance(value, ModelConfig):
|
||||
|
||||
@@ -449,7 +449,7 @@ def _calculate_edit_dimensions(target_area, ratio):
|
||||
"""Output size for LongCat-Image-Edit: fit `target_area`, ceil to /16.
|
||||
|
||||
Copied from diffusers pipeline_longcat_image_edit.calculate_dimensions.
|
||||
Note this intentionally differs from sglang.multimodal_gen.utils
|
||||
Note this intentionally differs from the Qwen-Image pipeline config
|
||||
calculate_dimensions (which rounds to /32).
|
||||
"""
|
||||
width = math.sqrt(target_area * ratio)
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
|
||||
|
||||
import math
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Callable
|
||||
|
||||
@@ -29,7 +30,16 @@ from sglang.multimodal_gen.runtime.utils.condition_expansion import (
|
||||
PromptToSampleBatchExpander,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.utils.vision import resize
|
||||
from sglang.multimodal_gen.utils import calculate_dimensions
|
||||
|
||||
|
||||
def _calculate_dimensions(target_area, ratio):
|
||||
width = math.sqrt(target_area * ratio)
|
||||
height = width / ratio
|
||||
|
||||
width = round(width / 32) * 32
|
||||
height = round(height / 32) * 32
|
||||
|
||||
return width, height, None
|
||||
|
||||
|
||||
def _extract_masked_hidden(hidden_states: torch.Tensor, mask: torch.Tensor):
|
||||
@@ -496,7 +506,7 @@ class QwenImageEditPipelineConfig(QwenImagePipelineConfig):
|
||||
height = batch.height
|
||||
width = batch.width
|
||||
image_size = batch.original_condition_image_size
|
||||
edit_width, edit_height, _ = calculate_dimensions(
|
||||
edit_width, edit_height, _ = _calculate_dimensions(
|
||||
1024 * 1024, image_size[0] / image_size[1]
|
||||
)
|
||||
vae_scale_factor = self.get_vae_scale_factor()
|
||||
@@ -599,7 +609,7 @@ class QwenImageEditPipelineConfig(QwenImagePipelineConfig):
|
||||
)
|
||||
|
||||
def calculate_condition_image_size(self, image, width, height) -> tuple[int, int]:
|
||||
calculated_width, calculated_height, _ = calculate_dimensions(
|
||||
calculated_width, calculated_height, _ = _calculate_dimensions(
|
||||
1024 * 1024, width / height
|
||||
)
|
||||
return calculated_width, calculated_height
|
||||
@@ -626,7 +636,7 @@ class QwenImageEditPlusPipelineConfig(QwenImageEditPipelineConfig):
|
||||
condition_image_sizes = []
|
||||
for img in image:
|
||||
image_width, image_height = img.size
|
||||
edit_width, edit_height, _ = calculate_dimensions(
|
||||
edit_width, edit_height, _ = _calculate_dimensions(
|
||||
VAE_IMAGE_SIZE, image_width / image_height
|
||||
)
|
||||
condition_image_sizes.append((edit_width, edit_height))
|
||||
@@ -674,13 +684,13 @@ class QwenImageEditPlusPipelineConfig(QwenImageEditPipelineConfig):
|
||||
return new_images
|
||||
|
||||
def calculate_condition_image_size(self, image, width, height) -> tuple[int, int]:
|
||||
calculated_width, calculated_height, _ = calculate_dimensions(
|
||||
calculated_width, calculated_height, _ = _calculate_dimensions(
|
||||
CONDITION_IMAGE_SIZE, width / height
|
||||
)
|
||||
return calculated_width, calculated_height
|
||||
|
||||
def calculate_vae_image_size(self, image, width, height) -> tuple[int, int]:
|
||||
calculated_width, calculated_height, _ = calculate_dimensions(
|
||||
calculated_width, calculated_height, _ = _calculate_dimensions(
|
||||
VAE_IMAGE_SIZE, width / height
|
||||
)
|
||||
return calculated_width, calculated_height
|
||||
|
||||
@@ -9,7 +9,7 @@ import math
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Callable
|
||||
|
||||
from sglang.multimodal_gen.utils import StoreBoolean
|
||||
from sglang.multimodal_gen.runtime.utils.argparse import StoreBoolean
|
||||
|
||||
_VALID_ROLLOUT_SDE_TYPES = ("sde", "cps", "ode")
|
||||
|
||||
|
||||
@@ -14,8 +14,8 @@ from sglang.multimodal_gen.runtime.layers.quantization.configs.nunchaku_config i
|
||||
is_nunchaku_available,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||
from sglang.multimodal_gen.runtime.utils.argparse import StoreBoolean
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.multimodal_gen.utils import StoreBoolean
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@@ -10,7 +10,8 @@ from sglang.multimodal_gen.configs.sample.sampling_params import (
|
||||
DataType,
|
||||
_sanitize_filename,
|
||||
)
|
||||
from sglang.multimodal_gen.utils import StoreBoolean, expand_path_fields
|
||||
from sglang.multimodal_gen.configs.utils import expand_path_fields
|
||||
from sglang.multimodal_gen.runtime.utils.argparse import StoreBoolean
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
|
||||
@@ -17,8 +17,9 @@ from enum import Enum, auto
|
||||
from typing import TYPE_CHECKING, Any, ClassVar
|
||||
|
||||
from sglang.multimodal_gen.configs.post_training import RLRolloutArgs
|
||||
from sglang.multimodal_gen.configs.utils import expand_path_fields
|
||||
from sglang.multimodal_gen.runtime.utils.argparse import StoreBoolean
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.multimodal_gen.utils import StoreBoolean, expand_path_fields
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@@ -1,9 +1,35 @@
|
||||
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
|
||||
|
||||
import argparse
|
||||
import os
|
||||
from dataclasses import fields
|
||||
from operator import attrgetter
|
||||
from typing import Any
|
||||
|
||||
|
||||
def expand_path_fields(obj) -> None:
|
||||
"""Expand paths in dataclass configuration without modifying other fields."""
|
||||
for f in fields(obj):
|
||||
value = attrgetter(f.name)(obj)
|
||||
if f.name.endswith("_path") and isinstance(value, str):
|
||||
setattr(obj, f.name, os.path.expanduser(value))
|
||||
elif f.name.endswith("_path") and isinstance(value, list):
|
||||
setattr(
|
||||
obj,
|
||||
f.name,
|
||||
[os.path.expanduser(v) if isinstance(v, str) else v for v in value],
|
||||
)
|
||||
elif f.name.endswith("_paths") and isinstance(value, dict):
|
||||
setattr(
|
||||
obj,
|
||||
f.name,
|
||||
{
|
||||
k: os.path.expanduser(v) if isinstance(v, str) else v
|
||||
for k, v in value.items()
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def update_config_from_args(
|
||||
config: Any, args_dict: dict[str, Any], prefix: str = "", pop_args: bool = False
|
||||
) -> bool:
|
||||
|
||||
@@ -50,6 +50,7 @@ from sglang.multimodal_gen.runtime.disaggregation.transport.protocol import (
|
||||
encode_transfer_msg,
|
||||
is_transfer_message,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.distributed.utils import broadcast_pyobj
|
||||
from sglang.multimodal_gen.runtime.entrypoints.utils import expand_request_outputs
|
||||
from sglang.multimodal_gen.runtime.pipelines_core import Req
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.diffusion_scheduler_utils import (
|
||||
@@ -57,7 +58,6 @@ from sglang.multimodal_gen.runtime.pipelines_core.diffusion_scheduler_utils impo
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||
from sglang.multimodal_gen.runtime.utils.common import get_zmq_socket
|
||||
from sglang.multimodal_gen.runtime.utils.distributed import broadcast_pyobj
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.multimodal_gen.runtime.utils.trace_wrapper import DiffStage, trace_slice
|
||||
from sglang.srt.observability.trace import TraceReqContext
|
||||
|
||||
@@ -11,6 +11,7 @@ import torch
|
||||
import torch.distributed as dist
|
||||
from torch.distributed import ProcessGroup, ReduceOp
|
||||
|
||||
from sglang.multimodal_gen.runtime import platforms
|
||||
from sglang.multimodal_gen.runtime.distributed.device_communicators.pynccl_wrapper import (
|
||||
NCCLLibrary,
|
||||
buffer_type,
|
||||
@@ -22,11 +23,42 @@ from sglang.multimodal_gen.runtime.distributed.device_communicators.pynccl_wrapp
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.distributed.utils import StatelessProcessGroup
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.multimodal_gen.utils import current_stream
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
_previous_set_stream = torch.cuda.set_stream
|
||||
|
||||
_current_stream = None
|
||||
|
||||
|
||||
def _patched_set_stream(stream: torch.cuda.Stream | None) -> None:
|
||||
global _current_stream
|
||||
_current_stream = stream
|
||||
if stream is not None:
|
||||
_previous_set_stream(stream)
|
||||
|
||||
|
||||
torch.cuda.set_stream = _patched_set_stream
|
||||
|
||||
|
||||
def _get_current_stream() -> torch.cuda.Stream | None:
|
||||
# cache the stream object to avoid constructing it for every collective;
|
||||
# callers must change streams through torch.cuda.set_stream
|
||||
if not platforms.current_platform.is_cuda_alike():
|
||||
return None
|
||||
|
||||
global _current_stream
|
||||
if _current_stream is None:
|
||||
# RCCL performs better on a dedicated stream than the default stream
|
||||
_current_stream = (
|
||||
torch.cuda.Stream()
|
||||
if platforms.current_platform.is_rocm()
|
||||
else torch.cuda.current_stream()
|
||||
)
|
||||
return _current_stream
|
||||
|
||||
|
||||
class PyNcclCommunicator:
|
||||
def __init__(
|
||||
self,
|
||||
@@ -110,7 +142,7 @@ class PyNcclCommunicator:
|
||||
self.world_size, self.unique_id, self.rank
|
||||
)
|
||||
|
||||
stream = current_stream()
|
||||
stream = _get_current_stream()
|
||||
# A small all_reduce for warmup.
|
||||
data = torch.zeros(1, device=device)
|
||||
self.all_reduce(data)
|
||||
@@ -134,7 +166,7 @@ class PyNcclCommunicator:
|
||||
out_tensor = torch.empty_like(in_tensor)
|
||||
|
||||
if stream is None:
|
||||
stream = current_stream()
|
||||
stream = _get_current_stream()
|
||||
self.nccl.ncclAllReduce(
|
||||
buffer_type(in_tensor.data_ptr()),
|
||||
buffer_type(out_tensor.data_ptr()),
|
||||
@@ -159,7 +191,7 @@ class PyNcclCommunicator:
|
||||
f"but the input tensor is on {input_tensor.device}"
|
||||
)
|
||||
if stream is None:
|
||||
stream = current_stream()
|
||||
stream = _get_current_stream()
|
||||
self.nccl.ncclAllGather(
|
||||
buffer_type(input_tensor.data_ptr()),
|
||||
buffer_type(output_tensor.data_ptr()),
|
||||
@@ -186,7 +218,7 @@ class PyNcclCommunicator:
|
||||
f"but the input tensor is on {input_tensor.device}"
|
||||
)
|
||||
if stream is None:
|
||||
stream = current_stream()
|
||||
stream = _get_current_stream()
|
||||
self.nccl.ncclReduceScatter(
|
||||
buffer_type(input_tensor.data_ptr()),
|
||||
buffer_type(output_tensor.data_ptr()),
|
||||
@@ -205,7 +237,7 @@ class PyNcclCommunicator:
|
||||
f"but the input tensor is on {tensor.device}"
|
||||
)
|
||||
if stream is None:
|
||||
stream = current_stream()
|
||||
stream = _get_current_stream()
|
||||
self.nccl.ncclSend(
|
||||
buffer_type(tensor.data_ptr()),
|
||||
tensor.numel(),
|
||||
@@ -223,7 +255,7 @@ class PyNcclCommunicator:
|
||||
f"but the input tensor is on {tensor.device}"
|
||||
)
|
||||
if stream is None:
|
||||
stream = current_stream()
|
||||
stream = _get_current_stream()
|
||||
self.nccl.ncclRecv(
|
||||
buffer_type(tensor.data_ptr()),
|
||||
tensor.numel(),
|
||||
@@ -272,7 +304,7 @@ class PyNcclCommunicator:
|
||||
f"got {input_.numel()} elements over {self.world_size} ranks"
|
||||
)
|
||||
if stream is None:
|
||||
stream = current_stream()
|
||||
stream = _get_current_stream()
|
||||
# dist.all_to_all_single defines split sizes along dim 0; convert rows
|
||||
# to element counts so n-D tensors split identically to torch
|
||||
in_row = input_.numel() // input_.size(0) if input_.dim() else 1
|
||||
@@ -341,7 +373,7 @@ class PyNcclCommunicator:
|
||||
f"but the input tensor is on {tensor.device}"
|
||||
)
|
||||
if stream is None:
|
||||
stream = current_stream()
|
||||
stream = _get_current_stream()
|
||||
if src == self.rank:
|
||||
sendbuff = buffer_type(tensor.data_ptr())
|
||||
# NCCL requires the sender also to have a receive buffer
|
||||
|
||||
+30
-2
@@ -35,8 +35,8 @@ from typing import Any
|
||||
import torch
|
||||
from torch.distributed import ReduceOp
|
||||
|
||||
from sglang.multimodal_gen import envs
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.multimodal_gen.utils import find_nccl_library
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -48,6 +48,34 @@ ncclResult_t = ctypes.c_int
|
||||
ncclComm_t = ctypes.c_void_p
|
||||
|
||||
|
||||
def _find_nccl_library() -> str:
|
||||
"""
|
||||
We either use the library file specified by the `SGLANG_DIFFUSION_NCCL_SO_PATH`
|
||||
environment variable, or we find the library file brought by PyTorch.
|
||||
After importing `torch`, `libnccl.so.2`, `librccl.so.1` or `libmccl.so.2`
|
||||
can be found by `ctypes` automatically.
|
||||
"""
|
||||
so_file = envs.SGLANG_DIFFUSION_NCCL_SO_PATH
|
||||
|
||||
# manually load the nccl library
|
||||
if so_file:
|
||||
logger.info(
|
||||
"Found nccl from environment variable SGLANG_DIFFUSION_NCCL_SO_PATH=%s",
|
||||
so_file,
|
||||
)
|
||||
else:
|
||||
if torch.version.cuda is not None:
|
||||
so_file = "libnccl.so.2"
|
||||
elif torch.version.hip is not None:
|
||||
so_file = "librccl.so.1"
|
||||
elif vars(torch.version).get("musa") is not None:
|
||||
so_file = "libmccl.so.2"
|
||||
else:
|
||||
raise ValueError("NCCL only supports CUDA, ROCm and MUSA backends.")
|
||||
logger.info("Found nccl from library %s", so_file)
|
||||
return str(so_file)
|
||||
|
||||
|
||||
class ncclUniqueId(ctypes.Structure):
|
||||
_fields_ = [("internal", ctypes.c_byte * 128)]
|
||||
|
||||
@@ -274,7 +302,7 @@ class NCCLLibrary:
|
||||
|
||||
def __init__(self, so_file: str | None = None):
|
||||
|
||||
so_file = so_file or find_nccl_library()
|
||||
so_file = so_file or _find_nccl_library()
|
||||
|
||||
try:
|
||||
if so_file not in NCCLLibrary.path_to_dict_mapping:
|
||||
|
||||
@@ -50,7 +50,6 @@ 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 ..utils.distributed import RankGenerator
|
||||
from .group_coordinator import (
|
||||
GroupCoordinator,
|
||||
PipelineGroupCoordinator,
|
||||
@@ -58,6 +57,7 @@ from .group_coordinator import (
|
||||
get_local_torch_device,
|
||||
new_device_group,
|
||||
)
|
||||
from .utils import RankGenerator
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@@ -13,11 +13,15 @@ import pickle
|
||||
import time
|
||||
from collections import deque
|
||||
from collections.abc import Sequence
|
||||
from typing import Any
|
||||
from typing import Any, List, Optional
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from torch.distributed import TCPStore
|
||||
|
||||
from sglang.multimodal_gen.runtime import platforms
|
||||
|
||||
try:
|
||||
from torch.distributed import all_gather_single as _all_gather_single
|
||||
except ImportError:
|
||||
@@ -201,3 +205,227 @@ class StatelessProcessGroup:
|
||||
store=store,
|
||||
data_expiration_seconds=data_expiration_seconds,
|
||||
)
|
||||
|
||||
|
||||
def broadcast_pyobj(
|
||||
data: List[Any],
|
||||
rank: int,
|
||||
dist_group: Optional[torch.distributed.ProcessGroup] = None,
|
||||
src: int = 0,
|
||||
force_cpu_device: bool = True,
|
||||
):
|
||||
"""Broadcast inputs from src rank to all other ranks with torch.dist backend.
|
||||
The `rank` here refer to the source rank on global process group (regardless
|
||||
of dist_group argument).
|
||||
"""
|
||||
|
||||
device = torch.device(
|
||||
platforms.current_platform.device_type if not force_cpu_device else "cpu"
|
||||
)
|
||||
|
||||
if rank == src:
|
||||
if data is None or len(data) == 0:
|
||||
tensor_size = torch.tensor([0], dtype=torch.long, device=device)
|
||||
dist.broadcast(tensor_size, src=src, group=dist_group)
|
||||
else:
|
||||
serialized_data = pickle.dumps(data)
|
||||
size = len(serialized_data)
|
||||
|
||||
tensor_data = torch.ByteTensor(
|
||||
np.frombuffer(serialized_data, dtype=np.uint8).copy()
|
||||
).to(device)
|
||||
tensor_size = torch.tensor([size], dtype=torch.long, device=device)
|
||||
|
||||
dist.broadcast(tensor_size, src=src, group=dist_group)
|
||||
dist.broadcast(tensor_data, src=src, group=dist_group)
|
||||
return data
|
||||
else:
|
||||
tensor_size = torch.tensor([0], dtype=torch.long, device=device)
|
||||
dist.broadcast(tensor_size, src=src, group=dist_group)
|
||||
size = tensor_size.item()
|
||||
|
||||
if size == 0:
|
||||
return []
|
||||
|
||||
tensor_data = torch.empty(size, dtype=torch.uint8, device=device)
|
||||
dist.broadcast(tensor_data, src=src, group=dist_group)
|
||||
|
||||
serialized_data = bytes(tensor_data.cpu().numpy())
|
||||
data = pickle.loads(serialized_data)
|
||||
return data
|
||||
|
||||
|
||||
def generate_masked_orthogonal_rank_groups(
|
||||
world_size: int, parallel_size: list[int], mask: list[bool]
|
||||
) -> list[list[int]]:
|
||||
"""Generate orthogonal parallel groups based on the parallel size and mask.
|
||||
|
||||
Arguments:
|
||||
world_size (int): world size
|
||||
|
||||
parallel_size (List[int]):
|
||||
The parallel size of each orthogonal parallel type. For example, if
|
||||
tensor_parallel_size = 2, pipeline_model_parallel_group = 3, data_parallel_size = 4,
|
||||
and the parallel mapping order is tp-pp-dp, then the parallel_size = [2, 3, 4].
|
||||
|
||||
mask (List[bool]):
|
||||
The mask controls which parallel methods the generated groups represent. If mask[i] is
|
||||
True, it means the generated group contains the i-th parallelism method. For example,
|
||||
if parallel_size = [tp_size, pp_size, dp_size], and mask = [True, False , True], then
|
||||
the generated group is the `tp-dp` group, if the mask = [False, True, False], then the
|
||||
generated group is the `pp` group.
|
||||
|
||||
Algorithm:
|
||||
For orthogonal parallelism, such as tp/dp/pp/cp, the global_rank and
|
||||
|
||||
If we want to get the `dp_group` (tp_size * pp_size groups of dp_size ranks each.
|
||||
For example, if the gpu size is 8 and order is 'tp-pp-dp', size is '2-2-2', and the
|
||||
dp_group here is [[0, 4], [1, 5], [2, 6], [3, 7]].)
|
||||
The tp_rank and pp_rank will be combined to form the `dp_group_index`.
|
||||
dp_group_index = tp_rank + pp_rank * tp_size (2)
|
||||
|
||||
So, Given that tp_rank and pp_rank satisfy equation (2), and dp_rank in
|
||||
range(0, dp_size), the ranks in dp_group[dp_group_index] satisfies the
|
||||
equation (1).
|
||||
|
||||
This function solve this math problem.
|
||||
|
||||
For example, if the parallel_size = [tp_size, dp_size, pp_size] = [2, 3, 4],
|
||||
and the mask = [False, True, False]. Then,
|
||||
dp_group_index(0) = tp_rank(0) + pp_rank(0) * 2
|
||||
dp_group_index(1) = tp_rank(1) + pp_rank(0) * 2
|
||||
...
|
||||
dp_group_index(7) = tp_rank(1) + pp_rank(3) * 2
|
||||
|
||||
dp_group[0] = 0 + range(0, 3) * 2 + 0 = [0, 2, 4]
|
||||
dp_group[1] = 1 + range(0, 3) * 2 + 0 = [1, 3, 5]
|
||||
...
|
||||
dp_group[7] = 1 + range(0, 3) * 2 + 3 * 2 * 3 = [19, 21, 23]
|
||||
"""
|
||||
|
||||
def prefix_product(a: List[int], init=1) -> List[int]:
|
||||
r = [init]
|
||||
for v in a:
|
||||
init = init * v
|
||||
r.append(init)
|
||||
return r
|
||||
|
||||
def inner_product(a: List[int], b: List[int]) -> int:
|
||||
return sum([x * y for x, y in zip(a, b)])
|
||||
|
||||
def decompose(index, shape, stride=None):
|
||||
"""
|
||||
This function solve the math problem below:
|
||||
There is an equation:
|
||||
index = sum(idx[i] * stride[i])
|
||||
And given the value of index, stride.
|
||||
Return the idx.
|
||||
This function will used to get the pp/dp/pp_rank
|
||||
from group_index and rank_in_group.
|
||||
"""
|
||||
if stride is None:
|
||||
stride = prefix_product(shape)
|
||||
idx = [(index // d) % s for s, d in zip(shape, stride)]
|
||||
# stride is a prefix_product result. And the value of stride[-1]
|
||||
# is not used.
|
||||
assert sum([x * y for x, y in zip(idx, stride[:-1])]) == index, (
|
||||
"idx {} with shape {} mismatch the return idx {}".format(index, shape, idx)
|
||||
)
|
||||
return idx
|
||||
|
||||
masked_shape = [s for s, m in zip(parallel_size, mask) if m]
|
||||
unmasked_shape = [s for s, m in zip(parallel_size, mask) if not m]
|
||||
|
||||
global_stride = prefix_product(parallel_size)
|
||||
masked_stride = [d for d, m in zip(global_stride, mask) if m]
|
||||
unmasked_stride = [d for d, m in zip(global_stride, mask) if not m]
|
||||
|
||||
group_size = prefix_product(masked_shape)[-1]
|
||||
num_of_group = world_size // group_size
|
||||
|
||||
ranks = []
|
||||
for group_index in range(num_of_group):
|
||||
# get indices from unmaksed for group_index.
|
||||
decomposed_group_idx = decompose(group_index, unmasked_shape)
|
||||
rank = []
|
||||
for rank_in_group in range(group_size):
|
||||
# get indices from masked for rank_in_group.
|
||||
decomposed_rank_idx = decompose(rank_in_group, masked_shape)
|
||||
rank.append(
|
||||
inner_product(decomposed_rank_idx, masked_stride)
|
||||
+ inner_product(decomposed_group_idx, unmasked_stride)
|
||||
)
|
||||
ranks.append(rank)
|
||||
return ranks
|
||||
|
||||
|
||||
class RankGenerator(object):
|
||||
def __init__(
|
||||
self,
|
||||
tp: int,
|
||||
sp: int,
|
||||
pp: int,
|
||||
cfg: int,
|
||||
dp: int,
|
||||
order: str,
|
||||
rank_offset: int = 0,
|
||||
) -> None:
|
||||
self.tp = tp
|
||||
self.sp = sp
|
||||
self.pp = pp
|
||||
self.cfg = cfg
|
||||
self.dp = dp
|
||||
self.rank_offset = rank_offset
|
||||
self.world_size = tp * sp * pp * cfg * dp
|
||||
|
||||
self.name_to_size = {
|
||||
"tp": self.tp,
|
||||
"sp": self.sp,
|
||||
"pp": self.pp,
|
||||
"cfg": self.cfg,
|
||||
"dp": self.dp,
|
||||
}
|
||||
order = order.lower()
|
||||
|
||||
for name in self.name_to_size.keys():
|
||||
if name not in order and self.name_to_size[name] != 1:
|
||||
raise RuntimeError(
|
||||
f"The size of ({name}) is ({self.name_to_size[name]}), but you haven't specified the order ({self.order})."
|
||||
)
|
||||
elif name not in order:
|
||||
order = order + "-" + name
|
||||
|
||||
self.order = order
|
||||
self.ordered_size = []
|
||||
|
||||
for token in order.split("-"):
|
||||
self.ordered_size.append(self.name_to_size[token])
|
||||
|
||||
def get_mask(self, order: str, token: str):
|
||||
ordered_token = order.split("-")
|
||||
token = token.split("-")
|
||||
mask = [False] * len(ordered_token)
|
||||
for t in token:
|
||||
mask[ordered_token.index(t)] = True
|
||||
return mask
|
||||
|
||||
def get_ranks(self, token):
|
||||
"""Get rank group by input token.
|
||||
|
||||
Arguments:
|
||||
token (str):
|
||||
Specify the ranks type that want to get. If we want
|
||||
to obtain multiple parallel types, we can use a hyphen
|
||||
'-' to separate them. For example, if we want to obtain
|
||||
the TP_DP group, the token should be 'tp-dp'.
|
||||
|
||||
"""
|
||||
mask = self.get_mask(self.order, token)
|
||||
ranks = generate_masked_orthogonal_rank_groups(
|
||||
self.world_size, self.ordered_size, mask
|
||||
)
|
||||
if self.rank_offset > 0:
|
||||
for rank_group in ranks:
|
||||
for i in range(len(rank_group)):
|
||||
rank_group[i] += self.rank_offset
|
||||
return ranks
|
||||
|
||||
@@ -5,7 +5,7 @@
|
||||
|
||||
import argparse
|
||||
|
||||
from sglang.multimodal_gen.utils import FlexibleArgumentParser
|
||||
from sglang.multimodal_gen.runtime.utils.argparse import FlexibleArgumentParser
|
||||
|
||||
|
||||
class CLISubcommand:
|
||||
|
||||
@@ -20,13 +20,13 @@ from sglang.multimodal_gen.runtime.entrypoints.cli.utils import (
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.entrypoints.utils import GenerationResult
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
from sglang.multimodal_gen.runtime.utils.argparse import FlexibleArgumentParser
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.multimodal_gen.runtime.utils.perf_logger import (
|
||||
MemorySnapshot,
|
||||
PerformanceLogger,
|
||||
RequestMetrics,
|
||||
)
|
||||
from sglang.multimodal_gen.utils import FlexibleArgumentParser
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@@ -6,7 +6,7 @@
|
||||
from sglang.multimodal_gen.runtime.entrypoints.cli.cli_types import CLISubcommand
|
||||
from sglang.multimodal_gen.runtime.entrypoints.cli.generate import GenerateSubcommand
|
||||
from sglang.multimodal_gen.runtime.entrypoints.cli.serve import ServeSubcommand
|
||||
from sglang.multimodal_gen.utils import FlexibleArgumentParser
|
||||
from sglang.multimodal_gen.runtime.utils.argparse import FlexibleArgumentParser
|
||||
|
||||
|
||||
def generate_cmd_init() -> list[CLISubcommand]:
|
||||
|
||||
@@ -12,7 +12,7 @@ from sglang.multimodal_gen.runtime.launch_server import (
|
||||
dispatch_launch,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
from sglang.multimodal_gen.utils import FlexibleArgumentParser
|
||||
from sglang.multimodal_gen.runtime.utils.argparse import FlexibleArgumentParser
|
||||
|
||||
|
||||
def add_multimodal_gen_serve_args(parser: argparse.ArgumentParser):
|
||||
|
||||
+4
-4
@@ -14,15 +14,15 @@ from sglang.multimodal_gen.runtime.entrypoints.openai.realtime.realtime_adapter
|
||||
build_realtime_sampling_params,
|
||||
save_realtime_first_frame,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.lingbot_world.constants import (
|
||||
LINGBOT_CAMERA_ACTIONS_CONDITION,
|
||||
LINGBOT_PROMPT_UPDATED_CONDITION,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.realtime.control_signals import (
|
||||
ControlSignalQueue,
|
||||
ParsedControlEventPayload,
|
||||
parse_control_event_payload,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.realtime.lingbot_world import (
|
||||
LINGBOT_CAMERA_ACTIONS_CONDITION,
|
||||
LINGBOT_PROMPT_UPDATED_CONDITION,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.realtime.states import (
|
||||
RealtimeCameraControlState,
|
||||
)
|
||||
|
||||
+1
-1
@@ -14,7 +14,7 @@ from PIL import Image
|
||||
from sglang.multimodal_gen.runtime.entrypoints.openai.realtime.timer import (
|
||||
RealtimeStageTimer,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.utils.realtime_video import (
|
||||
from sglang.multimodal_gen.runtime.realtime.video import (
|
||||
JPEG_FRAME_CONTENT_TYPE,
|
||||
RAW_RGB_CHANNELS,
|
||||
RAW_RGB_CONTENT_TYPE,
|
||||
|
||||
@@ -21,13 +21,13 @@ from sglang.multimodal_gen.runtime.server_args import (
|
||||
prepare_server_args,
|
||||
set_global_server_args,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.utils.common import (
|
||||
is_port_available,
|
||||
from sglang.multimodal_gen.runtime.utils.common import is_port_available
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import configure_logger, logger
|
||||
from sglang.multimodal_gen.runtime.utils.process import (
|
||||
kill_itself_when_parent_died,
|
||||
kill_process_tree,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import configure_logger, logger
|
||||
from sglang.multimodal_gen.runtime.utils.trace_wrapper import init_diffusion_tracing
|
||||
from sglang.multimodal_gen.utils import kill_itself_when_parent_died
|
||||
|
||||
_SCHEDULER_SHUTDOWN_TIMEOUT_MS = 5000
|
||||
_WORKER_JOIN_TIMEOUT_S = 10
|
||||
|
||||
@@ -8,7 +8,7 @@ from typing import Any
|
||||
|
||||
import numpy as np
|
||||
|
||||
from sglang.multimodal_gen.utils import dict_to_3d_list
|
||||
from sglang.multimodal_gen.runtime.layers.attention.mask_strategy import dict_to_3d_list
|
||||
|
||||
|
||||
def configure_sta(
|
||||
|
||||
@@ -15,6 +15,7 @@ from sglang.multimodal_gen.runtime.layers.attention.backends.attention_backend i
|
||||
AttentionMetadata,
|
||||
AttentionMetadataBuilder,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.layers.attention.mask_strategy import dict_to_3d_list
|
||||
from sglang.multimodal_gen.runtime.managers.forward_context import (
|
||||
ForwardContext,
|
||||
get_forward_context,
|
||||
@@ -22,7 +23,6 @@ from sglang.multimodal_gen.runtime.managers.forward_context import (
|
||||
from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum
|
||||
from sglang.multimodal_gen.runtime.server_args import get_global_server_args
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.multimodal_gen.utils import dict_to_3d_list
|
||||
|
||||
try:
|
||||
from st_attn import sliding_tile_attention
|
||||
|
||||
@@ -65,7 +65,7 @@ from sglang.multimodal_gen.runtime.managers.forward_context import (
|
||||
get_forward_context,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum
|
||||
from sglang.multimodal_gen.utils import get_compute_dtype
|
||||
from sglang.multimodal_gen.runtime.utils.precision import get_compute_dtype
|
||||
from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph import (
|
||||
eager_on_graph,
|
||||
is_in_breakable_cuda_graph,
|
||||
|
||||
@@ -0,0 +1,67 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/utils.py
|
||||
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
def dict_to_3d_list(
|
||||
mask_strategy: dict[str, Any] | None = None,
|
||||
t_max: int | None = None,
|
||||
l_max: int | None = None,
|
||||
h_max: int | None = None,
|
||||
) -> list[list[list[torch.Tensor | None]]]:
|
||||
"""
|
||||
Convert a dictionary of mask indices to a 3D list of tensors.
|
||||
Args:
|
||||
mask_strategy: keys are "t_l_h", values are torch.Tensor masks.
|
||||
t_max, l_max, h_max: if provided (all three), force the output shape to (t_max, l_max, h_max).
|
||||
If all three are None, infer shape from the data.
|
||||
"""
|
||||
# Case 1: no data, but fixed shape requested
|
||||
if mask_strategy is None:
|
||||
assert t_max is not None and l_max is not None and h_max is not None, (
|
||||
"If mask_strategy is None, you must provide t_max, l_max, and h_max"
|
||||
)
|
||||
return [
|
||||
[[None for _ in range(h_max)] for _ in range(l_max)] for _ in range(t_max)
|
||||
]
|
||||
|
||||
# Parse all keys into integer tuples
|
||||
indices = [tuple(map(int, key.split("_"))) for key in mask_strategy]
|
||||
|
||||
# Decide on dimensions
|
||||
if t_max is None and l_max is None and h_max is None:
|
||||
# fully dynamic: infer from data
|
||||
max_timesteps_idx = max(t for t, _, _ in indices) + 1
|
||||
max_layer_idx = max(l for _, l, _ in indices) + 1 # noqa: E741
|
||||
max_head_idx = max(h for _, _, h in indices) + 1
|
||||
else:
|
||||
# require all three to be provided
|
||||
assert t_max is not None and l_max is not None and h_max is not None, (
|
||||
"Either supply none of (t_max, l_max, h_max) to infer dimensions, "
|
||||
"or supply all three to fix the shape."
|
||||
)
|
||||
max_timesteps_idx = t_max
|
||||
max_layer_idx = l_max
|
||||
max_head_idx = h_max
|
||||
|
||||
# Preallocate
|
||||
result = [
|
||||
[[None for _ in range(max_head_idx)] for _ in range(max_layer_idx)]
|
||||
for _ in range(max_timesteps_idx)
|
||||
]
|
||||
|
||||
# Fill in, skipping any out-of-bounds entries
|
||||
for key, value in mask_strategy.items():
|
||||
t, l, h = map(int, key.split("_")) # noqa: E741
|
||||
if (
|
||||
0 <= t < max_timesteps_idx
|
||||
and 0 <= l < max_layer_idx
|
||||
and 0 <= h < max_head_idx
|
||||
):
|
||||
result[t][l][h] = value
|
||||
# else: silently ignore any key that doesn't fit
|
||||
|
||||
return result
|
||||
@@ -8,6 +8,7 @@ from collections.abc import Generator
|
||||
from contextlib import contextmanager
|
||||
from contextvars import ContextVar
|
||||
from functools import cache
|
||||
from pkgutil import resolve_name
|
||||
from typing import NamedTuple, cast
|
||||
|
||||
import torch
|
||||
@@ -19,10 +20,11 @@ from sglang.multimodal_gen.runtime.layers.attention.backends.attention_backend i
|
||||
from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs, get_global_server_args
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.multimodal_gen.utils import STR_BACKEND_ENV_VAR, resolve_obj_by_qualname
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
STR_BACKEND_ENV_VAR = "SGLANG_DIFFUSION_ATTENTION_BACKEND"
|
||||
|
||||
|
||||
def backend_name_to_enum(backend_name: str) -> AttentionBackendEnum | None:
|
||||
"""
|
||||
@@ -398,7 +400,7 @@ def _cached_get_attn_backend(
|
||||
raise ValueError(
|
||||
f"Invalid attention backend for {current_platform.device_name}"
|
||||
)
|
||||
return cast(type[AttentionBackend], resolve_obj_by_qualname(attention_cls))
|
||||
return cast(type[AttentionBackend], resolve_name(attention_cls))
|
||||
|
||||
|
||||
def _is_backend_supported(
|
||||
|
||||
@@ -20,7 +20,7 @@ from sglang.multimodal_gen.runtime.managers.forward_context import (
|
||||
from sglang.multimodal_gen.runtime.platforms.interface import AttentionBackendEnum
|
||||
from sglang.multimodal_gen.runtime.server_args import get_global_server_args
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.multimodal_gen.utils import get_compute_dtype
|
||||
from sglang.multimodal_gen.runtime.utils.precision import get_compute_dtype
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@@ -32,7 +32,7 @@ from sglang.multimodal_gen.runtime.layers.linear import (
|
||||
from sglang.multimodal_gen.runtime.layers.vocab_parallel_embedding import (
|
||||
VocabParallelEmbedding,
|
||||
)
|
||||
from sglang.multimodal_gen.utils import get_mixed_precision_state
|
||||
from sglang.multimodal_gen.runtime.utils.precision import get_mixed_precision_state
|
||||
|
||||
torch._dynamo.config.recompile_limit = 64
|
||||
|
||||
|
||||
+1
-1
@@ -88,6 +88,7 @@ from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import (
|
||||
load_dict,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.multimodal_gen.runtime.utils.precision_types import PRECISION_TO_TYPE
|
||||
from sglang.multimodal_gen.runtime.utils.quantization_utils import (
|
||||
get_quant_config,
|
||||
get_quant_config_from_safetensors_metadata,
|
||||
@@ -95,7 +96,6 @@ from sglang.multimodal_gen.runtime.utils.quantization_utils import (
|
||||
process_model_weights_after_loading,
|
||||
resolve_comfy_checkpoint_quantization,
|
||||
)
|
||||
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
|
||||
from sglang.srt.layers.linear import LinearBase as SrtLinearBase
|
||||
from sglang.srt.layers.quantization.fp8 import Fp8Config as SrtFp8Config
|
||||
from sglang.srt.layers.quantization.unquant import (
|
||||
|
||||
@@ -45,10 +45,10 @@ from sglang.multimodal_gen.runtime.utils.precision import (
|
||||
resolve_component_precision,
|
||||
resolve_decode_precision,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.utils.precision_types import PRECISION_TO_TYPE
|
||||
from sglang.multimodal_gen.runtime.weights.source import (
|
||||
filter_duplicate_precision_variant_safetensors,
|
||||
)
|
||||
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
|
||||
from sglang.srt.model_loader.checkpoint_quantization import (
|
||||
resolve_checkpoint_quant_spec,
|
||||
)
|
||||
|
||||
@@ -54,10 +54,10 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.host_memory_budget i
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.multimodal_gen.runtime.utils.precision import set_mixed_precision_policy
|
||||
from sglang.multimodal_gen.runtime.utils.quantization_utils import (
|
||||
process_model_weights_after_loading,
|
||||
)
|
||||
from sglang.multimodal_gen.utils import set_mixed_precision_policy
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@@ -80,6 +80,10 @@ from sglang.multimodal_gen.runtime.post_training.gpu_worker_post_training_mixin
|
||||
GPUWorkerPostTrainingMixin,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.realtime.session import RealtimeSessionCache
|
||||
from sglang.multimodal_gen.runtime.realtime.video import (
|
||||
RAW_RGB_CONTENT_TYPE,
|
||||
build_raw_rgb_frame_batches,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.server_args import PortArgs, ServerArgs
|
||||
from sglang.multimodal_gen.runtime.utils.common import set_cuda_arch, set_musa_arch
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import (
|
||||
@@ -90,17 +94,13 @@ from sglang.multimodal_gen.runtime.utils.perf_logger import (
|
||||
PerformanceLogger,
|
||||
capture_memory_snapshot,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.utils.process import kill_itself_when_parent_died
|
||||
from sglang.multimodal_gen.runtime.utils.profiler import maybe_record_function
|
||||
from sglang.multimodal_gen.runtime.utils.realtime_video import (
|
||||
RAW_RGB_CONTENT_TYPE,
|
||||
build_raw_rgb_frame_batches,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.utils.trace_wrapper import (
|
||||
DiffStage,
|
||||
init_diffusion_tracing,
|
||||
trace_slice,
|
||||
)
|
||||
from sglang.multimodal_gen.utils import kill_itself_when_parent_died
|
||||
from sglang.srt.environ import third_party_cache_defaults
|
||||
from sglang.srt.utils.network import NetworkAddress
|
||||
|
||||
|
||||
@@ -16,6 +16,7 @@ from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
|
||||
from sglang.multimodal_gen.runtime.disaggregation.scheduler_mixin import (
|
||||
SchedulerDisaggMixin,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.distributed.utils import broadcast_pyobj
|
||||
from sglang.multimodal_gen.runtime.entrypoints.control_requests import (
|
||||
GetDisaggStatsReq,
|
||||
ListLorasReq,
|
||||
@@ -62,7 +63,6 @@ from sglang.multimodal_gen.runtime.server_warmup import (
|
||||
should_return_warmup_result,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.utils.common import get_zmq_socket
|
||||
from sglang.multimodal_gen.runtime.utils.distributed import broadcast_pyobj
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.multimodal_gen.runtime.utils.profiler import maybe_record_function
|
||||
from sglang.multimodal_gen.runtime.utils.trace_wrapper import DiffStage, trace_slice
|
||||
|
||||
@@ -78,18 +78,18 @@ from sglang.multimodal_gen.runtime.models.dits.wanvideo import (
|
||||
WanTimeTextImageEmbedding,
|
||||
WanTransformer3DModel,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.lingbot_world.constants import (
|
||||
from sglang.multimodal_gen.runtime.platforms import (
|
||||
AttentionBackendEnum,
|
||||
current_platform,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.platforms.aiter import USE_AITER
|
||||
from sglang.multimodal_gen.runtime.realtime.lingbot_world import (
|
||||
LINGBOT_C2WS_PLUCKER_EMB_CACHE,
|
||||
LINGBOT_CAM_CONDITIONER_CACHE,
|
||||
LINGBOT_ROPE_CACHE,
|
||||
LINGBOT_SEQUENCE_SHARD_ROPE_CACHE,
|
||||
LINGBOT_TIME_EMBEDDINGS_CACHE,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.platforms import (
|
||||
AttentionBackendEnum,
|
||||
current_platform,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.platforms.aiter import USE_AITER
|
||||
from sglang.multimodal_gen.runtime.realtime.states import (
|
||||
get_realtime_causal_dit_state,
|
||||
)
|
||||
|
||||
@@ -13,6 +13,7 @@ from sglang.multimodal_gen.configs.models.fsdp import (
|
||||
from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import (
|
||||
LayerwiseOffloadableModuleMixin,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.models.dits import sana_wm_parity as parity_probe
|
||||
from sglang.multimodal_gen.runtime.models.dits.base import CachableDiT
|
||||
|
||||
# Re-exported for back-compat: callers import these names from this module path.
|
||||
@@ -78,9 +79,6 @@ from sglang.multimodal_gen.runtime.models.dits.sana_wm_components import ( # no
|
||||
compute_chunk_plucker,
|
||||
process_camera_conditions_ucpe,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.sana_wm import (
|
||||
parity_probe,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -12,7 +12,7 @@ from sglang.multimodal_gen.runtime.models import ( # noqa: F401
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
|
||||
from sglang.multimodal_gen.runtime.utils.precision_types import PRECISION_TO_TYPE
|
||||
|
||||
|
||||
def load_model_and_tokenizer(
|
||||
|
||||
@@ -38,6 +38,7 @@ from sglang.multimodal_gen.runtime.models.vlas.pi05_core import Pi05CoreModel
|
||||
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import maybe_download_model
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.multimodal_gen.runtime.utils.precision import set_mixed_precision_policy
|
||||
from sglang.multimodal_gen.runtime.vla.cuda_graph import (
|
||||
VLADenoiseGraphRunner,
|
||||
VLADenoiseGraphSignature,
|
||||
@@ -62,7 +63,6 @@ from sglang.multimodal_gen.runtime.vla.prompt_bucketing import (
|
||||
effective_token_length,
|
||||
select_prompt_token_bucket,
|
||||
)
|
||||
from sglang.multimodal_gen.utils import set_mixed_precision_policy
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@@ -35,8 +35,10 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages import (
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.multimodal_gen.runtime.utils.precision import resolve_precision
|
||||
from sglang.multimodal_gen.utils import set_mixed_precision_policy
|
||||
from sglang.multimodal_gen.runtime.utils.precision import (
|
||||
resolve_precision,
|
||||
set_mixed_precision_policy,
|
||||
)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@@ -38,8 +38,10 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages import (
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.multimodal_gen.runtime.utils.precision import resolve_precision
|
||||
from sglang.multimodal_gen.utils import set_mixed_precision_policy
|
||||
from sglang.multimodal_gen.runtime.utils.precision import (
|
||||
resolve_precision,
|
||||
set_mixed_precision_policy,
|
||||
)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@@ -64,7 +64,7 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.h
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
|
||||
from sglang.multimodal_gen.runtime.utils.precision_types import PRECISION_TO_TYPE
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@@ -15,7 +15,7 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.l
|
||||
LongCatImageEditTextEncodingStage,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
|
||||
from sglang.multimodal_gen.runtime.utils.precision_types import PRECISION_TO_TYPE
|
||||
|
||||
|
||||
def _prepare_mu(batch, server_args):
|
||||
|
||||
@@ -19,7 +19,7 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.progressive_resolution.
|
||||
QwenImageProgressiveDenoisingStage,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
|
||||
from sglang.multimodal_gen.runtime.utils.precision_types import PRECISION_TO_TYPE
|
||||
|
||||
|
||||
def prepare_mu(batch: Req, server_args: ServerArgs):
|
||||
|
||||
@@ -10,6 +10,7 @@ from sglang.multimodal_gen.runtime.distributed.parallel_state import (
|
||||
get_world_group,
|
||||
get_world_rank,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.distributed.utils import broadcast_pyobj
|
||||
from sglang.multimodal_gen.runtime.pipelines_core import Req
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.executors.pipeline_executor import (
|
||||
PipelineExecutor,
|
||||
@@ -20,7 +21,6 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.base import (
|
||||
StageParallelismType,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
from sglang.multimodal_gen.runtime.utils.distributed import broadcast_pyobj
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -12,6 +12,7 @@ in a functional manner, reducing the need for explicit parameter passing.
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import math
|
||||
import os
|
||||
import pprint
|
||||
from collections import Counter
|
||||
@@ -38,7 +39,6 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import (
|
||||
init_logger,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.utils.perf_logger import RequestMetrics
|
||||
from sglang.multimodal_gen.utils import align_to
|
||||
from sglang.srt.observability.trace import TraceNullContext, TraceReqContext
|
||||
|
||||
logger = init_logger(__name__)
|
||||
@@ -46,6 +46,10 @@ logger = init_logger(__name__)
|
||||
SAMPLING_PARAMS_FIELDS = {f.name for f in fields(SamplingParams)}
|
||||
|
||||
|
||||
def _align_to(value: int, alignment: int) -> int:
|
||||
return int(math.ceil(value / alignment) * alignment)
|
||||
|
||||
|
||||
@dataclass
|
||||
class BatchMetricsWindow:
|
||||
"""Counters accumulated between dynamic batching metric logs.
|
||||
@@ -425,11 +429,11 @@ class Req:
|
||||
|
||||
# TODO: in some cases (e.g., TI2I), height and weight might be undecided at this moment
|
||||
if self.height:
|
||||
target_height = align_to(self.height, 16)
|
||||
target_height = _align_to(self.height, 16)
|
||||
else:
|
||||
target_height = -1
|
||||
if self.width:
|
||||
target_width = align_to(self.width, 16)
|
||||
target_width = _align_to(self.width, 16)
|
||||
else:
|
||||
target_width = -1
|
||||
|
||||
|
||||
@@ -126,6 +126,10 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload im
|
||||
LayerwiseOffloadableModuleMixin,
|
||||
is_layerwise_offloaded_module,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.component_loading import (
|
||||
load_transformer_if_needed,
|
||||
register_loaded_transformer,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.stages.base import (
|
||||
PipelineStage,
|
||||
@@ -152,10 +156,6 @@ from sglang.multimodal_gen.runtime.post_training.rollout_denoising_mixin import
|
||||
RolloutDenoisingMixin,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
from sglang.multimodal_gen.runtime.utils.component_load import (
|
||||
load_transformer_if_needed,
|
||||
register_loaded_transformer,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.multimodal_gen.runtime.utils.nvtx_pytorch_hooks import maybe_nvtx_range
|
||||
from sglang.multimodal_gen.runtime.utils.perf_logger import StageProfiler
|
||||
|
||||
@@ -28,7 +28,6 @@ from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.multimodal_gen.runtime.utils.vision import load_image, load_video
|
||||
from sglang.multimodal_gen.utils import best_output_size
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -39,6 +38,31 @@ V = StageValidators
|
||||
# TODO: since this might change sampling params after logging, should be do this beforehand?
|
||||
|
||||
|
||||
def _best_output_size(w, h, dw, dh, expected_area):
|
||||
# float output size
|
||||
ratio = w / h
|
||||
ow = (expected_area * ratio) ** 0.5
|
||||
oh = expected_area / ow
|
||||
|
||||
# process width first
|
||||
ow1 = int(ow // dw * dw)
|
||||
oh1 = int(expected_area / ow1 // dh * dh)
|
||||
assert ow1 % dw == 0 and oh1 % dh == 0 and ow1 * oh1 <= expected_area
|
||||
ratio1 = ow1 / oh1
|
||||
|
||||
# process height first
|
||||
oh2 = int(oh // dh * dh)
|
||||
ow2 = int(expected_area / oh2 // dw * dw)
|
||||
assert oh2 % dh == 0 and ow2 % dw == 0 and ow2 * oh2 <= expected_area
|
||||
ratio2 = ow2 / oh2
|
||||
|
||||
# compare ratios
|
||||
if max(ratio / ratio1, ratio1 / ratio) < max(ratio / ratio2, ratio2 / ratio):
|
||||
return ow1, oh1
|
||||
else:
|
||||
return ow2, oh2
|
||||
|
||||
|
||||
class InputValidationStage(PipelineStage):
|
||||
"""
|
||||
Stage for validating and preparing inputs for diffusion pipelines.
|
||||
@@ -224,7 +248,7 @@ class InputValidationStage(PipelineStage):
|
||||
)
|
||||
dh, dw = patch_size[1] * vae_stride, patch_size[2] * vae_stride
|
||||
max_area = 704 * 1280
|
||||
ow, oh = best_output_size(iw, ih, dw, dh, max_area)
|
||||
ow, oh = _best_output_size(iw, ih, dw, dh, max_area)
|
||||
|
||||
scale = max(ow / iw, oh / ih)
|
||||
img = img.resize((round(iw * scale), round(ih * scale)), Image.LANCZOS)
|
||||
|
||||
+1
-1
@@ -14,7 +14,7 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.decoding import (
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
|
||||
from sglang.multimodal_gen.runtime.utils.precision_types import PRECISION_TO_TYPE
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
+1
-1
@@ -36,8 +36,8 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.base import (
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.multimodal_gen.runtime.utils.perf_logger import StageProfiler
|
||||
from sglang.multimodal_gen.runtime.utils.precision_types import PRECISION_TO_TYPE
|
||||
from sglang.multimodal_gen.runtime.utils.profiler import SGLDiffusionProfiler
|
||||
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
+4
-4
@@ -20,6 +20,10 @@ from sglang.multimodal_gen.runtime.managers.forward_context import set_forward_c
|
||||
from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager import (
|
||||
ComponentUse,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.component_loading import (
|
||||
load_transformer_if_needed,
|
||||
register_loaded_transformer,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch, Req
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.stages.base import PipelineStage
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.stages.denoising import (
|
||||
@@ -33,10 +37,6 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import (
|
||||
VerificationResult,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
from sglang.multimodal_gen.runtime.utils.component_load import (
|
||||
load_transformer_if_needed,
|
||||
register_loaded_transformer,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.multimodal_gen.runtime.utils.mesh3d_utils import export_to_trimesh
|
||||
|
||||
|
||||
+1
-1
@@ -39,7 +39,7 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import (
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.multimodal_gen.runtime.utils.nvtx_pytorch_hooks import maybe_nvtx_range
|
||||
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
|
||||
from sglang.multimodal_gen.runtime.utils.precision_types import PRECISION_TO_TYPE
|
||||
|
||||
SEQUENCE_PADDING_INDICATOR = -1
|
||||
OUTPUT_IMAGE_INDICATOR = 2
|
||||
|
||||
+7
-7
@@ -24,13 +24,6 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.causal_denoising import
|
||||
CausalDMDRealtimeCacheContext,
|
||||
CausalKVCache,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.lingbot_world.constants import (
|
||||
LINGBOT_C2WS_PLUCKER_EMB_CACHE,
|
||||
LINGBOT_CAM_CONDITIONER_CACHE,
|
||||
LINGBOT_CAMERA_ACTIONS_CONDITION,
|
||||
LINGBOT_INTERACTIVE_KV_WINDOW_CACHE,
|
||||
LINGBOT_PROMPT_UPDATED_CONDITION,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import (
|
||||
StageValidators as V,
|
||||
)
|
||||
@@ -38,6 +31,13 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import (
|
||||
VerificationResult,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||
from sglang.multimodal_gen.runtime.realtime.lingbot_world import (
|
||||
LINGBOT_C2WS_PLUCKER_EMB_CACHE,
|
||||
LINGBOT_CAM_CONDITIONER_CACHE,
|
||||
LINGBOT_CAMERA_ACTIONS_CONDITION,
|
||||
LINGBOT_INTERACTIVE_KV_WINDOW_CACHE,
|
||||
LINGBOT_PROMPT_UPDATED_CONDITION,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
|
||||
|
||||
+1
-1
@@ -71,8 +71,8 @@ from sglang.multimodal_gen.runtime.utils.perf_logger import StageProfiler
|
||||
from sglang.multimodal_gen.runtime.utils.precision import (
|
||||
autocast_context as precision_autocast_context,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.utils.precision_types import PRECISION_TO_TYPE
|
||||
from sglang.multimodal_gen.runtime.utils.profiler import SGLDiffusionProfiler
|
||||
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
|
||||
from sglang.srt.utils.common import get_compiler_backend
|
||||
|
||||
_is_npu = current_platform.is_npu()
|
||||
|
||||
+1
-1
@@ -47,7 +47,7 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.text_encoding import (
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
|
||||
from sglang.multimodal_gen.runtime.utils.precision_types import PRECISION_TO_TYPE
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
+2
-2
@@ -26,6 +26,7 @@ import torch
|
||||
|
||||
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
|
||||
from sglang.multimodal_gen.runtime.managers.forward_context import set_forward_context
|
||||
from sglang.multimodal_gen.runtime.models.dits import sana_wm_parity as parity_probe
|
||||
from sglang.multimodal_gen.runtime.models.dits.sana_wm_components import (
|
||||
compute_chunk_plucker,
|
||||
)
|
||||
@@ -36,9 +37,8 @@ from sglang.multimodal_gen.runtime.realtime.states import (
|
||||
get_realtime_causal_dit_state,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
|
||||
from sglang.multimodal_gen.runtime.utils.precision_types import PRECISION_TO_TYPE
|
||||
|
||||
from . import parity_probe
|
||||
from .base import (
|
||||
_SANA_WM_DEFAULT_ROTATION_SPEED_DEG,
|
||||
_SANA_WM_DEFAULT_TRANSLATION_SPEED,
|
||||
|
||||
+1
-1
@@ -37,7 +37,7 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.base import (
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
|
||||
from sglang.multimodal_gen.runtime.utils.precision_types import PRECISION_TO_TYPE
|
||||
|
||||
from .base import (
|
||||
SanaWMDecodingStage,
|
||||
|
||||
+2
-2
@@ -27,6 +27,7 @@ from sglang.multimodal_gen.runtime.managers.forward_context import set_forward_c
|
||||
from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager import (
|
||||
ComponentUse,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.models.dits import sana_wm_parity as parity_probe
|
||||
from sglang.multimodal_gen.runtime.models.dits.sana_wm import (
|
||||
_NUM_STREAM_CACHE_SLOTS,
|
||||
)
|
||||
@@ -47,9 +48,8 @@ from sglang.multimodal_gen.runtime.realtime.states import (
|
||||
get_realtime_causal_dit_state,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
|
||||
from sglang.multimodal_gen.runtime.utils.precision_types import PRECISION_TO_TYPE
|
||||
|
||||
from . import parity_probe
|
||||
from .base import (
|
||||
_align_sana_wm_cfg_text_conditions,
|
||||
_cat_optional_tensors,
|
||||
|
||||
+1
-1
@@ -19,6 +19,7 @@ import torch
|
||||
from torch import nn
|
||||
|
||||
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
|
||||
from sglang.multimodal_gen.runtime.models.dits import sana_wm_parity as parity_probe
|
||||
from sglang.multimodal_gen.runtime.models.dits.sana_wm_refiner_transformer import (
|
||||
pack_latents,
|
||||
unpack_latents,
|
||||
@@ -26,7 +27,6 @@ from sglang.multimodal_gen.runtime.models.dits.sana_wm_refiner_transformer impor
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
|
||||
from . import parity_probe
|
||||
from .refiner import (
|
||||
STAGE_2_DISTILLED_SIGMA_VALUES,
|
||||
SanaWMLTX2RefinerStage,
|
||||
|
||||
+2
-2
@@ -17,7 +17,6 @@ from sglang.multimodal_gen.runtime.distributed import (
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
from sglang.multimodal_gen.utils import masks_like
|
||||
|
||||
|
||||
def should_apply_wan_ti2v(batch: Req, server_args: ServerArgs) -> bool:
|
||||
@@ -65,7 +64,8 @@ def prepare_wan_ti2v_latents(
|
||||
assert latent_model_input.ndim == 5
|
||||
|
||||
latent_for_mask = latent_model_input.squeeze(0)
|
||||
_, reserved_frames_masks = masks_like([latent_for_mask], zero=True)
|
||||
reserved_frames_masks = [torch.ones_like(latent_for_mask)]
|
||||
reserved_frames_masks[0][:, 0] = 0
|
||||
reserved_frames_mask = reserved_frames_masks[0].unsqueeze(0)
|
||||
|
||||
latents = (
|
||||
|
||||
@@ -20,7 +20,7 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.decoding import (
|
||||
scale_and_shift_latents,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
|
||||
from sglang.multimodal_gen.runtime.utils.precision_types import PRECISION_TO_TYPE
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
|
||||
@@ -21,7 +21,7 @@ from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
from sglang.multimodal_gen.runtime.utils.precision import (
|
||||
autocast_context as precision_autocast_context,
|
||||
)
|
||||
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
|
||||
from sglang.multimodal_gen.runtime.utils.precision_types import PRECISION_TO_TYPE
|
||||
|
||||
|
||||
class RealtimeVAEState(BaseRealtimeState):
|
||||
|
||||
@@ -5,6 +5,7 @@
|
||||
|
||||
import os
|
||||
import traceback
|
||||
from pkgutil import resolve_name
|
||||
|
||||
# imported by other files, do not remove
|
||||
from sglang.multimodal_gen.runtime.platforms.interface import ( # noqa: F401
|
||||
@@ -13,7 +14,7 @@ from sglang.multimodal_gen.runtime.platforms.interface import ( # noqa: F401
|
||||
PlatformEnum,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.multimodal_gen.utils import resolve_obj_by_qualname
|
||||
from sglang.multimodal_gen.third_party import pynvml
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -22,9 +23,6 @@ def cuda_platform_plugin() -> str | None:
|
||||
is_cuda = False
|
||||
|
||||
try:
|
||||
from sglang.multimodal_gen.utils import import_pynvml
|
||||
|
||||
pynvml = import_pynvml() # type: ignore[no-untyped-call]
|
||||
pynvml.nvmlInit()
|
||||
try:
|
||||
# NOTE: Edge case: sgl_diffusion cpu build on a GPU machine.
|
||||
@@ -267,7 +265,7 @@ def __getattr__(name: str):
|
||||
global _current_platform
|
||||
if _current_platform is None:
|
||||
platform_cls_qualname = resolve_current_platform_cls_qualname()
|
||||
_current_platform = resolve_obj_by_qualname(platform_cls_qualname)()
|
||||
_current_platform = resolve_name(platform_cls_qualname)()
|
||||
global _init_trace
|
||||
_init_trace = "".join(traceback.format_stack())
|
||||
return _current_platform
|
||||
|
||||
@@ -23,7 +23,7 @@ from sglang.multimodal_gen.runtime.platforms.interface import (
|
||||
PlatformEnum,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.multimodal_gen.utils import import_pynvml
|
||||
from sglang.multimodal_gen.third_party import pynvml
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -38,8 +38,6 @@ _DYNAMIC_CUDNN_SDPA_BACKEND_CLS_STR = "sglang.multimodal_gen.runtime.layers.atte
|
||||
_P = ParamSpec("_P")
|
||||
_R = TypeVar("_R")
|
||||
|
||||
pynvml = import_pynvml() # type: ignore[no-untyped-call]
|
||||
|
||||
# pytorch 2.5 uses cudnn sdpa by default, which will cause crash on some models
|
||||
# see https://github.com/huggingface/diffusers/issues/9704 for details
|
||||
torch.backends.cuda.enable_cudnn_sdp(False)
|
||||
|
||||
@@ -8,13 +8,13 @@ import enum
|
||||
import random
|
||||
from collections.abc import Callable
|
||||
from functools import lru_cache
|
||||
from pkgutil import resolve_name
|
||||
from typing import TYPE_CHECKING, Any, NamedTuple
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.multimodal_gen.utils import resolve_obj_by_qualname
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.multimodal_gen.runtime.layers.attention.backends.attention_backend import (
|
||||
@@ -451,7 +451,7 @@ class Platform:
|
||||
|
||||
def get_attn_backend(self, *args, **kwargs) -> AttentionImpl:
|
||||
attention_cls_str = self.get_attn_backend_cls_str(*args, **kwargs)
|
||||
return resolve_obj_by_qualname(attention_cls_str)
|
||||
return resolve_name(attention_cls_str)
|
||||
|
||||
def tensor_on_device(self, t: torch.Tensor) -> bool:
|
||||
"""Check if a tensor is on the current platform's device."""
|
||||
|
||||
@@ -6,11 +6,11 @@ from __future__ import annotations
|
||||
from typing import ClassVar, Literal
|
||||
|
||||
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
|
||||
from sglang.multimodal_gen.runtime.utils.argparse import FlexibleArgumentParser
|
||||
from sglang.multimodal_gen.runtime.utils.common import (
|
||||
format_tcp_endpoint,
|
||||
parse_tcp_host_port,
|
||||
)
|
||||
from sglang.multimodal_gen.utils import FlexibleArgumentParser
|
||||
|
||||
|
||||
class DisaggServerArgsMixin:
|
||||
|
||||
@@ -27,6 +27,7 @@ from sglang.multimodal_gen.configs.pipeline_configs.ltx_2 import (
|
||||
)
|
||||
from sglang.multimodal_gen.configs.quantization.nunchaku import NunchakuSVDQuantArgs
|
||||
from sglang.multimodal_gen.configs.quantization.qvg_kv import QVGKVQuantArgs
|
||||
from sglang.multimodal_gen.configs.utils import expand_path_fields
|
||||
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
|
||||
from sglang.multimodal_gen.runtime.layers.quantization.configs.nunchaku_config import (
|
||||
NunchakuConfig,
|
||||
@@ -68,6 +69,10 @@ from sglang.multimodal_gen.runtime.server_args.auto_tune import (
|
||||
ServerArgsAutoTuner,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.server_args.disagg import DisaggServerArgsMixin
|
||||
from sglang.multimodal_gen.runtime.utils.argparse import (
|
||||
FlexibleArgumentParser,
|
||||
StoreBoolean,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.utils.common import (
|
||||
is_port_available,
|
||||
is_valid_ipv6_address,
|
||||
@@ -78,15 +83,10 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import (
|
||||
configure_logger,
|
||||
init_logger,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.utils.precision_types import PRECISION_TO_TYPE
|
||||
from sglang.multimodal_gen.runtime.weights.source import (
|
||||
is_explicit_weight_file_reference,
|
||||
)
|
||||
from sglang.multimodal_gen.utils import (
|
||||
PRECISION_TO_TYPE,
|
||||
FlexibleArgumentParser,
|
||||
StoreBoolean,
|
||||
expand_path_fields,
|
||||
)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@@ -0,0 +1,263 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/utils.py
|
||||
|
||||
import argparse
|
||||
import sys
|
||||
from typing import Any
|
||||
|
||||
import yaml
|
||||
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import (
|
||||
SortedHelpFormatter,
|
||||
init_logger,
|
||||
)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class StoreBoolean(argparse.Action):
|
||||
def __init__(self, option_strings, dest, default=False, required=False, help=None):
|
||||
super().__init__(
|
||||
option_strings=option_strings,
|
||||
dest=dest,
|
||||
nargs="?",
|
||||
const=True,
|
||||
default=default,
|
||||
required=required,
|
||||
help=help,
|
||||
)
|
||||
|
||||
def __call__(self, parser, namespace, values, option_string=None):
|
||||
if values is None:
|
||||
setattr(namespace, self.dest, True)
|
||||
elif isinstance(values, str):
|
||||
if values.lower() == "true":
|
||||
setattr(namespace, self.dest, True)
|
||||
elif values.lower() == "false":
|
||||
setattr(namespace, self.dest, False)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Invalid boolean value: {values}. Expected 'true' or 'false'."
|
||||
)
|
||||
else:
|
||||
setattr(namespace, self.dest, bool(values))
|
||||
|
||||
|
||||
class FlexibleArgumentParser(argparse.ArgumentParser):
|
||||
"""ArgumentParser that allows both underscore and dash in names."""
|
||||
|
||||
def __init__(self, *args, **kwargs) -> None:
|
||||
# Set the default 'formatter_class' to SortedHelpFormatter
|
||||
if "formatter_class" not in kwargs:
|
||||
kwargs["formatter_class"] = SortedHelpFormatter
|
||||
super().__init__(*args, **kwargs)
|
||||
|
||||
def parse_args( # type: ignore[override]
|
||||
self, args=None, namespace=None
|
||||
) -> argparse.Namespace:
|
||||
if args is None:
|
||||
args = sys.argv[1:]
|
||||
|
||||
if any(arg.startswith("--config") for arg in args):
|
||||
args = self._pull_args_from_config(args)
|
||||
|
||||
# Convert underscores to dashes and vice versa in argument names
|
||||
processed_args = []
|
||||
for arg in args:
|
||||
if arg.startswith("--"):
|
||||
if "=" in arg:
|
||||
key, value = arg.split("=", 1)
|
||||
key = "--" + key[len("--") :].replace("_", "-")
|
||||
processed_args.append(f"{key}={value}")
|
||||
else:
|
||||
processed_args.append("--" + arg[len("--") :].replace("_", "-"))
|
||||
elif arg.startswith("-O") and arg != "-O" and len(arg) == 2:
|
||||
# allow -O flag to be used without space, e.g. -O3
|
||||
processed_args.append("-O")
|
||||
processed_args.append(arg[2:])
|
||||
else:
|
||||
processed_args.append(arg)
|
||||
|
||||
namespace = super().parse_args(processed_args, namespace)
|
||||
|
||||
# Track which arguments were explicitly provided
|
||||
namespace._provided = set()
|
||||
|
||||
i = 0
|
||||
while i < len(args):
|
||||
arg = args[i]
|
||||
if arg.startswith("--"):
|
||||
# Handle --key=value format
|
||||
if "=" in arg:
|
||||
key = arg.split("=")[0][2:].replace("-", "_")
|
||||
namespace._provided.add(key)
|
||||
i += 1
|
||||
# Handle --key value format
|
||||
else:
|
||||
key = arg[2:].replace("-", "_")
|
||||
namespace._provided.add(key)
|
||||
# Skip the value if there is one
|
||||
if i + 1 < len(args) and not args[i + 1].startswith("-"):
|
||||
i += 2
|
||||
else:
|
||||
i += 1
|
||||
else:
|
||||
i += 1
|
||||
|
||||
return namespace # type: ignore[no-any-return]
|
||||
|
||||
def _pull_args_from_config(self, args: list[str]) -> list[str]:
|
||||
"""Method to pull arguments specified in the config file
|
||||
into the command-line args variable.
|
||||
|
||||
The arguments in config file will be inserted between
|
||||
the argument list.
|
||||
|
||||
example:
|
||||
```yaml
|
||||
port: 12323
|
||||
tensor-parallel-size: 4
|
||||
```
|
||||
```python
|
||||
$: vllm {serve,chat,complete} "facebook/opt-12B" \
|
||||
--config config.yaml -tp 2
|
||||
$: args = [
|
||||
"serve,chat,complete",
|
||||
"facebook/opt-12B",
|
||||
'--config', 'config.yaml',
|
||||
'-tp', '2'
|
||||
]
|
||||
$: args = [
|
||||
"serve,chat,complete",
|
||||
"facebook/opt-12B",
|
||||
'--port', '12323',
|
||||
'--tp-size', '4',
|
||||
'-tp', '2'
|
||||
]
|
||||
```
|
||||
|
||||
Please note how the config args are inserted after the sub command.
|
||||
this way the order of priorities is maintained when these are args
|
||||
parsed by super().
|
||||
"""
|
||||
index = -1
|
||||
config_arg = None
|
||||
for i, arg in enumerate(args):
|
||||
if arg.startswith("--config"):
|
||||
if index != -1:
|
||||
raise ValueError("More than one config file specified!")
|
||||
index = i
|
||||
config_arg = arg
|
||||
|
||||
if config_arg is None:
|
||||
return args
|
||||
args_before_config = args[:index]
|
||||
if "=" in config_arg:
|
||||
file_path = config_arg.split("=", 1)[1]
|
||||
args_after_config = args[index + 1 :]
|
||||
else:
|
||||
if index == len(args) - 1:
|
||||
raise ValueError(
|
||||
"No config file specified! "
|
||||
"Please check your command-line arguments."
|
||||
)
|
||||
file_path = args[index + 1]
|
||||
args_after_config = args[index + 2 :]
|
||||
|
||||
config_args = self._load_config_file(file_path)
|
||||
|
||||
# 0th index is for {serve,chat,complete}
|
||||
# followed by model_tag (only for serve)
|
||||
# followed by config args
|
||||
# followed by rest of cli args.
|
||||
# maintaining this order will enforce the precedence
|
||||
# of cli > config > defaults
|
||||
if args[0] == "serve":
|
||||
if index == 1:
|
||||
raise ValueError(
|
||||
"No model_tag specified! Please check your command-line arguments."
|
||||
)
|
||||
command = args_before_config[0]
|
||||
model_tag = args_before_config[1]
|
||||
other_args_before = args_before_config[2:]
|
||||
args = (
|
||||
[command, model_tag]
|
||||
+ config_args
|
||||
+ other_args_before
|
||||
+ args_after_config
|
||||
)
|
||||
else:
|
||||
command = args_before_config[0]
|
||||
other_args_before = args_before_config[1:]
|
||||
args = [command] + config_args + other_args_before + args_after_config
|
||||
|
||||
return args
|
||||
|
||||
def _load_config_file(self, file_path: str) -> list[str]:
|
||||
"""Loads a yaml file and returns the key value pairs as a
|
||||
flattened list with argparse like pattern
|
||||
```yaml
|
||||
port: 12323
|
||||
tensor-parallel-size: 4
|
||||
vae_config:
|
||||
load_encoder: false
|
||||
load_decoder: true
|
||||
```
|
||||
returns:
|
||||
processed_args: list[str] = [
|
||||
'--port': '12323',
|
||||
'--tp-size': '4',
|
||||
'--vae-config.load-encoder': 'false',
|
||||
'--vae-config.load-decoder': 'true'
|
||||
]
|
||||
"""
|
||||
|
||||
extension: str = file_path.split(".")[-1]
|
||||
if extension not in ("yaml", "yml", "json"):
|
||||
raise ValueError(
|
||||
"Config file must be of a yaml/yml/json type.\
|
||||
%s supplied",
|
||||
extension,
|
||||
)
|
||||
|
||||
processed_args: list[str] = []
|
||||
|
||||
config: dict[str, Any] = {}
|
||||
try:
|
||||
with open(file_path) as config_file:
|
||||
config = yaml.safe_load(config_file)
|
||||
except Exception as ex:
|
||||
logger.error(
|
||||
"Unable to read the config file at %s. \
|
||||
Make sure path is correct",
|
||||
file_path,
|
||||
)
|
||||
raise ex
|
||||
|
||||
store_boolean_arguments = [
|
||||
action.dest for action in self._actions if isinstance(action, StoreBoolean)
|
||||
]
|
||||
|
||||
def process_dict(prefix: str, d: dict[str, Any]):
|
||||
for key, value in d.items():
|
||||
full_key = f"{prefix}.{key}" if prefix else key
|
||||
|
||||
if isinstance(value, bool) and full_key not in store_boolean_arguments:
|
||||
if value:
|
||||
processed_args.append("--" + full_key)
|
||||
else:
|
||||
processed_args.append("--" + full_key)
|
||||
processed_args.append("false")
|
||||
elif isinstance(value, list):
|
||||
processed_args.append("--" + full_key)
|
||||
for item in value:
|
||||
processed_args.append(str(item))
|
||||
elif isinstance(value, dict):
|
||||
process_dict(full_key, value)
|
||||
else:
|
||||
processed_args.append("--" + full_key)
|
||||
processed_args.append(str(value))
|
||||
|
||||
process_dict("", config)
|
||||
|
||||
return processed_args
|
||||
@@ -4,10 +4,7 @@ import ipaddress
|
||||
import logging
|
||||
import os
|
||||
import platform
|
||||
import signal
|
||||
import socket
|
||||
import sys
|
||||
import threading
|
||||
from functools import lru_cache
|
||||
from typing import Any
|
||||
|
||||
@@ -19,45 +16,6 @@ import zmq
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def kill_process_tree(parent_pid, include_parent: bool = True, skip_pid: int = None):
|
||||
"""Kill the process and all its child processes."""
|
||||
# Remove sigchld handler to avoid spammy logs.
|
||||
if threading.current_thread() is threading.main_thread():
|
||||
signal.signal(signal.SIGCHLD, signal.SIG_DFL)
|
||||
|
||||
if parent_pid is None:
|
||||
parent_pid = os.getpid()
|
||||
include_parent = False
|
||||
|
||||
try:
|
||||
itself = psutil.Process(parent_pid)
|
||||
except psutil.NoSuchProcess:
|
||||
return
|
||||
|
||||
children = itself.children(recursive=True)
|
||||
for child in children:
|
||||
if child.pid == skip_pid:
|
||||
continue
|
||||
try:
|
||||
child.kill()
|
||||
except psutil.NoSuchProcess:
|
||||
pass
|
||||
|
||||
if include_parent:
|
||||
try:
|
||||
if parent_pid == os.getpid():
|
||||
itself.kill()
|
||||
sys.exit(0)
|
||||
|
||||
itself.kill()
|
||||
|
||||
# Sometime processes cannot be killed with SIGKILL (e.g, PID=1 launched by kubernetes),
|
||||
# so we send an additional signal to kill them.
|
||||
itself.send_signal(signal.SIGQUIT)
|
||||
except psutil.NoSuchProcess:
|
||||
pass
|
||||
|
||||
|
||||
def add_prefix(name: str, prefix: str) -> str:
|
||||
"""Add a weight path prefix to a module name.
|
||||
|
||||
|
||||
@@ -1,234 +0,0 @@
|
||||
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
|
||||
|
||||
import pickle
|
||||
from typing import Any, List, Optional
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
|
||||
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||
|
||||
|
||||
def broadcast_pyobj(
|
||||
data: List[Any],
|
||||
rank: int,
|
||||
dist_group: Optional[torch.distributed.ProcessGroup] = None,
|
||||
src: int = 0,
|
||||
force_cpu_device: bool = True,
|
||||
):
|
||||
"""Broadcast inputs from src rank to all other ranks with torch.dist backend.
|
||||
The `rank` here refer to the source rank on global process group (regardless
|
||||
of dist_group argument).
|
||||
"""
|
||||
|
||||
device = torch.device(
|
||||
current_platform.device_type if not force_cpu_device else "cpu"
|
||||
)
|
||||
|
||||
if rank == src:
|
||||
if data is None or len(data) == 0:
|
||||
tensor_size = torch.tensor([0], dtype=torch.long, device=device)
|
||||
dist.broadcast(tensor_size, src=src, group=dist_group)
|
||||
else:
|
||||
serialized_data = pickle.dumps(data)
|
||||
size = len(serialized_data)
|
||||
|
||||
tensor_data = torch.ByteTensor(
|
||||
np.frombuffer(serialized_data, dtype=np.uint8).copy()
|
||||
).to(device)
|
||||
tensor_size = torch.tensor([size], dtype=torch.long, device=device)
|
||||
|
||||
dist.broadcast(tensor_size, src=src, group=dist_group)
|
||||
dist.broadcast(tensor_data, src=src, group=dist_group)
|
||||
return data
|
||||
else:
|
||||
tensor_size = torch.tensor([0], dtype=torch.long, device=device)
|
||||
dist.broadcast(tensor_size, src=src, group=dist_group)
|
||||
size = tensor_size.item()
|
||||
|
||||
if size == 0:
|
||||
return []
|
||||
|
||||
tensor_data = torch.empty(size, dtype=torch.uint8, device=device)
|
||||
dist.broadcast(tensor_data, src=src, group=dist_group)
|
||||
|
||||
serialized_data = bytes(tensor_data.cpu().numpy())
|
||||
data = pickle.loads(serialized_data)
|
||||
return data
|
||||
|
||||
|
||||
def generate_masked_orthogonal_rank_groups(
|
||||
world_size: int, parallel_size: list[int], mask: list[bool]
|
||||
) -> list[list[int]]:
|
||||
"""Generate orthogonal parallel groups based on the parallel size and mask.
|
||||
|
||||
Arguments:
|
||||
world_size (int): world size
|
||||
|
||||
parallel_size (List[int]):
|
||||
The parallel size of each orthogonal parallel type. For example, if
|
||||
tensor_parallel_size = 2, pipeline_model_parallel_group = 3, data_parallel_size = 4,
|
||||
and the parallel mapping order is tp-pp-dp, then the parallel_size = [2, 3, 4].
|
||||
|
||||
mask (List[bool]):
|
||||
The mask controls which parallel methods the generated groups represent. If mask[i] is
|
||||
True, it means the generated group contains the i-th parallelism method. For example,
|
||||
if parallel_size = [tp_size, pp_size, dp_size], and mask = [True, False , True], then
|
||||
the generated group is the `tp-dp` group, if the mask = [False, True, False], then the
|
||||
generated group is the `pp` group.
|
||||
|
||||
Algorithm:
|
||||
For orthogonal parallelism, such as tp/dp/pp/cp, the global_rank and
|
||||
|
||||
If we want to get the `dp_group` (tp_size * pp_size groups of dp_size ranks each.
|
||||
For example, if the gpu size is 8 and order is 'tp-pp-dp', size is '2-2-2', and the
|
||||
dp_group here is [[0, 4], [1, 5], [2, 6], [3, 7]].)
|
||||
The tp_rank and pp_rank will be combined to form the `dp_group_index`.
|
||||
dp_group_index = tp_rank + pp_rank * tp_size (2)
|
||||
|
||||
So, Given that tp_rank and pp_rank satisfy equation (2), and dp_rank in
|
||||
range(0, dp_size), the ranks in dp_group[dp_group_index] satisfies the
|
||||
equation (1).
|
||||
|
||||
This function solve this math problem.
|
||||
|
||||
For example, if the parallel_size = [tp_size, dp_size, pp_size] = [2, 3, 4],
|
||||
and the mask = [False, True, False]. Then,
|
||||
dp_group_index(0) = tp_rank(0) + pp_rank(0) * 2
|
||||
dp_group_index(1) = tp_rank(1) + pp_rank(0) * 2
|
||||
...
|
||||
dp_group_index(7) = tp_rank(1) + pp_rank(3) * 2
|
||||
|
||||
dp_group[0] = 0 + range(0, 3) * 2 + 0 = [0, 2, 4]
|
||||
dp_group[1] = 1 + range(0, 3) * 2 + 0 = [1, 3, 5]
|
||||
...
|
||||
dp_group[7] = 1 + range(0, 3) * 2 + 3 * 2 * 3 = [19, 21, 23]
|
||||
"""
|
||||
|
||||
def prefix_product(a: List[int], init=1) -> List[int]:
|
||||
r = [init]
|
||||
for v in a:
|
||||
init = init * v
|
||||
r.append(init)
|
||||
return r
|
||||
|
||||
def inner_product(a: List[int], b: List[int]) -> int:
|
||||
return sum([x * y for x, y in zip(a, b)])
|
||||
|
||||
def decompose(index, shape, stride=None):
|
||||
"""
|
||||
This function solve the math problem below:
|
||||
There is an equation:
|
||||
index = sum(idx[i] * stride[i])
|
||||
And given the value of index, stride.
|
||||
Return the idx.
|
||||
This function will used to get the pp/dp/pp_rank
|
||||
from group_index and rank_in_group.
|
||||
"""
|
||||
if stride is None:
|
||||
stride = prefix_product(shape)
|
||||
idx = [(index // d) % s for s, d in zip(shape, stride)]
|
||||
# stride is a prefix_product result. And the value of stride[-1]
|
||||
# is not used.
|
||||
assert sum([x * y for x, y in zip(idx, stride[:-1])]) == index, (
|
||||
"idx {} with shape {} mismatch the return idx {}".format(index, shape, idx)
|
||||
)
|
||||
return idx
|
||||
|
||||
masked_shape = [s for s, m in zip(parallel_size, mask) if m]
|
||||
unmasked_shape = [s for s, m in zip(parallel_size, mask) if not m]
|
||||
|
||||
global_stride = prefix_product(parallel_size)
|
||||
masked_stride = [d for d, m in zip(global_stride, mask) if m]
|
||||
unmasked_stride = [d for d, m in zip(global_stride, mask) if not m]
|
||||
|
||||
group_size = prefix_product(masked_shape)[-1]
|
||||
num_of_group = world_size // group_size
|
||||
|
||||
ranks = []
|
||||
for group_index in range(num_of_group):
|
||||
# get indices from unmaksed for group_index.
|
||||
decomposed_group_idx = decompose(group_index, unmasked_shape)
|
||||
rank = []
|
||||
for rank_in_group in range(group_size):
|
||||
# get indices from masked for rank_in_group.
|
||||
decomposed_rank_idx = decompose(rank_in_group, masked_shape)
|
||||
rank.append(
|
||||
inner_product(decomposed_rank_idx, masked_stride)
|
||||
+ inner_product(decomposed_group_idx, unmasked_stride)
|
||||
)
|
||||
ranks.append(rank)
|
||||
return ranks
|
||||
|
||||
|
||||
class RankGenerator(object):
|
||||
def __init__(
|
||||
self,
|
||||
tp: int,
|
||||
sp: int,
|
||||
pp: int,
|
||||
cfg: int,
|
||||
dp: int,
|
||||
order: str,
|
||||
rank_offset: int = 0,
|
||||
) -> None:
|
||||
self.tp = tp
|
||||
self.sp = sp
|
||||
self.pp = pp
|
||||
self.cfg = cfg
|
||||
self.dp = dp
|
||||
self.rank_offset = rank_offset
|
||||
self.world_size = tp * sp * pp * cfg * dp
|
||||
|
||||
self.name_to_size = {
|
||||
"tp": self.tp,
|
||||
"sp": self.sp,
|
||||
"pp": self.pp,
|
||||
"cfg": self.cfg,
|
||||
"dp": self.dp,
|
||||
}
|
||||
order = order.lower()
|
||||
|
||||
for name in self.name_to_size.keys():
|
||||
if name not in order and self.name_to_size[name] != 1:
|
||||
raise RuntimeError(
|
||||
f"The size of ({name}) is ({self.name_to_size[name]}), but you haven't specified the order ({self.order})."
|
||||
)
|
||||
elif name not in order:
|
||||
order = order + "-" + name
|
||||
|
||||
self.order = order
|
||||
self.ordered_size = []
|
||||
|
||||
for token in order.split("-"):
|
||||
self.ordered_size.append(self.name_to_size[token])
|
||||
|
||||
def get_mask(self, order: str, token: str):
|
||||
ordered_token = order.split("-")
|
||||
token = token.split("-")
|
||||
mask = [False] * len(ordered_token)
|
||||
for t in token:
|
||||
mask[ordered_token.index(t)] = True
|
||||
return mask
|
||||
|
||||
def get_ranks(self, token):
|
||||
"""Get rank group by input token.
|
||||
|
||||
Arguments:
|
||||
token (str):
|
||||
Specify the ranks type that want to get. If we want
|
||||
to obtain multiple parallel types, we can use a hyphen
|
||||
'-' to separate them. For example, if we want to obtain
|
||||
the TP_DP group, the token should be 'tp-dp'.
|
||||
|
||||
"""
|
||||
mask = self.get_mask(self.order, token)
|
||||
ranks = generate_masked_orthogonal_rank_groups(
|
||||
self.world_size, self.ordered_size, mask
|
||||
)
|
||||
if self.rank_offset > 0:
|
||||
for rank_group in ranks:
|
||||
for i in range(len(rank_group)):
|
||||
rank_group[i] += self.rank_offset
|
||||
return ranks
|
||||
@@ -1,7 +1,10 @@
|
||||
import threading
|
||||
from contextlib import contextmanager, nullcontext
|
||||
from dataclasses import dataclass
|
||||
from typing import Iterator, Optional, Union
|
||||
|
||||
import torch
|
||||
from torch.distributed.fsdp import MixedPrecisionPolicy
|
||||
|
||||
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||
from sglang.multimodal_gen.runtime.utils.precision_types import PRECISION_TO_TYPE
|
||||
@@ -202,3 +205,55 @@ def temporary_module_dtype(
|
||||
yield module
|
||||
finally:
|
||||
module.to(dtype=original_dtype)
|
||||
|
||||
|
||||
@dataclass
|
||||
class MixedPrecisionState:
|
||||
param_dtype: torch.dtype | None = None
|
||||
reduce_dtype: torch.dtype | None = None
|
||||
output_dtype: torch.dtype | None = None
|
||||
compute_dtype: torch.dtype | None = None
|
||||
mp_policy: MixedPrecisionPolicy | None = None
|
||||
|
||||
|
||||
class _MixedPrecisionContext(threading.local):
|
||||
state: MixedPrecisionState | None = None
|
||||
|
||||
|
||||
_mixed_precision_state = _MixedPrecisionContext()
|
||||
|
||||
|
||||
def get_mixed_precision_state() -> MixedPrecisionState:
|
||||
"""Get the current mixed precision state."""
|
||||
state = _mixed_precision_state.state
|
||||
if state is None:
|
||||
raise ValueError("Mixed precision state not set")
|
||||
return state
|
||||
|
||||
|
||||
def set_mixed_precision_policy(
|
||||
param_dtype: torch.dtype,
|
||||
reduce_dtype: torch.dtype,
|
||||
output_dtype: torch.dtype | None = None,
|
||||
mp_policy: MixedPrecisionPolicy | None = None,
|
||||
):
|
||||
"""Set mixed precision policy for the current thread.
|
||||
|
||||
Args:
|
||||
param_dtype: Parameter dtype used for training
|
||||
reduce_dtype: Reduction dtype used for gradients
|
||||
output_dtype: Optional output dtype
|
||||
"""
|
||||
state = MixedPrecisionState(
|
||||
param_dtype=param_dtype,
|
||||
reduce_dtype=reduce_dtype,
|
||||
output_dtype=output_dtype,
|
||||
mp_policy=mp_policy,
|
||||
)
|
||||
_mixed_precision_state.state = state
|
||||
|
||||
|
||||
def get_compute_dtype() -> torch.dtype:
|
||||
"""Get the current compute dtype from mixed precision policy."""
|
||||
state = _mixed_precision_state.state
|
||||
return torch.get_default_dtype() if state is None else state.param_dtype
|
||||
|
||||
@@ -0,0 +1,69 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/utils.py
|
||||
|
||||
import ctypes
|
||||
import os
|
||||
import signal
|
||||
import sys
|
||||
import threading
|
||||
|
||||
import psutil
|
||||
|
||||
|
||||
def kill_itself_when_parent_died() -> None:
|
||||
if sys.platform != "linux":
|
||||
return
|
||||
|
||||
# keep GPU workers tied to the CLI process even if the parent is SIGKILLed
|
||||
PR_SET_PDEATHSIG = 1
|
||||
# Capture parent before arming PDEATHSIG: if the parent already died in the
|
||||
# fork->prctl window, PDEATHSIG won't fire, so detect the reparent explicitly.
|
||||
parent_pid = os.getppid()
|
||||
libc = ctypes.CDLL("libc.so.6", use_errno=True)
|
||||
if libc.prctl(PR_SET_PDEATHSIG, signal.SIGKILL) != 0:
|
||||
err = ctypes.get_errno()
|
||||
raise OSError(err, os.strerror(err))
|
||||
# getppid() changing means we were reparented (parent gone). Comparing to the
|
||||
# captured pid instead of "== 1" avoids self-killing when PID 1 is the real
|
||||
# parent (e.g. running as a container's init process).
|
||||
if os.getppid() != parent_pid:
|
||||
os.kill(os.getpid(), signal.SIGKILL)
|
||||
|
||||
|
||||
def kill_process_tree(parent_pid, include_parent: bool = True, skip_pid: int = None):
|
||||
"""Kill the process and all its child processes."""
|
||||
# Remove sigchld handler to avoid spammy logs.
|
||||
if threading.current_thread() is threading.main_thread():
|
||||
signal.signal(signal.SIGCHLD, signal.SIG_DFL)
|
||||
|
||||
if parent_pid is None:
|
||||
parent_pid = os.getpid()
|
||||
include_parent = False
|
||||
|
||||
try:
|
||||
itself = psutil.Process(parent_pid)
|
||||
except psutil.NoSuchProcess:
|
||||
return
|
||||
|
||||
children = itself.children(recursive=True)
|
||||
for child in children:
|
||||
if child.pid == skip_pid:
|
||||
continue
|
||||
try:
|
||||
child.kill()
|
||||
except psutil.NoSuchProcess:
|
||||
pass
|
||||
|
||||
if include_parent:
|
||||
try:
|
||||
if parent_pid == os.getpid():
|
||||
itself.kill()
|
||||
sys.exit(0)
|
||||
|
||||
itself.kill()
|
||||
|
||||
# Sometime processes cannot be killed with SIGKILL (e.g, PID=1 launched by kubernetes),
|
||||
# so we send an additional signal to kill them.
|
||||
itself.send_signal(signal.SIGQUIT)
|
||||
except psutil.NoSuchProcess:
|
||||
pass
|
||||
@@ -16,7 +16,7 @@ import numpy as np
|
||||
import pytest
|
||||
from openai import Client
|
||||
|
||||
from sglang.multimodal_gen.runtime.utils.realtime_video import (
|
||||
from sglang.multimodal_gen.runtime.realtime.video import (
|
||||
RAW_RGB_CONTENT_TYPE,
|
||||
RAW_RGB_DELTA_GZIP_CONTENT_TYPE,
|
||||
RAW_RGBA_DELTA_GZIP_CONTENT_TYPE,
|
||||
|
||||
@@ -24,12 +24,12 @@ from openai import Client
|
||||
|
||||
from sglang.multimodal_gen.benchmarks.compare_perf import calculate_upper_bound
|
||||
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||
from sglang.multimodal_gen.runtime.utils.common import kill_process_tree
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import (
|
||||
globally_suppress_loggers,
|
||||
init_logger,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.utils.perf_logger import RequestPerfRecord
|
||||
from sglang.multimodal_gen.runtime.utils.process import kill_process_tree
|
||||
from sglang.multimodal_gen.test.server.common.slack import upload_file_to_slack
|
||||
from sglang.multimodal_gen.test.server.realtime_consistency import (
|
||||
build_realtime_init_payload,
|
||||
|
||||
@@ -24,15 +24,15 @@ from sglang.multimodal_gen.runtime.models.dits.lingbot_world import (
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.stages.causal_denoising import (
|
||||
CausalDMDCachePolicy,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.lingbot_world.constants import (
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.lingbot_world.lingbot_world_causal_denoising import (
|
||||
LingBotWorldCausalDMDDenoisingStage,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.realtime.lingbot_world import (
|
||||
LINGBOT_C2WS_PLUCKER_EMB_CACHE,
|
||||
LINGBOT_CAM_CONDITIONER_CACHE,
|
||||
LINGBOT_PROMPT_UPDATED_CONDITION,
|
||||
LINGBOT_ROPE_CACHE,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.lingbot_world.lingbot_world_causal_denoising import (
|
||||
LingBotWorldCausalDMDDenoisingStage,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.realtime.states import RealtimeCausalDiTState
|
||||
|
||||
LINGBOT_INTERACTIVE_KV_WINDOW_ENV = "SGLANG_LINGBOT_ENABLE_INTERACTIVE_KV_WINDOW"
|
||||
|
||||
@@ -10,9 +10,7 @@ from sglang.multimodal_gen.runtime.entrypoints.utils import (
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.managers.gpu_worker import GPUWorker
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch
|
||||
from sglang.multimodal_gen.runtime.utils.realtime_video import (
|
||||
build_raw_rgb_frame_batches,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.realtime.video import build_raw_rgb_frame_batches
|
||||
|
||||
|
||||
def test_materialize_output_sample_converts_tensor_to_uint8_frames():
|
||||
|
||||
+1
-1
@@ -8,7 +8,7 @@ import msgspec.msgpack
|
||||
import numpy as np
|
||||
import pytest
|
||||
|
||||
from sglang.multimodal_gen.runtime.utils.realtime_video import (
|
||||
from sglang.multimodal_gen.runtime.realtime.video import (
|
||||
RAW_RGB_CONTENT_TYPE,
|
||||
RAW_RGB_DELTA_GZIP_CONTENT_TYPE,
|
||||
RAW_RGBA_DELTA_GZIP_CONTENT_TYPE,
|
||||
|
||||
@@ -14,7 +14,7 @@ from sglang.multimodal_gen.runtime.entrypoints.openai.realtime.realtime_output_a
|
||||
RawRGBRealtimeOutputAdapter,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch
|
||||
from sglang.multimodal_gen.runtime.utils.realtime_video import (
|
||||
from sglang.multimodal_gen.runtime.realtime.video import (
|
||||
JPEG_FRAME_CONTENT_TYPE,
|
||||
RAW_RGB_CONTENT_TYPE,
|
||||
WEBP_FRAME_CONTENT_TYPE,
|
||||
|
||||
@@ -58,9 +58,7 @@ from sglang.multimodal_gen.runtime.realtime.session import (
|
||||
from sglang.multimodal_gen.runtime.realtime.states import (
|
||||
RealtimeCausalDecodeState,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.utils.realtime_video import (
|
||||
RAW_RGB_CONTENT_TYPE,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.realtime.video import RAW_RGB_CONTENT_TYPE
|
||||
|
||||
|
||||
class _Req(SimpleNamespace):
|
||||
|
||||
@@ -142,7 +142,7 @@ class TestAttentionBackendFallback(unittest.TestCase):
|
||||
_FakePlatform,
|
||||
),
|
||||
patch(
|
||||
f"{_SELECTOR}.resolve_obj_by_qualname",
|
||||
f"{_SELECTOR}.resolve_name",
|
||||
side_effect=_FAKE_BACKENDS.__getitem__,
|
||||
),
|
||||
component_attn_backend_context_manager(
|
||||
|
||||
@@ -217,7 +217,7 @@ class TestCudaAttentionBackendSelection(unittest.TestCase):
|
||||
FakeCudaPlatform,
|
||||
),
|
||||
patch(
|
||||
"sglang.multimodal_gen.runtime.layers.attention.selector.resolve_obj_by_qualname",
|
||||
"sglang.multimodal_gen.runtime.layers.attention.selector.resolve_name",
|
||||
return_value=FakeAITERBackend,
|
||||
),
|
||||
):
|
||||
|
||||
@@ -26,7 +26,7 @@ class TestCudaPlatformDetection(unittest.TestCase):
|
||||
with (
|
||||
self.subTest(hip_version=hip_version),
|
||||
patch(
|
||||
"sglang.multimodal_gen.utils.import_pynvml",
|
||||
"sglang.multimodal_gen.runtime.platforms.pynvml.nvmlInit",
|
||||
side_effect=NVMLUnavailableError,
|
||||
),
|
||||
patch.object(platforms.os.path, "isfile", return_value=False),
|
||||
|
||||
@@ -1,69 +1,11 @@
|
||||
import importlib.util
|
||||
import sys
|
||||
import types
|
||||
import unittest
|
||||
from contextlib import nullcontext
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.multimodal_gen.runtime.utils import precision
|
||||
|
||||
def _load_precision_module():
|
||||
package_names = (
|
||||
"sglang",
|
||||
"sglang.multimodal_gen",
|
||||
"sglang.multimodal_gen.runtime",
|
||||
"sglang.multimodal_gen.runtime.utils",
|
||||
)
|
||||
stub_names = (
|
||||
*package_names,
|
||||
"sglang.multimodal_gen.runtime.platforms",
|
||||
"sglang.multimodal_gen.utils",
|
||||
)
|
||||
missing = object()
|
||||
previous_modules = {name: sys.modules.get(name, missing) for name in stub_names}
|
||||
|
||||
try:
|
||||
utils_module = types.ModuleType("sglang.multimodal_gen.utils")
|
||||
utils_module.PRECISION_TO_TYPE = {
|
||||
"fp16": torch.float16,
|
||||
"bf16": torch.bfloat16,
|
||||
"fp32": torch.float32,
|
||||
}
|
||||
platforms_module = types.ModuleType("sglang.multimodal_gen.runtime.platforms")
|
||||
platforms_module.current_platform = SimpleNamespace(
|
||||
device_type="cpu",
|
||||
is_mps=lambda: False,
|
||||
is_amp_supported=lambda: True,
|
||||
)
|
||||
for package_name in package_names:
|
||||
package = types.ModuleType(package_name)
|
||||
package.__path__ = []
|
||||
sys.modules[package_name] = package
|
||||
sys.modules["sglang.multimodal_gen.runtime.platforms"] = platforms_module
|
||||
sys.modules["sglang.multimodal_gen.utils"] = utils_module
|
||||
|
||||
precision_path = (
|
||||
Path(__file__).resolve().parents[2] / "runtime/utils/precision.py"
|
||||
)
|
||||
spec = importlib.util.spec_from_file_location(
|
||||
"_diffusion_precision_under_test", precision_path
|
||||
)
|
||||
precision = importlib.util.module_from_spec(spec)
|
||||
sys.modules[spec.name] = precision
|
||||
spec.loader.exec_module(precision)
|
||||
finally:
|
||||
for module_name, previous_module in previous_modules.items():
|
||||
if previous_module is missing:
|
||||
sys.modules.pop(module_name, None)
|
||||
else:
|
||||
sys.modules[module_name] = previous_module
|
||||
|
||||
return precision
|
||||
|
||||
|
||||
precision = _load_precision_module()
|
||||
align_tensor_to_module_dtype = precision.align_tensor_to_module_dtype
|
||||
autocast_context = precision.autocast_context
|
||||
autocast_enabled = precision.autocast_enabled
|
||||
|
||||
@@ -110,9 +110,7 @@ def test_qwen3vl_auxiliary_component_falls_back_from_global_backend(monkeypatch)
|
||||
monkeypatch.setattr(
|
||||
"sglang.multimodal_gen.runtime.platforms.current_platform", _FakePlatform
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
f"{_SELECTOR}.resolve_obj_by_qualname", lambda _name: _FakeFABackend
|
||||
)
|
||||
monkeypatch.setattr(f"{_SELECTOR}.resolve_name", lambda _name: _FakeFABackend)
|
||||
_cached_get_attn_backend.cache_clear()
|
||||
config = SimpleNamespace(
|
||||
head_dim=8,
|
||||
|
||||
@@ -90,7 +90,7 @@ from sglang.multimodal_gen.runtime.server_args import (
|
||||
MAX_SCHEDULER_RPC_TIMEOUT_S,
|
||||
ServerArgs,
|
||||
)
|
||||
from sglang.multimodal_gen.utils import FlexibleArgumentParser
|
||||
from sglang.multimodal_gen.runtime.utils.argparse import FlexibleArgumentParser
|
||||
|
||||
|
||||
@contextmanager
|
||||
|
||||
@@ -13,10 +13,10 @@ from sglang.multimodal_gen.configs.models.encoders import TextEncoderConfig
|
||||
from sglang.multimodal_gen.runtime.distributed.parallel_state import (
|
||||
_get_encoder_data_parallel_group_ranks,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.distributed.utils import RankGenerator
|
||||
from sglang.multimodal_gen.runtime.models.encoders.base import TextEncoder
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.stages import text_encoding as _te_mod
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
from sglang.multimodal_gen.runtime.utils.distributed import RankGenerator
|
||||
|
||||
|
||||
def _enc(hidden=4096, heads=64, inter=10240, folding_mode=None):
|
||||
|
||||
@@ -0,0 +1,113 @@
|
||||
# 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
|
||||
import torch
|
||||
|
||||
from sglang.multimodal_gen.configs.utils import expand_path_fields
|
||||
from sglang.multimodal_gen.runtime.layers.attention.mask_strategy import dict_to_3d_list
|
||||
from sglang.multimodal_gen.runtime.utils.argparse import (
|
||||
FlexibleArgumentParser,
|
||||
StoreBoolean,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.utils.precision import (
|
||||
get_compute_dtype,
|
||||
get_mixed_precision_state,
|
||||
set_mixed_precision_policy,
|
||||
)
|
||||
|
||||
|
||||
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")
|
||||
parser = FlexibleArgumentParser()
|
||||
parser.add_argument("command")
|
||||
parser.add_argument("--num-gpus", type=int, default=1)
|
||||
parser.add_argument("--use-cache", action=StoreBoolean, default=False)
|
||||
args = parser.parse_args(
|
||||
["generate", "--config", str(config), "--num_gpus=4", "--use-cache", "false"]
|
||||
)
|
||||
assert args.num_gpus == 4
|
||||
assert args.use_cache is False
|
||||
assert args._provided == {"num_gpus", "use_cache"}
|
||||
|
||||
|
||||
def test_expand_paths_preserves_slots_and_non_path_fields():
|
||||
@dataclass(slots=True)
|
||||
class Config:
|
||||
model_path: str = "~/model"
|
||||
image_path: list = field(default_factory=lambda: ["~/image.png", None])
|
||||
model_paths: dict = field(
|
||||
default_factory=lambda: {"vae": "~/vae", "other": None}
|
||||
)
|
||||
prompt: str = "~/not-a-path"
|
||||
|
||||
config = Config()
|
||||
expand_path_fields(config)
|
||||
assert config.model_path == os.path.expanduser("~/model")
|
||||
assert config.image_path == [os.path.expanduser("~/image.png"), None]
|
||||
assert config.model_paths == {"vae": os.path.expanduser("~/vae"), "other": None}
|
||||
assert config.prompt == "~/not-a-path"
|
||||
|
||||
|
||||
def test_mixed_precision_state_is_thread_local():
|
||||
barrier = Barrier(2)
|
||||
|
||||
def worker(dtype):
|
||||
assert get_compute_dtype() == torch.get_default_dtype()
|
||||
with pytest.raises(ValueError, match="Mixed precision state not set"):
|
||||
get_mixed_precision_state()
|
||||
set_mixed_precision_policy(dtype, torch.float32, output_dtype=dtype)
|
||||
barrier.wait(timeout=10)
|
||||
assert get_mixed_precision_state().output_dtype == dtype
|
||||
return get_compute_dtype()
|
||||
|
||||
with ThreadPoolExecutor(max_workers=2) as executor:
|
||||
assert list(executor.map(worker, (torch.float16, torch.bfloat16))) == [
|
||||
torch.float16,
|
||||
torch.bfloat16,
|
||||
]
|
||||
|
||||
|
||||
def test_attention_mask_strategy_preserves_tensor_identity():
|
||||
mask = torch.tensor([True, False])
|
||||
strategy = {"1_0_2": mask}
|
||||
inferred = dict_to_3d_list(strategy)
|
||||
assert len(inferred) == 2
|
||||
assert inferred[1][0][2] is mask
|
||||
assert inferred[0][0][2] is None
|
||||
assert dict_to_3d_list(strategy, 1, 1, 1) == [[[None]]]
|
||||
assert dict_to_3d_list(None, 2, 1, 1) == [[[None]], [[None]]]
|
||||
@@ -2,7 +2,7 @@ import signal
|
||||
import unittest
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
from sglang.multimodal_gen import utils
|
||||
from sglang.multimodal_gen.runtime.utils import process as utils
|
||||
|
||||
|
||||
class TestKillItselfWhenParentDied(unittest.TestCase):
|
||||
|
||||
@@ -24,6 +24,7 @@ from sglang.multimodal_gen.configs.models.vaes.stablediffusion3 import (
|
||||
from sglang.multimodal_gen.configs.models.vaes.wanvae import WanVAEConfig
|
||||
from sglang.multimodal_gen.configs.utils import update_config_from_args
|
||||
from sglang.multimodal_gen.runtime.distributed import parallel_state
|
||||
from sglang.multimodal_gen.runtime.distributed.utils import RankGenerator
|
||||
from sglang.multimodal_gen.runtime.layers.parallel_conv import (
|
||||
SpatialParallelCausalConv3d,
|
||||
SpatialParallelConv2d,
|
||||
@@ -52,8 +53,7 @@ from sglang.multimodal_gen.runtime.models.vaes.wanvae import (
|
||||
WanDecoder3d,
|
||||
WanDistAttentionBlock,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.utils.distributed import RankGenerator
|
||||
from sglang.multimodal_gen.utils import FlexibleArgumentParser
|
||||
from sglang.multimodal_gen.runtime.utils.argparse import FlexibleArgumentParser
|
||||
|
||||
|
||||
class _DispatchProbeVAE(ParallelTiledVAE):
|
||||
|
||||
@@ -5,13 +5,56 @@ from unittest.mock import patch
|
||||
import torch
|
||||
|
||||
from sglang.multimodal_gen.configs.pipeline_configs.base import maybe_unpad_latents
|
||||
from sglang.multimodal_gen.configs.pipeline_configs.wan import Wan2_2_TI2V_5B_Config
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.wan_ti2v import (
|
||||
expand_wan_ti2v_timestep,
|
||||
prepare_wan_ti2v_latents,
|
||||
prepare_wan_ti2v_sp_inputs,
|
||||
)
|
||||
|
||||
|
||||
class TestWanTI2VHelpers(unittest.TestCase):
|
||||
def test_conditioning_mask_preserves_first_frame_only(self):
|
||||
latents = torch.arange(1 * 2 * 3 * 4 * 4).reshape(1, 2, 3, 4, 4).float()
|
||||
image_latent = torch.full((1, 2, 1, 4, 4), 0.25)
|
||||
vae = SimpleNamespace(
|
||||
encode=lambda image: SimpleNamespace(mean=image_latent),
|
||||
scaling_factor=1.0,
|
||||
shift_factor=None,
|
||||
)
|
||||
config = Wan2_2_TI2V_5B_Config()
|
||||
batch = SimpleNamespace(
|
||||
image_latent=None,
|
||||
condition_image=torch.zeros(1, 3, 1, 64, 64),
|
||||
num_frames=9,
|
||||
height=64,
|
||||
width=64,
|
||||
)
|
||||
with (
|
||||
patch(
|
||||
"sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.wan_ti2v.get_local_torch_device",
|
||||
return_value="cpu",
|
||||
),
|
||||
patch(
|
||||
"sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.wan_ti2v.get_sp_world_size",
|
||||
return_value=1,
|
||||
),
|
||||
):
|
||||
_, _, masks = prepare_wan_ti2v_latents(
|
||||
vae,
|
||||
latents,
|
||||
torch.bfloat16,
|
||||
torch.float32,
|
||||
batch,
|
||||
SimpleNamespace(pipeline_config=config),
|
||||
)
|
||||
self.assertEqual(masks[0].dtype, torch.bfloat16)
|
||||
self.assertEqual(masks[0].shape, (2, 3, 4, 4))
|
||||
self.assertTrue(torch.all(masks[0][:, 0] == 0))
|
||||
self.assertTrue(torch.all(masks[0][:, 1:] == 1))
|
||||
torch.testing.assert_close(batch.latents[:, :, :1], image_latent)
|
||||
torch.testing.assert_close(batch.latents[:, :, 1:], latents[:, :, 1:])
|
||||
|
||||
def test_sp_mask_is_padded_before_sharding(self):
|
||||
mask = torch.ones(1, 21, 4, 4)
|
||||
mask[:, 0] = 0
|
||||
|
||||
@@ -1,794 +0,0 @@
|
||||
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
|
||||
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/utils.py
|
||||
|
||||
import argparse
|
||||
import ctypes
|
||||
import importlib
|
||||
import importlib.util
|
||||
import inspect
|
||||
import math
|
||||
import os
|
||||
import signal
|
||||
import sys
|
||||
import threading
|
||||
import traceback
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, fields, is_dataclass
|
||||
from functools import lru_cache, partial, wraps
|
||||
from typing import Any, TypeVar, cast
|
||||
|
||||
import cloudpickle
|
||||
import torch
|
||||
import yaml
|
||||
from torch.distributed.fsdp import MixedPrecisionPolicy
|
||||
|
||||
import sglang.multimodal_gen.envs as envs
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import (
|
||||
SortedHelpFormatter,
|
||||
init_logger,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.utils.precision_types import (
|
||||
PRECISION_TO_TYPE as PRECISION_TO_TYPE,
|
||||
)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
T = TypeVar("T")
|
||||
|
||||
|
||||
def expand_path_fields(obj) -> None:
|
||||
"""In-place expanduser on all dataclass fields whose name ends with '_path' or '_paths'."""
|
||||
eu = os.path.expanduser
|
||||
for f in fields(obj):
|
||||
v = getattr(obj, f.name)
|
||||
if f.name.endswith("_path") and isinstance(v, str):
|
||||
setattr(obj, f.name, eu(v))
|
||||
elif f.name.endswith("_path") and isinstance(v, list):
|
||||
setattr(obj, f.name, [eu(x) if isinstance(x, str) else x for x in v])
|
||||
elif f.name.endswith("_paths") and isinstance(v, dict):
|
||||
setattr(
|
||||
obj,
|
||||
f.name,
|
||||
{k: eu(p) if isinstance(p, str) else p for k, p in v.items()},
|
||||
)
|
||||
|
||||
|
||||
STR_BACKEND_ENV_VAR: str = "SGLANG_DIFFUSION_ATTENTION_BACKEND"
|
||||
STR_ATTN_CONFIG_ENV_VAR: str = "SGLANG_DIFFUSION_ATTENTION_CONFIG"
|
||||
|
||||
|
||||
def find_nccl_library() -> str:
|
||||
"""
|
||||
We either use the library file specified by the `VLLM_NCCL_SO_PATH`
|
||||
environment variable, or we find the library file brought by PyTorch.
|
||||
After importing `torch`, `libnccl.so.2`, `librccl.so.1` or `libmccl.so.2`
|
||||
can be found by `ctypes` automatically.
|
||||
"""
|
||||
so_file = envs.SGLANG_DIFFUSION_NCCL_SO_PATH
|
||||
|
||||
# manually load the nccl library
|
||||
if so_file:
|
||||
logger.info(
|
||||
"Found nccl from environment variable SGLANG_DIFFUSION_NCCL_SO_PATH=%s",
|
||||
so_file,
|
||||
)
|
||||
else:
|
||||
if torch.version.cuda is not None:
|
||||
so_file = "libnccl.so.2"
|
||||
elif torch.version.hip is not None:
|
||||
so_file = "librccl.so.1"
|
||||
elif hasattr(torch.version, "musa") and torch.version.musa is not None:
|
||||
so_file = "libmccl.so.2"
|
||||
else:
|
||||
raise ValueError("NCCL only supports CUDA, ROCm and MUSA backends.")
|
||||
logger.info("Found nccl from library %s", so_file)
|
||||
return str(so_file)
|
||||
|
||||
|
||||
prev_set_stream = torch.cuda.set_stream
|
||||
|
||||
_current_stream = None
|
||||
|
||||
|
||||
def _patched_set_stream(stream: torch.cuda.Stream | None) -> None:
|
||||
global _current_stream
|
||||
_current_stream = stream
|
||||
if stream is not None:
|
||||
prev_set_stream(stream)
|
||||
|
||||
|
||||
torch.cuda.set_stream = _patched_set_stream
|
||||
|
||||
|
||||
def current_stream() -> torch.cuda.Stream | None:
|
||||
"""
|
||||
replace `torch.cuda.current_stream()` with `sglang.multimodal_gen.utils.current_stream()`.
|
||||
it turns out that `torch.cuda.current_stream()` is quite expensive,
|
||||
as it will construct a new stream object at each call.
|
||||
here we patch `torch.cuda.set_stream` to keep track of the current stream
|
||||
directly, so that we can avoid calling `torch.cuda.current_stream()`.
|
||||
|
||||
the underlying hypothesis is that we do not call `torch._C._cuda_setStream`
|
||||
from C/C++ code.
|
||||
"""
|
||||
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||
|
||||
# For non-CUDA platforms, return None
|
||||
if not current_platform.is_cuda_alike():
|
||||
return None
|
||||
|
||||
global _current_stream
|
||||
if _current_stream is None:
|
||||
# when this function is called before any stream is set,
|
||||
# we return the default stream.
|
||||
# On ROCm using the default 0 stream in combination with RCCL
|
||||
# is hurting performance. Therefore creating a dedicated stream
|
||||
# per process
|
||||
_current_stream = (
|
||||
torch.cuda.Stream()
|
||||
if current_platform.is_rocm()
|
||||
else torch.cuda.current_stream()
|
||||
)
|
||||
return _current_stream
|
||||
|
||||
|
||||
class StoreBoolean(argparse.Action):
|
||||
def __init__(self, option_strings, dest, default=False, required=False, help=None):
|
||||
super().__init__(
|
||||
option_strings=option_strings,
|
||||
dest=dest,
|
||||
nargs="?",
|
||||
const=True,
|
||||
default=default,
|
||||
required=required,
|
||||
help=help,
|
||||
)
|
||||
|
||||
def __call__(self, parser, namespace, values, option_string=None):
|
||||
if values is None:
|
||||
setattr(namespace, self.dest, True)
|
||||
elif isinstance(values, str):
|
||||
if values.lower() == "true":
|
||||
setattr(namespace, self.dest, True)
|
||||
elif values.lower() == "false":
|
||||
setattr(namespace, self.dest, False)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Invalid boolean value: {values}. Expected 'true' or 'false'."
|
||||
)
|
||||
else:
|
||||
setattr(namespace, self.dest, bool(values))
|
||||
|
||||
|
||||
class FlexibleArgumentParser(argparse.ArgumentParser):
|
||||
"""ArgumentParser that allows both underscore and dash in names."""
|
||||
|
||||
def __init__(self, *args, **kwargs) -> None:
|
||||
# Set the default 'formatter_class' to SortedHelpFormatter
|
||||
if "formatter_class" not in kwargs:
|
||||
kwargs["formatter_class"] = SortedHelpFormatter
|
||||
super().__init__(*args, **kwargs)
|
||||
|
||||
def parse_args( # type: ignore[override]
|
||||
self, args=None, namespace=None
|
||||
) -> argparse.Namespace:
|
||||
if args is None:
|
||||
args = sys.argv[1:]
|
||||
|
||||
if any(arg.startswith("--config") for arg in args):
|
||||
args = self._pull_args_from_config(args)
|
||||
|
||||
# Convert underscores to dashes and vice versa in argument names
|
||||
processed_args = []
|
||||
for arg in args:
|
||||
if arg.startswith("--"):
|
||||
if "=" in arg:
|
||||
key, value = arg.split("=", 1)
|
||||
key = "--" + key[len("--") :].replace("_", "-")
|
||||
processed_args.append(f"{key}={value}")
|
||||
else:
|
||||
processed_args.append("--" + arg[len("--") :].replace("_", "-"))
|
||||
elif arg.startswith("-O") and arg != "-O" and len(arg) == 2:
|
||||
# allow -O flag to be used without space, e.g. -O3
|
||||
processed_args.append("-O")
|
||||
processed_args.append(arg[2:])
|
||||
else:
|
||||
processed_args.append(arg)
|
||||
|
||||
namespace = super().parse_args(processed_args, namespace)
|
||||
|
||||
# Track which arguments were explicitly provided
|
||||
namespace._provided = set()
|
||||
|
||||
i = 0
|
||||
while i < len(args):
|
||||
arg = args[i]
|
||||
if arg.startswith("--"):
|
||||
# Handle --key=value format
|
||||
if "=" in arg:
|
||||
key = arg.split("=")[0][2:].replace("-", "_")
|
||||
namespace._provided.add(key)
|
||||
i += 1
|
||||
# Handle --key value format
|
||||
else:
|
||||
key = arg[2:].replace("-", "_")
|
||||
namespace._provided.add(key)
|
||||
# Skip the value if there is one
|
||||
if i + 1 < len(args) and not args[i + 1].startswith("-"):
|
||||
i += 2
|
||||
else:
|
||||
i += 1
|
||||
else:
|
||||
i += 1
|
||||
|
||||
return namespace # type: ignore[no-any-return]
|
||||
|
||||
def _pull_args_from_config(self, args: list[str]) -> list[str]:
|
||||
"""Method to pull arguments specified in the config file
|
||||
into the command-line args variable.
|
||||
|
||||
The arguments in config file will be inserted between
|
||||
the argument list.
|
||||
|
||||
example:
|
||||
```yaml
|
||||
port: 12323
|
||||
tensor-parallel-size: 4
|
||||
```
|
||||
```python
|
||||
$: vllm {serve,chat,complete} "facebook/opt-12B" \
|
||||
--config config.yaml -tp 2
|
||||
$: args = [
|
||||
"serve,chat,complete",
|
||||
"facebook/opt-12B",
|
||||
'--config', 'config.yaml',
|
||||
'-tp', '2'
|
||||
]
|
||||
$: args = [
|
||||
"serve,chat,complete",
|
||||
"facebook/opt-12B",
|
||||
'--port', '12323',
|
||||
'--tp-size', '4',
|
||||
'-tp', '2'
|
||||
]
|
||||
```
|
||||
|
||||
Please note how the config args are inserted after the sub command.
|
||||
this way the order of priorities is maintained when these are args
|
||||
parsed by super().
|
||||
"""
|
||||
index = -1
|
||||
config_arg = None
|
||||
for i, arg in enumerate(args):
|
||||
if arg.startswith("--config"):
|
||||
if index != -1:
|
||||
raise ValueError("More than one config file specified!")
|
||||
index = i
|
||||
config_arg = arg
|
||||
|
||||
if config_arg is None:
|
||||
return args
|
||||
args_before_config = args[:index]
|
||||
if "=" in config_arg:
|
||||
file_path = config_arg.split("=", 1)[1]
|
||||
args_after_config = args[index + 1 :]
|
||||
else:
|
||||
if index == len(args) - 1:
|
||||
raise ValueError(
|
||||
"No config file specified! "
|
||||
"Please check your command-line arguments."
|
||||
)
|
||||
file_path = args[index + 1]
|
||||
args_after_config = args[index + 2 :]
|
||||
|
||||
config_args = self._load_config_file(file_path)
|
||||
|
||||
# 0th index is for {serve,chat,complete}
|
||||
# followed by model_tag (only for serve)
|
||||
# followed by config args
|
||||
# followed by rest of cli args.
|
||||
# maintaining this order will enforce the precedence
|
||||
# of cli > config > defaults
|
||||
if args[0] == "serve":
|
||||
if index == 1:
|
||||
raise ValueError(
|
||||
"No model_tag specified! Please check your command-line arguments."
|
||||
)
|
||||
command = args_before_config[0]
|
||||
model_tag = args_before_config[1]
|
||||
other_args_before = args_before_config[2:]
|
||||
args = (
|
||||
[command, model_tag]
|
||||
+ config_args
|
||||
+ other_args_before
|
||||
+ args_after_config
|
||||
)
|
||||
else:
|
||||
command = args_before_config[0]
|
||||
other_args_before = args_before_config[1:]
|
||||
args = [command] + config_args + other_args_before + args_after_config
|
||||
|
||||
return args
|
||||
|
||||
def _load_config_file(self, file_path: str) -> list[str]:
|
||||
"""Loads a yaml file and returns the key value pairs as a
|
||||
flattened list with argparse like pattern
|
||||
```yaml
|
||||
port: 12323
|
||||
tensor-parallel-size: 4
|
||||
vae_config:
|
||||
load_encoder: false
|
||||
load_decoder: true
|
||||
```
|
||||
returns:
|
||||
processed_args: list[str] = [
|
||||
'--port': '12323',
|
||||
'--tp-size': '4',
|
||||
'--vae-config.load-encoder': 'false',
|
||||
'--vae-config.load-decoder': 'true'
|
||||
]
|
||||
"""
|
||||
|
||||
extension: str = file_path.split(".")[-1]
|
||||
if extension not in ("yaml", "yml", "json"):
|
||||
raise ValueError(
|
||||
"Config file must be of a yaml/yml/json type.\
|
||||
%s supplied",
|
||||
extension,
|
||||
)
|
||||
|
||||
processed_args: list[str] = []
|
||||
|
||||
config: dict[str, Any] = {}
|
||||
try:
|
||||
with open(file_path) as config_file:
|
||||
config = yaml.safe_load(config_file)
|
||||
except Exception as ex:
|
||||
logger.error(
|
||||
"Unable to read the config file at %s. \
|
||||
Make sure path is correct",
|
||||
file_path,
|
||||
)
|
||||
raise ex
|
||||
|
||||
store_boolean_arguments = [
|
||||
action.dest for action in self._actions if isinstance(action, StoreBoolean)
|
||||
]
|
||||
|
||||
def process_dict(prefix: str, d: dict[str, Any]):
|
||||
for key, value in d.items():
|
||||
full_key = f"{prefix}.{key}" if prefix else key
|
||||
|
||||
if isinstance(value, bool) and full_key not in store_boolean_arguments:
|
||||
if value:
|
||||
processed_args.append("--" + full_key)
|
||||
else:
|
||||
processed_args.append("--" + full_key)
|
||||
processed_args.append("false")
|
||||
elif isinstance(value, list):
|
||||
processed_args.append("--" + full_key)
|
||||
for item in value:
|
||||
processed_args.append(str(item))
|
||||
elif isinstance(value, dict):
|
||||
process_dict(full_key, value)
|
||||
else:
|
||||
processed_args.append("--" + full_key)
|
||||
processed_args.append(str(value))
|
||||
|
||||
process_dict("", config)
|
||||
|
||||
return processed_args
|
||||
|
||||
|
||||
def warn_for_unimplemented_methods(cls: type[T]) -> type[T]:
|
||||
"""
|
||||
A replacement for `abc.ABC`.
|
||||
When we use `abc.ABC`, subclasses will fail to instantiate
|
||||
if they do not implement all abstract methods.
|
||||
Here, we only require `raise NotImplementedError` in the
|
||||
base class, and log a warning if the method is not implemented
|
||||
in the subclass.
|
||||
"""
|
||||
|
||||
original_init = cls.__init__
|
||||
|
||||
def find_unimplemented_methods(self: object):
|
||||
unimplemented_methods = []
|
||||
for attr_name in dir(self):
|
||||
# bypass inner method
|
||||
if attr_name.startswith("_"):
|
||||
continue
|
||||
|
||||
try:
|
||||
attr = getattr(self, attr_name)
|
||||
# get the func of callable method
|
||||
if callable(attr):
|
||||
attr_func = attr.__func__
|
||||
except AttributeError:
|
||||
continue
|
||||
src = inspect.getsource(attr_func)
|
||||
if "NotImplementedError" in src:
|
||||
unimplemented_methods.append(attr_name)
|
||||
if unimplemented_methods:
|
||||
method_names = ",".join(unimplemented_methods)
|
||||
msg = f"Methods {method_names} not implemented in {self}"
|
||||
logger.warning(msg)
|
||||
|
||||
@wraps(original_init)
|
||||
def wrapped_init(self, *args, **kwargs) -> None:
|
||||
original_init(self, *args, **kwargs)
|
||||
find_unimplemented_methods(self)
|
||||
|
||||
type.__setattr__(cls, "__init__", wrapped_init)
|
||||
return cls
|
||||
|
||||
|
||||
def align_to(value: int, alignment: int) -> int:
|
||||
"""align height, width according to alignment
|
||||
|
||||
Args:
|
||||
value (int): height or width
|
||||
alignment (int): target alignment factor
|
||||
|
||||
Returns:
|
||||
int: the aligned value
|
||||
"""
|
||||
return int(math.ceil(value / alignment) * alignment)
|
||||
|
||||
|
||||
def resolve_obj_by_qualname(qualname: str) -> Any:
|
||||
"""
|
||||
Resolve an object by its fully qualified name.
|
||||
"""
|
||||
module_name, obj_name = qualname.rsplit(".", 1)
|
||||
module = importlib.import_module(module_name)
|
||||
return getattr(module, obj_name)
|
||||
|
||||
|
||||
# From vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/utils.py
|
||||
def import_pynvml():
|
||||
"""
|
||||
Historical comments:
|
||||
|
||||
libnvml.so is the library behind nvidia-smi, and
|
||||
pynvml is a Python wrapper around it. We use it to get GPU
|
||||
status without initializing CUDA context in the current process.
|
||||
Historically, there are two packages that provide pynvml:
|
||||
- `nvidia-ml-py` (https://pypi.org/project/nvidia-ml-py/): The official
|
||||
wrapper. It is a dependency of sglang-diffusion, and is installed when users
|
||||
install sglang-diffusion. It provides a Python module named `pynvml`.
|
||||
- `pynvml` (https://pypi.org/project/pynvml/): An unofficial wrapper.
|
||||
Prior to version 12.0, it also provides a Python module `pynvml`,
|
||||
and therefore conflicts with the official one which is a standalone Python file.
|
||||
This causes errors when both of them are installed.
|
||||
Starting from version 12.0, it migrates to a new module
|
||||
named `pynvml_utils` to avoid the conflict.
|
||||
It is so confusing that many packages in the community use the
|
||||
unofficial one by mistake, and we have to handle this case.
|
||||
For example, `nvcr.io/nvidia/pytorch:24.12-py3` uses the unofficial
|
||||
one, and it will cause errors, see the issue
|
||||
https://github.com/vllm-project/vllm/issues/12847 for example.
|
||||
After all the troubles, we decide to copy the official `pynvml`
|
||||
module to our codebase, and use it directly.
|
||||
"""
|
||||
import sglang.multimodal_gen.third_party.pynvml as pynvml
|
||||
|
||||
return pynvml
|
||||
|
||||
|
||||
def update_environment_variables(envs: dict[str, str]):
|
||||
for k, v in envs.items():
|
||||
if k in os.environ and os.environ[k] != v:
|
||||
logger.warning(
|
||||
"Overwriting environment variable %s from '%s' to '%s'",
|
||||
k,
|
||||
os.environ[k],
|
||||
v,
|
||||
)
|
||||
os.environ[k] = v
|
||||
|
||||
|
||||
def run_method(
|
||||
obj: Any, method: str | bytes | Callable, args: tuple[Any], kwargs: dict[str, Any]
|
||||
) -> Any:
|
||||
"""
|
||||
Run a method of an object with the given arguments and keyword arguments.
|
||||
If the method is string, it will be converted to a method using getattr.
|
||||
If the method is serialized bytes and will be deserialized using
|
||||
cloudpickle.
|
||||
If the method is a callable, it will be called directly.
|
||||
"""
|
||||
if isinstance(method, bytes):
|
||||
func = partial(cloudpickle.loads(method), obj)
|
||||
elif isinstance(method, str):
|
||||
try:
|
||||
func = getattr(obj, method)
|
||||
except AttributeError:
|
||||
raise NotImplementedError(
|
||||
f"Method {method!r} is not implemented."
|
||||
) from None
|
||||
else:
|
||||
func = partial(method, obj) # type: ignore
|
||||
return func(*args, **kwargs)
|
||||
|
||||
|
||||
def shallow_asdict(obj) -> dict[str, Any]:
|
||||
if not is_dataclass(obj):
|
||||
raise TypeError("Expected dataclass instance")
|
||||
return {f.name: getattr(obj, f.name) for f in fields(obj)}
|
||||
|
||||
|
||||
def kill_itself_when_parent_died() -> None:
|
||||
if sys.platform != "linux":
|
||||
return
|
||||
|
||||
# keep GPU workers tied to the CLI process even if the parent is SIGKILLed
|
||||
PR_SET_PDEATHSIG = 1
|
||||
# Capture parent before arming PDEATHSIG: if the parent already died in the
|
||||
# fork->prctl window, PDEATHSIG won't fire, so detect the reparent explicitly.
|
||||
parent_pid = os.getppid()
|
||||
libc = ctypes.CDLL("libc.so.6", use_errno=True)
|
||||
if libc.prctl(PR_SET_PDEATHSIG, signal.SIGKILL) != 0:
|
||||
err = ctypes.get_errno()
|
||||
raise OSError(err, os.strerror(err))
|
||||
# getppid() changing means we were reparented (parent gone). Comparing to the
|
||||
# captured pid instead of "== 1" avoids self-killing when PID 1 is the real
|
||||
# parent (e.g. running as a container's init process).
|
||||
if os.getppid() != parent_pid:
|
||||
os.kill(os.getpid(), signal.SIGKILL)
|
||||
|
||||
|
||||
def get_exception_traceback() -> str:
|
||||
etype, value, tb = sys.exc_info()
|
||||
err_str = "".join(traceback.format_exception(etype, value, tb))
|
||||
return err_str
|
||||
|
||||
|
||||
class TypeBasedDispatcher:
|
||||
def __init__(self, mapping: list[tuple[type, Callable]]):
|
||||
self._mapping = mapping
|
||||
|
||||
def __call__(self, obj: Any):
|
||||
for ty, fn in self._mapping:
|
||||
if isinstance(obj, ty):
|
||||
return fn(obj)
|
||||
raise ValueError(f"Invalid object: {obj}")
|
||||
|
||||
|
||||
@dataclass
|
||||
class MixedPrecisionState:
|
||||
param_dtype: torch.dtype | None = None
|
||||
reduce_dtype: torch.dtype | None = None
|
||||
output_dtype: torch.dtype | None = None
|
||||
compute_dtype: torch.dtype | None = None
|
||||
mp_policy: MixedPrecisionPolicy | None = None
|
||||
|
||||
|
||||
# Thread-local storage for mixed precision state
|
||||
_mixed_precision_state = threading.local()
|
||||
|
||||
|
||||
def get_mixed_precision_state() -> MixedPrecisionState:
|
||||
"""Get the current mixed precision state."""
|
||||
if not hasattr(_mixed_precision_state, "state"):
|
||||
raise ValueError("Mixed precision state not set")
|
||||
return cast(MixedPrecisionState, _mixed_precision_state.state)
|
||||
|
||||
|
||||
def set_mixed_precision_policy(
|
||||
param_dtype: torch.dtype,
|
||||
reduce_dtype: torch.dtype,
|
||||
output_dtype: torch.dtype | None = None,
|
||||
mp_policy: MixedPrecisionPolicy | None = None,
|
||||
):
|
||||
"""Set mixed precision policy globally.
|
||||
|
||||
Args:
|
||||
param_dtype: Parameter dtype used for training
|
||||
reduce_dtype: Reduction dtype used for gradients
|
||||
output_dtype: Optional output dtype
|
||||
"""
|
||||
state = MixedPrecisionState(
|
||||
param_dtype=param_dtype,
|
||||
reduce_dtype=reduce_dtype,
|
||||
output_dtype=output_dtype,
|
||||
mp_policy=mp_policy,
|
||||
)
|
||||
_mixed_precision_state.state = state
|
||||
|
||||
|
||||
def get_compute_dtype() -> torch.dtype:
|
||||
"""Get the current compute dtype from mixed precision policy."""
|
||||
if not hasattr(_mixed_precision_state, "state"):
|
||||
return torch.get_default_dtype()
|
||||
else:
|
||||
state = get_mixed_precision_state()
|
||||
return state.param_dtype
|
||||
|
||||
|
||||
def dict_to_3d_list(
|
||||
mask_strategy: dict[str, Any] | None = None,
|
||||
t_max: int | None = None,
|
||||
l_max: int | None = None,
|
||||
h_max: int | None = None,
|
||||
) -> list[list[list[torch.Tensor | None]]]:
|
||||
"""
|
||||
Convert a dictionary of mask indices to a 3D list of tensors.
|
||||
Args:
|
||||
mask_strategy: keys are "t_l_h", values are torch.Tensor masks.
|
||||
t_max, l_max, h_max: if provided (all three), force the output shape to (t_max, l_max, h_max).
|
||||
If all three are None, infer shape from the data.
|
||||
"""
|
||||
# Case 1: no data, but fixed shape requested
|
||||
if mask_strategy is None:
|
||||
assert t_max is not None and l_max is not None and h_max is not None, (
|
||||
"If mask_strategy is None, you must provide t_max, l_max, and h_max"
|
||||
)
|
||||
return [
|
||||
[[None for _ in range(h_max)] for _ in range(l_max)] for _ in range(t_max)
|
||||
]
|
||||
|
||||
# Parse all keys into integer tuples
|
||||
indices = [tuple(map(int, key.split("_"))) for key in mask_strategy]
|
||||
|
||||
# Decide on dimensions
|
||||
if t_max is None and l_max is None and h_max is None:
|
||||
# fully dynamic: infer from data
|
||||
max_timesteps_idx = max(t for t, _, _ in indices) + 1
|
||||
max_layer_idx = max(l for _, l, _ in indices) + 1 # noqa: E741
|
||||
max_head_idx = max(h for _, _, h in indices) + 1
|
||||
else:
|
||||
# require all three to be provided
|
||||
assert t_max is not None and l_max is not None and h_max is not None, (
|
||||
"Either supply none of (t_max, l_max, h_max) to infer dimensions, "
|
||||
"or supply all three to fix the shape."
|
||||
)
|
||||
max_timesteps_idx = t_max
|
||||
max_layer_idx = l_max
|
||||
max_head_idx = h_max
|
||||
|
||||
# Preallocate
|
||||
result = [
|
||||
[[None for _ in range(max_head_idx)] for _ in range(max_layer_idx)]
|
||||
for _ in range(max_timesteps_idx)
|
||||
]
|
||||
|
||||
# Fill in, skipping any out-of-bounds entries
|
||||
for key, value in mask_strategy.items():
|
||||
t, l, h = map(int, key.split("_")) # noqa: E741
|
||||
if (
|
||||
0 <= t < max_timesteps_idx
|
||||
and 0 <= l < max_layer_idx
|
||||
and 0 <= h < max_head_idx
|
||||
):
|
||||
result[t][l][h] = value
|
||||
# else: silently ignore any key that doesn't fit
|
||||
|
||||
return result
|
||||
|
||||
|
||||
def set_random_seed(seed: int) -> None:
|
||||
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||
|
||||
current_platform.seed_everything(seed)
|
||||
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def is_vsa_available() -> bool:
|
||||
return importlib.util.find_spec("vsa") is not None
|
||||
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def is_vmoba_available() -> bool:
|
||||
if importlib.util.find_spec("kernel.csrc.attn.vmoba_attn.vmoba") is None:
|
||||
return False
|
||||
try:
|
||||
import flash_attn
|
||||
|
||||
return flash_attn.__version__ >= "2.7.4"
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
# adapted from: https://github.com/Wan-Video/Wan2.2/blob/main/wan/utils/utils.py
|
||||
def masks_like(
|
||||
tensors, zero=False, generator=None, p=0.2
|
||||
) -> tuple[list[torch.Tensor], list[torch.Tensor]]:
|
||||
"""
|
||||
Generate binary masks for Text-to-Image-to-Video (TI2V) tasks.
|
||||
|
||||
Creates masks to control which frames should be preserved vs replaced.
|
||||
Primarily used to fix the first frame to the input image while generating other frames.
|
||||
|
||||
Args:
|
||||
tensors: List of tensors with shape [C, T, H, W]
|
||||
zero: If True, set first frame (dim 1, index 0) to zero. Default: False
|
||||
generator: Optional random generator for stochastic masking
|
||||
p: Probability of applying special noise when generator is provided. Default: 0.2
|
||||
|
||||
Returns:
|
||||
Tuple of two lists of tensors:
|
||||
- When zero=False: Both lists contain all-ones tensors
|
||||
- When zero=True (no generator): First frame set to 0, others to 1
|
||||
- When zero=True (with generator): First frame set to small random values with probability p
|
||||
|
||||
Example:
|
||||
>>> latent = torch.randn(48, 69, 96, 160) # [C, T, H, W]
|
||||
>>> _, mask = masks_like([latent], zero=True)
|
||||
>>> # mask[0][:, 0] == 0 (first frame)
|
||||
>>> # mask[0][:, 1:] == 1 (other frames)
|
||||
>>> blended = (1.0 - mask[0]) * image + mask[0] * latent
|
||||
>>> # Result: first frame = image, other frames = latent
|
||||
"""
|
||||
assert isinstance(tensors, list)
|
||||
out1 = [torch.ones(u.shape, dtype=u.dtype, device=u.device) for u in tensors]
|
||||
|
||||
out2 = [torch.ones(u.shape, dtype=u.dtype, device=u.device) for u in tensors]
|
||||
|
||||
if zero:
|
||||
if generator is not None:
|
||||
for u, v in zip(out1, out2, strict=False):
|
||||
random_num = torch.rand(
|
||||
1, generator=generator, device=generator.device
|
||||
).item()
|
||||
if random_num < p:
|
||||
u[:, 0] = (
|
||||
torch.normal(
|
||||
mean=-3.5,
|
||||
std=0.5,
|
||||
size=(1,),
|
||||
device=u.device,
|
||||
generator=generator,
|
||||
)
|
||||
.expand_as(u[:, 0])
|
||||
.exp()
|
||||
)
|
||||
v[:, 0] = torch.zeros_like(v[:, 0])
|
||||
else:
|
||||
u[:, 0] = u[:, 0]
|
||||
v[:, 0] = v[:, 0]
|
||||
|
||||
else:
|
||||
for u, v in zip(out1, out2, strict=False):
|
||||
u[:, 0] = torch.zeros_like(u[:, 0])
|
||||
v[:, 0] = torch.zeros_like(v[:, 0])
|
||||
|
||||
return out1, out2
|
||||
|
||||
|
||||
# adapted from: https://github.com/Wan-Video/Wan2.2/blob/main/wan/utils/utils.py
|
||||
def best_output_size(w, h, dw, dh, expected_area):
|
||||
# float output size
|
||||
ratio = w / h
|
||||
ow = (expected_area * ratio) ** 0.5
|
||||
oh = expected_area / ow
|
||||
|
||||
# process width first
|
||||
ow1 = int(ow // dw * dw)
|
||||
oh1 = int(expected_area / ow1 // dh * dh)
|
||||
assert ow1 % dw == 0 and oh1 % dh == 0 and ow1 * oh1 <= expected_area
|
||||
ratio1 = ow1 / oh1
|
||||
|
||||
# process height first
|
||||
oh2 = int(oh // dh * dh)
|
||||
ow2 = int(expected_area / oh2 // dw * dw)
|
||||
assert oh2 % dh == 0 and ow2 % dw == 0 and ow2 * oh2 <= expected_area
|
||||
ratio2 = ow2 / oh2
|
||||
|
||||
# compare ratios
|
||||
if max(ratio / ratio1, ratio1 / ratio) < max(ratio / ratio2, ratio2 / ratio):
|
||||
return ow1, oh1
|
||||
else:
|
||||
return ow2, oh2
|
||||
|
||||
|
||||
def calculate_dimensions(target_area, ratio):
|
||||
width = math.sqrt(target_area * ratio)
|
||||
height = width / ratio
|
||||
|
||||
width = round(width / 32) * 32
|
||||
height = round(height / 32) * 32
|
||||
|
||||
return width, height, None
|
||||
Reference in New Issue
Block a user