diff --git a/docs/cookbook/diffusion/SANA-WM/SANA-WM.mdx b/docs/cookbook/diffusion/SANA-WM/SANA-WM.mdx index 970894f4f..72ac8b1dc 100644 --- a/docs/cookbook/diffusion/SANA-WM/SANA-WM.mdx +++ b/docs/cookbook/diffusion/SANA-WM/SANA-WM.mdx @@ -351,7 +351,7 @@ Header fields: **Encodings.** `application/x-raw-rgb` is uncompressed RGB24 (3 × uint8, `bytes_per_frame = width*height*3`). `application/x-raw-rgb-delta-gzip` is the zlib-compressed **per-frame XOR delta** against the preceding frame (each frame in the batch is XOR'd against the previous one; sent by default). `realtime_output_format: "raw"` forces uncompressed RGB; `"webp"` / `"jpeg"` send preview-encoded frames. -delta-gzip must be restored **frame-by-frame**: decompress the payload, then for each frame XOR it against the already-restored previous frame (the first frame of a batch references the last frame of the previous batch). See `restore_delta_gzip_raw_rgb_payload` in `runtime/utils/realtime_video.py`. The `"raw"` format below avoids this. +delta-gzip must be restored **frame-by-frame**: decompress the payload, then for each frame XOR it against the already-restored previous frame (the first frame of a batch references the last frame of the previous batch). See `restore_delta_gzip_raw_rgb_payload` in `runtime/realtime/video.py`. The `"raw"` format below avoids this. ### Minimal client example diff --git a/docs/docs/sglang-diffusion/contributing.mdx b/docs/docs/sglang-diffusion/contributing.mdx index f447518f6..1b2130a06 100644 --- a/docs/docs/sglang-diffusion/contributing.mdx +++ b/docs/docs/sglang-diffusion/contributing.mdx @@ -8,7 +8,8 @@ This guide outlines the requirements for contributing to the SGLang Diffusion mo ## Contributor Guides -- [Support New Models](./support_new_models): implementation guide for adding new diffusion pipelines +- [Support New Models](/docs/sglang-diffusion/support_new_models): implementation guide for adding new diffusion pipelines +- [Helper ownership](/docs/sglang-diffusion/support_new_models#place-helpers-with-their-owners): where to put shared and model-specific utilities - [CI Performance](./ci_perf): update and regenerate perf baselines diff --git a/docs/docs/sglang-diffusion/support_new_models.mdx b/docs/docs/sglang-diffusion/support_new_models.mdx index 1dabf8da5..c40f3401b 100644 --- a/docs/docs/sglang-diffusion/support_new_models.mdx +++ b/docs/docs/sglang-diffusion/support_new_models.mdx @@ -33,6 +33,30 @@ utilities, and common action-policy helpers. Model packages may call these helpers. Keep ownership in shared runtime folders unless the code is truly architecture-specific. +## Place helpers with their owners + +Use the narrowest existing owner before adding a utility module: + +| Scope | Location | +| --- | --- | +| Used by one file, or specific to one operation | A private helper in the consuming file; inline short one-off expressions | +| Shared within a domain | A purpose-named module in that domain, such as `runtime/realtime/video.py` or `runtime/layers/attention/mask_strategy.py` | +| Shared across domains, without model or pipeline semantics | An existing focused module under `runtime/utils/`, such as `argparse.py`, `process.py`, or `precision.py` | + +Do not create a top-level `utils/` package or grow a catch-all `utils.py` or +`common.py`. Split large mixed-responsibility files along ownership boundaries, +not arbitrary line counts. A helper folder is warranted only when several +cohesive modules need it, not for a single function or hypothetical reuse. + +Model code must not import pipeline stages. Put contracts shared by models and +stages in a lower-level domain module; for example, realtime cache keys belong +under `runtime/realtime/`. Keep GPU initialization, monkey patches, and model +loading out of generic utility imports. + +When moving internal helpers, update all callers, tests, and cookbook examples +together. Preserve documented registration and serving entry points; do not +add re-export chains just to retain obsolete internal utility paths. + ## Out-of-Tree Models and Pipelines An installed package can register native component models and a pipeline diff --git a/python/sglang/multimodal_gen/configs/models/vaes/base.py b/python/sglang/multimodal_gen/configs/models/vaes/base.py index 5d55e92ee..934a02219 100644 --- a/python/sglang/multimodal_gen/configs/models/vaes/base.py +++ b/python/sglang/multimodal_gen/configs/models/vaes/base.py @@ -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") diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/base.py b/python/sglang/multimodal_gen/configs/pipeline_configs/base.py index 762465ee4..4c51a5397 100644 --- a/python/sglang/multimodal_gen/configs/pipeline_configs/base.py +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/base.py @@ -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): diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/longcat_image.py b/python/sglang/multimodal_gen/configs/pipeline_configs/longcat_image.py index 2d0daf4a9..098d0f29e 100644 --- a/python/sglang/multimodal_gen/configs/pipeline_configs/longcat_image.py +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/longcat_image.py @@ -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) diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/qwen_image.py b/python/sglang/multimodal_gen/configs/pipeline_configs/qwen_image.py index d076d17b9..db62d1e3b 100644 --- a/python/sglang/multimodal_gen/configs/pipeline_configs/qwen_image.py +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/qwen_image.py @@ -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 diff --git a/python/sglang/multimodal_gen/configs/post_training/rl_rollout.py b/python/sglang/multimodal_gen/configs/post_training/rl_rollout.py index 91038329c..57c088883 100644 --- a/python/sglang/multimodal_gen/configs/post_training/rl_rollout.py +++ b/python/sglang/multimodal_gen/configs/post_training/rl_rollout.py @@ -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") diff --git a/python/sglang/multimodal_gen/configs/quantization/nunchaku.py b/python/sglang/multimodal_gen/configs/quantization/nunchaku.py index 3f89ea55f..ab4b6f74c 100644 --- a/python/sglang/multimodal_gen/configs/quantization/nunchaku.py +++ b/python/sglang/multimodal_gen/configs/quantization/nunchaku.py @@ -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__) diff --git a/python/sglang/multimodal_gen/configs/sample/action.py b/python/sglang/multimodal_gen/configs/sample/action.py index 1062a0379..24fb8914c 100644 --- a/python/sglang/multimodal_gen/configs/sample/action.py +++ b/python/sglang/multimodal_gen/configs/sample/action.py @@ -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 diff --git a/python/sglang/multimodal_gen/configs/sample/sampling_params.py b/python/sglang/multimodal_gen/configs/sample/sampling_params.py index 9fd211300..e6950716e 100644 --- a/python/sglang/multimodal_gen/configs/sample/sampling_params.py +++ b/python/sglang/multimodal_gen/configs/sample/sampling_params.py @@ -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__) diff --git a/python/sglang/multimodal_gen/configs/utils.py b/python/sglang/multimodal_gen/configs/utils.py index 11565db01..53ecc5bc5 100644 --- a/python/sglang/multimodal_gen/configs/utils.py +++ b/python/sglang/multimodal_gen/configs/utils.py @@ -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: diff --git a/python/sglang/multimodal_gen/runtime/disaggregation/scheduler_mixin.py b/python/sglang/multimodal_gen/runtime/disaggregation/scheduler_mixin.py index 998241307..9673156cd 100644 --- a/python/sglang/multimodal_gen/runtime/disaggregation/scheduler_mixin.py +++ b/python/sglang/multimodal_gen/runtime/disaggregation/scheduler_mixin.py @@ -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 diff --git a/python/sglang/multimodal_gen/runtime/distributed/device_communicators/pynccl.py b/python/sglang/multimodal_gen/runtime/distributed/device_communicators/pynccl.py index 64d551b5c..18e8e8bec 100644 --- a/python/sglang/multimodal_gen/runtime/distributed/device_communicators/pynccl.py +++ b/python/sglang/multimodal_gen/runtime/distributed/device_communicators/pynccl.py @@ -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 diff --git a/python/sglang/multimodal_gen/runtime/distributed/device_communicators/pynccl_wrapper.py b/python/sglang/multimodal_gen/runtime/distributed/device_communicators/pynccl_wrapper.py index 048f117dd..f218ff06e 100644 --- a/python/sglang/multimodal_gen/runtime/distributed/device_communicators/pynccl_wrapper.py +++ b/python/sglang/multimodal_gen/runtime/distributed/device_communicators/pynccl_wrapper.py @@ -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: diff --git a/python/sglang/multimodal_gen/runtime/distributed/parallel_state.py b/python/sglang/multimodal_gen/runtime/distributed/parallel_state.py index b6584805e..acc4bafc6 100644 --- a/python/sglang/multimodal_gen/runtime/distributed/parallel_state.py +++ b/python/sglang/multimodal_gen/runtime/distributed/parallel_state.py @@ -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__) diff --git a/python/sglang/multimodal_gen/runtime/distributed/utils.py b/python/sglang/multimodal_gen/runtime/distributed/utils.py index a3577610d..0f185e691 100644 --- a/python/sglang/multimodal_gen/runtime/distributed/utils.py +++ b/python/sglang/multimodal_gen/runtime/distributed/utils.py @@ -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 diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/cli/cli_types.py b/python/sglang/multimodal_gen/runtime/entrypoints/cli/cli_types.py index 16b9dd44b..104109069 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/cli/cli_types.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/cli/cli_types.py @@ -5,7 +5,7 @@ import argparse -from sglang.multimodal_gen.utils import FlexibleArgumentParser +from sglang.multimodal_gen.runtime.utils.argparse import FlexibleArgumentParser class CLISubcommand: diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/cli/generate.py b/python/sglang/multimodal_gen/runtime/entrypoints/cli/generate.py index f2c885843..dedcae21c 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/cli/generate.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/cli/generate.py @@ -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__) diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/cli/main.py b/python/sglang/multimodal_gen/runtime/entrypoints/cli/main.py index 26b6a4f67..af5be8722 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/cli/main.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/cli/main.py @@ -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]: diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/cli/serve.py b/python/sglang/multimodal_gen/runtime/entrypoints/cli/serve.py index 49e4743e9..f41c599d8 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/cli/serve.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/cli/serve.py @@ -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): diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/openai/realtime/adapters/lingbot_world_realtime_adapter.py b/python/sglang/multimodal_gen/runtime/entrypoints/openai/realtime/adapters/lingbot_world_realtime_adapter.py index 5313459aa..2e1e1be7f 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/openai/realtime/adapters/lingbot_world_realtime_adapter.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/openai/realtime/adapters/lingbot_world_realtime_adapter.py @@ -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, ) diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/openai/realtime/realtime_output_adapter.py b/python/sglang/multimodal_gen/runtime/entrypoints/openai/realtime/realtime_output_adapter.py index ba567c8a2..70904a704 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/openai/realtime/realtime_output_adapter.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/openai/realtime/realtime_output_adapter.py @@ -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, diff --git a/python/sglang/multimodal_gen/runtime/launch_server.py b/python/sglang/multimodal_gen/runtime/launch_server.py index 187313de4..a0d8d6f1e 100644 --- a/python/sglang/multimodal_gen/runtime/launch_server.py +++ b/python/sglang/multimodal_gen/runtime/launch_server.py @@ -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 diff --git a/python/sglang/multimodal_gen/runtime/layers/attention/STA_configuration.py b/python/sglang/multimodal_gen/runtime/layers/attention/STA_configuration.py index 5c17d0405..e747fe0b5 100644 --- a/python/sglang/multimodal_gen/runtime/layers/attention/STA_configuration.py +++ b/python/sglang/multimodal_gen/runtime/layers/attention/STA_configuration.py @@ -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( diff --git a/python/sglang/multimodal_gen/runtime/layers/attention/backends/sliding_tile_attn.py b/python/sglang/multimodal_gen/runtime/layers/attention/backends/sliding_tile_attn.py index 0db59e292..17d7131a4 100644 --- a/python/sglang/multimodal_gen/runtime/layers/attention/backends/sliding_tile_attn.py +++ b/python/sglang/multimodal_gen/runtime/layers/attention/backends/sliding_tile_attn.py @@ -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 diff --git a/python/sglang/multimodal_gen/runtime/layers/attention/layer.py b/python/sglang/multimodal_gen/runtime/layers/attention/layer.py index b04fcc17f..4b289356a 100644 --- a/python/sglang/multimodal_gen/runtime/layers/attention/layer.py +++ b/python/sglang/multimodal_gen/runtime/layers/attention/layer.py @@ -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, diff --git a/python/sglang/multimodal_gen/runtime/layers/attention/mask_strategy.py b/python/sglang/multimodal_gen/runtime/layers/attention/mask_strategy.py new file mode 100644 index 000000000..003626580 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/layers/attention/mask_strategy.py @@ -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 diff --git a/python/sglang/multimodal_gen/runtime/layers/attention/selector.py b/python/sglang/multimodal_gen/runtime/layers/attention/selector.py index f86cb9b94..f4a0ab145 100644 --- a/python/sglang/multimodal_gen/runtime/layers/attention/selector.py +++ b/python/sglang/multimodal_gen/runtime/layers/attention/selector.py @@ -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( diff --git a/python/sglang/multimodal_gen/runtime/layers/attention/turbo_layer.py b/python/sglang/multimodal_gen/runtime/layers/attention/turbo_layer.py index 31e2c8103..0d09c5941 100644 --- a/python/sglang/multimodal_gen/runtime/layers/attention/turbo_layer.py +++ b/python/sglang/multimodal_gen/runtime/layers/attention/turbo_layer.py @@ -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__) diff --git a/python/sglang/multimodal_gen/runtime/layers/lora/linear.py b/python/sglang/multimodal_gen/runtime/layers/lora/linear.py index b0ccd2af0..1e135b3e8 100644 --- a/python/sglang/multimodal_gen/runtime/layers/lora/linear.py +++ b/python/sglang/multimodal_gen/runtime/layers/lora/linear.py @@ -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 diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/text_encoder_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/text_encoder_loader.py index 147b0e9fe..cfa7555ec 100644 --- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/text_encoder_loader.py +++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/text_encoder_loader.py @@ -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 ( diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/vae_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/vae_loader.py index ddd427f42..58f796fd4 100644 --- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/vae_loader.py +++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/vae_loader.py @@ -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, ) diff --git a/python/sglang/multimodal_gen/runtime/loader/fsdp_load.py b/python/sglang/multimodal_gen/runtime/loader/fsdp_load.py index a275d1e1e..710202111 100644 --- a/python/sglang/multimodal_gen/runtime/loader/fsdp_load.py +++ b/python/sglang/multimodal_gen/runtime/loader/fsdp_load.py @@ -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__) diff --git a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py index fa2797d9e..20be558a4 100644 --- a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py +++ b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py @@ -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 diff --git a/python/sglang/multimodal_gen/runtime/managers/scheduler.py b/python/sglang/multimodal_gen/runtime/managers/scheduler.py index 7cc308fb3..9afafe829 100644 --- a/python/sglang/multimodal_gen/runtime/managers/scheduler.py +++ b/python/sglang/multimodal_gen/runtime/managers/scheduler.py @@ -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 diff --git a/python/sglang/multimodal_gen/runtime/models/dits/lingbot_world.py b/python/sglang/multimodal_gen/runtime/models/dits/lingbot_world.py index e1ea98790..0659d16ce 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/lingbot_world.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/lingbot_world.py @@ -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, ) diff --git a/python/sglang/multimodal_gen/runtime/models/dits/sana_wm.py b/python/sglang/multimodal_gen/runtime/models/dits/sana_wm.py index 7a06fe8e0..2365fe91f 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/sana_wm.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/sana_wm.py @@ -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__) diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/sana_wm/parity_probe.py b/python/sglang/multimodal_gen/runtime/models/dits/sana_wm_parity.py similarity index 100% rename from python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/sana_wm/parity_probe.py rename to python/sglang/multimodal_gen/runtime/models/dits/sana_wm_parity.py diff --git a/python/sglang/multimodal_gen/runtime/models/sensenova_u1/loader.py b/python/sglang/multimodal_gen/runtime/models/sensenova_u1/loader.py index 4c5e136e7..08c7211e8 100644 --- a/python/sglang/multimodal_gen/runtime/models/sensenova_u1/loader.py +++ b/python/sglang/multimodal_gen/runtime/models/sensenova_u1/loader.py @@ -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( diff --git a/python/sglang/multimodal_gen/runtime/models/vlas/pi05_policy.py b/python/sglang/multimodal_gen/runtime/models/vlas/pi05_policy.py index 94faba358..3e7c4d08b 100644 --- a/python/sglang/multimodal_gen/runtime/models/vlas/pi05_policy.py +++ b/python/sglang/multimodal_gen/runtime/models/vlas/pi05_policy.py @@ -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__) diff --git a/python/sglang/multimodal_gen/runtime/pipelines/comfyui_qwen_image_pipeline.py b/python/sglang/multimodal_gen/runtime/pipelines/comfyui_qwen_image_pipeline.py index de708e278..84d898234 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines/comfyui_qwen_image_pipeline.py +++ b/python/sglang/multimodal_gen/runtime/pipelines/comfyui_qwen_image_pipeline.py @@ -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__) diff --git a/python/sglang/multimodal_gen/runtime/pipelines/comfyui_zimage_pipeline.py b/python/sglang/multimodal_gen/runtime/pipelines/comfyui_zimage_pipeline.py index 3be2131d1..ac3700e7e 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines/comfyui_zimage_pipeline.py +++ b/python/sglang/multimodal_gen/runtime/pipelines/comfyui_zimage_pipeline.py @@ -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__) diff --git a/python/sglang/multimodal_gen/runtime/pipelines/hunyuan3d_pipeline.py b/python/sglang/multimodal_gen/runtime/pipelines/hunyuan3d_pipeline.py index 1f8d87bd5..af216a3e2 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines/hunyuan3d_pipeline.py +++ b/python/sglang/multimodal_gen/runtime/pipelines/hunyuan3d_pipeline.py @@ -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__) diff --git a/python/sglang/multimodal_gen/runtime/pipelines/longcat_image.py b/python/sglang/multimodal_gen/runtime/pipelines/longcat_image.py index b843026ad..915d2d995 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines/longcat_image.py +++ b/python/sglang/multimodal_gen/runtime/pipelines/longcat_image.py @@ -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): diff --git a/python/sglang/multimodal_gen/runtime/pipelines/qwen_image.py b/python/sglang/multimodal_gen/runtime/pipelines/qwen_image.py index 5b08bf644..737321334 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines/qwen_image.py +++ b/python/sglang/multimodal_gen/runtime/pipelines/qwen_image.py @@ -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): diff --git a/python/sglang/multimodal_gen/runtime/utils/component_load.py b/python/sglang/multimodal_gen/runtime/pipelines_core/component_loading.py similarity index 100% rename from python/sglang/multimodal_gen/runtime/utils/component_load.py rename to python/sglang/multimodal_gen/runtime/pipelines_core/component_loading.py diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/executors/parallel_executor.py b/python/sglang/multimodal_gen/runtime/pipelines_core/executors/parallel_executor.py index f3ff019ee..3f80a1280 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/executors/parallel_executor.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/executors/parallel_executor.py @@ -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__) diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/schedule_batch.py b/python/sglang/multimodal_gen/runtime/pipelines_core/schedule_batch.py index bc41754db..d931264e1 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/schedule_batch.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/schedule_batch.py @@ -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 diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py index d85370940..64e43b569 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py @@ -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 diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/input_validation.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/input_validation.py index 29814d22b..e15d600e1 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/input_validation.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/input_validation.py @@ -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) diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/helios_decoding.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/helios_decoding.py index 5c2cebdab..a9c7074ba 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/helios_decoding.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/helios_decoding.py @@ -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__) diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/helios_denoising.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/helios_denoising.py index 27521849e..2c04b8b85 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/helios_denoising.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/helios_denoising.py @@ -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__) diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/hunyuan3d/shape.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/hunyuan3d/shape.py index 99ea8f4fd..e2c7c8fb0 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/hunyuan3d/shape.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/hunyuan3d/shape.py @@ -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 diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/ideogram.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/ideogram.py index bddce267b..51f942acf 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/ideogram.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/ideogram.py @@ -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 diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/lingbot_world/lingbot_world_causal_denoising.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/lingbot_world/lingbot_world_causal_denoising.py index 4434f23c0..a86369d06 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/lingbot_world/lingbot_world_causal_denoising.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/lingbot_world/lingbot_world_causal_denoising.py @@ -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 diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/mova.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/mova.py index 41bad0364..df50ae9ff 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/mova.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/mova.py @@ -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() diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/sana_wm/base.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/sana_wm/base.py index 7689c5ddb..d3d051243 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/sana_wm/base.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/sana_wm/base.py @@ -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__) diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/sana_wm/realtime_chain.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/sana_wm/realtime_chain.py index 3a0a06268..a271a427b 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/sana_wm/realtime_chain.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/sana_wm/realtime_chain.py @@ -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, diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/sana_wm/refiner.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/sana_wm/refiner.py index 87e29cf64..282ae050f 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/sana_wm/refiner.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/sana_wm/refiner.py @@ -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, diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/sana_wm/streaming.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/sana_wm/streaming.py index de6d9f8db..f1614b8e4 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/sana_wm/streaming.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/sana_wm/streaming.py @@ -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, diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/sana_wm/streaming_refiner.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/sana_wm/streaming_refiner.py index 46190432e..445247e85 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/sana_wm/streaming_refiner.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/sana_wm/streaming_refiner.py @@ -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, diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/wan_ti2v.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/wan_ti2v.py index 4556bf0f6..792302496 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/wan_ti2v.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/wan_ti2v.py @@ -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 = ( diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/realtime/base.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/realtime/base.py index df38eee2b..1f975ed39 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/realtime/base.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/realtime/base.py @@ -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) diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/realtime/vae.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/realtime/vae.py index d83ea67a3..0461c480f 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/realtime/vae.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/realtime/vae.py @@ -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): diff --git a/python/sglang/multimodal_gen/runtime/platforms/__init__.py b/python/sglang/multimodal_gen/runtime/platforms/__init__.py index 31ccebc9e..ce46524ab 100644 --- a/python/sglang/multimodal_gen/runtime/platforms/__init__.py +++ b/python/sglang/multimodal_gen/runtime/platforms/__init__.py @@ -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 diff --git a/python/sglang/multimodal_gen/runtime/platforms/cuda.py b/python/sglang/multimodal_gen/runtime/platforms/cuda.py index 6acf931e6..b3b5f52d4 100644 --- a/python/sglang/multimodal_gen/runtime/platforms/cuda.py +++ b/python/sglang/multimodal_gen/runtime/platforms/cuda.py @@ -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) diff --git a/python/sglang/multimodal_gen/runtime/platforms/interface.py b/python/sglang/multimodal_gen/runtime/platforms/interface.py index ee4702660..2b30835a2 100644 --- a/python/sglang/multimodal_gen/runtime/platforms/interface.py +++ b/python/sglang/multimodal_gen/runtime/platforms/interface.py @@ -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.""" diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/lingbot_world/constants.py b/python/sglang/multimodal_gen/runtime/realtime/lingbot_world.py similarity index 100% rename from python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/lingbot_world/constants.py rename to python/sglang/multimodal_gen/runtime/realtime/lingbot_world.py diff --git a/python/sglang/multimodal_gen/runtime/utils/realtime_video.py b/python/sglang/multimodal_gen/runtime/realtime/video.py similarity index 100% rename from python/sglang/multimodal_gen/runtime/utils/realtime_video.py rename to python/sglang/multimodal_gen/runtime/realtime/video.py diff --git a/python/sglang/multimodal_gen/runtime/server_args/disagg.py b/python/sglang/multimodal_gen/runtime/server_args/disagg.py index 56223eb66..7453385a1 100644 --- a/python/sglang/multimodal_gen/runtime/server_args/disagg.py +++ b/python/sglang/multimodal_gen/runtime/server_args/disagg.py @@ -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: diff --git a/python/sglang/multimodal_gen/runtime/server_args/server_args.py b/python/sglang/multimodal_gen/runtime/server_args/server_args.py index 045d6d342..b3c00ca5a 100644 --- a/python/sglang/multimodal_gen/runtime/server_args/server_args.py +++ b/python/sglang/multimodal_gen/runtime/server_args/server_args.py @@ -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__) diff --git a/python/sglang/multimodal_gen/runtime/utils/argparse.py b/python/sglang/multimodal_gen/runtime/utils/argparse.py new file mode 100644 index 000000000..0324a6049 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/utils/argparse.py @@ -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 diff --git a/python/sglang/multimodal_gen/runtime/utils/common.py b/python/sglang/multimodal_gen/runtime/utils/common.py index b3826e685..6be3219a9 100644 --- a/python/sglang/multimodal_gen/runtime/utils/common.py +++ b/python/sglang/multimodal_gen/runtime/utils/common.py @@ -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. diff --git a/python/sglang/multimodal_gen/runtime/utils/distributed.py b/python/sglang/multimodal_gen/runtime/utils/distributed.py deleted file mode 100644 index 70f9a4e47..000000000 --- a/python/sglang/multimodal_gen/runtime/utils/distributed.py +++ /dev/null @@ -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 diff --git a/python/sglang/multimodal_gen/runtime/utils/precision.py b/python/sglang/multimodal_gen/runtime/utils/precision.py index 0cb673505..919a1cf36 100644 --- a/python/sglang/multimodal_gen/runtime/utils/precision.py +++ b/python/sglang/multimodal_gen/runtime/utils/precision.py @@ -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 diff --git a/python/sglang/multimodal_gen/runtime/utils/process.py b/python/sglang/multimodal_gen/runtime/utils/process.py new file mode 100644 index 000000000..90c318949 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/utils/process.py @@ -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 diff --git a/python/sglang/multimodal_gen/test/server/realtime_consistency.py b/python/sglang/multimodal_gen/test/server/realtime_consistency.py index bd64fe27c..c40a36545 100644 --- a/python/sglang/multimodal_gen/test/server/realtime_consistency.py +++ b/python/sglang/multimodal_gen/test/server/realtime_consistency.py @@ -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, diff --git a/python/sglang/multimodal_gen/test/server/test_server_utils.py b/python/sglang/multimodal_gen/test/server/test_server_utils.py index ae9ffe31a..513edd5e9 100644 --- a/python/sglang/multimodal_gen/test/server/test_server_utils.py +++ b/python/sglang/multimodal_gen/test/server/test_server_utils.py @@ -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, diff --git a/python/sglang/multimodal_gen/test/unit/realtime/test_lingbot_causal_denoising.py b/python/sglang/multimodal_gen/test/unit/realtime/test_lingbot_causal_denoising.py index 890f3afb9..16d40fdca 100644 --- a/python/sglang/multimodal_gen/test/unit/realtime/test_lingbot_causal_denoising.py +++ b/python/sglang/multimodal_gen/test/unit/realtime/test_lingbot_causal_denoising.py @@ -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" diff --git a/python/sglang/multimodal_gen/test/unit/realtime/test_output_materialization.py b/python/sglang/multimodal_gen/test/unit/realtime/test_output_materialization.py index 5fdaf3d06..4d42e0239 100644 --- a/python/sglang/multimodal_gen/test/unit/realtime/test_output_materialization.py +++ b/python/sglang/multimodal_gen/test/unit/realtime/test_output_materialization.py @@ -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(): diff --git a/python/sglang/multimodal_gen/test/unit/realtime/test_realtime_consistency_harness.py b/python/sglang/multimodal_gen/test/unit/realtime/test_realtime_consistency_harness.py index 44d6a8604..a17d54e1f 100644 --- a/python/sglang/multimodal_gen/test/unit/realtime/test_realtime_consistency_harness.py +++ b/python/sglang/multimodal_gen/test/unit/realtime/test_realtime_consistency_harness.py @@ -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, diff --git a/python/sglang/multimodal_gen/test/unit/realtime/test_realtime_output_transport.py b/python/sglang/multimodal_gen/test/unit/realtime/test_realtime_output_transport.py index a8b2e2808..d52dd9956 100644 --- a/python/sglang/multimodal_gen/test/unit/realtime/test_realtime_output_transport.py +++ b/python/sglang/multimodal_gen/test/unit/realtime/test_realtime_output_transport.py @@ -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, diff --git a/python/sglang/multimodal_gen/test/unit/realtime/test_realtime_runtime.py b/python/sglang/multimodal_gen/test/unit/realtime/test_realtime_runtime.py index 46d22bfea..d9b05dccf 100644 --- a/python/sglang/multimodal_gen/test/unit/realtime/test_realtime_runtime.py +++ b/python/sglang/multimodal_gen/test/unit/realtime/test_realtime_runtime.py @@ -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): diff --git a/python/sglang/multimodal_gen/test/unit/test_attention_backend_selector.py b/python/sglang/multimodal_gen/test/unit/test_attention_backend_selector.py index b0185febe..ff7590f8a 100644 --- a/python/sglang/multimodal_gen/test/unit/test_attention_backend_selector.py +++ b/python/sglang/multimodal_gen/test/unit/test_attention_backend_selector.py @@ -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( diff --git a/python/sglang/multimodal_gen/test/unit/test_cuda_attention_backend.py b/python/sglang/multimodal_gen/test/unit/test_cuda_attention_backend.py index b53cd246b..1964cde61 100644 --- a/python/sglang/multimodal_gen/test/unit/test_cuda_attention_backend.py +++ b/python/sglang/multimodal_gen/test/unit/test_cuda_attention_backend.py @@ -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, ), ): diff --git a/python/sglang/multimodal_gen/test/unit/test_platform_detection.py b/python/sglang/multimodal_gen/test/unit/test_platform_detection.py index 8d198742e..23cf65c51 100644 --- a/python/sglang/multimodal_gen/test/unit/test_platform_detection.py +++ b/python/sglang/multimodal_gen/test/unit/test_platform_detection.py @@ -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), diff --git a/python/sglang/multimodal_gen/test/unit/test_precision_consistency.py b/python/sglang/multimodal_gen/test/unit/test_precision_consistency.py index ca6efaa80..732bb5d57 100644 --- a/python/sglang/multimodal_gen/test/unit/test_precision_consistency.py +++ b/python/sglang/multimodal_gen/test/unit/test_precision_consistency.py @@ -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 diff --git a/python/sglang/multimodal_gen/test/unit/test_qwen3vl_text.py b/python/sglang/multimodal_gen/test/unit/test_qwen3vl_text.py index 908e99d4e..9a7e68236 100644 --- a/python/sglang/multimodal_gen/test/unit/test_qwen3vl_text.py +++ b/python/sglang/multimodal_gen/test/unit/test_qwen3vl_text.py @@ -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, diff --git a/python/sglang/multimodal_gen/test/unit/test_server_args.py b/python/sglang/multimodal_gen/test/unit/test_server_args.py index fa8cad5e0..ca942cca0 100644 --- a/python/sglang/multimodal_gen/test/unit/test_server_args.py +++ b/python/sglang/multimodal_gen/test/unit/test_server_args.py @@ -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 diff --git a/python/sglang/multimodal_gen/test/unit/test_text_encode_dp_gate.py b/python/sglang/multimodal_gen/test/unit/test_text_encode_dp_gate.py index be020137e..57d7f0b52 100644 --- a/python/sglang/multimodal_gen/test/unit/test_text_encode_dp_gate.py +++ b/python/sglang/multimodal_gen/test/unit/test_text_encode_dp_gate.py @@ -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): diff --git a/python/sglang/multimodal_gen/test/unit/test_utility_ownership.py b/python/sglang/multimodal_gen/test/unit/test_utility_ownership.py new file mode 100644 index 000000000..1371a244a --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/test_utility_ownership.py @@ -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]]] diff --git a/python/sglang/multimodal_gen/test/unit/test_utils_parent_death.py b/python/sglang/multimodal_gen/test/unit/test_utils_parent_death.py index d7e6f2980..b1dd1f231 100644 --- a/python/sglang/multimodal_gen/test/unit/test_utils_parent_death.py +++ b/python/sglang/multimodal_gen/test/unit/test_utils_parent_death.py @@ -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): diff --git a/python/sglang/multimodal_gen/test/unit/test_vae_spatial_parallel_decode.py b/python/sglang/multimodal_gen/test/unit/test_vae_spatial_parallel_decode.py index 0df66e936..19dd4cd45 100644 --- a/python/sglang/multimodal_gen/test/unit/test_vae_spatial_parallel_decode.py +++ b/python/sglang/multimodal_gen/test/unit/test_vae_spatial_parallel_decode.py @@ -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): diff --git a/python/sglang/multimodal_gen/test/unit/test_wan_ti2v_helpers.py b/python/sglang/multimodal_gen/test/unit/test_wan_ti2v_helpers.py index a150005f9..814dfe7cb 100644 --- a/python/sglang/multimodal_gen/test/unit/test_wan_ti2v_helpers.py +++ b/python/sglang/multimodal_gen/test/unit/test_wan_ti2v_helpers.py @@ -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 diff --git a/python/sglang/multimodal_gen/utils.py b/python/sglang/multimodal_gen/utils.py deleted file mode 100644 index be12a1551..000000000 --- a/python/sglang/multimodal_gen/utils.py +++ /dev/null @@ -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