[diffusion] refactor: reorganize runtime utility and server_args modules (#30447)

This commit is contained in:
Mick
2026-07-10 15:53:00 +08:00
committed by GitHub
parent 1d8e3c248b
commit 4a8e1b07a2
36 changed files with 217 additions and 235 deletions
@@ -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,
@@ -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
@@ -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
@@ -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
@@ -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):
@@ -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()
@@ -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:
@@ -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
@@ -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,
@@ -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__)
@@ -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"
@@ -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
@@ -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)
@@ -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,
@@ -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
@@ -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()
@@ -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()
@@ -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
@@ -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)
@@ -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__)
@@ -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)
@@ -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
@@ -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
@@ -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 = []
@@ -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__)
@@ -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__)
@@ -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__)
@@ -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__)
@@ -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()
@@ -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)
@@ -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__)
@@ -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,
@@ -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
@@ -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,
)