From 4a8e1b07a2ea85d65ae4b8f74ffa361bc627df41 Mon Sep 17 00:00:00 2001 From: Mick Date: Fri, 10 Jul 2026 15:53:00 +0800 Subject: [PATCH] [diffusion] refactor: reorganize runtime utility and server_args modules (#30447) --- .../configs/pipeline_configs/base.py | 2 +- .../configs/pipeline_configs/qwen_image.py | 2 +- .../runtime/disaggregation/disagg_args.py | 2 +- .../layers/attention/backends/aiter.py | 4 +- .../runtime/layers/layernorm.py | 7 +- .../multimodal_gen/runtime/layers/linear.py | 6 +- .../layers/quantization/bitsandbytes.py | 2 +- .../runtime/layers/quantization/fp8.py | 4 +- .../layers/quantization/modelopt_quant.py | 2 +- .../layers/quantization/nunchaku_linear.py | 2 +- .../layers/quantization/weight_only_fp8.py | 2 +- .../layers/vocab_parallel_embedding.py | 2 +- .../runtime/models/dits/common.py | 18 ++ .../runtime/models/dits/hunyuanvideo.py | 2 +- .../runtime/models/dits/joy_image.py | 2 +- .../runtime/models/dits/lingbot_world.py | 6 +- .../runtime/models/dits/wanvideo.py | 10 +- .../runtime/models/parameter.py | 4 +- .../multimodal_gen/runtime/models/utils.py | 156 ------------------ .../runtime/pipelines/diffusers_pipeline.py | 4 +- .../diffusion_scheduler_utils.py | 40 +++++ .../pipelines_core/stages/causal_denoising.py | 2 +- .../pipelines_core/stages/denoising_dmd.py | 4 +- .../pipelines_core/stages/image_encoding.py | 12 +- .../pipelines_core/stages/input_validation.py | 2 +- .../stages/model_specific_stages/cosmos3.py | 2 +- .../stages/model_specific_stages/glm_image.py | 2 +- .../qwen_image_layered.py | 2 +- .../multimodal_gen/runtime/platforms/aiter.py | 10 ++ .../runtime/server_args/__init__.py | 40 +++++ .../auto_tune.py} | 2 +- .../disagg.py} | 0 .../runtime/{ => server_args}/server_args.py | 4 +- .../vision_utils.py => utils/vision.py} | 0 .../runtime/utils/weight_attrs.py | 33 ++++ .../test/unit/test_server_args.py | 58 +++---- 36 files changed, 217 insertions(+), 235 deletions(-) create mode 100644 python/sglang/multimodal_gen/runtime/models/dits/common.py delete mode 100644 python/sglang/multimodal_gen/runtime/models/utils.py create mode 100644 python/sglang/multimodal_gen/runtime/platforms/aiter.py create mode 100644 python/sglang/multimodal_gen/runtime/server_args/__init__.py rename python/sglang/multimodal_gen/runtime/{server_args_auto_tune.py => server_args/auto_tune.py} (99%) rename python/sglang/multimodal_gen/runtime/{server_args_disagg.py => server_args/disagg.py} (100%) rename python/sglang/multimodal_gen/runtime/{ => server_args}/server_args.py (99%) rename python/sglang/multimodal_gen/runtime/{models/vision_utils.py => utils/vision.py} (100%) create mode 100644 python/sglang/multimodal_gen/runtime/utils/weight_attrs.py diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/base.py b/python/sglang/multimodal_gen/configs/pipeline_configs/base.py index 9610e8b5f..09a2068ec 100644 --- a/python/sglang/multimodal_gen/configs/pipeline_configs/base.py +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/base.py @@ -35,8 +35,8 @@ from sglang.multimodal_gen.runtime.distributed.parallel_state import ( get_sp_parallel_rank, get_sp_world_size, ) -from sglang.multimodal_gen.runtime.models.vision_utils import get_default_height_width 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 ( FlexibleArgumentParser, StoreBoolean, 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 9bdf486c3..1bfb41e96 100644 --- a/python/sglang/multimodal_gen/configs/pipeline_configs/qwen_image.py +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/qwen_image.py @@ -22,7 +22,7 @@ from sglang.multimodal_gen.configs.pipeline_configs.base import ( from sglang.multimodal_gen.configs.post_training.pipeline_configs import ( QwenImageRolloutPipelineMixin, ) -from sglang.multimodal_gen.runtime.models.vision_utils import resize +from sglang.multimodal_gen.runtime.utils.vision import resize from sglang.multimodal_gen.utils import calculate_dimensions diff --git a/python/sglang/multimodal_gen/runtime/disaggregation/disagg_args.py b/python/sglang/multimodal_gen/runtime/disaggregation/disagg_args.py index 15167d7b0..5e6c35922 100644 --- a/python/sglang/multimodal_gen/runtime/disaggregation/disagg_args.py +++ b/python/sglang/multimodal_gen/runtime/disaggregation/disagg_args.py @@ -6,7 +6,7 @@ from __future__ import annotations import argparse from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType -from sglang.multimodal_gen.runtime.server_args_disagg import DisaggServerArgsMixin +from sglang.multimodal_gen.runtime.server_args.disagg import DisaggServerArgsMixin # Keep the historical disagg_args import path working. DISAGG_RESULT_PORT_OFFSETS = DisaggServerArgsMixin.DISAGG_RESULT_PORT_OFFSETS diff --git a/python/sglang/multimodal_gen/runtime/layers/attention/backends/aiter.py b/python/sglang/multimodal_gen/runtime/layers/attention/backends/aiter.py index 024c4d63b..2b7b8eea5 100755 --- a/python/sglang/multimodal_gen/runtime/layers/attention/backends/aiter.py +++ b/python/sglang/multimodal_gen/runtime/layers/attention/backends/aiter.py @@ -15,7 +15,7 @@ from sglang.multimodal_gen.runtime.layers.attention.backends.attention_backend i AttentionMetadataBuilder, ) from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum -from sglang.srt.models.deepseek_common.utils import _use_aiter_gfx95 +from sglang.multimodal_gen.runtime.platforms.aiter import USE_AITER_GFX95 logger = logging.getLogger(__name__) @@ -38,7 +38,7 @@ def _can_use_fmha_fp8_prefill( num_kv_heads: int, ) -> bool: """True if MHA q/k/v head_dim==128 on a gfx950-class arch.""" - if not _use_aiter_gfx95: + if not USE_AITER_GFX95: return False if num_kv_heads != num_heads: return False diff --git a/python/sglang/multimodal_gen/runtime/layers/layernorm.py b/python/sglang/multimodal_gen/runtime/layers/layernorm.py index d7fde1f32..795972bc0 100755 --- a/python/sglang/multimodal_gen/runtime/layers/layernorm.py +++ b/python/sglang/multimodal_gen/runtime/layers/layernorm.py @@ -25,16 +25,15 @@ from sglang.multimodal_gen.runtime.distributed.parallel_state import ( ) from sglang.multimodal_gen.runtime.layers.custom_op import CustomOp from sglang.multimodal_gen.runtime.platforms import current_platform +from sglang.multimodal_gen.runtime.platforms.aiter import USE_AITER from sglang.multimodal_gen.runtime.utils.common import get_bool_env_var _is_cuda = current_platform.is_cuda() -_is_hip = current_platform.is_hip() _is_npu = current_platform.is_npu() _is_musa = current_platform.is_musa() _is_cpu = current_platform.is_cpu() _is_xpu = current_platform.is_xpu() _use_rocm_flydsl = get_bool_env_var("SGLANG_USE_ROCM_FLYDSL") -_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip if _is_cuda or _is_xpu: from sgl_kernel import fused_add_rmsnorm, rmsnorm @@ -48,7 +47,7 @@ if _is_npu: if _is_musa: from sgl_kernel import fused_add_rmsnorm -if _use_aiter: +if USE_AITER: from aiter import rmsnorm2d_fwd as rms_norm from aiter import rmsnorm2d_fwd_with_add as fused_add_rms_norm @@ -81,7 +80,7 @@ class RMSNorm(CustomOp): ) if get_bool_env_var("SGLANG_ENABLE_DETERMINISTIC_INFERENCE"): self._forward_method = self.forward_native - elif _use_aiter: + elif USE_AITER: self._forward_method = self.forward_aiter def forward_triton(self, x: torch.Tensor, residual: Optional[torch.Tensor] = None): diff --git a/python/sglang/multimodal_gen/runtime/layers/linear.py b/python/sglang/multimodal_gen/runtime/layers/linear.py index 94affa99b..19007ecd0 100644 --- a/python/sglang/multimodal_gen/runtime/layers/linear.py +++ b/python/sglang/multimodal_gen/runtime/layers/linear.py @@ -33,12 +33,12 @@ from sglang.multimodal_gen.runtime.models.parameter import ( PerTensorScaleParameter, RowvLLMParameter, ) - -# yapf: enable -from sglang.multimodal_gen.runtime.models.utils import set_weight_attrs from sglang.multimodal_gen.runtime.platforms import current_platform from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger +# yapf: enable +from sglang.multimodal_gen.runtime.utils.weight_attrs import set_weight_attrs + logger = init_logger(__name__) IS_AMP_SUPPORTED = current_platform.is_amp_supported() diff --git a/python/sglang/multimodal_gen/runtime/layers/quantization/bitsandbytes.py b/python/sglang/multimodal_gen/runtime/layers/quantization/bitsandbytes.py index 6895ad000..36482f0bc 100644 --- a/python/sglang/multimodal_gen/runtime/layers/quantization/bitsandbytes.py +++ b/python/sglang/multimodal_gen/runtime/layers/quantization/bitsandbytes.py @@ -17,7 +17,7 @@ from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config impor QuantizationConfig, QuantizeMethodBase, ) -from sglang.multimodal_gen.runtime.models.utils import set_weight_attrs +from sglang.multimodal_gen.runtime.utils.weight_attrs import set_weight_attrs def _require_bitsandbytes() -> None: diff --git a/python/sglang/multimodal_gen/runtime/layers/quantization/fp8.py b/python/sglang/multimodal_gen/runtime/layers/quantization/fp8.py index 06ad6aefa..9312d7100 100644 --- a/python/sglang/multimodal_gen/runtime/layers/quantization/fp8.py +++ b/python/sglang/multimodal_gen/runtime/layers/quantization/fp8.py @@ -24,6 +24,7 @@ from sglang.multimodal_gen.runtime.models.parameter import ( PerTensorScaleParameter, ) from sglang.multimodal_gen.runtime.platforms import current_platform +from sglang.multimodal_gen.runtime.platforms.aiter import USE_AITER from sglang.multimodal_gen.runtime.utils.common import ( cpu_has_amx_support, get_bool_env_var, @@ -63,9 +64,8 @@ _is_cpu_amx_available = cpu_has_amx_support() _is_cpu = current_platform.is_cpu() _is_fp8_fnuz = is_fp8_fnuz() _use_hip_int4 = get_bool_env_var("SGLANG_INT4_WEIGHT") and _is_hip -_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip -if _use_aiter or _use_hip_int4: +if USE_AITER or _use_hip_int4: pass diff --git a/python/sglang/multimodal_gen/runtime/layers/quantization/modelopt_quant.py b/python/sglang/multimodal_gen/runtime/layers/quantization/modelopt_quant.py index f539c1ed5..9f91f06a8 100755 --- a/python/sglang/multimodal_gen/runtime/layers/quantization/modelopt_quant.py +++ b/python/sglang/multimodal_gen/runtime/layers/quantization/modelopt_quant.py @@ -20,8 +20,8 @@ from sglang.multimodal_gen.runtime.models.parameter import ( ModelWeightParameter, PerTensorScaleParameter, ) -from sglang.multimodal_gen.runtime.models.utils import set_weight_attrs from sglang.multimodal_gen.runtime.platforms import current_platform +from sglang.multimodal_gen.runtime.utils.weight_attrs import set_weight_attrs from sglang.srt.layers.quantization.fp8_utils import ( apply_fp8_linear, cutlass_fp8_supported, diff --git a/python/sglang/multimodal_gen/runtime/layers/quantization/nunchaku_linear.py b/python/sglang/multimodal_gen/runtime/layers/quantization/nunchaku_linear.py index 516f76699..e79ed3709 100644 --- a/python/sglang/multimodal_gen/runtime/layers/quantization/nunchaku_linear.py +++ b/python/sglang/multimodal_gen/runtime/layers/quantization/nunchaku_linear.py @@ -6,8 +6,8 @@ import torch.nn as nn from torch.nn.parameter import Parameter from sglang.multimodal_gen.runtime.layers.linear import LinearMethodBase -from sglang.multimodal_gen.runtime.models.utils import set_weight_attrs from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger +from sglang.multimodal_gen.runtime.utils.weight_attrs import set_weight_attrs logger = init_logger(__name__) diff --git a/python/sglang/multimodal_gen/runtime/layers/quantization/weight_only_fp8.py b/python/sglang/multimodal_gen/runtime/layers/quantization/weight_only_fp8.py index 2d540eb5a..0e19ee07b 100644 --- a/python/sglang/multimodal_gen/runtime/layers/quantization/weight_only_fp8.py +++ b/python/sglang/multimodal_gen/runtime/layers/quantization/weight_only_fp8.py @@ -15,7 +15,7 @@ from sglang.multimodal_gen.runtime.distributed import ( tensor_model_parallel_all_reduce, ) from sglang.multimodal_gen.runtime.layers.utils import get_group_rank, get_group_size -from sglang.multimodal_gen.runtime.models.utils import set_weight_attrs +from sglang.multimodal_gen.runtime.utils.weight_attrs import set_weight_attrs FP8_WEIGHT_DTYPE = torch.float8_e4m3fn W8A8_FP8_GEMM_ENV = "SGLANG_DIFFUSION_ENABLE_W8A8_FP8_GEMM" diff --git a/python/sglang/multimodal_gen/runtime/layers/vocab_parallel_embedding.py b/python/sglang/multimodal_gen/runtime/layers/vocab_parallel_embedding.py index fecb4245f..30950f013 100644 --- a/python/sglang/multimodal_gen/runtime/layers/vocab_parallel_embedding.py +++ b/python/sglang/multimodal_gen/runtime/layers/vocab_parallel_embedding.py @@ -22,8 +22,8 @@ from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config impor ) from sglang.multimodal_gen.runtime.layers.utils import get_group_rank, get_group_size from sglang.multimodal_gen.runtime.models.parameter import BasevLLMParameter -from sglang.multimodal_gen.runtime.models.utils import set_weight_attrs from sglang.multimodal_gen.runtime.platforms import current_platform +from sglang.multimodal_gen.runtime.utils.weight_attrs import set_weight_attrs DEFAULT_VOCAB_PADDING_SIZE = 64 diff --git a/python/sglang/multimodal_gen/runtime/models/dits/common.py b/python/sglang/multimodal_gen/runtime/models/dits/common.py new file mode 100644 index 000000000..4f98ef813 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/models/dits/common.py @@ -0,0 +1,18 @@ +# SPDX-License-Identifier: Apache-2.0 + +import torch + + +def modulate( + x: torch.Tensor, + shift: torch.Tensor | None = None, + scale: torch.Tensor | None = None, +) -> torch.Tensor: + """Modulate by shift and scale.""" + if scale is None and shift is None: + return x + if shift is None: + return x * (1 + scale.unsqueeze(1)) # type: ignore[union-attr] + if scale is None: + return x + shift.unsqueeze(1) # type: ignore[union-attr] + return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1) diff --git a/python/sglang/multimodal_gen/runtime/models/dits/hunyuanvideo.py b/python/sglang/multimodal_gen/runtime/models/dits/hunyuanvideo.py index 30d764b8d..2f329c564 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/hunyuanvideo.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/hunyuanvideo.py @@ -50,7 +50,7 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload im LayerwiseOffloadableModuleMixin, ) from sglang.multimodal_gen.runtime.models.dits.base import CachableDiT -from sglang.multimodal_gen.runtime.models.utils import modulate +from sglang.multimodal_gen.runtime.models.dits.common import modulate from sglang.multimodal_gen.runtime.platforms import ( AttentionBackendEnum, current_platform, diff --git a/python/sglang/multimodal_gen/runtime/models/dits/joy_image.py b/python/sglang/multimodal_gen/runtime/models/dits/joy_image.py index a3a87cbd7..d7d56ff50 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/joy_image.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/joy_image.py @@ -38,11 +38,11 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload im ) from sglang.multimodal_gen.runtime.models.dits.base import CachableDiT from sglang.multimodal_gen.runtime.models.dits.wanvideo import WanTimeTextImageEmbedding -from sglang.multimodal_gen.runtime.models.utils import set_weight_attrs from sglang.multimodal_gen.runtime.platforms import ( AttentionBackendEnum, ) from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger +from sglang.multimodal_gen.runtime.utils.weight_attrs import set_weight_attrs logger = init_logger(__name__) _MODULATION_FACTOR = 6 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 849438704..ed0f2303f 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/lingbot_world.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/lingbot_world.py @@ -78,7 +78,6 @@ from sglang.multimodal_gen.runtime.models.dits.wanvideo import ( WanTimeTextImageEmbedding, WanTransformer3DModel, ) -from sglang.multimodal_gen.runtime.models.utils import _use_aiter from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.lingbot_world.constants import ( LINGBOT_C2WS_PLUCKER_EMB_CACHE, LINGBOT_CAM_CONDITIONER_CACHE, @@ -90,6 +89,7 @@ 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, ) @@ -111,7 +111,7 @@ def _safe_tensor_version(tensor: torch.Tensor) -> int: return 0 if tensor.is_inference() else tensor._version -if _use_aiter: +if USE_AITER: from aiter.ops.rope import rope_cached_2c_fwd_inplace @@ -515,7 +515,7 @@ class LingBotWorldTransformerBlock(nn.Module): query, key = apply_flashinfer_rope_qk_inplace( query, key, cos_sin_cache, is_neox=False ) - elif _use_aiter: + elif USE_AITER: query_shape = query.shape key_shape = key.shape num_tokens = query.shape[:-2].numel() diff --git a/python/sglang/multimodal_gen/runtime/models/dits/wanvideo.py b/python/sglang/multimodal_gen/runtime/models/dits/wanvideo.py index 09de378b1..907871973 100755 --- a/python/sglang/multimodal_gen/runtime/models/dits/wanvideo.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/wanvideo.py @@ -53,13 +53,11 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload im LayerwiseOffloadableModuleMixin, ) from sglang.multimodal_gen.runtime.models.dits.base import CachableDiT -from sglang.multimodal_gen.runtime.models.utils import ( - _use_aiter, -) 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.server_args import get_global_server_args from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.srt.utils import add_prefix @@ -67,7 +65,7 @@ from sglang.srt.utils import add_prefix logger = init_logger(__name__) _is_cuda = current_platform.is_cuda() -if _use_aiter: +if USE_AITER: from aiter.ops.rope import rope_cached_2c_fwd_inplace @@ -555,7 +553,7 @@ class WanTransformerBlock(nn.Module): query, key = apply_flashinfer_rope_qk_inplace( query, key, cos_sin_cache, is_neox=False ) - elif _use_aiter: + elif USE_AITER: query_shape = query.shape key_shape = key.shape num_tokens = query.shape[:-2].numel() @@ -802,7 +800,7 @@ class WanTransformerBlock_VSA(nn.Module): query, key = apply_flashinfer_rope_qk_inplace( query, key, cos_sin_cache, is_neox=False ) - elif _use_aiter: + elif USE_AITER: query_shape = query.shape key_shape = key.shape num_tokens = query.shape[:-2].numel() diff --git a/python/sglang/multimodal_gen/runtime/models/parameter.py b/python/sglang/multimodal_gen/runtime/models/parameter.py index f3f171a69..4f72feb74 100644 --- a/python/sglang/multimodal_gen/runtime/models/parameter.py +++ b/python/sglang/multimodal_gen/runtime/models/parameter.py @@ -12,8 +12,8 @@ import torch from torch.nn import Parameter from sglang.multimodal_gen.runtime.distributed import get_tp_rank -from sglang.multimodal_gen.runtime.models.utils import _make_synced_weight_loader from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger +from sglang.multimodal_gen.runtime.utils.weight_attrs import make_synced_weight_loader logger = init_logger(__name__) @@ -50,7 +50,7 @@ class BasevLLMParameter(Parameter): from sglang.multimodal_gen.runtime.platforms import current_platform if current_platform.is_tpu(): - weight_loader = _make_synced_weight_loader(weight_loader) + weight_loader = make_synced_weight_loader(weight_loader) self._weight_loader = weight_loader diff --git a/python/sglang/multimodal_gen/runtime/models/utils.py b/python/sglang/multimodal_gen/runtime/models/utils.py deleted file mode 100644 index 7628e5922..000000000 --- a/python/sglang/multimodal_gen/runtime/models/utils.py +++ /dev/null @@ -1,156 +0,0 @@ -# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo - -# SPDX-License-Identifier: Apache-2.0 -# SPDX-FileCopyrightText: Copyright contributors to the vLLM project -# Adapted from: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/utils.py -"""Utils for model executor.""" - -from typing import Any - -import torch - -from sglang.multimodal_gen.runtime.platforms import current_platform -from sglang.srt.utils import ( - get_bool_env_var, - is_gfx95_supported, - is_hip, -) - -_is_hip = is_hip() -_is_gfx95_supported = is_gfx95_supported() -_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip -_use_aiter_gfx95 = _use_aiter and _is_gfx95_supported - - -def set_weight_attrs( - weight: torch.Tensor, - weight_attrs: dict[str, Any] | None, -): - """Set attributes on a weight tensor. - - This method is used to set attributes on a weight tensor. This method - will not overwrite existing attributes. - - Args: - weight: The weight tensor. - weight_attrs: A dictionary of attributes to set on the weight tensor. - """ - if weight_attrs is None: - return - for key, value in weight_attrs.items(): - assert not hasattr(weight, key), f"Overwriting existing tensor attribute: {key}" - - # NOTE(woosuk): During weight loading, we often do something like: - # narrowed_tensor = param.data.narrow(0, offset, len) - # narrowed_tensor.copy_(real_weight) - # expecting narrowed_tensor and param.data to share the same storage. - # However, on TPUs, narrowed_tensor will lazily propagate to the base - # tensor, which is param.data, leading to the redundant memory usage. - # This sometimes causes OOM errors during model loading. To avoid this, - # we sync the param tensor after its weight loader is called. - # TODO(woosuk): Remove this hack once we have a better solution. - from sglang.multimodal_gen.runtime.platforms import current_platform - - if current_platform.is_tpu() and key == "weight_loader": - value = _make_synced_weight_loader(value) - setattr(weight, key, value) - - -def _make_synced_weight_loader(original_weight_loader) -> Any: - - def _synced_weight_loader(param, *args, **kwargs): - original_weight_loader(param, *args, **kwargs) - torch._sync(param) - - return _synced_weight_loader - - -def extract_layer_index(layer_name: str) -> int: - """ - Extract the layer index from the module name. - Examples: - - "encoder.layers.0" -> 0 - - "encoder.layers.1.self_attn" -> 1 - - "2.self_attn" -> 2 - - "model.encoder.layers.0.sub.1" -> ValueError - """ - subnames = layer_name.split(".") - int_vals: list[int] = [] - for subname in subnames: - try: - int_vals.append(int(subname)) - except ValueError: - continue - assert len(int_vals) == 1, ( - f"layer name {layer_name} should" " only contain one integer" - ) - return int_vals[0] - - -def modulate( - x: torch.Tensor, - shift: torch.Tensor | None = None, - scale: torch.Tensor | None = None, -) -> torch.Tensor: - """modulate by shift and scale""" - if scale is None and shift is None: - return x - elif shift is None: - return x * (1 + scale.unsqueeze(1)) # type: ignore[union-attr] - elif scale is None: - return x + shift.unsqueeze(1) # type: ignore[union-attr] - else: - return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze( - 1 - ) # type: ignore[union-attr] - - -def pred_noise_to_pred_video( - pred_noise: torch.Tensor, - noise_input_latent: torch.Tensor, - timestep: torch.Tensor, - scheduler: Any, -) -> torch.Tensor: - """ - Convert predicted noise to clean latent. - - Args: - pred_noise: the predicted noise with shape [B, C, H, W] - where B is batch_size or batch_size * num_frames - noise_input_latent: the noisy latent with shape [B, C, H, W], - timestep: the timestep with shape [1] or [bs * num_frames] or [bs, num_frames] - scheduler: the scheduler - - Returns: - the predicted video with shape [B, C, H, W] - """ - # If timestep is [bs, num_frames] - if timestep.ndim == 2: - timestep = timestep.flatten(0, 1) - assert timestep.numel() == noise_input_latent.shape[0] - elif timestep.ndim == 1: - # If timestep is [1] - if timestep.shape[0] == 1: - timestep = timestep.expand(noise_input_latent.shape[0]) - else: - assert timestep.numel() == noise_input_latent.shape[0] - else: - raise ValueError( - f"[pred_noise_to_pred_video] Invalid timestep shape: {timestep.shape}" - ) - # timestep shape should be [B] - dtype = pred_noise.dtype - device = pred_noise.device - pred_noise = pred_noise.double().to(device) - noise_input_latent = noise_input_latent.double().to(device) - sigmas = scheduler.sigmas.double().to(device) - high_dtype = ( - torch.float64 if current_platform.is_float64_supported() else torch.float32 - ) - timesteps = scheduler.timesteps.to(high_dtype).to(device) - timestep_id = torch.argmin( - (timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1 - ) - sigma_t = sigmas[timestep_id].reshape(-1, 1, 1, 1) - pred_video = noise_input_latent - sigma_t * pred_noise - return pred_video.to(dtype) diff --git a/python/sglang/multimodal_gen/runtime/pipelines/diffusers_pipeline.py b/python/sglang/multimodal_gen/runtime/pipelines/diffusers_pipeline.py index e66c695e7..c84e74e19 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines/diffusers_pipeline.py +++ b/python/sglang/multimodal_gen/runtime/pipelines/diffusers_pipeline.py @@ -24,9 +24,6 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager im ComponentResidencyStrategy, get_global_component_residency_manager, ) -from sglang.multimodal_gen.runtime.models.vision_utils import ( - load_image as load_vision_image, -) from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import ( ComposedPipelineBase, ) @@ -46,6 +43,7 @@ from sglang.multimodal_gen.runtime.server_args import ServerArgs 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 resolve_precision +from sglang.multimodal_gen.runtime.utils.vision import load_image as load_vision_image logger = init_logger(__name__) diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/diffusion_scheduler_utils.py b/python/sglang/multimodal_gen/runtime/pipelines_core/diffusion_scheduler_utils.py index 04d3628f0..1fea8c0e8 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/diffusion_scheduler_utils.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/diffusion_scheduler_utils.py @@ -5,7 +5,10 @@ from __future__ import annotations from copy import deepcopy from typing import Any +import torch + from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req +from sglang.multimodal_gen.runtime.platforms import current_platform def clone_scheduler_runtime(scheduler: Any) -> Any: @@ -30,3 +33,40 @@ def get_or_create_request_scheduler( else scheduler_template ) return batch.scheduler + + +def pred_noise_to_pred_video( + pred_noise: torch.Tensor, + noise_input_latent: torch.Tensor, + timestep: torch.Tensor, + scheduler: Any, +) -> torch.Tensor: + """Convert predicted noise to clean latent.""" + if timestep.ndim == 2: + timestep = timestep.flatten(0, 1) + assert timestep.numel() == noise_input_latent.shape[0] + elif timestep.ndim == 1: + if timestep.shape[0] == 1: + timestep = timestep.expand(noise_input_latent.shape[0]) + else: + assert timestep.numel() == noise_input_latent.shape[0] + else: + raise ValueError( + f"[pred_noise_to_pred_video] Invalid timestep shape: {timestep.shape}" + ) + + dtype = pred_noise.dtype + device = pred_noise.device + pred_noise = pred_noise.double().to(device) + noise_input_latent = noise_input_latent.double().to(device) + sigmas = scheduler.sigmas.double().to(device) + high_dtype = ( + torch.float64 if current_platform.is_float64_supported() else torch.float32 + ) + timesteps = scheduler.timesteps.to(high_dtype).to(device) + timestep_id = torch.argmin( + (timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1 + ) + sigma_t = sigmas[timestep_id].reshape(-1, 1, 1, 1) + pred_video = noise_input_latent - sigma_t * pred_noise + return pred_video.to(dtype) diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/causal_denoising.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/causal_denoising.py index da9c5207f..af1eb5f39 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/causal_denoising.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/causal_denoising.py @@ -13,9 +13,9 @@ from sglang.multimodal_gen.runtime.layers.kvcache.causal_attention_cache import CrossAttentionKVCache, ) from sglang.multimodal_gen.runtime.managers.forward_context import set_forward_context -from sglang.multimodal_gen.runtime.models.utils import pred_noise_to_pred_video from sglang.multimodal_gen.runtime.pipelines_core.diffusion_scheduler_utils import ( get_or_create_request_scheduler, + pred_noise_to_pred_video, ) from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req from sglang.multimodal_gen.runtime.pipelines_core.stages.denoising import DenoisingStage diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising_dmd.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising_dmd.py index 8b12b6b28..ef182c778 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising_dmd.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising_dmd.py @@ -9,7 +9,9 @@ from sglang.multimodal_gen.runtime.managers.forward_context import set_forward_c from sglang.multimodal_gen.runtime.models.schedulers.scheduling_flow_match_euler_discrete import ( FlowMatchEulerDiscreteScheduler, ) -from sglang.multimodal_gen.runtime.models.utils import pred_noise_to_pred_video +from sglang.multimodal_gen.runtime.pipelines_core.diffusion_scheduler_utils import ( + pred_noise_to_pred_video, +) from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req from sglang.multimodal_gen.runtime.pipelines_core.stages import DenoisingStage from sglang.multimodal_gen.runtime.platforms import current_platform diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/image_encoding.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/image_encoding.py index 91508caa6..7afce8967 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/image_encoding.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/image_encoding.py @@ -28,11 +28,6 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload im configure_layerwise_offload_modules, ) from sglang.multimodal_gen.runtime.models.vaes.common import ParallelTiledVAE -from sglang.multimodal_gen.runtime.models.vision_utils import ( - normalize, - numpy_to_pt, - pil_to_numpy, -) from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req from sglang.multimodal_gen.runtime.pipelines_core.stages.base import PipelineStage from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import ( @@ -50,6 +45,11 @@ from sglang.multimodal_gen.runtime.utils.precision import ( resolve_precision, temporary_module_dtype, ) +from sglang.multimodal_gen.runtime.utils.vision import ( + normalize, + numpy_to_pt, + pil_to_numpy, +) logger = init_logger(__name__) @@ -710,7 +710,7 @@ class LTX2ImageEncodingStage(PipelineStage): if self.vae is None: raise ValueError("VAE must be provided for LTX-2 TI2V.") - from sglang.multimodal_gen.runtime.models.vision_utils import load_image + from sglang.multimodal_gen.runtime.utils.vision import load_image # 1. Load images, apply codec compression, resize for condition_image conditioned_imgs = [] 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 bdd1cbc2c..56dd05431 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 @@ -13,7 +13,6 @@ from PIL import Image from sglang.multimodal_gen.configs.pipeline_configs import WanI2V480PConfig from sglang.multimodal_gen.configs.pipeline_configs.base import ModelTaskType from sglang.multimodal_gen.configs.pipeline_configs.mova import MOVAPipelineConfig -from sglang.multimodal_gen.runtime.models.vision_utils import load_image, load_video from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req from sglang.multimodal_gen.runtime.pipelines_core.stages.base import PipelineStage from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import ( @@ -23,6 +22,7 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import ( 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__) diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/cosmos3.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/cosmos3.py index 1492e41f3..169a79459 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/cosmos3.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/cosmos3.py @@ -30,7 +30,6 @@ from sglang.multimodal_gen.runtime.distributed.parallel_state import ( get_sp_world_size, ) from sglang.multimodal_gen.runtime.managers.forward_context import set_forward_context -from sglang.multimodal_gen.runtime.models.vision_utils import load_video from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req from sglang.multimodal_gen.runtime.pipelines_core.stages.base import ( PipelineStage, @@ -57,6 +56,7 @@ 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.profiler import SGLDiffusionProfiler +from sglang.multimodal_gen.runtime.utils.vision import load_video from sglang.srt.utils.common import get_compiler_backend logger = init_logger(__name__) diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/glm_image.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/glm_image.py index 5c3647b85..1a6ce9e00 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/glm_image.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/glm_image.py @@ -16,7 +16,6 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager im ComponentUse, ) from sglang.multimodal_gen.runtime.models.dits.glm_image import GlmImageKVCache -from sglang.multimodal_gen.runtime.models.vision_utils import load_image from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req from sglang.multimodal_gen.runtime.pipelines_core.stages.base import ( PipelineStage, @@ -28,6 +27,7 @@ from sglang.multimodal_gen.runtime.utils.precision import ( align_tensor_to_module_dtype, get_module_dtype, ) +from sglang.multimodal_gen.runtime.utils.vision import load_image logger = init_logger(__name__) diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/qwen_image_layered.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/qwen_image_layered.py index 9d850259a..c76e678cb 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/qwen_image_layered.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/qwen_image_layered.py @@ -12,12 +12,12 @@ 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.vision_utils import load_image from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req from sglang.multimodal_gen.runtime.pipelines_core.stages.base import PipelineStage 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 align_tensor_to_module_dtype +from sglang.multimodal_gen.runtime.utils.vision import load_image logger = init_logger(__name__) diff --git a/python/sglang/multimodal_gen/runtime/platforms/aiter.py b/python/sglang/multimodal_gen/runtime/platforms/aiter.py new file mode 100644 index 000000000..d15240964 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/platforms/aiter.py @@ -0,0 +1,10 @@ +# SPDX-License-Identifier: Apache-2.0 + +from sglang.srt.utils import ( + get_bool_env_var, + is_gfx95_supported, + is_hip, +) + +USE_AITER = get_bool_env_var("SGLANG_USE_AITER") and is_hip() +USE_AITER_GFX95 = USE_AITER and is_gfx95_supported() diff --git a/python/sglang/multimodal_gen/runtime/server_args/__init__.py b/python/sglang/multimodal_gen/runtime/server_args/__init__.py new file mode 100644 index 000000000..cdcf73a7a --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/server_args/__init__.py @@ -0,0 +1,40 @@ +# SPDX-License-Identifier: Apache-2.0 + +from sglang.multimodal_gen.runtime.server_args import server_args as _server_args +from sglang.multimodal_gen.runtime.server_args.server_args import ( + BREAKABLE_CUDA_GRAPH_SUPPORTED_MODEL_IDS, + BREAKABLE_CUDA_GRAPH_SUPPORTED_PIPELINE_CONFIGS, + DEFAULT_BCG_TEXT_BUCKETS, + LORA_MERGE_MODES, + LTX2_TWO_STAGE_DEVICE_MODE_CHOICES, + Backend, + PortArgs, + ServerArgs, + _normalize_ltx2_two_stage_device_mode, + get_global_server_args, + is_ltx2_two_stage_pipeline_name, + prepare_server_args, + set_global_server_args, +) + +__all__ = [ + "Backend", + "BREAKABLE_CUDA_GRAPH_SUPPORTED_MODEL_IDS", + "BREAKABLE_CUDA_GRAPH_SUPPORTED_PIPELINE_CONFIGS", + "DEFAULT_BCG_TEXT_BUCKETS", + "LORA_MERGE_MODES", + "LTX2_TWO_STAGE_DEVICE_MODE_CHOICES", + "PortArgs", + "ServerArgs", + "_normalize_ltx2_two_stage_device_mode", + "get_global_server_args", + "is_ltx2_two_stage_pipeline_name", + "prepare_server_args", + "set_global_server_args", +] + + +def __getattr__(name: str): + if name == "_global_server_args": + return _server_args._global_server_args + raise AttributeError(name) diff --git a/python/sglang/multimodal_gen/runtime/server_args_auto_tune.py b/python/sglang/multimodal_gen/runtime/server_args/auto_tune.py similarity index 99% rename from python/sglang/multimodal_gen/runtime/server_args_auto_tune.py rename to python/sglang/multimodal_gen/runtime/server_args/auto_tune.py index 8059c9532..df005777a 100644 --- a/python/sglang/multimodal_gen/runtime/server_args_auto_tune.py +++ b/python/sglang/multimodal_gen/runtime/server_args/auto_tune.py @@ -20,7 +20,7 @@ from sglang.multimodal_gen.runtime.platforms import current_platform from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger if TYPE_CHECKING: - from sglang.multimodal_gen.runtime.server_args import ServerArgs + from sglang.multimodal_gen.runtime.server_args.server_args import ServerArgs logger = init_logger(__name__) diff --git a/python/sglang/multimodal_gen/runtime/server_args_disagg.py b/python/sglang/multimodal_gen/runtime/server_args/disagg.py similarity index 100% rename from python/sglang/multimodal_gen/runtime/server_args_disagg.py rename to python/sglang/multimodal_gen/runtime/server_args/disagg.py diff --git a/python/sglang/multimodal_gen/runtime/server_args.py b/python/sglang/multimodal_gen/runtime/server_args/server_args.py similarity index 99% rename from python/sglang/multimodal_gen/runtime/server_args.py rename to python/sglang/multimodal_gen/runtime/server_args/server_args.py index f032e0c0b..eae08df88 100644 --- a/python/sglang/multimodal_gen/runtime/server_args.py +++ b/python/sglang/multimodal_gen/runtime/server_args/server_args.py @@ -42,11 +42,11 @@ from sglang.multimodal_gen.runtime.platforms import ( AttentionBackendEnum, current_platform, ) -from sglang.multimodal_gen.runtime.server_args_auto_tune import ( +from sglang.multimodal_gen.runtime.server_args.auto_tune import ( PERFORMANCE_MODES, ServerArgsAutoTuner, ) -from sglang.multimodal_gen.runtime.server_args_disagg import DisaggServerArgsMixin +from sglang.multimodal_gen.runtime.server_args.disagg import DisaggServerArgsMixin from sglang.multimodal_gen.runtime.utils.common import ( is_port_available, is_valid_ipv6_address, diff --git a/python/sglang/multimodal_gen/runtime/models/vision_utils.py b/python/sglang/multimodal_gen/runtime/utils/vision.py similarity index 100% rename from python/sglang/multimodal_gen/runtime/models/vision_utils.py rename to python/sglang/multimodal_gen/runtime/utils/vision.py diff --git a/python/sglang/multimodal_gen/runtime/utils/weight_attrs.py b/python/sglang/multimodal_gen/runtime/utils/weight_attrs.py new file mode 100644 index 000000000..685cb0005 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/utils/weight_attrs.py @@ -0,0 +1,33 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +# Adapted from: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/utils.py + +from typing import Any + +import torch + +from sglang.multimodal_gen.runtime.platforms import current_platform + + +def set_weight_attrs( + weight: torch.Tensor, + weight_attrs: dict[str, Any] | None, +): + """Set attributes on a weight tensor without overwriting existing ones.""" + if weight_attrs is None: + return + for key, value in weight_attrs.items(): + assert not hasattr(weight, key), f"Overwriting existing tensor attribute: {key}" + + if current_platform.is_tpu() and key == "weight_loader": + value = make_synced_weight_loader(value) + setattr(weight, key, value) + + +def make_synced_weight_loader(original_weight_loader) -> Any: + + def _synced_weight_loader(param, *args, **kwargs): + original_weight_loader(param, *args, **kwargs) + torch._sync(param) + + return _synced_weight_loader 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 5f43756dc..b4f326198 100644 --- a/python/sglang/multimodal_gen/test/unit/test_server_args.py +++ b/python/sglang/multimodal_gen/test/unit/test_server_args.py @@ -190,23 +190,23 @@ class TestServerArgsPathExpansion(unittest.TestCase): PipelineConfig, "from_kwargs", return_value=QwenImagePipelineConfig() ), patch( - "sglang.multimodal_gen.runtime.server_args.current_platform.is_cpu", + "sglang.multimodal_gen.runtime.platforms.current_platform.is_cpu", return_value=False, ), patch( - "sglang.multimodal_gen.runtime.server_args.current_platform.is_mps", + "sglang.multimodal_gen.runtime.platforms.current_platform.is_mps", return_value=False, ), patch( - "sglang.multimodal_gen.runtime.server_args.current_platform.is_cuda", + "sglang.multimodal_gen.runtime.platforms.current_platform.is_cuda", return_value=True, ), patch( - "sglang.multimodal_gen.runtime.server_args.current_platform.get_device_total_memory", + "sglang.multimodal_gen.runtime.platforms.current_platform.get_device_total_memory", return_value=80 * 1024**3, ), patch( - "sglang.multimodal_gen.runtime.server_args.current_platform.get_available_gpu_memory", + "sglang.multimodal_gen.runtime.platforms.current_platform.get_available_gpu_memory", return_value=80, ), ): @@ -366,11 +366,11 @@ class TestServerArgsPathExpansion(unittest.TestCase): return_value=None, ), patch( - "sglang.multimodal_gen.runtime.server_args.current_platform.get_device_total_memory", + "sglang.multimodal_gen.runtime.platforms.current_platform.get_device_total_memory", return_value=80 * 1024**3, ), patch( - "sglang.multimodal_gen.runtime.server_args.current_platform.get_available_gpu_memory", + "sglang.multimodal_gen.runtime.platforms.current_platform.get_available_gpu_memory", return_value=80, ), ): @@ -706,27 +706,27 @@ class TestOffloadDefaults(unittest.TestCase): with ( patch.object(PipelineConfig, "from_kwargs", return_value=pipeline_config), patch( - "sglang.multimodal_gen.runtime.server_args.current_platform.is_cpu", + "sglang.multimodal_gen.runtime.platforms.current_platform.is_cpu", return_value=False, ), patch( - "sglang.multimodal_gen.runtime.server_args.current_platform.is_mps", + "sglang.multimodal_gen.runtime.platforms.current_platform.is_mps", return_value=False, ), patch( - "sglang.multimodal_gen.runtime.server_args.current_platform.is_cuda", + "sglang.multimodal_gen.runtime.platforms.current_platform.is_cuda", return_value=True, ), patch( - "sglang.multimodal_gen.runtime.server_args.current_platform.enable_dit_layerwise_offload_for_wan_by_default", + "sglang.multimodal_gen.runtime.platforms.current_platform.enable_dit_layerwise_offload_for_wan_by_default", return_value=True, ), patch( - "sglang.multimodal_gen.runtime.server_args.current_platform.get_device_total_memory", + "sglang.multimodal_gen.runtime.platforms.current_platform.get_device_total_memory", return_value=memory_gb * 1024**3, ), patch( - "sglang.multimodal_gen.runtime.server_args.current_platform.get_available_gpu_memory", + "sglang.multimodal_gen.runtime.platforms.current_platform.get_available_gpu_memory", side_effect=get_available_gpu_memory, ), ): @@ -744,19 +744,19 @@ class TestOffloadDefaults(unittest.TestCase): with ( patch.object(PipelineConfig, "from_kwargs", return_value=pipeline_config), patch( - "sglang.multimodal_gen.runtime.server_args.current_platform.is_cpu", + "sglang.multimodal_gen.runtime.platforms.current_platform.is_cpu", return_value=False, ), patch( - "sglang.multimodal_gen.runtime.server_args.current_platform.is_cuda", + "sglang.multimodal_gen.runtime.platforms.current_platform.is_cuda", return_value=True, ), patch( - "sglang.multimodal_gen.runtime.server_args.current_platform.get_device_total_memory", + "sglang.multimodal_gen.runtime.platforms.current_platform.get_device_total_memory", return_value=memory_gb * 1024**3, ), patch( - "sglang.multimodal_gen.runtime.server_args.current_platform.get_available_gpu_memory", + "sglang.multimodal_gen.runtime.platforms.current_platform.get_available_gpu_memory", return_value=memory_gb, ), ): @@ -1613,23 +1613,23 @@ class TestOffloadDefaults(unittest.TestCase): PipelineConfig, "from_kwargs", return_value=QwenImagePipelineConfig() ), patch( - "sglang.multimodal_gen.runtime.server_args.current_platform.is_cpu", + "sglang.multimodal_gen.runtime.platforms.current_platform.is_cpu", return_value=False, ), patch( - "sglang.multimodal_gen.runtime.server_args.current_platform.is_mps", + "sglang.multimodal_gen.runtime.platforms.current_platform.is_mps", return_value=False, ), patch( - "sglang.multimodal_gen.runtime.server_args.current_platform.is_cuda", + "sglang.multimodal_gen.runtime.platforms.current_platform.is_cuda", return_value=True, ), patch( - "sglang.multimodal_gen.runtime.server_args.current_platform.get_device_total_memory", + "sglang.multimodal_gen.runtime.platforms.current_platform.get_device_total_memory", return_value=80 * 1024**3, ), patch( - "sglang.multimodal_gen.runtime.server_args.current_platform.get_available_gpu_memory", + "sglang.multimodal_gen.runtime.platforms.current_platform.get_available_gpu_memory", return_value=80, ), ): @@ -1657,23 +1657,23 @@ class TestOffloadDefaults(unittest.TestCase): PipelineConfig, "from_kwargs", return_value=LTX2PipelineConfig() ), patch( - "sglang.multimodal_gen.runtime.server_args.current_platform.is_cpu", + "sglang.multimodal_gen.runtime.platforms.current_platform.is_cpu", return_value=False, ), patch( - "sglang.multimodal_gen.runtime.server_args.current_platform.is_mps", + "sglang.multimodal_gen.runtime.platforms.current_platform.is_mps", return_value=False, ), patch( - "sglang.multimodal_gen.runtime.server_args.current_platform.is_cuda", + "sglang.multimodal_gen.runtime.platforms.current_platform.is_cuda", return_value=True, ), patch( - "sglang.multimodal_gen.runtime.server_args.current_platform.get_device_total_memory", + "sglang.multimodal_gen.runtime.platforms.current_platform.get_device_total_memory", return_value=140 * 1024**3, ), patch( - "sglang.multimodal_gen.runtime.server_args.current_platform.get_available_gpu_memory", + "sglang.multimodal_gen.runtime.platforms.current_platform.get_available_gpu_memory", return_value=134, ), ): @@ -1855,9 +1855,9 @@ class TestPerRoleParallelism(unittest.TestCase): self.assertEqual(args.get_role_parallelism(RoleType.DENOISER)["tp_size"], 2) self.assertEqual(args.get_role_parallelism(RoleType.DECODER)["sp_degree"], 4) - def test_disagg_args_import_path_stays_compatible(self): + def test_disagg_args_import_path_matches_server_args_package(self): from sglang.multimodal_gen.runtime.disaggregation import disagg_args - from sglang.multimodal_gen.runtime.server_args_disagg import ( + from sglang.multimodal_gen.runtime.server_args.disagg import ( DisaggServerArgsMixin, )