diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/base.py b/python/sglang/multimodal_gen/configs/pipeline_configs/base.py index e86ad7bde..72d70d721 100644 --- a/python/sglang/multimodal_gen/configs/pipeline_configs/base.py +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/base.py @@ -201,6 +201,9 @@ class PipelineConfig: native_only_components: ClassVar[tuple[str, ...]] = () task_type: ModelTaskType = ModelTaskType.I2I skip_input_image_preprocess: bool = False + # False when changing component placement after a calibration request is + # known to alter the pipeline's numerical path. + supports_auto_residency: bool = True # Components that cannot fall back to a native Transformers/Diffusers # implementation because their pipeline requires SGLang-specific behavior. native_only_components: tuple[str, ...] = () diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/ltx_2.py b/python/sglang/multimodal_gen/configs/pipeline_configs/ltx_2.py index ccc4d2715..240c23023 100644 --- a/python/sglang/multimodal_gen/configs/pipeline_configs/ltx_2.py +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/ltx_2.py @@ -723,3 +723,6 @@ class LTX2PipelineConfig(PipelineConfig): @dataclasses.dataclass class LTX23PipelineConfig(LTX2PipelineConfig): """Configuration overrides for LTX-2.3.""" + + # original-mode lora swaps invalidate post-warmup timing calibration + supports_auto_residency: bool = False 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 91169c686..d076d17b9 100644 --- a/python/sglang/multimodal_gen/configs/pipeline_configs/qwen_image.py +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/qwen_image.py @@ -755,6 +755,8 @@ class QwenImageEditPlus_2511_PipelineConfig(QwenImageEditPlusPipelineConfig): class QwenImageLayeredPipelineConfig(QwenImageEditPipelineConfig): resolution: int = 640 vae_precision: str = "bf16" + # promoting the auxiliary components regresses first-request latency + supports_auto_residency: bool = False def get_model_deployment_config(self) -> ModelDeploymentConfig: return ModelDeploymentConfig( diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/sana_wm.py b/python/sglang/multimodal_gen/configs/pipeline_configs/sana_wm.py index b0e013f99..afd29e82d 100644 --- a/python/sglang/multimodal_gen/configs/pipeline_configs/sana_wm.py +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/sana_wm.py @@ -56,6 +56,9 @@ class SanaWMPipelineConfig(PipelineConfig): optional 6-DoF camera trajectory, optional Stage-2 LTX-2 refiner).""" task_type: ModelTaskType = ModelTaskType.TI2V + # The current two-stage path is not numerically invariant when its primary + # transformer and connectors are promoted after warmup. + supports_auto_residency: bool = False # SanaWMBeforeDenoisingStage._splice_first_frame handles condition-image # resize + VAE-encode itself, so bypass the framework's generic TI2V diff --git a/python/sglang/multimodal_gen/configs/sample/sampling_params.py b/python/sglang/multimodal_gen/configs/sample/sampling_params.py index bdd68652c..9fd211300 100644 --- a/python/sglang/multimodal_gen/configs/sample/sampling_params.py +++ b/python/sglang/multimodal_gen/configs/sample/sampling_params.py @@ -136,6 +136,58 @@ def _sanitize_filename(name: str, replacement: str = "_", max_length: int = 150) return ascii_name +_SEQUENCE_SHARD_PIPELINE_FAMILIES = ("wan", "helios", "joy", "cosmos3") + + +def resolve_sequence_shard( + pipeline_config: Any, enable_sequence_shard: bool | None +) -> bool: + """Whether this pipeline shards the sequence dim instead of aligning frames. + + Shared by ``SamplingParams._adjust_visual_fields`` and the synthetic + warmup builder so warmup requests follow the same frame contract as real + requests. + """ + pipeline_name_lower = pipeline_config.__class__.__name__.lower() + return any( + family in pipeline_name_lower for family in _SEQUENCE_SHARD_PIPELINE_FAMILIES + ) and (enable_sequence_shard is None or enable_sequence_shard) + + +def align_num_frames_for_num_gpus( + num_frames: int, + *, + num_gpus: int, + vae_config: Any, + round_down: bool, +) -> int: + """Align the latent frame count to be divisible by ``num_gpus``.""" + if num_gpus <= 1: + return num_frames + use_temporal_scaling_frames = vae_config.use_temporal_scaling_frames + temporal_scale_factor = vae_config.arch_config.temporal_compression_ratio + + if use_temporal_scaling_frames: + orig_latent_num_frames = (num_frames - 1) // temporal_scale_factor + 1 + else: + orig_latent_num_frames = num_frames + + if orig_latent_num_frames % num_gpus == 0: + return num_frames + + if round_down: + # Ensure we have at least 1 batch per GPU + new_latent_num_frames = max(1, (orig_latent_num_frames // num_gpus)) * num_gpus + else: + new_latent_num_frames = math.ceil(orig_latent_num_frames / num_gpus) * num_gpus + + if use_temporal_scaling_frames: + # Convert back to frames, keeping num_frames-1 a multiple of the + # temporal scale factor + return (new_latent_num_frames - 1) * temporal_scale_factor + 1 + return new_latent_num_frames + + class DataType(Enum): IMAGE = auto() VIDEO = auto() @@ -792,14 +844,9 @@ class SamplingParams: ) logger.warning(error_msg) - pipeline_name_lower = server_args.pipeline_config.__class__.__name__.lower() - - if ( - "wan" in pipeline_name_lower - or "helios" in pipeline_name_lower - or "joy" in pipeline_name_lower - or "cosmos3" in pipeline_name_lower - ) and (self.enable_sequence_shard is None or self.enable_sequence_shard): + if resolve_sequence_shard( + server_args.pipeline_config, self.enable_sequence_shard + ): self.enable_sequence_shard = True logger.debug("Automatically enabled enable_sequence_shard") else: @@ -832,43 +879,13 @@ class SamplingParams: ) if self.adjust_frames: - # Adjust number of frames based on number of GPUs for video task - use_temporal_scaling_frames = ( - pipeline_config.vae_config.use_temporal_scaling_frames + new_num_frames = align_num_frames_for_num_gpus( + self.num_frames, + num_gpus=server_args.num_gpus, + vae_config=pipeline_config.vae_config, + round_down=self.num_frames_round_down, ) - num_frames = self.num_frames - num_gpus = server_args.num_gpus - temporal_scale_factor = ( - pipeline_config.vae_config.arch_config.temporal_compression_ratio - ) - - if use_temporal_scaling_frames: - orig_latent_num_frames = ( - num_frames - 1 - ) // temporal_scale_factor + 1 - else: - orig_latent_num_frames = num_frames - - if orig_latent_num_frames % server_args.num_gpus != 0: - # Adjust latent frames to be divisible by number of GPUs - if self.num_frames_round_down: - # Ensure we have at least 1 batch per GPU - new_latent_num_frames = ( - max(1, (orig_latent_num_frames // num_gpus)) * num_gpus - ) - else: - new_latent_num_frames = ( - math.ceil(orig_latent_num_frames / num_gpus) * num_gpus - ) - - if use_temporal_scaling_frames: - # Convert back to number of frames, ensuring num_frames-1 is a multiple of temporal_scale_factor - new_num_frames = ( - new_latent_num_frames - 1 - ) * temporal_scale_factor + 1 - else: - new_num_frames = new_latent_num_frames - + if new_num_frames != self.num_frames: logger.info( "Adjusting number of frames from %s to %s based on number of GPUs (%s)", self.num_frames, diff --git a/python/sglang/multimodal_gen/envs.py b/python/sglang/multimodal_gen/envs.py index 2df4962d9..2941b8d92 100644 --- a/python/sglang/multimodal_gen/envs.py +++ b/python/sglang/multimodal_gen/envs.py @@ -37,6 +37,7 @@ if TYPE_CHECKING: SGLANG_DIFFUSION_TEST_FORCE_HOST_AVAILABLE_GIB: float | None = None SGLANG_DIFFUSION_TEST_CAP_DEVICE_MEMORY_GIB: float | None = None SGLANG_DIFFUSION_STAGE_LOGGING: bool = False + SGLANG_DIFFUSION_DISABLE_AUTO_RESIDENCY: bool = False SGLANG_DIFFUSION_MINIMAX_H3_ADALN_GPU_PLANS: int = 64 SGLANG_DIFFUSION_MINIMAX_H3_ADALN_FP32: bool = False SGLANG_DIFFUSION_CFG_GATE_STEP: float = 1.0 @@ -265,6 +266,12 @@ environment_variables: dict[str, Callable[[], Any]] = { # If set, sgl_diffusion will enable stage logging, which will print the time # taken for each stage "SGLANG_DIFFUSION_STAGE_LOGGING": _lazy_bool("SGLANG_DIFFUSION_STAGE_LOGGING"), + # Kill-switch for the warmup-calibrated auto residency promotion that runs + # under `--performance-mode auto` with server warmup. Set to disable the + # promotion without giving up the rest of the auto performance policy. + "SGLANG_DIFFUSION_DISABLE_AUTO_RESIDENCY": _lazy_bool( + "SGLANG_DIFFUSION_DISABLE_AUTO_RESIDENCY" + ), # Plan slots in the MiniMax-H3 --minimax-h3-adaln-online GPU slab # (9.25 MiB per slot-timestep; 64 x width 4 = 2.31 GiB). A request needs # up to num_inference_steps - 1 slots; the default covers the 50-step diff --git a/python/sglang/multimodal_gen/runtime/distributed/device_communicators/ipc_a2a.py b/python/sglang/multimodal_gen/runtime/distributed/device_communicators/ipc_a2a.py index ddda93f3f..90524471e 100644 --- a/python/sglang/multimodal_gen/runtime/distributed/device_communicators/ipc_a2a.py +++ b/python/sglang/multimodal_gen/runtime/distributed/device_communicators/ipc_a2a.py @@ -140,6 +140,18 @@ class IpcA2AState: """Drop mappings that belong to a model-parallel group being replaced.""" self.__init__() + def drop_staging(self) -> None: + """Release the cached staging buffers (both ranks call this at the same point). + + Staging is keyed by message size and only evicted by count, so a warmup + probe at the full serving shape leaves buffers sized for it behind. + """ + if not self.staging: + return + if torch.cuda.is_available(): + torch.cuda.synchronize() + self.staging.clear() + def _share(self, t, group): """Exchange `t` with the peer via torch IPC, re-opening the handle in the LOCAL device context (the mapping is only dereferenceable from the diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/control_requests.py b/python/sglang/multimodal_gen/runtime/entrypoints/control_requests.py new file mode 100644 index 000000000..9f142e7cd --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/entrypoints/control_requests.py @@ -0,0 +1,50 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Control-request protocol between the HTTP process and scheduler workers. + +These types are cross-process IPC contracts, not utilities: the HTTP side +constructs them (scheduler_client treats them as ``_CONTROL_REQ_TYPES`` and +fans them out to every replica) and each scheduler dispatches them through +``Scheduler.request_handlers``. Keep this module import-light -- both +processes import it, and the HTTP process must not drag in torch-heavy +worker modules through it. +""" + +from typing import List, Optional, Union + +import msgspec + + +class SetLoraReq(msgspec.Struct): + lora_nickname: Union[str, List[str]] + lora_path: Optional[Union[str, List[Optional[str]]]] = None + target: Union[str, List[str]] = "all" + strength: Union[float, List[float]] = 1.0 + merge_mode: Optional[str] = None + lora_alpha: Optional[Union[int, List[Optional[int]]]] = None + + +class MergeLoraWeightsReq(msgspec.Struct): + target: str = "all" + strength: float = 1.0 + + +class UnmergeLoraWeightsReq(msgspec.Struct): + target: str = "all" + + +class ListLorasReq(msgspec.Struct): + pass + + +class ShutdownReq(msgspec.Struct): + pass + + +class ReleaseRealtimeSessionReq(msgspec.Struct): + session_id: str + + +class GetDisaggStatsReq(msgspec.Struct): + """Request to get disagg pipeline metrics from the scheduler.""" + + pass diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py b/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py index 0523f5650..712ba7d1b 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py @@ -19,13 +19,15 @@ from sglang.multimodal_gen.configs.sample.sampling_params import ( DataType, SamplingParams, ) -from sglang.multimodal_gen.runtime.entrypoints.utils import ( - GenerationResult, +from sglang.multimodal_gen.runtime.entrypoints.control_requests import ( ListLorasReq, MergeLoraWeightsReq, SetLoraReq, ShutdownReq, UnmergeLoraWeightsReq, +) +from sglang.multimodal_gen.runtime.entrypoints.utils import ( + GenerationResult, expand_request_outputs, format_lora_message, prepare_request, diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/http_server.py b/python/sglang/multimodal_gen/runtime/entrypoints/http_server.py index 1b5098acb..8dc5385a4 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/http_server.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/http_server.py @@ -264,7 +264,9 @@ async def stats_endpoint(request: Request): Returns queue depth, request counts, latency, throughput, etc. Sends a GetDisaggStatsReq to the scheduler via ZMQ and returns the result. """ - from sglang.multimodal_gen.runtime.entrypoints.utils import GetDisaggStatsReq + from sglang.multimodal_gen.runtime.entrypoints.control_requests import ( + GetDisaggStatsReq, + ) server_args: ServerArgs = request.app.state.server_args response: dict = { diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/openai/common_api.py b/python/sglang/multimodal_gen/runtime/entrypoints/openai/common_api.py index 9a3e6d831..d3c9234a1 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/openai/common_api.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/openai/common_api.py @@ -5,11 +5,13 @@ from fastapi import APIRouter, Body, HTTPException from pydantic import BaseModel, Field from sglang.multimodal_gen.registry import get_model_info -from sglang.multimodal_gen.runtime.entrypoints.utils import ( +from sglang.multimodal_gen.runtime.entrypoints.control_requests import ( ListLorasReq, MergeLoraWeightsReq, SetLoraReq, UnmergeLoraWeightsReq, +) +from sglang.multimodal_gen.runtime.entrypoints.utils import ( format_lora_message, ) from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/openai/realtime/realtime_video_api.py b/python/sglang/multimodal_gen/runtime/entrypoints/openai/realtime/realtime_video_api.py index 305a40818..e8573a8ff 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/openai/realtime/realtime_video_api.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/openai/realtime/realtime_video_api.py @@ -8,6 +8,9 @@ from typing import TYPE_CHECKING import msgspec.msgpack from fastapi import APIRouter, WebSocket, WebSocketDisconnect +from sglang.multimodal_gen.runtime.entrypoints.control_requests import ( + ReleaseRealtimeSessionReq, +) from sglang.multimodal_gen.runtime.entrypoints.openai.protocol import ( RealtimeEvent, RealtimeVideoGenerationsRequest, @@ -28,9 +31,6 @@ from sglang.multimodal_gen.runtime.entrypoints.openai.realtime.timer import ( from sglang.multimodal_gen.runtime.entrypoints.openai.utils import ( process_generation_batch, ) -from sglang.multimodal_gen.runtime.entrypoints.utils import ( - ReleaseRealtimeSessionReq, -) from sglang.multimodal_gen.runtime.scheduler_client import async_scheduler_client from sglang.multimodal_gen.runtime.server_args import get_global_server_args from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/openai/realtime/registry.py b/python/sglang/multimodal_gen/runtime/entrypoints/openai/realtime/registry.py index 4c682d0d7..005131af4 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/openai/realtime/registry.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/openai/realtime/registry.py @@ -4,11 +4,10 @@ from __future__ import annotations from typing import TYPE_CHECKING -from sglang.multimodal_gen.runtime.entrypoints.openai.realtime.realtime_adapter import ( - BaseRealtimeModelAdapter, -) - if TYPE_CHECKING: + from sglang.multimodal_gen.runtime.entrypoints.openai.realtime.realtime_adapter import ( + BaseRealtimeModelAdapter, + ) from sglang.multimodal_gen.runtime.server_args import ServerArgs @@ -67,3 +66,12 @@ def get_realtime_model_adapter( "Realtime video is not supported for pipeline config " f"{type(pipeline_config).__name__}; no realtime adapter is registered." ) + + +def has_realtime_model_adapter(server_args: ServerArgs) -> bool: + """Whether the resolved pipeline config has a registered adapter.""" + _register_builtin_realtime_model_adapters() + return any( + config_cls in _REALTIME_ADAPTER_REGISTRY + for config_cls in type(server_args.pipeline_config).__mro__ + ) diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/openai/utils.py b/python/sglang/multimodal_gen/runtime/entrypoints/openai/utils.py index bb2c3abe7..f5c966d80 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/openai/utils.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/openai/utils.py @@ -18,12 +18,14 @@ from sglang.multimodal_gen.configs.sample.sampling_params import ( DataType, SamplingParams, ) -from sglang.multimodal_gen.runtime.entrypoints.utils import ( +from sglang.multimodal_gen.runtime.entrypoints.control_requests import ( ListLorasReq, MergeLoraWeightsReq, SetLoraReq, ShutdownReq, UnmergeLoraWeightsReq, +) +from sglang.multimodal_gen.runtime.entrypoints.utils import ( format_lora_message, save_outputs, ) diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/utils.py b/python/sglang/multimodal_gen/runtime/entrypoints/utils.py index 842b8aa1b..a23fa92fb 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/utils.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/utils.py @@ -160,49 +160,6 @@ def _close_cached_cuda_video_buffer() -> None: atexit.register(_close_cached_cuda_video_buffer) -@dataclass -class SetLoraReq: - lora_nickname: Union[str, List[str]] - lora_path: Optional[Union[str, List[Optional[str]]]] = None - target: Union[str, List[str]] = "all" - strength: Union[float, List[float]] = 1.0 - merge_mode: Optional[str] = None - lora_alpha: Optional[Union[int, List[Optional[int]]]] = None - - -@dataclass -class MergeLoraWeightsReq: - target: str = "all" - strength: float = 1.0 - - -@dataclass -class UnmergeLoraWeightsReq: - target: str = "all" - - -@dataclass -class ListLorasReq: - pass - - -@dataclass -class ShutdownReq: - pass - - -@dataclass -class ReleaseRealtimeSessionReq: - session_id: str - - -@dataclass -class GetDisaggStatsReq: - """Request to get disagg pipeline metrics from the scheduler.""" - - pass - - def format_lora_message( lora_nickname: Union[str, List[str]], target: Union[str, List[str]], diff --git a/python/sglang/multimodal_gen/runtime/launch_server.py b/python/sglang/multimodal_gen/runtime/launch_server.py index f4b655ae5..187313de4 100644 --- a/python/sglang/multimodal_gen/runtime/launch_server.py +++ b/python/sglang/multimodal_gen/runtime/launch_server.py @@ -12,8 +12,8 @@ from sglang.multimodal_gen.runtime.disaggregation.orchestrator import ( DiffusionServer, ) from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType +from sglang.multimodal_gen.runtime.entrypoints.control_requests import ShutdownReq from sglang.multimodal_gen.runtime.entrypoints.http_server import create_app -from sglang.multimodal_gen.runtime.entrypoints.utils import ShutdownReq from sglang.multimodal_gen.runtime.managers.gpu_worker import run_scheduler_process from sglang.multimodal_gen.runtime.scheduler_client import SchedulerClient from sglang.multimodal_gen.runtime.server_args import ( diff --git a/python/sglang/multimodal_gen/runtime/layers/usp.py b/python/sglang/multimodal_gen/runtime/layers/usp.py index 8b2d66db1..80342c3cb 100644 --- a/python/sglang/multimodal_gen/runtime/layers/usp.py +++ b/python/sglang/multimodal_gen/runtime/layers/usp.py @@ -41,6 +41,20 @@ def _maybe_wait(tensor: torch.Tensor) -> torch.Tensor: _A2A_STAGING_BUFFERS: dict[tuple[str, torch.dtype, int], torch.Tensor] = {} +def drop_a2a_staging_buffers() -> None: + """Release the cached all-to-all staging buffers on this rank. + + The cache only ever grows to the largest message seen, so a warmup probe + at the full serving shape leaves buffers sized for it behind; the caller + releases them at a point every rank reaches together. + """ + if not _A2A_STAGING_BUFFERS: + return + if torch.cuda.is_available(): + torch.cuda.synchronize() + _A2A_STAGING_BUFFERS.clear() + + def _a2a_staging_buffer( role: str, shape: tuple[int, ...], dtype: torch.dtype, device: torch.device ) -> torch.Tensor: diff --git a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py index 6f098c452..de32a17a7 100644 --- a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py +++ b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py @@ -15,6 +15,8 @@ import numpy as np import torch from setproctitle import setproctitle +from sglang.multimodal_gen.runtime.warmup_request_builder import lighten_warmup_req + from sglang.multimodal_gen.runtime.utils.logging_utils import ( # isort: skip globally_suppress_loggers, ) @@ -49,7 +51,17 @@ from sglang.multimodal_gen.runtime.entrypoints.utils import ( post_process_sample, save_outputs, ) +from sglang.multimodal_gen.runtime.managers.memory_managers.auto_residency import ( + DefaultWorkload, + WarmupMemoryRecord, + estimate_default_workload_peak_bytes, + resolve_default_workload, +) +from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager import ( + peek_global_component_residency_manager, +) from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import ( + LayerwiseUsageTracker, configure_layerwise_offload_modules, ) from sglang.multimodal_gen.runtime.managers.memory_managers.memory_occupation_controller import ( @@ -93,14 +105,6 @@ from sglang.srt.utils.network import NetworkAddress logger = init_logger(__name__) -OFFLOAD_DISABLE_RECOMMENDATION_ORDER = ( - "vae", - "image_encoder", - "text_encoder", - "text_encoder_2", - "transformer", -) - @dataclass class _ExpandedOutputParts: @@ -129,6 +133,64 @@ def _worker_cpu_intra_op_threads(num_gpus: int) -> int | None: return max(1, min(16, cpu_count // max(1, num_gpus))) +OFFLOAD_DISABLE_RECOMMENDATION_ORDER = ( + "vae", + "image_encoder", + "text_encoder", + "text_encoder_2", + "transformer", +) + + +PROBE_FIT_MIN_MARGIN_BYTES = 1 << 30 + + +def _shape_label(req: Req) -> str: + return f"{req.width}x{req.height}x{req.num_frames or 1}f" + + +def fit_auto_residency_probe( + req: Req, + *, + records: list[WarmupMemoryRecord], + free_bytes: int, + total_bytes: int, + server_args: ServerArgs, +) -> tuple[Req, int | None, int]: + """Shrink a full-shape probe until its extrapolated peak fits the memory left. + + The probe measures the default workload under the load-safe placement, so + a probe the card cannot hold would only be found out by running out of + memory. The bounded warmup that runs before it gives one measurement to + extrapolate from; while that extrapolation exceeds free memory minus the + reserve, frames go first and then area, the ladder the OOM retry walks. + Returns the fitted request, its estimate and the number of shrink steps. + """ + # Only the probe has to fit, so the margin is allocator slack, not the + # planner's placement reserve (which held back 4 GiB of a 32 GiB card and + # shrank a probe that had 10 GiB to spare). + budget = free_bytes - max(PROBE_FIT_MIN_MARGIN_BYTES, total_bytes // 50) + # The bounded warmup already ran at the smallest measured shape; a probe + # below it measures nothing new and degenerate shapes fail inside models. + floor_units = min((record.workload_units() for record in records), default=0) + fitted, steps = req, 0 + while True: + units = ( + max(1, int(fitted.width or 1)) + * max(1, int(fitted.height or 1)) + * max(1, int(fitted.num_frames or 1)) + ) + estimate = estimate_default_workload_peak_bytes( + records=records, target_units=units + ) + if estimate is None or estimate <= budget or units <= floor_units: + return fitted, estimate, steps + lighter = lighten_warmup_req(server_args, fitted) + if lighter is None: + return fitted, estimate, steps + fitted, steps = lighter, steps + 1 + + class GPUWorker(GPUWorkerPostTrainingMixin): """ A worker that executes the model on a single GPU. @@ -163,6 +225,10 @@ class GPUWorker(GPUWorkerPostTrainingMixin): load_snapshot.peak_allocated_mb if load_snapshot is not None else 0.0 ) self._runtime_peak_reserved_mb = 0.0 + # Warmup probes run the default workload's full shape and may exceed any + # serving request; keep their peak out of the runtime figure. + self._warmup_peak_reserved_mb = 0.0 + self._release_warmup_pool_before_serving = False self._runtime_peak_allocated_mb = 0.0 self.sp_group = get_sp_group() self.sp_cpu_group = self.sp_group.cpu_group @@ -173,6 +239,26 @@ class GPUWorker(GPUWorkerPostTrainingMixin): self.cfg_cpu_group = self.cfg_group.cpu_group self._realtime_sessions = RealtimeSessionCache(max_sessions=1) self.memory_occupation: MemoryOccupationController | None = None + # per-rank memory measurements of server warmup forwards; consumed by + # the auto-residency placement decision before the server turns ready + self._auto_residency_warmup_records: list[WarmupMemoryRecord] = [] + # default workload resolved once for the per-request residency hint + self._cached_default_workload: DefaultWorkload | None = None + self._cached_default_workload_failed = False + + def _default_workload_for_hint(self) -> DefaultWorkload | None: + if ( + self._cached_default_workload is None + and not self._cached_default_workload_failed + ): + try: + self._cached_default_workload = resolve_default_workload( + self.server_args + ) + except Exception: + logger.debug("Default workload unresolvable", exc_info=True) + self._cached_default_workload_failed = True + return self._cached_default_workload def release_realtime_session(self, session_id: str) -> OutputBatch: """release the session of a realtime connection""" @@ -364,24 +450,37 @@ class GPUWorker(GPUWorkerPostTrainingMixin): remaining_gpu_mem_gb = ( current_platform.get_device_total_memory() / (1024**3) - peak_reserved_gb ) - can_stay_resident = self.get_can_stay_resident_components(remaining_gpu_mem_gb) + try: + can_stay_resident = self.get_can_stay_resident_components( + remaining_gpu_mem_gb + ) + except Exception: + # a debug-only hint must never fail a completed request + logger.debug("Residency hint unavailable", exc_info=True) + can_stay_resident = [] pool_overhead_gb = peak_reserved_gb - peak_allocated_gb pool_overhead_pct = ( pool_overhead_gb / peak_reserved_gb * 100 if peak_reserved_gb else 0.0 ) + residency_hint = ( + f" Components that can remain on GPU: {can_stay_resident}. " + "Make it explicit with --component-residency =resident; " + "--performance-mode auto with server warmup applies safe " + "adjustments automatically." + if can_stay_resident + else "" + ) logger.debug( "GPU memory: peak=%.2f GB, allocated=%.2f GB, pool=%.2f GB (%.1f%%), " - "headroom=%.2f GB. Components that can remain on GPU: %s. " - "Adjust --cpu-offload-components or --layerwise-offload-components " - "to change residency.", + "headroom=%.2f GB.%s", peak_reserved_gb, peak_allocated_gb, pool_overhead_gb, pool_overhead_pct, remaining_gpu_mem_gb, - can_stay_resident, + residency_hint, ) def execute_forward( @@ -409,6 +508,8 @@ class GPUWorker(GPUWorkerPostTrainingMixin): return self._execute_forward_batch(batch) req = batch[0] + if req.is_warmup and req.extra.get("auto_residency_full_shape_probe"): + self._fit_auto_residency_probe(req) return self._execute_forward_common( req, forward_fn=lambda: self.pipeline.forward(req, self.server_args), @@ -502,9 +603,57 @@ class GPUWorker(GPUWorkerPostTrainingMixin): """ output_batch = None forward_failed = False + # Prewarm reqs (is_warmup=False) run a different offload layout and + # must not contaminate the calibration records. Pipelines that cannot + # apply a residency plan also skip the temporary per-layer hooks. + measure_server_warmup = ( + req.is_warmup + and bool(req.extra.get("server_based_warmup")) + and self.server_args.pipeline_config.supports_auto_residency + and current_platform.is_cuda() + ) + warmup_workload = ( + ( + int(req.width or 0), + int(req.height or 0), + int(req.num_frames or 1), + max(1, int(req.num_inference_steps or 1)), + ) + if measure_server_warmup + else None + ) + warmup_baseline_allocated_bytes = 0 + layerwise_usage_tracker: LayerwiseUsageTracker | None = None + layerwise_layer_uses_by_stage: dict[ + str, dict[str, dict[str, tuple[int, ...]]] + ] = {} try: + if measure_server_warmup: + # Drop the previous request's allocator pool so each probe + # starts from the same placement and can return released + # component storage before its allocated peak is measured. + torch.get_device_module().empty_cache() + self._release_warmup_pool(req) if not current_platform.is_cpu() and not current_platform.is_mps(): torch.get_device_module().reset_peak_memory_stats() + if measure_server_warmup: + warmup_baseline_allocated_bytes = ( + torch.get_device_module().memory_allocated() + ) + if ( + self.server_args.performance_mode == "auto" + and self.pipeline is not None + ): + layerwise_usage_tracker = LayerwiseUsageTracker( + self.pipeline.modules, + stage_name_provider=( + lambda: ( + req.metrics.active_stage_name + if req.metrics is not None + else None + ) + ), + ) start_time = ( execution_start_time @@ -559,7 +708,7 @@ class GPUWorker(GPUWorkerPostTrainingMixin): req_label = req.request_id[:8] if req.request_id else "unnamed" with maybe_record_function(f"SAVE_OUTPUTS {req_label}"): self._materialize_output_transport(output_batch, req, save_output_paths) - self._record_output_peak_memory(output_batch) + self._record_output_peak_memory(output_batch, is_warmup=req.is_warmup) collect_perf = ( req.perf_dump_path is not None or envs.SGLANG_DIFFUSION_STAGE_LOGGING @@ -615,12 +764,115 @@ class GPUWorker(GPUWorkerPostTrainingMixin): if output_batch is None: output_batch = OutputBatch() output_batch.error = f"Error executing {error_context}: {e}" - self._record_output_peak_memory(output_batch) + self._record_output_peak_memory(output_batch, is_warmup=req.is_warmup) # clean cache if OOM if not current_platform.is_cpu(): torch.get_device_module().empty_cache() + finally: + # also runs on the propagate_forward_errors re-raise: a warmup + # forward that never completed must still leave a failed record, + # or the estimator would plan from the remaining partial data + if measure_server_warmup: + assert warmup_workload is not None + if layerwise_usage_tracker is not None: + ( + layerwise_layer_uses, + layerwise_layer_uses_by_stage, + ) = layerwise_usage_tracker.finish_with_stages() + else: + layerwise_layer_uses = {} + self._record_server_warmup_memory( + req=req, + workload=warmup_workload, + baseline_allocated_bytes=warmup_baseline_allocated_bytes, + succeeded=output_batch is not None and output_batch.error is None, + layerwise_layer_uses=layerwise_layer_uses, + layerwise_layer_uses_by_stage=layerwise_layer_uses_by_stage, + ) return output_batch + def _record_server_warmup_memory( + self, + *, + req: Req, + workload: tuple[int, int, int, int], + baseline_allocated_bytes: int, + succeeded: bool, + layerwise_layer_uses: dict[str, dict[str, tuple[int, ...]]] | None = None, + layerwise_layer_uses_by_stage: ( + dict[str, dict[str, dict[str, tuple[int, ...]]]] | None + ) = None, + ) -> None: + phase_allocated_peaks: dict[str, int] = {} + phase_components: dict[str, tuple[str, ...]] = {} + phase_used_components: dict[str, tuple[str, ...]] = {} + phase_full_weight_transition_components: dict[str, tuple[str, ...]] = {} + untracked_active_components: tuple[str, ...] = () + residency_manager = peek_global_component_residency_manager() + if residency_manager is not None: + for phase_name, peak in residency_manager.take_warmup_phase_peaks().items(): + phase_allocated_peaks[phase_name] = peak.allocated_bytes + phase_components[phase_name] = peak.active_components + phase_used_components[phase_name] = peak.used_components + phase_full_weight_transition_components[phase_name] = ( + peak.full_weight_transition_components + ) + untracked_active_components = residency_manager.current_device_components() + request_allocated_peak = max( + int(torch.get_device_module().max_memory_allocated()), + max(phase_allocated_peaks.values(), default=0), + ) + if request_allocated_peak > max(phase_allocated_peaks.values(), default=0): + # Work after the residency-managed stage timeline (for example, + # output materialization) must remain a placement constraint. + # A reserved-only increase is not a second live placement. Record + # it separately so post-placement validation can still require + # allocator headroom without charging cache to candidate deltas. + phase_allocated_peaks["request:untracked"] = request_allocated_peak + phase_components["request:untracked"] = untracked_active_components + phase_used_components["request:untracked"] = () + phase_full_weight_transition_components["request:untracked"] = () + metrics = req.metrics + width, height, num_frames, num_inference_steps = workload + self._auto_residency_warmup_records.append( + WarmupMemoryRecord( + width=width, + height=height, + num_frames=num_frames, + baseline_allocated_bytes=int(baseline_allocated_bytes), + peak_allocated_bytes=request_allocated_peak, + succeeded=succeeded, + peak_reserved_bytes=int( + torch.get_device_module().max_memory_reserved() + ), + phase_peak_allocated_bytes=phase_allocated_peaks, + phase_active_components=phase_components, + phase_used_components=phase_used_components, + phase_full_weight_transition_components=( + phase_full_weight_transition_components + ), + layerwise_layer_uses=layerwise_layer_uses or {}, + layerwise_layer_uses_by_stage=layerwise_layer_uses_by_stage or {}, + num_inference_steps=num_inference_steps, + total_duration_ms=( + float(metrics.total_duration_ms) if metrics is not None else 0.0 + ), + stage_duration_ms=(dict(metrics.stages) if metrics is not None else {}), + step_duration_ms=(tuple(metrics.steps) if metrics is not None else ()), + step_duration_ms_by_stage=( + { + stage_name: tuple(durations) + for stage_name, durations in metrics.steps_by_stage.items() + } + if metrics is not None + else {} + ), + stage_iterations=( + dict(metrics.stage_iterations) if metrics is not None else {} + ), + ) + ) + def _materialize_output_transport( self, output_batch: OutputBatch, @@ -733,16 +985,100 @@ class GPUWorker(GPUWorkerPostTrainingMixin): ) return np.asarray(materialized.frames) - def _record_output_peak_memory(self, output_batch: OutputBatch) -> None: + def _fit_auto_residency_probe(self, req: Req) -> None: + """Size the full-shape probe to what the card has left, on every rank alike.""" + records = [r for r in self._auto_residency_warmup_records if r.succeeded] + if not records or not current_platform.is_cuda(): + return + device = current_platform.get_device(self.local_rank) + free_bytes = int( + current_platform.get_available_gpu_memory(empty_cache=True) * (1 << 30) + ) + total_bytes = int(torch.cuda.get_device_properties(device).total_memory) + _, _, steps = fit_auto_residency_probe( + req, + records=records, + free_bytes=free_bytes, + total_bytes=total_bytes, + server_args=self.server_args, + ) + requested_units = ( + max(1, int(req.width or 1)) + * max(1, int(req.height or 1)) + * max(1, int(req.num_frames or 1)) + ) + estimate = estimate_default_workload_peak_bytes( + records=records, target_units=requested_units + ) + # Ranks see different free memory and hold different records; the + # forward must run one shape everywhere, so the most cautious rank wins. + agreed = torch.tensor([steps], dtype=torch.int64, device=device) + agreed = get_replica_group().all_reduce( + agreed, op=torch.distributed.ReduceOp.MAX + ) + steps = int(agreed.item()) + if steps == 0: + return + fitted = req + for _ in range(steps): + lighter = lighten_warmup_req(self.server_args, fitted) + if lighter is None: + break + fitted = lighter + if self.is_output_rank: + logger.warning( + "Auto residency probe %s would not fit: extrapolated peak %.1f GiB " + "against %.1f GiB free; probing at %s instead", + _shape_label(req), + (estimate or 0) / (1 << 30), + free_bytes / (1 << 30), + _shape_label(fitted), + ) + req.sampling_params = fitted.sampling_params + + def _release_warmup_pool(self, req: Req) -> None: + """Drop what the full-shape probe left behind before the next request. + + The probe runs a shape serving may never see. Its cached allocator + blocks would become the floor of every runtime peak measurement, and + the all-to-all staging buffers it created (IPC and Ulysses) stay + allocated at its message size. The request after the probe (the + bounded re-warm) regrows all of them at a serving-sized shape. + """ + if req.is_warmup and req.extra.get("auto_residency_full_shape_probe"): + self._release_warmup_pool_before_serving = True + return + if not self._release_warmup_pool_before_serving: + return + self._release_warmup_pool_before_serving = False + if current_platform.is_cpu() or current_platform.is_mps(): + return + from sglang.multimodal_gen.runtime.distributed.device_communicators.ipc_a2a import ( + IPC_A2A, + ) + from sglang.multimodal_gen.runtime.layers.usp import drop_a2a_staging_buffers + + IPC_A2A.drop_staging() + drop_a2a_staging_buffers() + torch.get_device_module().empty_cache() + + def _record_output_peak_memory( + self, output_batch: OutputBatch, *, is_warmup: bool = False + ) -> None: if current_platform.is_cpu(): return snapshot = capture_memory_snapshot() - self._runtime_peak_reserved_mb = max( - self._runtime_peak_reserved_mb, snapshot.peak_reserved_mb - ) - self._runtime_peak_allocated_mb = max( - self._runtime_peak_allocated_mb, snapshot.peak_allocated_mb - ) + if is_warmup: + self._warmup_peak_reserved_mb = max( + self._warmup_peak_reserved_mb, snapshot.peak_reserved_mb + ) + else: + self._runtime_peak_reserved_mb = max( + self._runtime_peak_reserved_mb, snapshot.peak_reserved_mb + ) + self._runtime_peak_allocated_mb = max( + self._runtime_peak_allocated_mb, snapshot.peak_allocated_mb + ) if self.is_output_rank: output_batch.peak_memory_mb = snapshot.peak_reserved_mb @@ -755,6 +1091,7 @@ class GPUWorker(GPUWorkerPostTrainingMixin): [ self._load_peak_reserved_mb, self._runtime_peak_reserved_mb, + self._warmup_peak_reserved_mb, self._load_peak_allocated_mb, self._runtime_peak_allocated_mb, ], @@ -769,6 +1106,7 @@ class GPUWorker(GPUWorkerPostTrainingMixin): ( load_peak_mb, runtime_peak_mb, + warmup_peak_mb, load_peak_allocated_mb, runtime_peak_allocated_mb, ) = peaks.tolist() @@ -789,6 +1127,9 @@ class GPUWorker(GPUWorkerPostTrainingMixin): peak_allocated_mb=runtime_peak_allocated_mb, ), ) + metrics.record_memory_snapshot( + "warmup_peak", replace(snapshot, peak_reserved_mb=warmup_peak_mb) + ) def _forward_group(self, batch: list[Req]) -> OutputBatch: assert self.pipeline is not None diff --git a/python/sglang/multimodal_gen/runtime/managers/memory_managers/auto_residency.py b/python/sglang/multimodal_gen/runtime/managers/memory_managers/auto_residency.py new file mode 100644 index 000000000..c80db14ee --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/managers/memory_managers/auto_residency.py @@ -0,0 +1,690 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Warmup-calibrated automatic component residency placement. + +Under ``--performance-mode auto`` with server warmup, each rank measures the +peak GPU memory of bounded synthetic warmup requests and a low-step probe at +the complete default serving shape, then selects a complete serving placement +for every eligible component under the measured memory constraints. + +When no full-shape measurement is available, the fallback estimate splits the +measured peak into persistent weights and workload-scaled activations. Scaling +the whole peak would multiply resident weights by the video frame/area cap +ratio (~16x for Wan-class defaults) and residency adjustment would never trigger. + +The planner targets the model default workload only (default resolution, +default frames, batch=1). Larger shapes, batches, or multi-image inputs need +explicit ``--component-residency``. + +Loading and serving are deliberately separate placement states. The existing +auto policy provides the initial state; when that state can complete loading +and calibration, this module optimizes the long-lived serving state and +validates the transition with a post-placement warmup. It does not force a +single placement to serve two different lifecycle objectives. +""" + +from __future__ import annotations + +import statistics +from typing import TYPE_CHECKING, Iterable, Mapping + +import msgspec + +from sglang.multimodal_gen.runtime.managers.memory_managers.component_residency import ( + COMPONENT_OFFLOAD, + LAYERWISE_OFFLOAD, + RESIDENT, +) +from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload_components import ( + is_dit_component_name, +) +from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger + +if TYPE_CHECKING: + from sglang.multimodal_gen.runtime.server_args import ServerArgs + +logger = init_logger(__name__) + +GIB_BYTES = 1024**3 + +# Activation memory rarely scales perfectly linearly with workload units; +# pad the extrapolated activation part before checking the budget. +ACTIVATION_EXTRAPOLATION_MARGIN = 1.2 +# A target-shape measurement plus the mandatory post-placement warmup justifies +# a tighter reserve than an extrapolated estimate. Both retain an absolute +# floor for allocator slack, shape variance, and CUDA graph or compile pools. +MEASURED_VRAM_RESERVE_FRACTION = 0.05 +EXTRAPOLATED_VRAM_RESERVE_FRACTION = 0.10 +MIN_VRAM_RESERVE_BYTES = 4 * GIB_BYTES +# The absolute floor is sized for datacenter cards, where either fraction can +# dominate. On a 12 GiB card a flat 4 GiB would fence off a third of the device, +# so cap the floor as a share of what is actually there. +MAX_VRAM_RESERVE_FRACTION = 0.20 + +# A feasible placement is not automatically useful. Predictions inside this +# interval are treated as latency-equivalent. The joint optimizer then avoids +# changing strategy, minimizes device memory, preserves the faster estimate, +# and finally minimizes HostPin. The raw estimate is already an upper bound: +# transfer time is capped by the measured request. +ESTIMATED_PINNED_H2D_BYTES_PER_SECOND = 24 * GIB_BYTES +MIN_LATENCY_EQUIVALENCE_NS = 50_000_000 +MAX_LATENCY_EQUIVALENCE_NS = 100_000_000 +LATENCY_EQUIVALENCE_FRACTION = 0.01 +# The transfer model ranks feasible placements; the mandatory warmup is the +# authority on whether a selected placement actually helped. Allow normal +# measurement noise, but undo a round whose calibrated request is materially +# slower than the original layout. +POST_ADJUSTMENT_REGRESSION_FRACTION = 0.05 + +PLACEMENT_STATUS_SKIPPED = "skipped" +PLACEMENT_STATUS_ADJUSTED = "adjusted" +PLACEMENT_STATUS_VALIDATED = "validated" +PLACEMENT_STATUS_ROLLED_BACK = "rolled_back" +PLACEMENT_STATUS_ROLLBACK_FAILED = "rollback_failed" + + +def describe_error(error: BaseException) -> str: + """Never-empty error text (str(AssertionError()) is "" and would be + dropped by any truthiness filter).""" + text = str(error) + return f"{type(error).__name__}: {text}" if text else type(error).__name__ + + +class WarmupMemoryRecord(msgspec.Struct, frozen=True): + """Per-rank memory measurement of one server warmup forward.""" + + width: int + height: int + num_frames: int + baseline_allocated_bytes: int + peak_allocated_bytes: int + succeeded: bool + peak_reserved_bytes: int = 0 + phase_peak_allocated_bytes: dict[str, int] = {} + phase_active_components: dict[str, tuple[str, ...]] = {} + phase_used_components: dict[str, tuple[str, ...]] = {} + phase_full_weight_transition_components: dict[str, tuple[str, ...]] = {} + layerwise_layer_uses: dict[str, dict[str, tuple[int, ...]]] = {} + layerwise_layer_uses_by_stage: dict[str, dict[str, dict[str, tuple[int, ...]]]] = {} + num_inference_steps: int = 1 + total_duration_ms: float = 0.0 + stage_duration_ms: dict[str, float] = {} + step_duration_ms: tuple[float, ...] = () + step_duration_ms_by_stage: dict[str, tuple[float, ...]] = {} + stage_iterations: dict[str, tuple[int, int]] = {} + + def workload_units(self) -> int: + return max(1, self.width) * max(1, self.height) * max(1, self.num_frames) + + +class ResidencyTarget(msgspec.Struct, frozen=True): + """One complete target state for an auto-managed component.""" + + component_name: str + residency_mode: str + target_resident_weight_bytes: int + # Estimated per-request host-to-device traffic this target removes. + h2d_bytes_per_request: int + # Layerwise candidates jointly choose stage-scoped GPU residency and host + # pinning. None is used by ordinary component placement. + target_layerwise_resident_layers: tuple[int, ...] | None = None + target_layerwise_pinned_layers: tuple[tuple[int, ...], ...] | None = None + pinned_host_delta_bytes: int = 0 + host_unpin_scratch_bytes: int = 0 + host_pin_scratch_bytes: int = 0 + host_materialize_scratch_bytes: int = 0 + # Signed device-memory delta while applying the placement before the + # validation warmup. Layerwise -> resident materializes every managed + # layer immediately; a demotion can release those bytes first and fund a + # later materialization in the same transaction. + device_transition_delta_bytes: int = 0 + permanent_residency: bool = False + # Device-memory delta relative to the measured placement. A component + # already loaded for its own phase has a different delta from phases where + # it is absent; keeping both avoids adding the same weights twice. + active_device_delta_bytes: int = 0 + # Delta when the component is already present because of async prefetch, + # but is not the semantic owner of this phase. + present_device_delta_bytes: int = 0 + inactive_device_delta_bytes: int = 0 + # None preserves the historical derived target for hand-built callers: + # partial layerwise targets remain layerwise, every other option is + # resident. Generated complete-state frontiers set this explicitly. + target_residency_mode: str | None = None + current_placement: bool = False + target_device_weight_bytes: int = 0 + target_pinned_host_bytes: int = 0 + + def target_mode(self) -> str: + if self.target_residency_mode is not None: + return self.target_residency_mode + if ( + self.target_layerwise_resident_layers is not None + and not self.permanent_residency + ): + return LAYERWISE_OFFLOAD + return RESIDENT + + def option_key(self) -> str: + target_mode = self.target_mode() + if target_mode == COMPONENT_OFFLOAD: + return f"{self.component_name}:component-offload" + if self.target_layerwise_resident_layers is None: + return f"{self.component_name}:resident" + layer_counts = ",".join( + str(count) for count in self.target_layerwise_resident_layers + ) + pinned = "|".join( + ",".join(str(index) for index in indices) or "-" + for indices in self.target_layerwise_pinned_layers or () + ) + permanence = "permanent" if self.permanent_residency else "stage" + return f"{self.component_name}:{permanence}:layers={layer_counts}:pins={pinned}" + + +class DefaultWorkload(msgspec.Struct, frozen=True): + """The model-default request shape the planner is calibrated for.""" + + width: int | None + height: int | None + num_frames: int + num_inference_steps: int + + def workload_units(self) -> int | None: + if self.width is None or self.height is None: + return None + return max(1, self.width) * max(1, self.height) * max(1, self.num_frames) + + def describe(self) -> str: + if self.width is None or self.height is None: + return "model-default" + return f"{self.width}x{self.height}x{self.num_frames}f" + + +class RankResidencyReport(msgspec.Struct, frozen=True): + """One rank's inputs to the replica-wide placement decision.""" + + rank: int + budget_bytes: int + estimated_peak_bytes: int | None + target_workload_measured: bool = False + observed_reserved_bytes: int = 0 + estimated_peak_bytes_by_phase: dict[str, int] = {} + active_components_by_phase: dict[str, tuple[str, ...]] = {} + used_components_by_phase: dict[str, tuple[str, ...]] = {} + full_weight_transition_components_by_phase: dict[str, tuple[str, ...]] = {} + current_device_weight_bytes_by_component: dict[str, int] = {} + node_rank: int = 0 + pinned_host_bytes: int = 0 + host_pin_capacity_bytes: int = 0 + host_transition_headroom_bytes: int = 0 + device_transition_allocated_bytes: int = 0 + estimated_request_duration_ns: int = 0 + measured_request_duration_ns: int = 0 + candidate_latency_savings_ns: dict[str, int] = {} + candidates: list[ResidencyTarget] = [] + skip_reason: str | None = None + + +def resolve_default_workload(server_args: ServerArgs) -> DefaultWorkload: + """Resolve the default request shape the planner is optimized for.""" + from sglang.multimodal_gen.runtime.warmup_request_builder import ( + get_model_sampling_defaults, + resolve_default_workload_shape, + ) + + defaults = get_model_sampling_defaults(server_args) + width, height, num_frames = resolve_default_workload_shape(server_args, defaults) + return DefaultWorkload( + width=width, + height=height, + num_frames=num_frames, + num_inference_steps=defaults.num_inference_steps or 1, + ) + + +def resolve_measured_default_workload( + workload: DefaultWorkload, records: Iterable[WarmupMemoryRecord] +) -> DefaultWorkload: + """Fill an implicit default resolution from the executed warmup. + + Image-edit pipelines can derive their output size from the input image, so + the sampling defaults legitimately omit width and height. The warmup record + is captured after input validation and therefore contains the effective + serving shape. Keep the model-default frame count because video warmup may + intentionally cap frames before measurement. + """ + if workload.workload_units() is not None: + return workload + measured = [ + record + for record in records + if record.succeeded and record.width > 0 and record.height > 0 + ] + if not measured: + return workload + representative = max(measured, key=lambda record: record.width * record.height) + return DefaultWorkload( + width=representative.width, + height=representative.height, + num_frames=workload.num_frames, + num_inference_steps=workload.num_inference_steps, + ) + + +def estimate_layerwise_layer_uses( + *, + records: Iterable[WarmupMemoryRecord], + target_units: int | None, + target_num_inference_steps: int, +) -> dict[str, dict[str, tuple[int, ...]]]: + """Estimate per-request layer calls from the same calibration forward. + + A full-shape memory probe deliberately runs only a few denoise steps. + Stage-attributed calls use that stage's measured and target iteration + counts, so independent shape, paint, refiner, and chunk loops are not all + multiplied by one request-wide ratio. Legacy records retain the repeated + DiT-layer heuristic. + """ + successful = [record for record in records if record.succeeded] + if target_units is not None: + covering = [ + record for record in successful if record.workload_units() >= target_units + ] + if covering: + successful = covering + + estimated: dict[str, dict[str, list[int]]] = {} + for record in successful: + source_steps = max(1, record.num_inference_steps) + component_stages = _component_stages(record) + repeated_stages = _repeated_stages(record, component_stages) + for component_name, groups in record.layerwise_layer_uses.items(): + component = estimated.setdefault(component_name, {}) + for layer_name, counts in groups.items(): + target = component.setdefault(layer_name, [0] * len(counts)) + if len(target) != len(counts): + continue + stage_counts = [ + ( + stage_name, + stage_components[component_name][layer_name], + ) + for stage_name, stage_components in ( + record.layerwise_layer_uses_by_stage.items() + ) + if component_name in stage_components + and layer_name in stage_components[component_name] + and len(stage_components[component_name][layer_name]) == len(counts) + ] + for layer_index, count in enumerate(counts): + if stage_counts: + measured_by_stage = sum( + per_layer_counts[layer_index] + for _, per_layer_counts in stage_counts + ) + scaled = max(0, count - measured_by_stage) + for stage_name, per_layer_counts in stage_counts: + measured_iterations, target_iterations = _stage_iterations( + record, + stage_name, + repeated_stages=repeated_stages, + target_num_inference_steps=(target_num_inference_steps), + ) + stage_count = per_layer_counts[layer_index] + if stage_count <= 1: + scaled += stage_count + else: + scaled += ( + stage_count * target_iterations + + measured_iterations + - 1 + ) // measured_iterations + else: + scaled = count + component_is_repeated = is_dit_component_name( + component_name + ) or any( + stage_name in repeated_stages + for stage_name in component_stages.get(component_name, ()) + ) + if ( + component_is_repeated + and count > 1 + and target_num_inference_steps > source_steps + ): + scaled = ( + count * target_num_inference_steps + source_steps - 1 + ) // source_steps + target[layer_index] = max(target[layer_index], scaled) + return { + component_name: { + layer_name: tuple(counts) for layer_name, counts in groups.items() + } + for component_name, groups in estimated.items() + } + + +def _component_stages( + record: WarmupMemoryRecord, + *, + timed_stage_names: set[str] | None = None, +) -> dict[str, set[str]]: + component_stages: dict[str, set[str]] = {} + phase_components = record.phase_used_components or record.phase_active_components + for phase_name, components in phase_components.items(): + fields = phase_name.split(":", 2) + if len(fields) < 2 or not fields[0].isdigit(): + continue + stage_name = fields[1] + if timed_stage_names is not None and stage_name not in timed_stage_names: + continue + for component_name in components: + component_stages.setdefault(component_name, set()).add(stage_name) + return component_stages + + +def _repeated_stages( + record: WarmupMemoryRecord, + component_stages: Mapping[str, set[str]], +) -> set[str]: + stages = { + stage_name + for component_name, stage_names in component_stages.items() + if is_dit_component_name(component_name) + for stage_name in stage_names + } + stages.update(record.stage_iterations) + stages.update( + stage_name + for stage_name in set(record.stage_duration_ms).union( + *(stage_names for stage_names in component_stages.values()) + ) + if stage_name.endswith("DenoisingStage") + and not stage_name.endswith("BeforeDenoisingStage") + ) + return stages + + +def _stage_iterations( + record: WarmupMemoryRecord, + stage_name: str, + *, + repeated_stages: set[str], + target_num_inference_steps: int, +) -> tuple[int, int]: + explicit = record.stage_iterations.get(stage_name) + if explicit is not None: + return max(1, explicit[0]), max(1, explicit[1]) + measured = max(1, record.num_inference_steps) + target = ( + max(1, target_num_inference_steps) + if stage_name in repeated_stages + else measured + ) + return measured, target + + +def estimate_default_workload_timing( + *, + records: Iterable[WarmupMemoryRecord], + target_units: int | None, + target_num_inference_steps: int, +) -> tuple[int, dict[str, int], dict[str, tuple[str, ...]]]: + """Estimate full-request and stage durations from the warmup workload. + + Repeated stages scale by their own measured/default iteration counts. The + full-shape probe intentionally executes only a few steps, so using its raw + total would make every one-shot encoder transfer look important relative + to a long video request. + """ + successful = [record for record in records if record.succeeded] + if not successful: + return 0, {}, {} + if target_units is not None: + at_target = [ + record for record in successful if record.workload_units() >= target_units + ] + if at_target: + successful = at_target + representative = max( + successful, + key=lambda record: ( + record.workload_units(), + record.total_duration_ms, + ), + ) + if representative.total_duration_ms <= 0 or not representative.stage_duration_ms: + return 0, {}, {} + + component_stages = _component_stages( + representative, + timed_stage_names=set(representative.stage_duration_ms), + ) + repeated_stages = _repeated_stages(representative, component_stages) + + stage_duration_ns: dict[str, int] = {} + for stage_name, duration_ms in representative.stage_duration_ms.items(): + measured_iterations, target_iterations = _stage_iterations( + representative, + stage_name, + repeated_stages=repeated_stages, + target_num_inference_steps=target_num_inference_steps, + ) + step_durations = representative.step_duration_ms_by_stage.get(stage_name, ()) + if not step_durations and len(repeated_stages) == 1: + step_durations = representative.step_duration_ms + if stage_name in repeated_stages and step_durations: + steady_steps = ( + step_durations[1:] if len(step_durations) > 1 else step_durations + ) + non_step_ms = max(0.0, duration_ms - sum(step_durations)) + target_duration_ms = non_step_ms + ( + statistics.median(steady_steps) * target_iterations + ) + else: + target_duration_ms = duration_ms * target_iterations / measured_iterations + stage_duration_ns[stage_name] = max(0, int(target_duration_ms * 1_000_000)) + + measured_stage_ms = sum(representative.stage_duration_ms.values()) + untracked_ms = max(0.0, representative.total_duration_ms - measured_stage_ms) + request_duration_ns = sum(stage_duration_ns.values()) + int( + untracked_ms * 1_000_000 + ) + + return ( + request_duration_ns, + stage_duration_ns, + { + component_name: tuple(sorted(stage_names)) + for component_name, stage_names in component_stages.items() + }, + ) + + +def estimate_default_workload_peak_bytes( + *, + records: Iterable[WarmupMemoryRecord], + target_units: int | None, + constant_weight_bytes: int = 0, +) -> int | None: + """Extrapolate live warmup memory to the default workload. + + Real measurements use allocated bytes. Cached allocator blocks are + reclaimable and are covered by the explicit VRAM reserve; treating them as + live memory would charge the same storage again when adding resident + weights. + + Preference order: + 1. A measurement at or above the target workload bounds the peak directly. + 2. Two distinct measured workload sizes fit ``peak = constant + slope * + units``; only the fitted linear part is extrapolated (and padded). + Under offload the pre-forward baseline is nearly empty, so a + single-point split cannot separate constant costs (streamed weights, + attention workspace, tiled VAE decode) from workload-linear + activations -- measured on Wan2.1-14B, single-point scaling + overestimated a ~30 GiB peak as ~183 GiB. + 3. One usable size: scale everything above the pre-forward allocated + baseline (conservative; may block adjustment but never over-allocates). + + Returns None when the estimate cannot be trusted: no successful records, + the target workload is unknown (an unknown target would silently equate the + area/frame-capped warmup peak with the real serving peak), or a probe at or + below the target ran out of memory. That last case is not missing data but + a measurement: the card could not hold the target as it is already + configured, so making more weights resident can only make it worse. A probe + that failed strictly above the target says nothing about the target and is + dropped instead. + """ + records = list(records) + if target_units is None: + return None + failed_units = [ + record.workload_units() for record in records if not record.succeeded + ] + if any(units <= target_units for units in failed_units): + return None + records = [record for record in records if record.succeeded] + if not records: + return None + + peak_by_units: dict[int, int] = {} + for record in records: + units = record.workload_units() + peak = record.peak_allocated_bytes + peak_by_units[units] = max(peak_by_units.get(units, 0), peak) + + covering_peaks = [ + peak for units, peak in peak_by_units.items() if units >= target_units + ] + if covering_peaks: + return max(covering_peaks) + + if len(peak_by_units) >= 2: + # fit on the two largest sizes: closest to the target, best local slope + (large_units, large_peak), (small_units, small_peak) = sorted( + peak_by_units.items(), reverse=True + )[:2] + slope = (large_peak - small_peak) / (large_units - small_units) + if slope >= 0: + constant = max( + large_peak - slope * large_units, + min(constant_weight_bytes, large_peak), + ) + return int( + constant + slope * target_units * ACTIVATION_EXTRAPOLATION_MARGIN + ) + # negative slope is measurement noise; fall through to the + # conservative single-point estimate + + estimates = [] + for record in records: + peak = record.peak_allocated_bytes + baseline = min( + max(record.baseline_allocated_bytes, constant_weight_bytes), peak + ) + activation = peak - baseline + ratio = target_units / record.workload_units() + estimates.append( + baseline + int(activation * ratio * ACTIVATION_EXTRAPOLATION_MARGIN) + ) + return max(estimates) + + +def estimate_workload_phase_peaks( + *, + records: Iterable[WarmupMemoryRecord], + target_units: int | None, + component_weight_bytes: Mapping[str, int], +) -> tuple[ + dict[str, int], + dict[str, tuple[str, ...]], + dict[str, tuple[str, ...]], + dict[str, tuple[str, ...]], +]: + """Estimate each measured execution phase at the target workload. + + A component already active in a phase is part of that phase's measured + peak. Keeping a component resident therefore adds no weight bytes to its + own component-offload phase, while it adds the full footprint to phases + where the component was absent. Measurements with different active layouts + remain separate constraints; combining one layout's peak with another + layout's component set would describe a state that never occurred. + """ + successful = [record for record in records if record.succeeded] + if target_units is not None: + covering = [ + record for record in successful if record.workload_units() >= target_units + ] + if covering: + successful = covering + + grouped: dict[ + tuple[str, tuple[str, ...], tuple[str, ...], tuple[str, ...]], + list[WarmupMemoryRecord], + ] = {} + for record in successful: + used_by_phase = record.phase_used_components or record.phase_active_components + for phase_name in record.phase_peak_allocated_bytes: + active = tuple(sorted(record.phase_active_components.get(phase_name, ()))) + used = tuple(sorted(used_by_phase.get(phase_name, ()))) + full_weight_transitions = tuple( + sorted( + record.phase_full_weight_transition_components.get(phase_name, ()) + ) + ) + grouped.setdefault( + (phase_name, active, used, full_weight_transitions), [] + ).append(record) + + layouts_per_phase: dict[str, int] = {} + for phase_name, _, _, _ in grouped: + layouts_per_phase[phase_name] = layouts_per_phase.get(phase_name, 0) + 1 + + estimated_peaks: dict[str, int] = {} + active_components: dict[str, tuple[str, ...]] = {} + used_components: dict[str, tuple[str, ...]] = {} + full_weight_transition_components: dict[str, tuple[str, ...]] = {} + layout_indices: dict[str, int] = {} + for ( + phase_name, + active, + used, + full_weight_transitions, + ), phase_records in sorted(grouped.items()): + output_name = phase_name + if layouts_per_phase[phase_name] > 1: + index = layout_indices.get(phase_name, 0) + layout_indices[phase_name] = index + 1 + output_name = f"{phase_name}:layout:{index}" + weight_floor = sum(component_weight_bytes.get(name, 0) for name in active) + phase_measurements = [ + WarmupMemoryRecord( + width=record.width, + height=record.height, + num_frames=record.num_frames, + baseline_allocated_bytes=min( + record.baseline_allocated_bytes, + record.phase_peak_allocated_bytes[phase_name], + ), + peak_allocated_bytes=record.phase_peak_allocated_bytes[phase_name], + succeeded=True, + ) + for record in phase_records + ] + estimate = estimate_default_workload_peak_bytes( + records=phase_measurements, + target_units=target_units, + constant_weight_bytes=weight_floor, + ) + if estimate is None: + continue + estimated_peaks[output_name] = estimate + active_components[output_name] = active + used_components[output_name] = used + full_weight_transition_components[output_name] = full_weight_transitions + return ( + estimated_peaks, + active_components, + used_components, + full_weight_transition_components, + ) diff --git a/python/sglang/multimodal_gen/runtime/managers/memory_managers/component_manager.py b/python/sglang/multimodal_gen/runtime/managers/memory_managers/component_manager.py index b45d774a9..7984fd51e 100644 --- a/python/sglang/multimodal_gen/runtime/managers/memory_managers/component_manager.py +++ b/python/sglang/multimodal_gen/runtime/managers/memory_managers/component_manager.py @@ -1,4 +1,4 @@ -from collections.abc import Iterator +from collections.abc import Iterable, Iterator from contextlib import contextmanager from dataclasses import dataclass from typing import Mapping, MutableMapping, Protocol, Sequence @@ -58,6 +58,14 @@ class ResidencyState: batch_is_warmup: bool = False +@dataclass(frozen=True, slots=True) +class WarmupPhasePeak: + active_components: tuple[str, ...] + allocated_bytes: int + used_components: tuple[str, ...] = () + full_weight_transition_components: tuple[str, ...] = () + + class ResidencyBatch(Protocol): is_warmup: bool @@ -129,6 +137,13 @@ class ComponentResidencyManager: ] = {} self._uses_seen: dict[str, ComponentUse] = {} self._modules_seen: dict[str, nn.Module] = {} + self._track_warmup_memory = False + self._warmup_phase_key: str | None = None + self._warmup_phase_components: tuple[str, ...] = () + self._warmup_phase_used_components: tuple[str, ...] = () + self._warmup_phase_full_weight_transition_components: tuple[str, ...] = () + self._warmup_phase_peaks: dict[str, WarmupPhasePeak] = {} + self._completed_warmup_phase_peaks: dict[str, WarmupPhasePeak] = {} def refresh_pipeline(self, pipeline: ComponentResidencyPipeline) -> None: custom_strategies = dict(pipeline.component_residency_strategies) @@ -178,6 +193,24 @@ class ComponentResidencyManager: self._ordered_uses = tuple( use for uses in self._stage_uses_by_index for use in uses ) + self._track_warmup_memory = ( + self.state.batch_is_warmup + and self.server_args.pipeline_config.supports_auto_residency + and current_platform.is_cuda() + and torch.get_device_module().is_available() + ) + self._warmup_phase_key = None + self._warmup_phase_components = () + self._warmup_phase_used_components = () + self._warmup_phase_full_weight_transition_components = () + self._warmup_phase_peaks = {} + self._completed_warmup_phase_peaks = {} + if self._track_warmup_memory: + # GPUWorker reset the request peak before entering the pipeline. + # Start the first interval without another reset so preprocessing + # before stage 0 remains part of the placement constraints. + self._warmup_phase_key = "request:before-stage" + self._warmup_phase_components = self._warmup_active_components() self._validate_explicit_nonresident_components() def _validate_explicit_nonresident_components(self) -> None: @@ -224,6 +257,49 @@ class ComponentResidencyManager: self.state.stage_index = stage_index self.state.stage_name = self.stage_name(stage) self.state.next_stage_name = self._next_stage_name(stage_index) + if self._track_warmup_memory: + self._begin_warmup_phase( + key=f"{stage_index}:{self.state.stage_name}:setup", + components=self._warmup_active_components(), + used_components=(), + ) + + @contextmanager + def full_weight_transition(self, component_names: Iterable[str]) -> Iterator[None]: + """Measure request logic that temporarily materializes complete weights.""" + names = tuple(sorted(set(component_names))) + if not self._track_warmup_memory or not names: + yield + return + previous_phase = ( + self._warmup_phase_key, + self._warmup_phase_components, + self._warmup_phase_used_components, + self._warmup_phase_full_weight_transition_components, + ) + self._begin_warmup_phase( + key=( + f"{self.state.stage_index}:{self.state.stage_name}:" + f"full-weight-transition:{','.join(names)}" + ), + components=self._warmup_active_components(), + used_components=(), + full_weight_transition_components=names, + ) + try: + yield + finally: + previous_key, components, used_components, transitions = previous_phase + if previous_key is None: + self._record_warmup_phase_peak() + self._warmup_phase_key = None + else: + self._begin_warmup_phase( + key=previous_key, + components=components, + used_components=used_components, + full_weight_transition_components=transitions, + ) def begin_stage(self) -> None: """Prepare a stage that declares one uninterrupted component use.""" @@ -233,6 +309,7 @@ class ComponentResidencyManager: def end_stage(self) -> None: """Close the component interval owned by the current stage.""" + self._record_warmup_phase_peak() if self._active_use is None: return if self._active_use.stage_name != self.state.stage_name: @@ -263,7 +340,13 @@ class ComponentResidencyManager: self._active_use_module is not None and active_module is not self._active_use_module ) + requires_prepare = active_module is not None and ( + self._active_use_module is None + or module_changed + or use.target_dtype != previous_use.target_dtype + ) if module_changed: + self._begin_warmup_transition(previous_use, None) self._disable_active_nvtx() self._finish_use( previous_use, @@ -271,18 +354,19 @@ class ComponentResidencyManager: keep_on_warmup=False, force=True, ) - if active_module is not None and ( - self._active_use_module is None - or module_changed - or use.target_dtype != previous_use.target_dtype - ): + self._begin_warmup_transition(None, use) + elif requires_prepare: + self._begin_warmup_transition(previous_use, use) + if requires_prepare: active_module = self._prepare_forward_use(use, module=active_module) + self._begin_warmup_use(use) self._active_use = use self._active_use_module = active_module self.state.current_use = use self._enable_nvtx_for_use(use, active_module) return if self._active_use is not None: + self._begin_warmup_transition(self._active_use, None) self._disable_active_nvtx() self._finish_use( self._active_use, @@ -292,8 +376,10 @@ class ComponentResidencyManager: self._active_use = None self._active_use_module = None self.state.current_use = None + self._begin_warmup_transition(None, use) self._mark_current_use(use) module = self._prepare_forward_use(use, module=module) + self._begin_warmup_use(use) self._active_use = use self._active_use_module = module self._enable_nvtx_for_use(use, module) @@ -303,6 +389,7 @@ class ComponentResidencyManager: """End one sequential component use interval.""" if self._active_use is None or not self._same_use(self._active_use, use): return + self._begin_warmup_transition(self._active_use, None) self._disable_active_nvtx() self._finish_use( self._active_use, @@ -316,6 +403,7 @@ class ComponentResidencyManager: self._active_use = None self._active_use_module = None self.state.current_use = None + self._begin_warmup_between_uses() self._prefetch_next_memory_intensive_use() @contextmanager @@ -378,6 +466,7 @@ class ComponentResidencyManager: if self._active_use is None: return active_use = self._active_use + self._begin_warmup_transition(active_use, None) self._disable_active_nvtx() self._finish_use( active_use, @@ -387,6 +476,7 @@ class ComponentResidencyManager: self._active_use = None self._active_use_module = None self.state.current_use = None + self._begin_warmup_between_uses() if prefetch_next: self._prefetch_next_memory_intensive_use() @@ -483,8 +573,11 @@ class ComponentResidencyManager: self._uses_seen[use.component_name] = use self._modules_seen[use.component_name] = module + self._begin_warmup_prefetch(use) if strategy.prefetch_for_use(module, use, self.state): self._prefetched_use_keys.add(self._use_key(use)) + else: + self._begin_warmup_between_uses() def _finish_use( self, @@ -534,11 +627,154 @@ class ComponentResidencyManager: not self._is_single_dit_component(component_name) or keep_single_dit ) strategy = self.strategy_for(component_name, module) + if self._track_warmup_memory: + will_prepare = self.state.batch_is_warmup and preferred + self._begin_warmup_phase( + key=f"request:cleanup:{component_name}", + components=self._warmup_active_components( + (use,) if will_prepare else () + ), + used_components=(component_name,) if will_prepare else (), + ) was_on_supported_device = self._module_on_supported_device(module) strategy.finish_request(module, use, self.state, preferred=preferred) self._empty_cache_after_large_release( use, strategy, module, was_on_supported_device ) + if self._track_warmup_memory: + self._record_warmup_phase_peak() + self._warmup_phase_peaks["idle"] = WarmupPhasePeak( + active_components=self._warmup_active_components(), + allocated_bytes=int(torch.get_device_module().memory_allocated()), + used_components=(), + ) + self._completed_warmup_phase_peaks = dict(self._warmup_phase_peaks) + self._track_warmup_memory = False + + def _begin_warmup_phase( + self, + *, + key: str, + components: tuple[str, ...], + used_components: tuple[str, ...], + full_weight_transition_components: tuple[str, ...] = (), + ) -> None: + if not self._track_warmup_memory: + return + self._record_warmup_phase_peak() + self._warmup_phase_key = key + self._warmup_phase_components = tuple(sorted(set(components))) + self._warmup_phase_used_components = tuple(sorted(set(used_components))) + self._warmup_phase_full_weight_transition_components = tuple( + sorted(set(full_weight_transition_components)) + ) + torch.get_device_module().reset_peak_memory_stats() + + def _begin_warmup_transition( + self, previous: ComponentUse | None, upcoming: ComponentUse | None + ) -> None: + if not self._track_warmup_memory: + return + previous_name = previous.component_name if previous is not None else "idle" + upcoming_name = upcoming.component_name if upcoming is not None else "idle" + self._begin_warmup_phase( + key=( + f"{self.state.stage_index}:{self.state.stage_name}:transition:" + f"{previous_name}->{upcoming_name}" + ), + components=self._warmup_active_components( + tuple(use for use in (previous, upcoming) if use is not None) + ), + used_components=tuple( + use.component_name for use in (previous, upcoming) if use is not None + ), + ) + + def _begin_warmup_use(self, use: ComponentUse) -> None: + phase = use.phase or use.component_name + self._begin_warmup_phase( + key=f"{self.state.stage_index}:{self.state.stage_name}:use:{phase}", + components=self._warmup_active_components((use,)), + used_components=(use.component_name,), + ) + + def _begin_warmup_between_uses(self) -> None: + self._begin_warmup_phase( + key=f"{self.state.stage_index}:{self.state.stage_name}:between", + components=self._warmup_active_components(), + used_components=(), + ) + + def _begin_warmup_prefetch(self, use: ComponentUse) -> None: + phase = use.phase or use.component_name + self._begin_warmup_phase( + key=f"{self.state.stage_index}:{self.state.stage_name}:prefetch:{phase}", + components=self._warmup_active_components((use,)), + used_components=(use.component_name,), + ) + + def _warmup_active_components( + self, active_uses: Sequence[ComponentUse] = () + ) -> tuple[str, ...]: + components = {use.component_name for use in active_uses} + for component_name, module in self.pipeline.modules.items(): + if not isinstance(module, nn.Module): + continue + if is_layerwise_offloaded_module(module): + continue + if self._module_on_supported_device(module): + components.add(component_name) + return tuple(sorted(components)) + + def _record_warmup_phase_peak(self) -> None: + if not self._track_warmup_memory or self._warmup_phase_key is None: + return + peak = WarmupPhasePeak( + active_components=self._warmup_phase_components, + allocated_bytes=int(torch.get_device_module().max_memory_allocated()), + used_components=self._warmup_phase_used_components, + full_weight_transition_components=( + self._warmup_phase_full_weight_transition_components + ), + ) + previous = self._warmup_phase_peaks.get(self._warmup_phase_key) + if previous is None: + self._warmup_phase_peaks[self._warmup_phase_key] = peak + else: + self._warmup_phase_peaks[self._warmup_phase_key] = WarmupPhasePeak( + active_components=tuple( + sorted( + set(previous.active_components) & set(peak.active_components) + ) + ), + allocated_bytes=max(previous.allocated_bytes, peak.allocated_bytes), + used_components=tuple( + sorted(set(previous.used_components) & set(peak.used_components)) + ), + full_weight_transition_components=tuple( + sorted( + set(previous.full_weight_transition_components) + & set(peak.full_weight_transition_components) + ) + ), + ) + + def take_warmup_phase_peaks( + self, + ) -> dict[str, WarmupPhasePeak]: + """Return and clear the most recently completed warmup phase peaks.""" + peaks = self._completed_warmup_phase_peaks + self._completed_warmup_phase_peaks = {} + return peaks + + def current_device_components(self) -> tuple[str, ...]: + """Components whose complete module is currently on the device. + + Dormant layerwise-managed modules are excluded because only their + resident window is present. Active layerwise uses are attributed by + the managed phase timeline instead. + """ + return self._warmup_active_components() def stage_name(self, stage: ComponentResidencyStage) -> str: return self._stage_names_by_id.get(id(stage), stage.__class__.__name__) @@ -703,6 +939,11 @@ class ComponentResidencyManager: _GLOBAL_COMPONENT_RESIDENCY_MANAGER: ComponentResidencyManager | None = None +def peek_global_component_residency_manager() -> ComponentResidencyManager | None: + """Return the process-global manager without creating one.""" + return _GLOBAL_COMPONENT_RESIDENCY_MANAGER + + def get_global_component_residency_manager( pipeline: ComponentResidencyPipeline, server_args: ServerArgs, diff --git a/python/sglang/multimodal_gen/runtime/managers/memory_managers/layerwise_offload.py b/python/sglang/multimodal_gen/runtime/managers/memory_managers/layerwise_offload.py index 67644a22c..68d43809a 100644 --- a/python/sglang/multimodal_gen/runtime/managers/memory_managers/layerwise_offload.py +++ b/python/sglang/multimodal_gen/runtime/managers/memory_managers/layerwise_offload.py @@ -5,7 +5,15 @@ import threading from collections.abc import Mapping, Sequence from contextlib import nullcontext from time import perf_counter -from typing import Any, Dict, List, Optional, Set, Tuple +from typing import ( + Any, + Callable, + Dict, + List, + Optional, + Set, + Tuple, +) import torch from torch.distributed.tensor import DTensor @@ -2205,3 +2213,99 @@ def configure_layerwise_offload_modules( elif warn_missing: logger.debug("No selected pipeline component enabled layerwise offload") return configured_component_names + + +class LayerwiseUsageTracker: + """Temporary per-layer call counters for one calibration request. + + Managers cannot provide this information when layerwise offload has not + been configured yet or was disabled by a resident placement. Observe the + declared layer groups directly and remove every hook before returning, so + ordinary serving requests pay no counter or hook-dispatch cost. + """ + + def __init__( + self, + modules: Mapping[str, object], + *, + stage_name_provider: Callable[[], str | None] | None = None, + ) -> None: + self._handles: list[Any] = [] + self._counts: dict[str, dict[str, list[int]]] = {} + self._counts_by_stage: dict[str, dict[str, dict[str, list[int]]]] = {} + self._stage_name_provider = stage_name_provider + for component_name, module in modules.items(): + if not isinstance(module, LayerwiseOffloadableModuleMixin): + continue + named_modules = dict(module.named_modules()) + component_counts: dict[str, list[int]] = {} + for layer_name in module.layer_names: + layers = named_modules.get(layer_name) + if not isinstance(layers, (torch.nn.ModuleList, torch.nn.Sequential)): + continue + counts = [0] * len(layers) + if not counts: + continue + component_counts[layer_name] = counts + for layer_index, layer in enumerate(layers): + + def record_use( + _module, + _inputs, + *, + target_counts=counts, + target_index=layer_index, + target_component_name=component_name, + target_layer_name=layer_name, + target_layer_count=len(layers), + ) -> None: + target_counts[target_index] += 1 + if self._stage_name_provider is None: + return + stage_name = self._stage_name_provider() + if stage_name is None: + return + stage_counts = self._counts_by_stage.setdefault(stage_name, {}) + component_stage_counts = stage_counts.setdefault( + target_component_name, {} + ) + layer_stage_counts = component_stage_counts.setdefault( + target_layer_name, [0] * target_layer_count + ) + layer_stage_counts[target_index] += 1 + + self._handles.append(layer.register_forward_pre_hook(record_use)) + if component_counts: + self._counts[component_name] = component_counts + + def finish(self) -> dict[str, dict[str, tuple[int, ...]]]: + counts, _ = self.finish_with_stages() + return counts + + def finish_with_stages( + self, + ) -> tuple[ + dict[str, dict[str, tuple[int, ...]]], + dict[str, dict[str, dict[str, tuple[int, ...]]]], + ]: + for handle in self._handles: + handle.remove() + self._handles.clear() + counts = { + component_name: { + layer_name: tuple(counts) + for layer_name, counts in component_counts.items() + } + for component_name, component_counts in self._counts.items() + } + counts_by_stage = { + stage_name: { + component_name: { + layer_name: tuple(counts) + for layer_name, counts in component_counts.items() + } + for component_name, component_counts in stage_counts.items() + } + for stage_name, stage_counts in self._counts_by_stage.items() + } + return counts, counts_by_stage diff --git a/python/sglang/multimodal_gen/runtime/managers/scheduler.py b/python/sglang/multimodal_gen/runtime/managers/scheduler.py index e1a31976f..7cc308fb3 100644 --- a/python/sglang/multimodal_gen/runtime/managers/scheduler.py +++ b/python/sglang/multimodal_gen/runtime/managers/scheduler.py @@ -16,15 +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.entrypoints.post_training.io_struct import ( - GetWeightsChecksumReqInput, - ReleaseMemoryOccupationReqInput, - ResumeMemoryOccupationReqInput, - UpdateWeightFromDiskReqInput, - UpdateWeightFromTensorCheckerReqInput, - UpdateWeightFromTensorReqInput, -) -from sglang.multimodal_gen.runtime.entrypoints.utils import ( +from sglang.multimodal_gen.runtime.entrypoints.control_requests import ( GetDisaggStatsReq, ListLorasReq, MergeLoraWeightsReq, @@ -33,6 +25,14 @@ from sglang.multimodal_gen.runtime.entrypoints.utils import ( ShutdownReq, UnmergeLoraWeightsReq, ) +from sglang.multimodal_gen.runtime.entrypoints.post_training.io_struct import ( + GetWeightsChecksumReqInput, + ReleaseMemoryOccupationReqInput, + ResumeMemoryOccupationReqInput, + UpdateWeightFromDiskReqInput, + UpdateWeightFromTensorCheckerReqInput, + UpdateWeightFromTensorReqInput, +) from sglang.multimodal_gen.runtime.ipc_array import ( is_local_endpoint, spill_large_arrays_to_file_refs, 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 0e217caeb..bc41754db 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/schedule_batch.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/schedule_batch.py @@ -341,6 +341,7 @@ class Req: self.suppress_logs = True self.metrics.suppress_stage_breakdown = True self.extra["cache_dit_num_inference_steps"] = self.num_inference_steps + self.extra["warmup_target_num_inference_steps"] = self.num_inference_steps self.num_inference_steps = warmup_steps def copy_as_warmup(self, warmup_steps: int = 1) -> Req: @@ -348,6 +349,34 @@ class Req: req.set_as_warmup(warmup_steps) return req + def record_stage_iterations( + self, + measured_iterations: int, + target_iterations: int | None = None, + ) -> None: + """Record a stage loop against its full default-request work. + + Most stages declare the count as a formula of the step count + (``PipelineStage.default_workload_iterations``) and never call this. + It is for loops whose length is only known inside them (chunked or + block-wise schedules); ``target_iterations`` defaults to scaling the + measured count from the probe's steps to the default workload's. + """ + if not self.is_warmup or self.metrics is None: + return + measured = max(1, int(measured_iterations)) + if target_iterations is None: + measured_request_steps = max(1, int(self.num_inference_steps)) + target_request_steps = int( + self.extra.get( + "warmup_target_num_inference_steps", measured_request_steps + ) + ) + target_iterations = ( + measured * max(1, target_request_steps) + measured_request_steps - 1 + ) // measured_request_steps + self.metrics.record_stage_iterations(measured, target_iterations) + def validate(self): """Initialize dependent fields after dataclass initialization.""" if getattr(self.sampling_params, "data_type", None) == DataType.ACTION: diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/base.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/base.py index 70d43d582..1defacad7 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/base.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/base.py @@ -55,6 +55,31 @@ class StageVerificationError(Exception): pass +def record_default_workload_iterations(stage, batch) -> None: + """Record how often a stage's repeated unit ran, once per warmup request. + + A stage whose loop length is a formula of the request's step count + declares it in ``default_workload_iterations``; the same formula at the + default workload's step count is the target. A stage that recorded + explicitly inside its loop (a count only known there) is left alone. + """ + metrics = batch.metrics + if metrics is None or metrics.active_stage_name is None: + return + if metrics.active_stage_name in metrics.stage_iterations: + return + measured = stage.default_workload_iterations(batch, int(batch.num_inference_steps)) + if measured is None: + return + target_steps = int( + batch.extra.get("warmup_target_num_inference_steps", batch.num_inference_steps) + ) + target = stage.default_workload_iterations(batch, target_steps) + metrics.record_stage_iterations( + max(1, int(measured)), max(1, int(measured if target is None else target)) + ) + + class PipelineStage(StageDedupMixin, ABC): """ Abstract base class for all pipeline stages. @@ -163,6 +188,14 @@ class PipelineStage(StageDedupMixin, ABC): def set_profile_stage_name(self, stage_name: str) -> None: self._profile_stage_name = stage_name + def default_workload_iterations( + self, batch: Req, num_inference_steps: int + ) -> int | None: + """How many times this stage's repeated unit runs for a request with + ``num_inference_steps`` steps; ``None`` means the stage is not repeated + (or records its count itself with ``batch.record_stage_iterations``).""" + return None + def _component_stage_name(self, stage_name: str | None = None) -> str: return stage_name or self._registered_stage_name or self.__class__.__name__ @@ -383,7 +416,14 @@ class PipelineStage(StageDedupMixin, ABC): # Execute the actual stage logic with unified profiling. previous_batch_is_warmup = self._current_batch_is_warmup + metrics = batch.metrics + warmup_metrics = metrics if batch.is_warmup else None + previous_active_stage = ( + warmup_metrics.active_stage_name if warmup_metrics is not None else None + ) self._current_batch_is_warmup = batch.is_warmup + if warmup_metrics is not None: + warmup_metrics.active_stage_name = self._component_stage_name() try: with StageProfiler( stage_name, @@ -394,7 +434,11 @@ class PipelineStage(StageDedupMixin, ABC): perf_dump_path_provided=batch.perf_dump_path is not None, ): result = self.forward(batch, server_args) + if warmup_metrics is not None: + record_default_workload_iterations(self, batch) finally: + if warmup_metrics is not None: + warmup_metrics.active_stage_name = previous_active_stage self._current_batch_is_warmup = previous_batch_is_warmup self._current_use_nvtx = False 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 5a225e731..ff9aeba24 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 @@ -145,6 +145,12 @@ class CausalDMDRealtimeCacheContext: class CausalDMDDenoisingStage(DenoisingStage): + def default_workload_iterations( + self, batch: Req, num_inference_steps: int + ) -> int | None: + # blocks x fixed DMD steps, known only once the block sizes are laid out + return None + """ Denoising stage for causal diffusion. """ @@ -218,9 +224,8 @@ class CausalDMDDenoisingStage(DenoisingStage): (scheduler.timesteps.cpu(), torch.tensor([0], dtype=torch.float32)) ) timesteps = scheduler_timesteps[1000 - timesteps] - timesteps = timesteps.to(device) logger.info("Using timesteps: %s", timesteps) - return timesteps + return timesteps.to(device) def _prepare_causal_dmd_image_kwargs( self, @@ -1254,6 +1259,9 @@ class CausalDMDDenoisingStage(DenoisingStage): block_sizes = [1] + [self.num_frames_per_block] * num_blocks start_index = 0 + total_iterations = len(block_sizes) * len(timesteps) + batch.record_stage_iterations(total_iterations, total_iterations) + def prepare_context_input(current_latents): return current_latents 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 5841ee462..277d59718 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py @@ -334,6 +334,11 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin): def role_affinity(self): return RoleType.DENOISER + def default_workload_iterations( + self, batch: Req, num_inference_steps: int + ) -> int | None: + return num_inference_steps + def __init__( self, transformer, scheduler, pipeline=None, transformer_2=None, vae=None ) -> None: 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 5cefe0a7f..2e7848571 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 @@ -25,6 +25,12 @@ logger = init_logger(__name__) class DmdDenoisingStage(DenoisingStage): + def default_workload_iterations( + self, batch: Req, num_inference_steps: int + ) -> int | None: + # a fixed distilled schedule: the same count at any requested step count + return len(self.server_args.pipeline_config.dmd_denoising_steps) + """ Denoising stage for DMD. """ 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 5989b21e5..a00599e69 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 @@ -1063,6 +1063,11 @@ class Cosmos3TimestepPreparationStage(PipelineStage): class Cosmos3DenoisingStage(PipelineStage, RolloutDenoisingMixin): + def default_workload_iterations( + self, batch: Req, num_inference_steps: int + ) -> int | None: + return num_inference_steps + """Cosmos3 denoise loop, including CFG and the parallelism modes. The UND pathway runs once and its K/V is cached per cache_key (``cond`` / 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 c8237318d..27521849e 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 @@ -776,5 +776,9 @@ class HeliosChunkedDenoisingStage(PipelineStage): # separately to avoid temporal artifacts at chunk boundaries. batch.latent_chunks = chunk_latents_list batch.latents = history_latents[:, :, -total_generated_latent_frames:] + batch.record_stage_iterations( + global_step_offset, + global_step_offset if is_enable_stage2 else None, + ) return batch 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 1065cee60..4434f23c0 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 @@ -691,6 +691,7 @@ class LingBotWorldCausalDMDDenoisingStage(CausalDMDDenoisingStage): "Ensure ImageVAEEncodingStage runs before this stage." ) ctx = self._prepare_causal_dmd_forward_context(batch, server_args) + batch.record_stage_iterations(len(ctx.timesteps), len(ctx.timesteps)) latents = ctx.latents cache_ctx = self._prepare_realtime_causal_caches(batch, server_args, ctx) diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/longlive2.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/longlive2.py index 3716d5320..3b22866dd 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/longlive2.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/longlive2.py @@ -421,6 +421,9 @@ class LongLive2CausalDenoisingStage(CausalDMDDenoisingStage): num_blocks = (t - 1) // self.num_frames_per_block block_sizes = [1] + [self.num_frames_per_block] * num_blocks + total_iterations = len(block_sizes) * len(timesteps) + batch.record_stage_iterations(total_iterations, total_iterations) + start_index = 0 self._validate_block_prompt_count(batch, block_sizes) diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/ltx_2/denoising_av.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/ltx_2/denoising_av.py index 57e09ddef..a5d59c203 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/ltx_2/denoising_av.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/ltx_2/denoising_av.py @@ -108,6 +108,12 @@ class LTX2AVDenoisingStage(LTX2DenoisingStage): class LTX2RefinementStage(LTX2AVDenoisingStage): + def default_workload_iterations( + self, batch: Req, num_inference_steps: int + ) -> int | None: + # the refiner runs its distilled sigma schedule regardless of the request's steps + return max(1, len(self.distilled_sigmas) - 1) + """Stage-2 refinement wrapper that re-noises distilled LTX latents once.""" def __init__( diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/minimax_h3/stages/denoising.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/minimax_h3/stages/denoising.py index d23e89258..a0211526f 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/minimax_h3/stages/denoising.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/minimax_h3/stages/denoising.py @@ -443,6 +443,12 @@ def _precompute_rope_cache( class MiniMaxH3DenoisingStage(DenoisingStage): + def default_workload_iterations( + self, batch: Req, num_inference_steps: int + ) -> int | None: + # one denoise per sigma interval: steps - 1 + return max(1, num_inference_steps - 1) + def __init__(self, transformer, pipeline=None) -> None: super().__init__( transformer=transformer, 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 e346b54c0..41bad0364 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 @@ -152,6 +152,11 @@ class MOVATimestepPreparationStage(PipelineStage): class MOVADenoisingStage(PipelineStage): + def default_workload_iterations( + self, batch: Req, num_inference_steps: int + ) -> int | None: + return num_inference_steps + """Run MOVA dual-tower denoising loop.""" def __init__(self, video_dit, video_dit_2, audio_dit, dual_tower_bridge, scheduler): 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 31e1f7213..87e29cf64 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 @@ -725,6 +725,8 @@ class SanaWMLTX2RefinerStage(PipelineStage): return batch batch_size = int(batch.latents.shape[0]) + total_iterations = batch_size * (len(STAGE_2_DISTILLED_SIGMA_VALUES) - 1) + batch.record_stage_iterations(total_iterations, total_iterations) prompts = self._prompts_for_batch(batch, batch_size) fps = float(getattr(batch, "fps", 16) or 16) 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 3b38547a4..de6d9f8db 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 @@ -630,6 +630,8 @@ class SanaWMStreamingDenoisingStage(CausalDMDDenoisingStage): len(explicit_sigmas), do_cfg, ) + total_iterations = num_chunks * len(explicit_sigmas) + batch.record_stage_iterations(total_iterations, total_iterations) for chunk_idx in self.progress_bar(range(num_chunks), batch=batch): chunk_kv, sink_num = self._accumulate_kv_cache( 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 6325f0dd7..46190432e 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 @@ -671,6 +671,8 @@ class SanaWMStreamingRefinerStage(SanaWMLTX2RefinerStage): ) return batch n_blocks = math.ceil(n_active / self.block_size) + total_iterations = n_blocks * (len(STAGE_2_DISTILLED_SIGMA_VALUES) - 1) + batch.record_stage_iterations(total_iterations, total_iterations) self.log_info( "SANA-WM streaming refiner: latent=%s, sink=%d, block=%d, blocks=%d, kv_max=%d, seed=%d", tuple(latents.shape), diff --git a/python/sglang/multimodal_gen/runtime/scheduler_client.py b/python/sglang/multimodal_gen/runtime/scheduler_client.py index aa0ef7c1e..27054bc4e 100644 --- a/python/sglang/multimodal_gen/runtime/scheduler_client.py +++ b/python/sglang/multimodal_gen/runtime/scheduler_client.py @@ -7,6 +7,13 @@ from typing import Any, Optional import zmq import zmq.asyncio +from sglang.multimodal_gen.runtime.entrypoints.control_requests import ( + ListLorasReq, + MergeLoraWeightsReq, + SetLoraReq, + ShutdownReq, + UnmergeLoraWeightsReq, +) from sglang.multimodal_gen.runtime.entrypoints.post_training.io_struct import ( GetWeightsChecksumReqInput, ReleaseMemoryOccupationReqInput, @@ -15,13 +22,6 @@ from sglang.multimodal_gen.runtime.entrypoints.post_training.io_struct import ( UpdateWeightFromTensorCheckerReqInput, UpdateWeightFromTensorReqInput, ) -from sglang.multimodal_gen.runtime.entrypoints.utils import ( - ListLorasReq, - MergeLoraWeightsReq, - SetLoraReq, - ShutdownReq, - UnmergeLoraWeightsReq, -) from sglang.multimodal_gen.runtime.ipc_array import materialize_file_refs from sglang.multimodal_gen.runtime.pipelines_core import Req from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch diff --git a/python/sglang/multimodal_gen/runtime/server_args/auto_tune.py b/python/sglang/multimodal_gen/runtime/server_args/auto_tune.py index 099762774..34abb57e6 100644 --- a/python/sglang/multimodal_gen/runtime/server_args/auto_tune.py +++ b/python/sglang/multimodal_gen/runtime/server_args/auto_tune.py @@ -10,6 +10,10 @@ from sglang.multimodal_gen import envs from sglang.multimodal_gen.configs.pipeline_configs.model_deployment_config import ( ModelDeploymentConfig, ) +from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType +from sglang.multimodal_gen.runtime.entrypoints.openai.realtime.registry import ( + has_realtime_model_adapter, +) from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload_components import ( LAYERWISE_OFFLOAD_ALL_COMPONENTS, LAYERWISE_OFFLOAD_DIT_GROUP, @@ -26,6 +30,49 @@ if TYPE_CHECKING: logger = init_logger(__name__) + +def auto_residency_args_skip_reason(server_args: ServerArgs) -> str | None: + """Return why args cannot use warmup-calibrated residency.""" + if envs.SGLANG_DIFFUSION_DISABLE_AUTO_RESIDENCY: + return "disabled via SGLANG_DIFFUSION_DISABLE_AUTO_RESIDENCY" + if server_args.performance_mode != "auto": + return f"performance_mode={server_args.performance_mode}" + if ( + server_args.pipeline_class_name == "LTX2TwoStagePipeline" + and server_args.ltx2_two_stage_device_mode is None + ): + return "legacy LTX-2 two-stage placement" + if server_args.ltx2_two_stage_device_mode == "original": + return "LTX-2 original two-stage placement" + if ( + server_args.warmup_mode != "server" + or server_args.disagg_role != RoleType.MONOLITHIC + ): + return "no synthetic server warmup to calibrate from" + task_type = server_args.pipeline_config.task_type + if not (task_type.is_visual_gen() or task_type.is_mesh_gen()): + return "no synthetic server warmup to calibrate from" + if not server_args.pipeline_config.supports_auto_residency: + return "pipeline does not support post-warmup residency changes" + if has_realtime_model_adapter(server_args): + return "realtime serving has no representative synthetic warmup" + if server_args.backend == "diffusers": + return "diffusers backend" + if server_args.enable_breakable_cuda_graph: + return "breakable CUDA graph captures during warmup" + if server_args.enable_torch_compile: + # Compile warmup temporarily evicts resident auxiliaries and may + # layerwise-offload the DiT, so its peak is not a serving peak. + return "torch.compile warmup uses a stripped memory layout" + if envs.SGLANG_CACHE_DIT_ENABLED: + return "cache-dit enabled" + if server_args.batching_max_size > 1: + return "dynamic batching enabled" + if not current_platform.is_cuda(): + return "requires CUDA" + return None + + PERFORMANCE_MODES = ("manual", "auto", "speed", "memory") DEFAULT_LAYERWISE_COMPONENT_ARG_NAMES = ( 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 e56433c50..c44474b99 100644 --- a/python/sglang/multimodal_gen/runtime/server_args/server_args.py +++ b/python/sglang/multimodal_gen/runtime/server_args/server_args.py @@ -454,6 +454,10 @@ class ServerArgs(DisaggServerArgsMixin): warmup_resolutions: list[str] = None warmup_num_frames: int | None = None warmup_steps: int = 1 + # JSON overrides for the representative request shape used by synthetic + # warmup and automatic residency planning. Execution remains bounded by + # warmup_steps and the server warmup frame/area caps. + warmup_sampling_params: dict[str, Any] | str | None = None disable_autocast: bool | None = None @@ -2384,6 +2388,18 @@ class ServerArgs(DisaggServerArgsMixin): default=ServerArgs.warmup_steps, help="The number of warmup steps to perform for each resolution.", ) + parser.add_argument( + "--warmup-sampling-params", + type=str, + default=ServerArgs.warmup_sampling_params, + help=( + "JSON object overriding model sampling defaults for synthetic " + "warmup and auto residency planning, for example " + '\'{"width":832,"height":480,"num_frames":9,' + '"num_inference_steps":4}\'. Warmup still applies its ' + "bounded execution caps." + ), + ) # component residency and legacy offload controls parser.add_argument( "--component-residency", diff --git a/python/sglang/multimodal_gen/runtime/server_warmup.py b/python/sglang/multimodal_gen/runtime/server_warmup.py index db5102b38..4723666db 100644 --- a/python/sglang/multimodal_gen/runtime/server_warmup.py +++ b/python/sglang/multimodal_gen/runtime/server_warmup.py @@ -7,15 +7,22 @@ from typing import Any, Awaitable, Callable from tqdm.auto import tqdm +from sglang.multimodal_gen.runtime.entrypoints.openai.realtime.registry import ( + has_realtime_model_adapter, +) from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import ( OutputBatch, Req, ) from sglang.multimodal_gen.runtime.server_args import ServerArgs +from sglang.multimodal_gen.runtime.server_args.auto_tune import ( + auto_residency_args_skip_reason, +) from sglang.multimodal_gen.runtime.utils.image_io import save_base64_image_to_path from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.warmup_request_builder import ( build_warmup_reqs, + lighten_warmup_req, should_include_warmup_image, supports_synthetic_warmup, ) @@ -74,15 +81,7 @@ def should_run_server_warmup(server_args: ServerArgs) -> bool: def is_realtime_serving(server_args: ServerArgs) -> bool: """Synthetic warmup has no realtime session state.""" - try: - from sglang.multimodal_gen.runtime.entrypoints.openai.realtime.registry import ( - get_realtime_model_adapter, - ) - - get_realtime_model_adapter(server_args) - return True - except Exception: - return False + return has_realtime_model_adapter(server_args) def should_run_synthetic_server_warmup(server_args: ServerArgs) -> bool: @@ -101,11 +100,69 @@ def should_run_explicit_client_warmup(server_args: ServerArgs) -> bool: ) +def auto_residency_skip_reason(server_args: ServerArgs) -> str | None: + """Final gate for warmup-calibrated residency placement. + + Only rules out paths the planner was not designed for; the workers + re-check per component (explicit placement, FSDP modules, custom + strategies, missing sizes) and per measurement. + """ + args_reason = auto_residency_args_skip_reason(server_args) + if args_reason is not None: + return args_reason + if not should_run_synthetic_server_warmup(server_args): + return "no synthetic server warmup to calibrate from" + return None + + +# Enough to clear a probe that overshot the card, few enough that a failure +# which is not about probe size gives up quickly instead of walking the +# workload down to nothing. +MAX_WARMUP_DEGRADE_ATTEMPTS = 3 + + +_OUT_OF_MEMORY_MARKERS = ( + "out of memory", + "outofmemory", + "cudaerrormemoryallocation", + "cublas_status_alloc_failed", + "cannot allocate memory", + "unable to allocate", +) + + +def _is_out_of_memory(error: Any) -> bool: + text = str(error).lower() + return any(marker in text for marker in _OUT_OF_MEMORY_MARKERS) + + +def _degrade_after_oom(server_args: ServerArgs, req: Req) -> Req | None: + """Next warmup probe to try after `req` ran the card out of memory. + + Only memory failures are worth retrying smaller; anything else fails the + same way at every size and should surface instead of being shrunk away. + """ + lighter = lighten_warmup_req(server_args, req) + if lighter is None: + return None + logger.warning( + "%s ran out of memory; retrying warmup at %s", + format_warmup_req(req), + format_warmup_req(lighter), + ) + return lighter + + def format_warmup_req(req_or_group: Any) -> str: req = get_first_generation_req(req_or_group) - prefix = ( - "server warmup req" if is_server_based_warmup(req_or_group) else "warmup req" - ) + if req is not None and req.extra.get("auto_residency_full_shape_probe"): + prefix = "auto residency probe" + else: + prefix = ( + "server warmup req" + if is_server_based_warmup(req_or_group) + else "warmup req" + ) if req is None: return prefix @@ -131,6 +188,8 @@ def build_client_warmup_reqs( server_args: ServerArgs, *, warmup_input_path: str | None = None, + rewarm: bool = False, + step_limit: int | None = None, ) -> list[Req]: warmup_reqs = build_warmup_reqs( server_args, @@ -143,6 +202,12 @@ def build_client_warmup_reqs( for req in warmup_reqs: if req.is_warmup: req.extra["warmup_total"] = warmup_total + if step_limit is not None: + req.num_inference_steps = min(req.num_inference_steps, step_limit) + if rewarm: + # a repeat pass after an auto-residency change: keep it out of + # the scheduler's warmup progress accounting (already at N/N) + req.extra["server_warmup_rewarm"] = True return warmup_reqs @@ -151,6 +216,8 @@ async def run_async_client_warmup( forward: Callable[[Req], Awaitable[OutputBatch]], *, fail_open: bool = False, + rewarm: bool = False, + step_limit: int | None = None, ) -> None: try: warmup_input_path = None @@ -158,9 +225,20 @@ async def run_async_client_warmup( warmup_input_path = prepare_warmup_image_path(server_args) for req in build_client_warmup_reqs( - server_args, warmup_input_path=warmup_input_path + server_args, + warmup_input_path=warmup_input_path, + rewarm=rewarm, + step_limit=step_limit, ): response = await forward(req) + for _ in range(MAX_WARMUP_DEGRADE_ATTEMPTS): + if response.error is None or not _is_out_of_memory(response.error): + break + lighter = _degrade_after_oom(server_args, req) + if lighter is None: + break + req = lighter + response = await forward(req) if response.error is not None: raise RuntimeError(response.error) except Exception: @@ -184,6 +262,14 @@ def run_sync_client_warmup( server_args, warmup_input_path=warmup_input_path ): response = forward(req) + for _ in range(MAX_WARMUP_DEGRADE_ATTEMPTS): + if response.error is None or not _is_out_of_memory(response.error): + break + lighter = _degrade_after_oom(server_args, req) + if lighter is None: + break + req = lighter + response = forward(req) if response.error is not None: raise RuntimeError(response.error) @@ -275,6 +361,18 @@ class SchedulerWarmupMixin: if not is_warmup: return + req = get_first_generation_req(req_or_group) + if req is not None and req.extra.get("server_warmup_rewarm"): + # auto-residency re-warm passes repeat already-counted requests; + # advancing the bar again would log N+1/N in CI + if output_batch.error is not None: + logger.warning( + "%s processing failed: %s", + self._format_warmup_req(req_or_group), + output_batch.error, + ) + return + server_based_warmup = is_server_based_warmup(req_or_group) self._warmup_processed += 1 self._advance_warmup_progress_bar(req_or_group, output_batch) diff --git a/python/sglang/multimodal_gen/runtime/utils/perf_logger.py b/python/sglang/multimodal_gen/runtime/utils/perf_logger.py index 297a06951..4f5feb7d4 100644 --- a/python/sglang/multimodal_gen/runtime/utils/perf_logger.py +++ b/python/sglang/multimodal_gen/runtime/utils/perf_logger.py @@ -58,6 +58,9 @@ class RequestMetrics: self.request_id = request_id self.stages: Dict[str, float] = {} self.steps: list[float] = [] + self.steps_by_stage: Dict[str, list[float]] = {} + self.stage_iterations: Dict[str, tuple[int, int]] = {} + self.active_stage_name: str | None = None self.total_duration_ms: float = 0.0 self.suppress_stage_breakdown: bool = False # memory tracking: {checkpoint_name: MemorySnapshot} @@ -77,7 +80,26 @@ class RequestMetrics: """Records the duration of a denoising step in execution order.""" if self.suppress_stage_breakdown: return - self.steps.append(duration_s * 1000) + duration_ms = duration_s * 1000 + self.steps.append(duration_ms) + if self.active_stage_name is not None: + self.steps_by_stage.setdefault(self.active_stage_name, []).append( + duration_ms + ) + + def record_stage_iterations( + self, measured_iterations: int, target_iterations: int + ) -> None: + """Record calibration and default-workload iterations for this stage.""" + if self.suppress_stage_breakdown or self.active_stage_name is None: + return + measured = max(1, int(measured_iterations)) + target = max(1, int(target_iterations)) + previous = self.stage_iterations.get(self.active_stage_name, (0, 0)) + self.stage_iterations[self.active_stage_name] = ( + previous[0] + measured, + previous[1] + target, + ) def record_memory_snapshot(self, checkpoint_name: str, snapshot: MemorySnapshot): if self.suppress_stage_breakdown: diff --git a/python/sglang/multimodal_gen/runtime/warmup_request_builder.py b/python/sglang/multimodal_gen/runtime/warmup_request_builder.py index 83e59f821..40e49eb53 100644 --- a/python/sglang/multimodal_gen/runtime/warmup_request_builder.py +++ b/python/sglang/multimodal_gen/runtime/warmup_request_builder.py @@ -6,6 +6,12 @@ first real request, without copying user traffic. It starts from the model's sampling defaults, then keeps startup bounded by choosing common low-cost resolution/frame buckets and trimming the denoising step count. +When warmup-calibrated auto residency is active, warmup uses the full default +serving shape but keeps the trimmed step count. Memory depends on the +activation shape rather than the number of repeated denoising steps, so this +directly measures placement headroom without running a full generation or an +extra stateful pipeline request. + Image models may run a tiny second step because first/last step paths often initialize different kernels or scheduler state. Video models cap frames and steps to keep startup bounded. Explicit warmup resolutions share this builder; @@ -13,17 +19,26 @@ callers send them through the scheduler client so warmup exercises the same request transport path as real generation. """ +import json from copy import copy +from dataclasses import fields, replace from typing import Any from sglang.multimodal_gen.configs.pipeline_configs.base import ModelTaskType -from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams +from sglang.multimodal_gen.configs.sample.sampling_params import ( + SamplingParams, + align_num_frames_for_num_gpus, + resolve_sequence_shard, +) from sglang.multimodal_gen.registry import get_pipeline_config_classes from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req from sglang.multimodal_gen.runtime.server_args import ( ServerArgs, is_ltx2_two_stage_pipeline_name, ) +from sglang.multimodal_gen.runtime.server_args.auto_tune import ( + auto_residency_args_skip_reason, +) from sglang.multimodal_gen.runtime.utils.common import parse_size from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger @@ -42,6 +57,11 @@ SERVER_WARMUP_MAX_VIDEO_FRAMES = 17 SERVER_WARMUP_LTX2_TWO_STAGE_MAX_VIDEO_FRAMES = 25 SERVER_WARMUP_IMAGE_STEPS = 2 SERVER_WARMUP_VIDEO_STEPS = 2 +# Two-step schedulers can have one compile-heavy step and one lower-order +# boundary step, leaving no representative steady-state timing sample. Auto +# residency extrapolates this timing to the default request, so collect four +# steps while retaining the shorter warmup for every non-planning path. +AUTO_RESIDENCY_TIMING_STEPS = 4 def get_model_sampling_defaults(server_args: ServerArgs) -> SamplingParams: @@ -50,13 +70,44 @@ def get_model_sampling_defaults(server_args: ServerArgs) -> SamplingParams: config_classes = get_pipeline_config_classes(pipeline_class_name) if config_classes is not None: _, sampling_params_cls = config_classes - return sampling_params_cls() + defaults = sampling_params_cls() + return _apply_warmup_sampling_overrides(server_args, defaults) - return SamplingParams.from_pretrained( + defaults = SamplingParams.from_pretrained( server_args.model_path, backend=server_args.backend, model_id=server_args.model_id, ) + return _apply_warmup_sampling_overrides(server_args, defaults) + + +def _apply_warmup_sampling_overrides( + server_args: ServerArgs, defaults: SamplingParams +) -> SamplingParams: + value = server_args.warmup_sampling_params + if value is None: + return defaults + if isinstance(value, str): + try: + value = json.loads(value) + except json.JSONDecodeError as error: + raise ValueError( + "--warmup-sampling-params must be a JSON object" + ) from error + if not isinstance(value, dict): + raise ValueError("--warmup-sampling-params must be a JSON object") + field_names = {item.name for item in fields(defaults)} + unknown = value.keys() - field_names + if unknown: + raise ValueError( + f"invalid --warmup-sampling-params fields: {', '.join(sorted(unknown))}" + ) + updated = copy(defaults) + for name, field_value in value.items(): + # Some model contracts intentionally expose fixed dataclass fields + # with init=False; a warmup workload still needs to mirror the request. + object.__setattr__(updated, name, field_value) + return updated def _resolve_default_warmup_resolution( @@ -221,6 +272,72 @@ def _fit_resolution_to_area( ) +def _halve_num_frames(server_args: ServerArgs, num_frames: int) -> int: + """Halve the latent frame count, keeping the model's frame arithmetic.""" + if num_frames <= 1: + return num_frames + ratio = ( + server_args.pipeline_config.vae_config.arch_config.temporal_compression_ratio + ) + if not ratio or ratio <= 1: + return max(1, num_frames // 2) + latent_frames = (num_frames - 1) // ratio + 1 + # round up: halving 5 latent frames should land on 3 (9 frames), not 2 (5) + return ((latent_frames + 1) // 2 - 1) * ratio + 1 + + +def _lighter_valid_num_frames(server_args: ServerArgs, num_frames: int) -> int: + """Largest frame count at or below half that the model's frame contract accepts. + + ``adjust_num_frames`` rounds up (LongLive2 maps 17 frames to 29), so the + halved count is walked down until it is a fixed point of the contract. + """ + halved = _halve_num_frames(server_args, num_frames) + adjust = getattr(server_args.pipeline_config, "adjust_num_frames", None) + for candidate in range(halved, 0, -1): + adjusted = adjust(candidate) if callable(adjust) else candidate + if not isinstance(adjusted, int) or isinstance(adjusted, bool): + adjusted = candidate + if adjusted == candidate: + return candidate if candidate < num_frames else num_frames + return num_frames + + +def lighten_warmup_req(server_args: ServerArgs, req: Req) -> Req | None: + """Roughly halve a warmup probe, or None once it cannot shrink further. + + Warmup peak memory tracks width * height * num_frames, so a card that could + not hold the full probe usually holds half of it while still walking the + same code path. Frames go first: they drive video activation size, and + cutting them leaves the spatial kernels at their serving shape. + """ + params = req.sampling_params + if params is None: + return None + + num_frames = params.num_frames or 1 + lighter_frames = _lighter_valid_num_frames(server_args, num_frames) + if lighter_frames < num_frames: + return _replace_warmup_workload(req, num_frames=lighter_frames) + + width = params.width + height = params.height + if not width or not height: + return None + alignment = _warmup_resolution_alignment(server_args) + lighter = _fit_resolution_to_area(width, height, width * height // 2, alignment) + # below the alignment floor the fit rounds back up; only take a real cut + if lighter[0] * lighter[1] >= width * height: + return None + return _replace_warmup_workload(req, width=lighter[0], height=lighter[1]) + + +def _replace_warmup_workload(req: Req, **changes: int) -> Req: + lighter = copy(req) + lighter.sampling_params = replace(req.sampling_params, **changes) + return lighter + + def _is_resolution_aligned(resolution: tuple[int, int], alignment: int) -> bool: width, height = resolution return width % alignment == 0 and height % alignment == 0 @@ -268,7 +385,87 @@ def _resolve_warmup_num_frames( ) warmup_num_frames = min(num_frames, frame_budget) - return server_args.pipeline_config.adjust_num_frames(warmup_num_frames) + return _apply_warmup_frame_contract( + server_args, sampling_defaults, num_frames=warmup_num_frames + ) + + +def _apply_warmup_frame_contract( + server_args: ServerArgs, sampling_defaults: SamplingParams, *, num_frames: int +) -> int: + """Re-apply the real-request frame contract to a capped warmup frame count. + + Warmup requests skip ``SamplingParams._adjust``, so the cap must run the + model frame contract itself -- without this, e.g. LongLive2's capped 17 + frames map to 5 latent frames (not divisible by its 8-frame causal block) + and every server warmup fails silently under fail-open. Pipelines that + align frames to ``num_gpus`` (rather than sharding the sequence dim) get + the same latent alignment real requests get. + """ + num_frames = server_args.pipeline_config.adjust_num_frames(num_frames) + if ( + sampling_defaults.adjust_frames + and not resolve_sequence_shard( + server_args.pipeline_config, sampling_defaults.enable_sequence_shard + ) + and server_args.num_gpus > 1 + ): + num_frames = align_num_frames_for_num_gpus( + num_frames, + num_gpus=server_args.num_gpus, + vae_config=server_args.pipeline_config.vae_config, + round_down=sampling_defaults.num_frames_round_down, + ) + return num_frames + + +def resolve_default_workload_shape( + server_args: ServerArgs, + sampling_defaults: SamplingParams, +) -> tuple[int | None, int | None, int]: + """Resolve the serving shape used by warmup-based memory planning.""" + width = sampling_defaults.width + height = sampling_defaults.height + if (width is None or height is None) and sampling_defaults.supported_resolutions: + width, height = max( + sampling_defaults.supported_resolutions, + key=lambda size: size[0] * size[1], + ) + num_frames = sampling_defaults.num_frames or 1 + if num_frames > 1: + num_frames = _apply_warmup_frame_contract( + server_args, sampling_defaults, num_frames=num_frames + ) + return width, height, num_frames + + +def _resolve_auto_residency_warmup_shape( + server_args: ServerArgs, + sampling_defaults: SamplingParams, + *, + warmup_shape: tuple[int, int, int | None], + server_based_warmup: bool, +) -> tuple[int, int, int] | None: + """Return the full serving shape when bounded warmup is smaller. + + The probe still runs only the bounded warmup step count. Denoising steps + repeat the same activation shape, so one full-shape forward measures the + placement constraints directly without paying for a full generation or + extrapolating a small-shape allocator peak. + """ + if not server_based_warmup: + return None + if auto_residency_args_skip_reason(server_args) is not None: + return None + width, height, num_frames = resolve_default_workload_shape( + server_args, sampling_defaults + ) + if width is None or height is None: + return None + target = (width, height, num_frames) + if target == warmup_shape: + return None + return target def _effective_cfg_scale(sampling_defaults: SamplingParams) -> float | None: @@ -381,11 +578,42 @@ def build_warmup_reqs( sampling_defaults, server_based_warmup=server_based_warmup, ) + auto_residency_warmup_shape = ( + _resolve_auto_residency_warmup_shape( + server_args, + sampling_defaults, + warmup_shape=(width, height, warmup_num_frames), + server_based_warmup=server_based_warmup, + ) + if warmup_resolutions is None + else None + ) + collect_auto_residency_metrics = ( + warmup_resolutions is None + and server_based_warmup + and auto_residency_args_skip_reason(server_args) is None + ) + if collect_auto_residency_metrics and sampling_defaults.num_inference_steps: + warmup_steps = min( + int(sampling_defaults.num_inference_steps), + max(warmup_steps, AUTO_RESIDENCY_TIMING_STEPS), + ) + shapes = [ + (width, height, warmup_num_frames, False) for width, height in resolutions + ] + if auto_residency_warmup_shape is not None: + # The bounded warmup runs first: its measurement lets the worker size + # the full-shape probe to the memory the card actually has left. The + # bounded shape then runs once more so the allocator pool and kernel + # caches serving starts from are shaped by a serving-sized request, + # not by the probe (the worker drops the probe's pool before it). + shapes.append((*auto_residency_warmup_shape, True)) + shapes.append(shapes[0]) # build warmup reqs warmup_reqs = [] include_warmup_image = should_include_warmup_image(server_args, server_based_warmup) - for width, height in resolutions: + for width, height, num_frames, is_probe in shapes: req_kwargs = dict( data_type=task_type.data_type(), width=width, @@ -399,7 +627,7 @@ def build_warmup_reqs( guidance_scale_2=sampling_defaults.guidance_scale_2, true_cfg_scale=sampling_defaults.true_cfg_scale, num_inference_steps=sampling_defaults.num_inference_steps, - num_frames=warmup_num_frames, + num_frames=num_frames, ) if include_warmup_image: if warmup_input_path is None: @@ -407,7 +635,13 @@ def build_warmup_reqs( "Warmup image path is required for image-input model" ) req_kwargs["prompt"] = DEFAULT_PLACEHOLDER_PROMPT - req_kwargs["image_path"] = [warmup_input_path] + default_image_path = sampling_defaults.image_path + image_count = ( + len(default_image_path) + if isinstance(default_image_path, (list, tuple)) + else 1 + ) + req_kwargs["image_path"] = [warmup_input_path] * max(1, image_count) if server_args.enable_cfg_parallel: if not req_kwargs.get("negative_prompt"): req_kwargs["negative_prompt"] = DEFAULT_PLACEHOLDER_PROMPT @@ -441,6 +675,13 @@ def build_warmup_reqs( req.extra["return_warmup_result"] = True if server_based_warmup: req.extra["server_based_warmup"] = True + if collect_auto_residency_metrics: + # Stage timers already synchronize around warmup stages. Keep + # their values for residency planning instead of discarding + # the measurements after paying that cost. + req.metrics.suppress_stage_breakdown = False + if is_probe: + req.extra["auto_residency_full_shape_probe"] = True warmup_reqs.append(req) return warmup_reqs diff --git a/python/sglang/multimodal_gen/test/server/perf_baselines/h100.json b/python/sglang/multimodal_gen/test/server/perf_baselines/h100.json index e07db5874..045e5a854 100644 --- a/python/sglang/multimodal_gen/test/server/perf_baselines/h100.json +++ b/python/sglang/multimodal_gen/test/server/perf_baselines/h100.json @@ -1568,6 +1568,7 @@ "expected_median_denoise_ms": 1010.9, "load_peak_vram_mb": 27320.0, "runtime_peak_vram_mb": 35926.0, + "warmup_peak_vram_mb": 56748.0, "load_peak_allocated_mb": 26241.0, "runtime_peak_allocated_mb": 31922.0, "estimated_full_test_time_s": 77.0 @@ -1701,6 +1702,7 @@ "expected_median_denoise_ms": 656.68, "load_peak_vram_mb": 35350.0, "runtime_peak_vram_mb": 42674.0, + "warmup_peak_vram_mb": 45782.0, "load_peak_allocated_mb": 34033.0, "runtime_peak_allocated_mb": 38569.0, "estimated_full_test_time_s": 243.0 @@ -1833,6 +1835,7 @@ "expected_median_denoise_ms": 2020.71, "load_peak_vram_mb": 6440.0, "runtime_peak_vram_mb": 21524.0, + "warmup_peak_vram_mb": 26902.0, "load_peak_allocated_mb": 6175.0, "runtime_peak_allocated_mb": 15472.0, "estimated_full_test_time_s": 204.3 @@ -2812,7 +2815,7 @@ "expected_avg_denoise_ms": 477.49, "expected_median_denoise_ms": 56.74, "load_peak_vram_mb": 34022.0, - "runtime_peak_vram_mb": 61190.0, + "runtime_peak_vram_mb": 62706.0, "load_peak_allocated_mb": 33866.0, "runtime_peak_allocated_mb": 57075.0, "estimated_full_test_time_s": 153.1 diff --git a/python/sglang/multimodal_gen/test/server/test_server_common.py b/python/sglang/multimodal_gen/test/server/test_server_common.py index 95ea8be3a..58011b481 100644 --- a/python/sglang/multimodal_gen/test/server/test_server_common.py +++ b/python/sglang/multimodal_gen/test/server/test_server_common.py @@ -465,6 +465,7 @@ class DiffusionServerBase: summary, expected_load_peak_vram_mb, expected_runtime_peak_vram_mb, + scenario.warmup_peak_vram_mb, expected_load_peak_allocated_mb=scenario.load_peak_allocated_mb, expected_runtime_peak_allocated_mb=( scenario.runtime_peak_allocated_mb @@ -539,10 +540,12 @@ class DiffusionServerBase: if os.environ.get("SGLANG_GEN_BASELINE", "0") == "1": logger.info( - "%s realtime peak VRAM baseline: load=%.0fMiB, runtime=%.0fMiB", + "%s realtime peak VRAM baseline: load=%.0fMiB, runtime=%.0fMiB, " + "warmup=%.0fMiB", case.id, summary.load_peak_vram_mb, summary.runtime_peak_vram_mb, + summary.warmup_peak_vram_mb, ) return @@ -559,6 +562,7 @@ class DiffusionServerBase: summary, scenario.load_peak_vram_mb, scenario.runtime_peak_vram_mb, + scenario.warmup_peak_vram_mb, expected_load_peak_allocated_mb=scenario.load_peak_allocated_mb, expected_runtime_peak_allocated_mb=scenario.runtime_peak_allocated_mb, ) @@ -579,6 +583,7 @@ class DiffusionServerBase: "median_denoise_ms": summary.median_denoise_ms, "load_peak_vram_mb": summary.load_peak_vram_mb, "runtime_peak_vram_mb": summary.runtime_peak_vram_mb, + "warmup_peak_vram_mb": summary.warmup_peak_vram_mb, "load_peak_allocated_mb": summary.load_peak_allocated_mb, "runtime_peak_allocated_mb": summary.runtime_peak_allocated_mb, "stage_metrics": summary.stage_metrics, @@ -680,6 +685,7 @@ class DiffusionServerBase: { "load_peak_vram_mb": round(summary.load_peak_vram_mb, 2), "runtime_peak_vram_mb": round(summary.runtime_peak_vram_mb, 2), + "warmup_peak_vram_mb": round(summary.warmup_peak_vram_mb, 2), "load_peak_allocated_mb": round(summary.load_peak_allocated_mb, 2), "runtime_peak_allocated_mb": round( summary.runtime_peak_allocated_mb, 2 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 683e5a443..086929e30 100644 --- a/python/sglang/multimodal_gen/test/server/test_server_utils.py +++ b/python/sglang/multimodal_gen/test/server/test_server_utils.py @@ -405,7 +405,7 @@ class ServerManager: "--log-level=debug", ] if self.extra_args.strip(): - command.extend(self.extra_args.strip().split()) + command.extend(shlex.split(self.extra_args)) access_log_exclude_flag = "--uvicorn-access-log-exclude-prefixes" if not any(arg.startswith(access_log_exclude_flag) for arg in command): command.extend(["--uvicorn-access-log-exclude-prefixes", "/health"]) @@ -590,6 +590,7 @@ class PerformanceValidator: summary: PerformanceSummary, expected_load_peak_vram_mb: float, expected_runtime_peak_vram_mb: float, + expected_warmup_peak_vram_mb: float | None = None, expected_load_peak_allocated_mb: float | None = None, expected_runtime_peak_allocated_mb: float | None = None, ) -> None: @@ -611,6 +612,16 @@ class PerformanceValidator: expected_allocated=expected_runtime_peak_allocated_mb, tolerance=self.tolerances.runtime_peak_vram, ) + # the full-shape warmup probe keeps its own budget, separate from serving + if expected_warmup_peak_vram_mb is not None and summary.warmup_peak_vram_mb > 0: + self._assert_le( + "Warmup Peak VRAM", + summary.warmup_peak_vram_mb, + expected_warmup_peak_vram_mb, + self.tolerances.runtime_peak_vram, + min_abs_tolerance=128.0, + unit=" MiB", + ) def _assert_peak_vram( self, diff --git a/python/sglang/multimodal_gen/test/server/testcase_configs.py b/python/sglang/multimodal_gen/test/server/testcase_configs.py index 8ac53328b..5c8cd4f3e 100644 --- a/python/sglang/multimodal_gen/test/server/testcase_configs.py +++ b/python/sglang/multimodal_gen/test/server/testcase_configs.py @@ -117,6 +117,9 @@ class ScenarioConfig: estimated_full_test_time_s: float | None = None load_peak_vram_mb: float | None = None runtime_peak_vram_mb: float | None = None + # Peak of the warmup calibration probe (the default workload's full shape + # under the load-safe placement); None skips the check until a baseline exists. + warmup_peak_vram_mb: float | None = None # Allocated peaks; when present they are the enforced VRAM figure and the # reserved peaks above are reported only (reserved tracks pool history). load_peak_allocated_mb: float | None = None @@ -140,6 +143,7 @@ class ScenarioConfig: estimated_full_test_time_s=optional_float("estimated_full_test_time_s"), load_peak_vram_mb=optional_float("load_peak_vram_mb"), runtime_peak_vram_mb=optional_float("runtime_peak_vram_mb"), + warmup_peak_vram_mb=optional_float("warmup_peak_vram_mb"), load_peak_allocated_mb=optional_float("load_peak_allocated_mb"), runtime_peak_allocated_mb=optional_float("runtime_peak_allocated_mb"), load_peak_host_anon_mb=optional_float("load_peak_host_anon_mb"), @@ -451,6 +455,7 @@ class PerformanceSummary: all_denoise_steps: dict[int, float] load_peak_vram_mb: float = 0.0 runtime_peak_vram_mb: float = 0.0 + warmup_peak_vram_mb: float = 0.0 load_peak_allocated_mb: float = 0.0 runtime_peak_allocated_mb: float = 0.0 load_peak_host_anon_mb: float = 0.0 @@ -490,6 +495,9 @@ class PerformanceSummary: runtime_peak_vram_mb = float( record.memory_snapshots.get("runtime_peak", {}).get("peak_reserved_mb", 0.0) ) + warmup_peak_vram_mb = float( + record.memory_snapshots.get("warmup_peak", {}).get("peak_reserved_mb", 0.0) + ) load_peak_allocated_mb = float( record.memory_snapshots.get("load_peak", {}).get("peak_allocated_mb", 0.0) ) @@ -517,6 +525,7 @@ class PerformanceSummary: all_denoise_steps=per_step, load_peak_vram_mb=load_peak_vram_mb, runtime_peak_vram_mb=runtime_peak_vram_mb, + warmup_peak_vram_mb=warmup_peak_vram_mb, load_peak_allocated_mb=load_peak_allocated_mb, runtime_peak_allocated_mb=runtime_peak_allocated_mb, load_peak_host_anon_mb=load_peak_host_anon_mb, 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 3f886a053..46d22bfea 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 @@ -35,6 +35,7 @@ from sglang.multimodal_gen.runtime.entrypoints.openai.realtime.realtime_output_a ) from sglang.multimodal_gen.runtime.entrypoints.openai.realtime.registry import ( get_realtime_model_adapter, + has_realtime_model_adapter, ) from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.lingbot_world.lingbot_world_causal_denoising import ( @@ -976,6 +977,7 @@ def test_realtime_input_validation_reuses_generator_across_chunks(): def test_realtime_registry_resolves_lingbot_adapter(): server_args = SimpleNamespace(pipeline_config=LingBotWorldCausalDMDConfig()) + assert has_realtime_model_adapter(server_args) adapter = get_realtime_model_adapter(server_args) assert isinstance(adapter, lingbot_realtime.LingBotWorldRealtimeAdapter) diff --git a/python/sglang/multimodal_gen/test/unit/test_auto_residency.py b/python/sglang/multimodal_gen/test/unit/test_auto_residency.py new file mode 100644 index 000000000..3c2014cf7 --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/test_auto_residency.py @@ -0,0 +1,921 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Unit tests for warmup-calibrated auto residency adjustment.""" + +from types import SimpleNamespace + +import pytest + +from sglang.multimodal_gen.configs.pipeline_configs.base import ModelTaskType +from sglang.multimodal_gen.configs.pipeline_configs.longlive2 import LongLive2T2VConfig +from sglang.multimodal_gen.runtime.managers.memory_managers.auto_residency import ( + ACTIVATION_EXTRAPOLATION_MARGIN, + GIB_BYTES, + DefaultWorkload, + WarmupMemoryRecord, + estimate_default_workload_peak_bytes, + estimate_default_workload_timing, + estimate_layerwise_layer_uses, + estimate_workload_phase_peaks, + resolve_measured_default_workload, +) +from sglang.multimodal_gen.runtime.warmup_request_builder import ( + SERVER_WARMUP_MAX_VIDEO_FRAMES, + _resolve_auto_residency_warmup_shape, + _resolve_warmup_num_frames, +) + + +def _record( + *, + width=832, + height=480, + num_frames=17, + baseline_gib=10, + peak_gib=12, + peak_reserved_gib=0, + succeeded=True, + num_inference_steps=1, + total_duration_ms=0.0, + stage_duration_ms=None, + step_duration_ms=(), + step_duration_ms_by_stage=None, + stage_iterations=None, + phase_active_components=None, + phase_used_components=None, + phase_full_weight_transition_components=None, + layerwise_layer_uses=None, + layerwise_layer_uses_by_stage=None, +) -> WarmupMemoryRecord: + return WarmupMemoryRecord( + width=width, + height=height, + num_frames=num_frames, + baseline_allocated_bytes=baseline_gib * GIB_BYTES, + peak_allocated_bytes=peak_gib * GIB_BYTES, + succeeded=succeeded, + peak_reserved_bytes=peak_reserved_gib * GIB_BYTES, + num_inference_steps=num_inference_steps, + total_duration_ms=total_duration_ms, + stage_duration_ms=stage_duration_ms or {}, + step_duration_ms=step_duration_ms, + step_duration_ms_by_stage=step_duration_ms_by_stage or {}, + stage_iterations=stage_iterations or {}, + phase_active_components=phase_active_components or {}, + phase_used_components=phase_used_components or {}, + phase_full_weight_transition_components=( + phase_full_weight_transition_components or {} + ), + layerwise_layer_uses=layerwise_layer_uses or {}, + layerwise_layer_uses_by_stage=layerwise_layer_uses_by_stage or {}, + ) + + +class TestEstimateDefaultWorkloadTiming: + def test_scales_only_denoising_with_target_steps(self): + record = _record( + num_inference_steps=2, + total_duration_ms=1_200, + stage_duration_ms={ + "TextEncodingStage": 100, + "DenoisingStage": 1_000, + }, + phase_active_components={ + "0:TextEncodingStage:use:text_encoder": ("text_encoder",), + "1:DenoisingStage:use:transformer": ("transformer",), + }, + ) + + request_ns, stage_ns, component_stages = estimate_default_workload_timing( + records=[record], + target_units=record.workload_units(), + target_num_inference_steps=40, + ) + + assert stage_ns == { + "TextEncodingStage": 100_000_000, + "DenoisingStage": 20_000_000_000, + } + assert request_ns == 20_200_000_000 + assert component_stages == { + "text_encoder": ("TextEncodingStage",), + "transformer": ("DenoisingStage",), + } + + def test_uses_steady_step_instead_of_scaling_first_step_setup(self): + record = _record( + num_inference_steps=2, + total_duration_ms=720, + stage_duration_ms={ + "TextEncodingStage": 100, + "DenoisingStage": 620, + }, + step_duration_ms=(500, 100), + phase_active_components={ + "0:TextEncodingStage:use:text_encoder": ("text_encoder",), + "1:DenoisingStage:use:transformer": ("transformer",), + }, + ) + + request_ns, stage_ns, _ = estimate_default_workload_timing( + records=[record], + target_units=record.workload_units(), + target_num_inference_steps=10, + ) + + assert stage_ns["DenoisingStage"] == 1_020_000_000 + assert request_ns == 1_120_000_000 + + def test_scales_standard_and_nonstandard_denoising_stages_together(self): + record = _record( + num_inference_steps=2, + total_duration_ms=1_200, + stage_duration_ms={ + "TextEncodingStage": 100, + "DenoisingStage": 500, + "CustomDenoisingStage": 600, + }, + phase_active_components={ + "0:TextEncodingStage:use:text_encoder": ("text_encoder",), + "1:DenoisingStage:use:transformer": ("transformer",), + "2:CustomDenoisingStage:use:custom_refiner": ("custom_refiner",), + }, + ) + + request_ns, stage_ns, _ = estimate_default_workload_timing( + records=[record], + target_units=record.workload_units(), + target_num_inference_steps=10, + ) + + assert stage_ns == { + "TextEncodingStage": 100_000_000, + "DenoisingStage": 2_500_000_000, + "CustomDenoisingStage": 3_000_000_000, + } + assert request_ns == 5_600_000_000 + + def test_uses_each_stage_iteration_target(self): + record = _record( + num_inference_steps=4, + total_duration_ms=1_100, + stage_duration_ms={ + "ShapeDenoisingStage": 500, + "PaintStage": 600, + }, + step_duration_ms_by_stage={ + "ShapeDenoisingStage": (200, 100, 100, 100), + }, + stage_iterations={ + "ShapeDenoisingStage": (4, 50), + "PaintStage": (4, 30), + }, + phase_active_components={ + "0:ShapeDenoisingStage:use:hy3dshape_model": ("hy3dshape_model",), + "1:PaintStage:use:paint_transformer": ("paint_transformer",), + }, + ) + + request_ns, stage_ns, _ = estimate_default_workload_timing( + records=[record], + target_units=record.workload_units(), + target_num_inference_steps=50, + ) + + assert stage_ns == { + "ShapeDenoisingStage": 5_000_000_000, + "PaintStage": 4_500_000_000, + } + assert request_ns == 9_500_000_000 + + +class TestEstimateLayerwiseLayerUses: + def test_scales_repeated_groups_but_not_one_shot_groups(self): + record = _record( + num_inference_steps=2, + layerwise_layer_uses={ + "transformer": { + "token_refiner.blocks": (1, 1), + "blocks": (2, 2, 2), + }, + "vae": { + "encoder.down_blocks": (0, 0), + "decoder.up_blocks": (1, 1), + }, + }, + ) + + uses = estimate_layerwise_layer_uses( + records=[record], + target_units=record.workload_units(), + target_num_inference_steps=10, + ) + + assert uses["transformer"] == { + "token_refiner.blocks": (1, 1), + "blocks": (10, 10, 10), + } + assert uses["vae"] == { + "encoder.down_blocks": (0, 0), + "decoder.up_blocks": (1, 1), + } + + def test_stage_role_scales_nonstandard_denoiser_but_not_encoder_repeats(self): + record = _record( + num_inference_steps=2, + phase_used_components={ + "0:TextEncodingStage:use:custom_encoder": ("custom_encoder",), + "1:CustomDenoisingStage:use:custom_refiner": ("custom_refiner",), + }, + layerwise_layer_uses={ + "custom_encoder": {"layers": (2, 2)}, + "custom_refiner": {"blocks": (2, 2)}, + }, + ) + + uses = estimate_layerwise_layer_uses( + records=[record], + target_units=record.workload_units(), + target_num_inference_steps=10, + ) + + assert uses["custom_encoder"]["layers"] == (2, 2) + assert uses["custom_refiner"]["blocks"] == (10, 10) + + def test_scales_each_stage_calls_with_its_own_iteration_target(self): + record = _record( + num_inference_steps=4, + phase_used_components={ + "0:ShapeStage:use:transformer": ("transformer",), + "1:PaintStage:use:transformer": ("transformer",), + }, + stage_iterations={ + "ShapeStage": (4, 50), + "PaintStage": (4, 30), + }, + layerwise_layer_uses={ + "transformer": {"blocks": (9, 9), "one_shot": (2, 2)} + }, + layerwise_layer_uses_by_stage={ + "ShapeStage": {"transformer": {"blocks": (4, 4), "one_shot": (1, 1)}}, + "PaintStage": {"transformer": {"blocks": (4, 4), "one_shot": (1, 1)}}, + }, + ) + + uses = estimate_layerwise_layer_uses( + records=[record], + target_units=record.workload_units(), + target_num_inference_steps=50, + ) + + # one untracked one-shot call + 50 shape calls + 30 paint calls + assert uses["transformer"]["blocks"] == (81, 81) + assert uses["transformer"]["one_shot"] == (2, 2) + + +class TestEstimateDefaultWorkloadPeak: + def test_same_shape_uses_measured_peak(self): + record = _record() + estimate = estimate_default_workload_peak_bytes( + records=[record], target_units=record.workload_units() + ) + assert estimate == record.peak_allocated_bytes + + def test_unknown_target_disables_estimation(self): + # An unknown target would silently equate the capped warmup peak with + # the serving peak and promote with no margin at all. + record = _record() + estimate = estimate_default_workload_peak_bytes( + records=[record], target_units=None + ) + assert estimate is None + + def test_single_capped_record_scales_only_the_activation_part(self): + # Warmup capped to 832x480x17; Wan-class default is 704x1280x121. + record = _record() + target_units = 704 * 1280 * 121 + ratio = target_units / record.workload_units() + assert ratio > 10 # the cap ratio this formula exists for + + estimate = estimate_default_workload_peak_bytes( + records=[record], target_units=target_units + ) + activation = record.peak_allocated_bytes - record.baseline_allocated_bytes + expected = record.baseline_allocated_bytes + int( + activation * ratio * ACTIVATION_EXTRAPOLATION_MARGIN + ) + assert estimate == expected + # Scaling the whole peak would inflate the estimate by the resident + # weights times the cap ratio and adjustment would never trigger. + naive = int(record.peak_allocated_bytes * ratio) + assert estimate < naive + + def test_two_point_fit_separates_constant_from_linear(self): + # Two calibration sizes let the estimator measure the slope instead + # of assuming everything above the baseline scales. Under offload the + # baseline is nearly empty, so the single-point formula would scale + # the whole peak (~x13 here); the fit extrapolates only the measured + # linear part. + small = _record(num_frames=9, baseline_gib=1, peak_gib=14) + large = _record(num_frames=17, baseline_gib=1, peak_gib=16) + target_units = 1280 * 720 * 81 + + estimate = estimate_default_workload_peak_bytes( + records=[small, large], target_units=target_units + ) + slope = (large.peak_allocated_bytes - small.peak_allocated_bytes) / ( + large.workload_units() - small.workload_units() + ) + constant = large.peak_allocated_bytes - slope * large.workload_units() + expected = int( + constant + slope * target_units * ACTIVATION_EXTRAPOLATION_MARGIN + ) + assert estimate == expected + single_point = estimate_default_workload_peak_bytes( + records=[large], target_units=target_units + ) + assert estimate < single_point + + def test_negative_slope_falls_back_to_single_point_formula(self): + small = _record(num_frames=9, peak_gib=16) + large = _record(num_frames=17, peak_gib=14) + target_units = 1280 * 720 * 81 + estimate = estimate_default_workload_peak_bytes( + records=[small, large], target_units=target_units + ) + fallback = max( + estimate_default_workload_peak_bytes( + records=[record], target_units=target_units + ) + for record in (small, large) + ) + assert estimate == fallback + + def test_weight_floor_prevents_scaling_constant_component_memory(self): + small = WarmupMemoryRecord( + width=832, + height=480, + num_frames=9, + baseline_allocated_bytes=GIB_BYTES, + peak_allocated_bytes=int(29.5 * GIB_BYTES), + succeeded=True, + ) + large = WarmupMemoryRecord( + width=832, + height=480, + num_frames=17, + baseline_allocated_bytes=GIB_BYTES, + peak_allocated_bytes=int(29.4 * GIB_BYTES), + succeeded=True, + ) + target_units = 1024 * 1024 * 81 + + estimate = estimate_default_workload_peak_bytes( + records=[small, large], + target_units=target_units, + constant_weight_bytes=int(29.4 * GIB_BYTES), + ) + fallback_without_weights = estimate_default_workload_peak_bytes( + records=[small, large], target_units=target_units + ) + + assert estimate < 50 * GIB_BYTES + assert estimate < fallback_without_weights + + def test_covering_measurement_bounds_the_target(self): + capped = _record(num_frames=17, peak_gib=12) + full = _record(width=1280, height=720, num_frames=81, peak_gib=30) + estimate = estimate_default_workload_peak_bytes( + records=[capped, full], target_units=1280 * 720 * 81 + ) + assert estimate == full.peak_allocated_bytes + + def test_multiple_records_take_the_max(self): + low = _record(peak_gib=12) + high = _record(peak_gib=20) + estimate = estimate_default_workload_peak_bytes( + records=[low, high], target_units=low.workload_units() + ) + assert estimate == high.peak_allocated_bytes + + def test_failure_at_the_target_size_disables_estimation(self): + good = _record(num_frames=9, peak_gib=8) + failed = _record(num_frames=17, succeeded=False) + assert ( + estimate_default_workload_peak_bytes( + records=[good, failed], target_units=failed.workload_units() + ) + is None + ) + + def test_failure_below_the_target_size_disables_estimation(self): + good = _record(num_frames=9, peak_gib=8) + failed = _record(num_frames=17, succeeded=False) + assert ( + estimate_default_workload_peak_bytes( + records=[good, failed], target_units=failed.workload_units() * 2 + ) + is None + ) + + def test_failure_above_the_target_size_is_dropped(self): + good = _record(num_frames=9, peak_gib=8) + failed = _record(num_frames=81, succeeded=False) + estimate = estimate_default_workload_peak_bytes( + records=[good, failed], target_units=good.workload_units() + ) + assert estimate == good.peak_allocated_bytes + + def test_no_records_disables_estimation(self): + assert ( + estimate_default_workload_peak_bytes(records=[], target_units=None) is None + ) + + def test_phase_estimation_preserves_active_component_membership(self): + small = _record(num_frames=9, peak_gib=30) + large = _record(num_frames=17, peak_gib=32) + small = WarmupMemoryRecord( + width=small.width, + height=small.height, + num_frames=small.num_frames, + baseline_allocated_bytes=small.baseline_allocated_bytes, + peak_allocated_bytes=small.peak_allocated_bytes, + succeeded=small.succeeded, + phase_peak_allocated_bytes={"denoise": 30 * GIB_BYTES}, + phase_active_components={"denoise": ("transformer",)}, + ) + large = WarmupMemoryRecord( + width=large.width, + height=large.height, + num_frames=large.num_frames, + baseline_allocated_bytes=large.baseline_allocated_bytes, + peak_allocated_bytes=large.peak_allocated_bytes, + succeeded=large.succeeded, + phase_peak_allocated_bytes={"denoise": 32 * GIB_BYTES}, + phase_active_components={"denoise": ("transformer",)}, + ) + + peaks, active, used, _ = estimate_workload_phase_peaks( + records=[small, large], + target_units=832 * 480 * 81, + component_weight_bytes={"transformer": 28 * GIB_BYTES}, + ) + + assert peaks["denoise"] >= 32 * GIB_BYTES + assert active == {"denoise": ("transformer",)} + assert used == active + + def test_phase_estimation_uses_allocated_peak(self): + record = WarmupMemoryRecord( + width=1024, + height=1024, + num_frames=1, + baseline_allocated_bytes=5 * GIB_BYTES, + peak_allocated_bytes=12 * GIB_BYTES, + succeeded=True, + phase_peak_allocated_bytes={"denoise": 11 * GIB_BYTES}, + phase_active_components={"denoise": ("transformer",)}, + ) + + peaks, _, _, _ = estimate_workload_phase_peaks( + records=[record], + target_units=record.workload_units(), + component_weight_bytes={"transformer": 10 * GIB_BYTES}, + ) + + assert peaks["denoise"] == 11 * GIB_BYTES + + def test_phase_estimation_preserves_full_weight_transition_components(self): + record = WarmupMemoryRecord( + width=1024, + height=1024, + num_frames=1, + baseline_allocated_bytes=2 * GIB_BYTES, + peak_allocated_bytes=4 * GIB_BYTES, + succeeded=True, + phase_peak_allocated_bytes={"lora_switch": 4 * GIB_BYTES}, + phase_full_weight_transition_components={"lora_switch": ("transformer",)}, + ) + + peaks, _, _, transitions = estimate_workload_phase_peaks( + records=[record], + target_units=record.workload_units(), + component_weight_bytes={"transformer": 2 * GIB_BYTES}, + ) + + assert peaks == {"lora_switch": 4 * GIB_BYTES} + assert transitions == {"lora_switch": ("transformer",)} + + def test_phase_estimation_prefers_target_layout_over_smaller_warmup(self): + small = WarmupMemoryRecord( + width=256, + height=256, + num_frames=9, + baseline_allocated_bytes=2 * GIB_BYTES, + peak_allocated_bytes=30 * GIB_BYTES, + succeeded=True, + phase_peak_allocated_bytes={"denoise": 30 * GIB_BYTES}, + phase_active_components={"denoise": ()}, + ) + target = WarmupMemoryRecord( + width=768, + height=512, + num_frames=25, + baseline_allocated_bytes=2 * GIB_BYTES, + peak_allocated_bytes=48 * GIB_BYTES, + succeeded=True, + phase_peak_allocated_bytes={"denoise": 48 * GIB_BYTES}, + phase_active_components={"denoise": ("transformer",)}, + ) + + peaks, active, used, _ = estimate_workload_phase_peaks( + records=[small, target], + target_units=target.workload_units(), + component_weight_bytes={"transformer": 40 * GIB_BYTES}, + ) + + assert peaks == {"denoise": 48 * GIB_BYTES} + assert active == {"denoise": ("transformer",)} + assert used == active + + def test_phase_estimation_keeps_distinct_active_layouts_separate(self): + transformer_phase = WarmupMemoryRecord( + width=768, + height=512, + num_frames=25, + baseline_allocated_bytes=2 * GIB_BYTES, + peak_allocated_bytes=48 * GIB_BYTES, + succeeded=True, + phase_peak_allocated_bytes={"denoise": 48 * GIB_BYTES}, + phase_active_components={"denoise": ("transformer",)}, + ) + encoder_phase = WarmupMemoryRecord( + width=768, + height=512, + num_frames=25, + baseline_allocated_bytes=2 * GIB_BYTES, + peak_allocated_bytes=20 * GIB_BYTES, + succeeded=True, + phase_peak_allocated_bytes={"denoise": 20 * GIB_BYTES}, + phase_active_components={"denoise": ("text_encoder",)}, + ) + + peaks, active, used, _ = estimate_workload_phase_peaks( + records=[transformer_phase, encoder_phase], + target_units=transformer_phase.workload_units(), + component_weight_bytes={ + "transformer": 40 * GIB_BYTES, + "text_encoder": 10 * GIB_BYTES, + }, + ) + + assert peaks == { + "denoise:layout:0": 20 * GIB_BYTES, + "denoise:layout:1": 48 * GIB_BYTES, + } + assert active == { + "denoise:layout:0": ("text_encoder",), + "denoise:layout:1": ("transformer",), + } + assert used == active + + +class TestResolveMeasuredDefaultWorkload: + def test_uses_effective_warmup_resolution_for_implicit_image_size(self): + workload = DefaultWorkload( + width=None, + height=None, + num_frames=1, + num_inference_steps=40, + ) + + resolved = resolve_measured_default_workload( + workload, + [ + _record(width=512, height=512), + _record(width=1024, height=1024), + ], + ) + + assert resolved == DefaultWorkload( + width=1024, + height=1024, + num_frames=1, + num_inference_steps=40, + ) + + def test_keeps_default_frames_when_warmup_caps_video(self): + workload = DefaultWorkload( + width=None, + height=None, + num_frames=81, + num_inference_steps=30, + ) + + resolved = resolve_measured_default_workload( + workload, [_record(width=832, height=480, num_frames=17)] + ) + + assert resolved.num_frames == 81 + + def test_does_not_replace_explicit_default_shape(self): + workload = DefaultWorkload( + width=1280, + height=720, + num_frames=81, + num_inference_steps=30, + ) + + assert ( + resolve_measured_default_workload( + workload, [_record(width=512, height=512)] + ) + is workload + ) + + +class TestWarmupFrameAdjustment: + def _server_args(self, *, bcg: bool = False, num_gpus: int = 1) -> SimpleNamespace: + return SimpleNamespace( + pipeline_config=LongLive2T2VConfig(), + pipeline_class_name=None, + enable_breakable_cuda_graph=bcg, + num_gpus=num_gpus, + ) + + def _defaults(self) -> SimpleNamespace: + return SimpleNamespace( + num_frames=61, + adjust_frames=True, + enable_sequence_shard=None, + num_frames_round_down=False, + ) + + def test_capped_frames_keep_the_model_frame_contract(self): + # LongLive2 default 61 frames is capped to 17, whose 5 latent frames + # break the 8-frame causal block; the builder must re-align to 29. + assert SERVER_WARMUP_MAX_VIDEO_FRAMES == 17 + num_frames = _resolve_warmup_num_frames( + self._server_args(), self._defaults(), server_based_warmup=True + ) + assert num_frames == 29 + + def test_bcg_keeps_full_serving_frames(self): + num_frames = _resolve_warmup_num_frames( + self._server_args(bcg=True), self._defaults(), server_based_warmup=True + ) + assert num_frames == 61 + + def test_non_server_warmup_keeps_default_frames(self): + num_frames = _resolve_warmup_num_frames( + self._server_args(), self._defaults(), server_based_warmup=False + ) + assert num_frames == 61 + + def test_capped_frames_get_the_gpu_alignment_real_requests_get(self): + # a frame-aligning pipeline (no sequence shard) on multiple GPUs: + # 17 frames -> 5 latent frames -> ceil to 6 latents on 2 GPUs -> 21 + args = SimpleNamespace( + pipeline_config=SimpleNamespace( + task_type=ModelTaskType.T2V, + adjust_num_frames=lambda n: n, + vae_config=SimpleNamespace( + use_temporal_scaling_frames=True, + arch_config=SimpleNamespace(temporal_compression_ratio=4), + ), + ), + pipeline_class_name=None, + enable_breakable_cuda_graph=False, + num_gpus=2, + ) + defaults = SimpleNamespace( + num_frames=81, + adjust_frames=True, + enable_sequence_shard=None, + num_frames_round_down=False, + ) + num_frames = _resolve_warmup_num_frames( + args, defaults, server_based_warmup=True + ) + assert num_frames == 21 + + +class TestAutoResidencyWarmupShape: + def _patch_gate(self, monkeypatch, reason: str | None = None) -> None: + monkeypatch.setattr( + "sglang.multimodal_gen.runtime.warmup_request_builder.auto_residency_args_skip_reason", + lambda _args: reason, + ) + + def _wan_like_args(self) -> SimpleNamespace: + return SimpleNamespace( + pipeline_config=SimpleNamespace( + task_type=ModelTaskType.T2V, + adjust_num_frames=lambda n: n, + ), + num_gpus=1, + ) + + def _defaults( + self, + num_frames: int, + *, + width: int | None = 1280, + height: int | None = 720, + supported_resolutions=None, + ) -> SimpleNamespace: + return SimpleNamespace( + width=width, + height=height, + num_frames=num_frames, + supported_resolutions=supported_resolutions, + adjust_frames=True, + enable_sequence_shard=None, + num_frames_round_down=False, + ) + + def test_capped_video_gets_a_full_shape_probe(self, monkeypatch): + self._patch_gate(monkeypatch) + probe = _resolve_auto_residency_warmup_shape( + self._wan_like_args(), + self._defaults(81), + warmup_shape=(832, 480, 17), + server_based_warmup=True, + ) + assert probe == (1280, 720, 81) + + def test_matching_warmup_needs_no_probe(self, monkeypatch): + self._patch_gate(monkeypatch) + probe = _resolve_auto_residency_warmup_shape( + self._wan_like_args(), + self._defaults(17, width=832, height=480), + warmup_shape=(832, 480, 17), + server_based_warmup=True, + ) + assert probe is None + + def test_unknown_target_resolution_skips_probe(self, monkeypatch): + self._patch_gate(monkeypatch) + probe = _resolve_auto_residency_warmup_shape( + self._wan_like_args(), + self._defaults(81, width=None, height=None), + warmup_shape=(832, 480, 17), + server_based_warmup=True, + ) + assert probe is None + + def test_skip_gate_disables_probe(self, monkeypatch): + # Full-shape calibration must share the adjustment's own gate (kill + # switch, quantized, manual, ...). + self._patch_gate(monkeypatch, reason="performance_mode=manual") + probe = _resolve_auto_residency_warmup_shape( + self._wan_like_args(), + self._defaults(81), + warmup_shape=(832, 480, 17), + server_based_warmup=True, + ) + assert probe is None + + def test_supported_resolution_fills_missing_target_size(self, monkeypatch): + self._patch_gate(monkeypatch) + probe = _resolve_auto_residency_warmup_shape( + self._wan_like_args(), + self._defaults( + 81, + width=None, + height=None, + supported_resolutions=[(832, 480), (1024, 1024)], + ), + warmup_shape=(832, 480, 17), + server_based_warmup=True, + ) + assert probe == (1024, 1024, 81) + + +class TestAutoResidencySkipReason: + def _base_args(self, **overrides) -> SimpleNamespace: + args = SimpleNamespace( + performance_mode="auto", + warmup_mode="server", + warmup_resolutions=None, + disagg_role="monolithic", + backend="sglang", + enable_breakable_cuda_graph=False, + enable_torch_compile=False, + batching_max_size=1, + dp_size=1, + ulysses_degree=1, + use_fsdp_inference=False, + quantization=None, + component_quantizations={}, + transformer_weights_path=None, + nunchaku_config=None, + direct_gpu_weight_loading=False, + ltx2_two_stage_device_mode=None, + pipeline_class_name=None, + pipeline_config=SimpleNamespace( + task_type=ModelTaskType.T2V, + supports_auto_residency=True, + ), + ) + for key, value in overrides.items(): + setattr(args, key, value) + return args + + def _skip_reason(self, args): + from sglang.multimodal_gen.runtime.server_warmup import ( + auto_residency_skip_reason, + ) + + return auto_residency_skip_reason(args) + + def test_env_kill_switch(self, monkeypatch): + monkeypatch.setenv("SGLANG_DIFFUSION_DISABLE_AUTO_RESIDENCY", "1") + reason = self._skip_reason(self._base_args()) + assert ( + reason is not None and "SGLANG_DIFFUSION_DISABLE_AUTO_RESIDENCY" in reason + ) + + def test_manual_performance_mode(self, monkeypatch): + monkeypatch.delenv("SGLANG_DIFFUSION_DISABLE_AUTO_RESIDENCY", raising=False) + reason = self._skip_reason(self._base_args(performance_mode="manual")) + assert reason == "performance_mode=manual" + + @pytest.mark.parametrize( + "overrides, expected_fragment", + [ + ({"warmup_mode": "request"}, "server warmup"), + ({"disagg_role": "denoiser"}, "server warmup"), + ({"backend": "diffusers"}, "diffusers"), + ( + {"ltx2_two_stage_device_mode": "original"}, + "LTX-2 original two-stage placement", + ), + ( + {"pipeline_class_name": "LTX2TwoStagePipeline"}, + "legacy LTX-2 two-stage placement", + ), + ({"enable_breakable_cuda_graph": True}, "CUDA graph"), + # compile warmup strips the memory layout (layerwise DiT + + # resident aux components on CPU): its peaks are not serving peaks + ({"enable_torch_compile": True}, "stripped memory layout"), + ({"batching_max_size": 4}, "batching"), + ( + { + "pipeline_config": SimpleNamespace( + task_type=ModelTaskType.T2V, + supports_auto_residency=False, + ) + }, + "post-warmup residency changes", + ), + ], + ) + def test_excluded_paths(self, monkeypatch, overrides, expected_fragment): + monkeypatch.delenv("SGLANG_DIFFUSION_DISABLE_AUTO_RESIDENCY", raising=False) + monkeypatch.delenv("SGLANG_CACHE_DIT_ENABLED", raising=False) + reason = self._skip_reason(self._base_args(**overrides)) + assert reason is not None and expected_fragment in reason + + @pytest.mark.parametrize( + "overrides", + [ + {"quantization": "fp8"}, + {"component_quantizations": {"image_encoder": "fp8"}}, + {"transformer_weights_path": "/x.safetensors"}, + {"direct_gpu_weight_loading": True}, + ], + ) + def test_fixed_loading_paths_still_calibrate_other_components( + self, monkeypatch, overrides + ): + monkeypatch.delenv("SGLANG_DIFFUSION_DISABLE_AUTO_RESIDENCY", raising=False) + monkeypatch.delenv("SGLANG_CACHE_DIT_ENABLED", raising=False) + reason = self._skip_reason(self._base_args(**overrides)) + assert reason is None or reason == "requires CUDA" + + def test_cache_dit_excluded(self, monkeypatch): + monkeypatch.delenv("SGLANG_DIFFUSION_DISABLE_AUTO_RESIDENCY", raising=False) + monkeypatch.setenv("SGLANG_CACHE_DIT_ENABLED", "true") + reason = self._skip_reason(self._base_args()) + assert reason is not None and "cache-dit" in reason + + @pytest.mark.parametrize( + "overrides", + [ + {"dp_size": 2}, + {"ulysses_degree": 2}, + {"use_fsdp_inference": True}, + ], + ) + def test_parallel_paths_reach_platform_gate(self, monkeypatch, overrides): + monkeypatch.delenv("SGLANG_DIFFUSION_DISABLE_AUTO_RESIDENCY", raising=False) + monkeypatch.delenv("SGLANG_CACHE_DIT_ENABLED", raising=False) + reason = self._skip_reason(self._base_args(**overrides)) + assert reason is None or reason == "requires CUDA" + + def test_eligible_path_reaches_platform_gate(self, monkeypatch): + monkeypatch.delenv("SGLANG_DIFFUSION_DISABLE_AUTO_RESIDENCY", raising=False) + monkeypatch.delenv("SGLANG_CACHE_DIT_ENABLED", raising=False) + reason = self._skip_reason(self._base_args()) + # on a CUDA host everything passes; CPU CI stops at the platform gate + assert reason is None or reason == "requires CUDA" diff --git a/python/sglang/multimodal_gen/test/unit/test_cfg_parallel_warmup.py b/python/sglang/multimodal_gen/test/unit/test_cfg_parallel_warmup.py index df57273e3..adfd31aeb 100644 --- a/python/sglang/multimodal_gen/test/unit/test_cfg_parallel_warmup.py +++ b/python/sglang/multimodal_gen/test/unit/test_cfg_parallel_warmup.py @@ -28,12 +28,13 @@ from sglang.multimodal_gen.configs.pipeline_configs.longlive2 import ( LongLive2T2VConfig, ) from sglang.multimodal_gen.configs.sample.longlive2 import LongLive2SamplingParams +from sglang.multimodal_gen.configs.sample.minimax_h3 import MiniMaxH3SamplingParams from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams -from sglang.multimodal_gen.runtime.entrypoints.diffusion_generator import DiffGenerator -from sglang.multimodal_gen.runtime.entrypoints.utils import ( +from sglang.multimodal_gen.runtime.entrypoints.control_requests import ( SetLoraReq, UnmergeLoraWeightsReq, ) +from sglang.multimodal_gen.runtime.entrypoints.diffusion_generator import DiffGenerator from sglang.multimodal_gen.runtime.managers.scheduler import Scheduler from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import ( OutputBatch, @@ -53,6 +54,7 @@ from sglang.multimodal_gen.runtime.server_warmup import ( from sglang.multimodal_gen.runtime.warmup_request_builder import ( DEFAULT_PLACEHOLDER_PROMPT, SERVER_WARMUP_IMAGE_FALLBACK_RESOLUTION, + _apply_warmup_sampling_overrides, _resolve_warmup_num_frames, build_warmup_reqs, should_include_warmup_image, @@ -108,6 +110,51 @@ def _make_validation_server_args(enable_cfg_parallel: bool) -> MagicMock: class TestWarmupReqCfgParallel(unittest.TestCase): """Warmup request construction and req-based warmup guards.""" + def test_sampling_workload_override_accepts_json(self): + defaults = SamplingParams( + width=1024, + height=1024, + num_frames=81, + num_inference_steps=35, + ) + server_args = SimpleNamespace( + warmup_sampling_params=( + '{"width":832,"height":480,"num_frames":9,"num_inference_steps":4}' + ) + ) + + overridden = _apply_warmup_sampling_overrides(server_args, defaults) + + self.assertEqual( + ( + overridden.width, + overridden.height, + overridden.num_frames, + overridden.num_inference_steps, + ), + (832, 480, 9, 4), + ) + self.assertEqual((defaults.width, defaults.height), (1024, 1024)) + + def test_sampling_workload_override_rejects_unknown_field(self): + server_args = SimpleNamespace( + warmup_sampling_params={"not_a_sampling_field": 1} + ) + + with self.assertRaisesRegex(ValueError, "invalid --warmup-sampling-params"): + _apply_warmup_sampling_overrides(server_args, SamplingParams()) + + def test_sampling_workload_override_supports_fixed_model_fields(self): + defaults = MiniMaxH3SamplingParams() + server_args = SimpleNamespace( + warmup_sampling_params={"num_frames": 49, "fps": 12} + ) + + overridden = _apply_warmup_sampling_overrides(server_args, defaults) + + self.assertEqual((overridden.num_frames, overridden.fps), (49, 12)) + self.assertEqual((defaults.num_frames, defaults.fps), (1, 24)) + def test_warmup_req_cfg_parallel_sets_do_cfg(self): server_args = _make_bare_scheduler(enable_cfg_parallel=True).server_args sampling_defaults = SamplingParams() @@ -270,6 +317,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase): server_args.warmup_steps = 1 server_args.enable_cfg_parallel = False server_args.enable_torch_compile = False + server_args.num_gpus = 1 server_args.pipeline_config.task_type = ModelTaskType.T2V server_args.pipeline_config.adjust_num_frames.side_effect = lambda value: value @@ -334,6 +382,68 @@ class TestWarmupReqCfgParallel(unittest.TestCase): self.assertTrue(req.extra["return_warmup_result"]) self.assertTrue(req.extra["server_based_warmup"]) + def test_auto_residency_uses_one_full_serving_shape_probe(self): + server_args = SimpleNamespace( + warmup_steps=1, + enable_cfg_parallel=False, + enable_torch_compile=False, + enable_breakable_cuda_graph=False, + pipeline_class_name=None, + num_gpus=1, + pipeline_config=SimpleNamespace( + task_type=ModelTaskType.T2V, + adjust_num_frames=lambda value: value, + vae_stride=None, + vae_scale_factor=None, + vae_config=SimpleNamespace(arch_config=None), + ), + is_arg_explicitly_set=lambda _name: False, + ) + sampling_defaults = SamplingParams( + width=1280, + height=720, + num_frames=81, + num_inference_steps=35, + adjust_frames=False, + supported_resolutions=[(1280, 720), (832, 480)], + ) + with ( + patch( + "sglang.multimodal_gen.runtime.warmup_request_builder.get_model_sampling_defaults", + return_value=sampling_defaults, + ), + patch( + "sglang.multimodal_gen.runtime.warmup_request_builder.auto_residency_args_skip_reason", + return_value=None, + ), + ): + reqs = build_warmup_reqs( + server_args, + warmup_resolutions=None, + server_based_warmup=True, + ) + + # the bounded warmup runs first so the worker can size the probe, and + # once more after it so serving starts from a serving-shaped pool + self.assertEqual(len(reqs), 3) + self.assertFalse(reqs[0].extra.get("auto_residency_full_shape_probe")) + self.assertFalse(reqs[2].extra.get("auto_residency_full_shape_probe")) + self.assertEqual( + (reqs[2].width, reqs[2].height, reqs[2].num_frames), + (reqs[0].width, reqs[0].height, reqs[0].num_frames), + ) + self.assertEqual( + (reqs[1].width, reqs[1].height, reqs[1].num_frames), + (1280, 720, 81), + ) + self.assertTrue(reqs[1].extra["auto_residency_full_shape_probe"]) + self.assertFalse(reqs[1].metrics.suppress_stage_breakdown) + self.assertEqual(reqs[1].num_inference_steps, 4) + self.assertIn( + "auto residency probe (1280x720x81f, 4/35 steps)", + format_warmup_req(reqs[1]), + ) + def test_server_based_warmup_uses_model_default_resolution(self): server_args = MagicMock() server_args.warmup_steps = 1 @@ -362,6 +472,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase): server_args.warmup_steps = 1 server_args.enable_cfg_parallel = False server_args.enable_torch_compile = False + server_args.num_gpus = 1 server_args.pipeline_config.task_type = ModelTaskType.T2V server_args.pipeline_config.adjust_num_frames.side_effect = lambda value: value @@ -474,6 +585,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase): server_args.warmup_steps = 1 server_args.enable_cfg_parallel = False server_args.enable_torch_compile = False + server_args.num_gpus = 1 server_args.pipeline_config.task_type = ModelTaskType.T2V server_args.pipeline_config.adjust_num_frames.side_effect = lambda value: value @@ -503,6 +615,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase): pipeline_config=pipeline_config, enable_breakable_cuda_graph=False, pipeline_class_name=None, + num_gpus=1, ) num_frames = _resolve_warmup_num_frames( @@ -545,6 +658,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase): server_args.warmup_steps = 1 server_args.enable_cfg_parallel = False server_args.enable_torch_compile = False + server_args.num_gpus = 1 server_args.pipeline_config.task_type = ModelTaskType.T2V server_args.pipeline_config.adjust_num_frames.side_effect = lambda value: value @@ -581,10 +695,15 @@ class TestWarmupReqCfgParallel(unittest.TestCase): server_args.warmup_steps = 1 server_args.enable_cfg_parallel = False server_args.enable_torch_compile = False + server_args.num_gpus = 1 server_args.pipeline_class_name = "LTX2TwoStageHQPipeline" server_args.pipeline_config.task_type = ModelTaskType.T2V server_args.pipeline_config.vae_scale_factor = 32 + server_args.pipeline_config.vae_config = SimpleNamespace( + use_temporal_scaling_frames=True, + arch_config=SimpleNamespace(temporal_compression_ratio=8), + ) server_args.pipeline_config.adjust_num_frames.return_value = 25 server_args.num_gpus = 2 @@ -725,6 +844,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase): server_args.is_arg_explicitly_set.return_value = False server_args.pipeline_config = SimpleNamespace( task_type=ModelTaskType.I2M, + supports_auto_residency=True, vae_stride=None, vae_scale_factor=None, vae_config=None, @@ -772,6 +892,33 @@ class TestWarmupReqCfgParallel(unittest.TestCase): self.assertEqual(reqs[0].image_path, ["/tmp/warmup.png"]) + def test_server_based_warmup_keeps_image_input_count(self): + server_args = MagicMock() + server_args.warmup_steps = 1 + server_args.enable_cfg_parallel = False + server_args.enable_torch_compile = False + server_args.pipeline_config.task_type = ModelTaskType.TI2I + + with patch( + "sglang.multimodal_gen.runtime.warmup_request_builder.get_model_sampling_defaults", + return_value=SamplingParams( + width=512, + height=512, + image_path=["first.png", "second.png"], + ), + ): + reqs = build_warmup_reqs( + server_args, + warmup_resolutions=None, + warmup_input_path="/tmp/warmup.png", + server_based_warmup=True, + ) + + self.assertEqual( + reqs[0].image_path, + ["/tmp/warmup.png", "/tmp/warmup.png"], + ) + def test_server_based_warmup_keeps_required_image_input(self): server_args = MagicMock() server_args.warmup_steps = 1 @@ -797,6 +944,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase): server_args.warmup_steps = 1 server_args.enable_cfg_parallel = False server_args.enable_torch_compile = False + server_args.num_gpus = 1 server_args.pipeline_config.task_type = ModelTaskType.TI2V with patch( diff --git a/python/sglang/multimodal_gen/test/unit/test_component_residency.py b/python/sglang/multimodal_gen/test/unit/test_component_residency.py index c6cd76686..3c51e8418 100644 --- a/python/sglang/multimodal_gen/test/unit/test_component_residency.py +++ b/python/sglang/multimodal_gen/test/unit/test_component_residency.py @@ -8,6 +8,7 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager im ComponentResidencyManager, ComponentUse, ResidencyState, + WarmupPhasePeak, ) from sglang.multimodal_gen.runtime.managers.memory_managers.component_residency import ( ComponentResidencyError, @@ -22,9 +23,19 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.image_encoding import ( from sglang.multimodal_gen.runtime.pipelines_core.stages.realtime.text_encoding import ( RealtimeTextEncodingStage, ) +from sglang.multimodal_gen.runtime.platforms import current_platform from sglang.multimodal_gen.runtime.server_args import ServerArgs +def _server_args(*, supports_auto_residency=True): + return SimpleNamespace( + enable_layerwise_nvtx_marker=False, + pipeline_config=SimpleNamespace( + supports_auto_residency=supports_auto_residency, + ), + ) + + def test_component_offload_releases_preferred_component_after_request(): strategy = ComponentOffloadStrategy() strategy.finish_use = Mock() @@ -143,7 +154,7 @@ def test_group_warmup_state_requires_every_batch_to_be_warmup(): _stage_name_mapping={}, component_residency_strategies={}, ) - server_args = SimpleNamespace(enable_layerwise_nvtx_marker=False) + server_args = _server_args() manager = ComponentResidencyManager(pipeline, server_args) manager.begin_request( @@ -166,6 +177,246 @@ class _Stage: return self.uses +def test_warmup_records_use_and_transition_peaks(monkeypatch): + device_module = SimpleNamespace( + is_available=lambda: True, + reset_peak_memory_stats=Mock(), + max_memory_allocated=lambda: 7, + memory_allocated=lambda: 2, + ) + monkeypatch.setattr(torch, "get_device_module", lambda: device_module) + monkeypatch.setattr(current_platform, "is_cuda", lambda: True) + + use = ComponentUse("denoise", "transformer") + stage = _Stage(use) + module = torch.nn.Linear(2, 2) + pipeline = SimpleNamespace( + modules={"transformer": module}, + _stage_name_mapping={"denoise": stage}, + component_residency_strategies={}, + ) + server_args = _server_args() + manager = ComponentResidencyManager(pipeline, server_args) + manager.strategy_for = Mock(return_value=Mock()) + manager.refresh_pipeline(pipeline) + manager.begin_request([stage], SimpleNamespace(is_warmup=True), server_args) + + manager.before_stage(stage, 0, SimpleNamespace(is_warmup=True), server_args) + manager.begin_stage() + manager.end_stage() + manager.finish_request() + + peaks = manager.take_warmup_phase_peaks() + inactive_peak = WarmupPhasePeak((), 7) + transformer_peak = WarmupPhasePeak( + ("transformer",), 7, used_components=("transformer",) + ) + assert peaks["request:before-stage"] == inactive_peak + assert peaks["0:denoise:setup"] == inactive_peak + assert peaks["0:denoise:transition:idle->transformer"] == transformer_peak + assert peaks["0:denoise:use:transformer"] == transformer_peak + assert peaks["0:denoise:transition:transformer->idle"] == transformer_peak + assert peaks["0:denoise:between"] == inactive_peak + # A non-preferred component is being released during cleanup, so it is no + # longer part of the placement that follows this transition. + assert peaks["request:cleanup:transformer"] == inactive_peak + assert peaks["idle"] == WarmupPhasePeak( + active_components=(), + allocated_bytes=2, + ) + + +def test_warmup_skips_memory_tracking_for_unsupported_pipeline(monkeypatch): + device_module = SimpleNamespace( + is_available=lambda: True, + reset_peak_memory_stats=Mock(), + ) + monkeypatch.setattr(torch, "get_device_module", lambda: device_module) + monkeypatch.setattr(current_platform, "is_cuda", lambda: True) + + stage = _Stage() + pipeline = SimpleNamespace( + modules={}, + _stage_name_mapping={"stage": stage}, + component_residency_strategies={}, + ) + server_args = _server_args(supports_auto_residency=False) + manager = ComponentResidencyManager(pipeline, server_args) + + manager.begin_request([stage], SimpleNamespace(is_warmup=True), server_args) + manager.before_stage(stage, 0, SimpleNamespace(is_warmup=True), server_args) + manager.finish_request() + + assert manager._track_warmup_memory is False + assert manager.take_warmup_phase_peaks() == {} + device_module.reset_peak_memory_stats.assert_not_called() + + +def test_warmup_records_full_weight_transition_without_preparing(monkeypatch): + device_module = SimpleNamespace( + is_available=lambda: True, + reset_peak_memory_stats=Mock(), + max_memory_allocated=lambda: 7, + memory_allocated=lambda: 2, + ) + monkeypatch.setattr(torch, "get_device_module", lambda: device_module) + monkeypatch.setattr(current_platform, "is_cuda", lambda: True) + + stage = _Stage() + module = torch.nn.Linear(2, 2) + pipeline = SimpleNamespace( + modules={"transformer": module}, + _stage_name_mapping={"lora_switch": stage}, + component_residency_strategies={}, + ) + server_args = _server_args() + manager = ComponentResidencyManager(pipeline, server_args) + manager.strategy_for = Mock() + manager.refresh_pipeline(pipeline) + manager.begin_request([stage], SimpleNamespace(is_warmup=True), server_args) + + manager.before_stage(stage, 0, SimpleNamespace(is_warmup=True), server_args) + with manager.full_weight_transition(("transformer",)): + pass + + assert manager._warmup_phase_peaks[ + "0:lora_switch:full-weight-transition:transformer" + ] == WarmupPhasePeak( + (), + 7, + full_weight_transition_components=("transformer",), + ) + assert manager._warmup_phase_key == "0:lora_switch:setup" + manager.strategy_for.assert_not_called() + + +def test_warmup_records_same_component_dtype_prepare_as_transition(monkeypatch): + device_module = SimpleNamespace( + is_available=lambda: True, + reset_peak_memory_stats=Mock(), + max_memory_allocated=lambda: 7, + memory_allocated=lambda: 2, + ) + monkeypatch.setattr(torch, "get_device_module", lambda: device_module) + monkeypatch.setattr(current_platform, "is_cuda", lambda: True) + + first = ComponentUse("stage", "transformer", target_dtype=torch.float16) + second = ComponentUse("stage", "transformer", target_dtype=torch.bfloat16) + stage = _Stage(first, second) + module = torch.nn.Linear(2, 2) + pipeline = SimpleNamespace( + modules={"transformer": module}, + _stage_name_mapping={"stage": stage}, + component_residency_strategies={}, + ) + server_args = _server_args() + manager = ComponentResidencyManager(pipeline, server_args) + strategy = Mock() + manager.strategy_for = Mock(return_value=strategy) + manager.refresh_pipeline(pipeline) + manager.begin_request([stage], SimpleNamespace(is_warmup=True), server_args) + manager.before_stage(stage, 0, SimpleNamespace(is_warmup=True), server_args) + + manager.begin_use(first, module=module) + manager.begin_use(second, module=module) + + manager._record_warmup_phase_peak() + assert manager._warmup_phase_peaks[ + "0:stage:transition:transformer->transformer" + ] == WarmupPhasePeak(("transformer",), 7, used_components=("transformer",)) + assert strategy.prepare_for_use.call_count == 2 + + +def test_warmup_attributes_prefetch_peak_to_prefetched_component(monkeypatch): + device_module = SimpleNamespace( + is_available=lambda: True, + reset_peak_memory_stats=Mock(), + max_memory_allocated=lambda: 7, + memory_allocated=lambda: 2, + ) + monkeypatch.setattr(torch, "get_device_module", lambda: device_module) + monkeypatch.setattr(current_platform, "is_cuda", lambda: True) + + encoder_use = ComponentUse("encode", "text_encoder") + transformer_use = ComponentUse("denoise", "transformer", memory_intensive=True) + encode_stage = _Stage(encoder_use) + denoise_stage = _Stage(transformer_use) + modules = { + "text_encoder": torch.nn.Linear(2, 2), + "transformer": torch.nn.Linear(2, 2), + } + pipeline = SimpleNamespace( + modules=modules, + _stage_name_mapping={ + "encode": encode_stage, + "denoise": denoise_stage, + }, + component_residency_strategies={}, + ) + server_args = _server_args() + manager = ComponentResidencyManager(pipeline, server_args) + strategy = Mock() + strategy.prefetch_for_use.return_value = True + manager.strategy_for = Mock(return_value=strategy) + manager.refresh_pipeline(pipeline) + manager.begin_request( + [encode_stage, denoise_stage], + SimpleNamespace(is_warmup=True), + server_args, + ) + + manager.before_stage(encode_stage, 0, SimpleNamespace(is_warmup=True), server_args) + manager.begin_stage() + manager.end_stage() + manager.before_stage(denoise_stage, 1, SimpleNamespace(is_warmup=True), server_args) + + assert manager._warmup_phase_peaks[ + "0:encode:prefetch:transformer" + ] == WarmupPhasePeak(("transformer",), 7, used_components=("transformer",)) + assert manager._warmup_phase_peaks["0:encode:between"] == WarmupPhasePeak((), 7) + + +def test_warmup_splits_sequential_component_transition(monkeypatch): + device_module = SimpleNamespace( + is_available=lambda: True, + reset_peak_memory_stats=Mock(), + max_memory_allocated=lambda: 7, + memory_allocated=lambda: 2, + ) + monkeypatch.setattr(torch, "get_device_module", lambda: device_module) + monkeypatch.setattr(current_platform, "is_cuda", lambda: True) + + first = ComponentUse("stage", "text_encoder") + second = ComponentUse("stage", "transformer") + stage = _Stage(first, second) + modules = { + "text_encoder": torch.nn.Linear(2, 2), + "transformer": torch.nn.Linear(2, 2), + } + pipeline = SimpleNamespace( + modules=modules, + _stage_name_mapping={"stage": stage}, + component_residency_strategies={}, + ) + server_args = _server_args() + manager = ComponentResidencyManager(pipeline, server_args) + manager.strategy_for = Mock(return_value=Mock()) + manager.refresh_pipeline(pipeline) + manager.begin_request([stage], SimpleNamespace(is_warmup=True), server_args) + manager.before_stage(stage, 0, SimpleNamespace(is_warmup=True), server_args) + + manager.begin_use(first) + manager.begin_use(second) + manager._record_warmup_phase_peak() + + assert manager._warmup_phase_peaks["0:stage:transition:text_encoder->idle"] == ( + WarmupPhasePeak(("text_encoder",), 7, used_components=("text_encoder",)) + ) + assert manager._warmup_phase_peaks["0:stage:transition:idle->transformer"] == ( + WarmupPhasePeak(("transformer",), 7, used_components=("transformer",)) + ) + + def _manager_for_stage(stage, modules): pipeline = SimpleNamespace( modules=modules, diff --git a/python/sglang/multimodal_gen/test/unit/test_diffusion_generator_shutdown.py b/python/sglang/multimodal_gen/test/unit/test_diffusion_generator_shutdown.py index e43f550dd..d00f9e436 100644 --- a/python/sglang/multimodal_gen/test/unit/test_diffusion_generator_shutdown.py +++ b/python/sglang/multimodal_gen/test/unit/test_diffusion_generator_shutdown.py @@ -3,8 +3,8 @@ from types import SimpleNamespace from unittest.mock import Mock, patch from sglang.multimodal_gen.runtime.entrypoints import diffusion_generator as dg +from sglang.multimodal_gen.runtime.entrypoints.control_requests import ShutdownReq from sglang.multimodal_gen.runtime.entrypoints.diffusion_generator import DiffGenerator -from sglang.multimodal_gen.runtime.entrypoints.utils import ShutdownReq class _FakeProcess: diff --git a/python/sglang/multimodal_gen/test/unit/test_dp_routing.py b/python/sglang/multimodal_gen/test/unit/test_dp_routing.py index ad0593faa..32bcd4399 100644 --- a/python/sglang/multimodal_gen/test/unit/test_dp_routing.py +++ b/python/sglang/multimodal_gen/test/unit/test_dp_routing.py @@ -6,7 +6,10 @@ and a realtime session always lands on the same replica it started on. import itertools -from sglang.multimodal_gen.runtime.entrypoints.utils import SetLoraReq, ShutdownReq +from sglang.multimodal_gen.runtime.entrypoints.control_requests import ( + SetLoraReq, + ShutdownReq, +) from sglang.multimodal_gen.runtime.pipelines_core import Req from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch from sglang.multimodal_gen.runtime.scheduler_client import ( diff --git a/python/sglang/multimodal_gen/test/unit/test_ipc_a2a_lifecycle.py b/python/sglang/multimodal_gen/test/unit/test_ipc_a2a_lifecycle.py index 031cd1c63..c6ed8ddc9 100644 --- a/python/sglang/multimodal_gen/test/unit/test_ipc_a2a_lifecycle.py +++ b/python/sglang/multimodal_gen/test/unit/test_ipc_a2a_lifecycle.py @@ -56,3 +56,38 @@ def test_peer_cuda_device_uses_the_ulysses_group_mapping(): patch(f"{_IPC}.dist.all_gather_object", side_effect=gather), ): assert _peer_cuda_device(group, rank=1, device=3) == 2 + + +def test_drop_staging_clears_cached_buffers_and_nothing_else(): + from collections import OrderedDict + + from sglang.multimodal_gen.runtime.distributed.device_communicators.ipc_a2a import ( + IpcA2AState, + ) + + state = IpcA2AState() + state.inited = True + state.calls = 7 + state.staging = OrderedDict( + {(4, 4, "bf16"): ("local", "peer"), (8, 8, "bf16"): ("l", "p")} + ) + + state.drop_staging() + + assert state.staging == OrderedDict() + assert state.inited is True and state.calls == 7 + state.drop_staging() # idempotent on an empty cache + + +def test_drop_a2a_staging_buffers_clears_the_ulysses_cache(): + import torch + + from sglang.multimodal_gen.runtime.layers import usp + + usp._A2A_STAGING_BUFFERS[("qkv", torch.float16, 0)] = torch.empty( + 8, dtype=torch.float16 + ) + with patch.object(torch.cuda, "is_available", return_value=False): + usp.drop_a2a_staging_buffers() + usp.drop_a2a_staging_buffers() # idempotent on an empty cache + assert usp._A2A_STAGING_BUFFERS == {} diff --git a/python/sglang/multimodal_gen/test/unit/test_launch_server_shutdown.py b/python/sglang/multimodal_gen/test/unit/test_launch_server_shutdown.py index 308ce5b50..0a3949aca 100644 --- a/python/sglang/multimodal_gen/test/unit/test_launch_server_shutdown.py +++ b/python/sglang/multimodal_gen/test/unit/test_launch_server_shutdown.py @@ -4,7 +4,7 @@ from unittest.mock import Mock, patch from sglang.multimodal_gen.runtime import launch_server as ls from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType -from sglang.multimodal_gen.runtime.entrypoints.utils import ShutdownReq +from sglang.multimodal_gen.runtime.entrypoints.control_requests import ShutdownReq class _FakeProcess: diff --git a/python/sglang/multimodal_gen/test/unit/test_performance_metrics.py b/python/sglang/multimodal_gen/test/unit/test_performance_metrics.py index 95b741270..ebfcd509d 100644 --- a/python/sglang/multimodal_gen/test/unit/test_performance_metrics.py +++ b/python/sglang/multimodal_gen/test/unit/test_performance_metrics.py @@ -1,11 +1,16 @@ import json +from types import SimpleNamespace from unittest.mock import Mock, patch import pytest import torch import sglang.multimodal_gen.runtime.managers.gpu_worker as gpu_worker_module +import sglang.multimodal_gen.runtime.managers.memory_managers.component_manager as component_manager_module from sglang.multimodal_gen.runtime.managers.gpu_worker import GPUWorker +from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager import ( + WarmupPhasePeak, +) from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch from sglang.multimodal_gen.runtime.platforms import current_platform from sglang.multimodal_gen.runtime.utils.perf_logger import ( @@ -23,6 +28,16 @@ from sglang.multimodal_gen.test.server.testcase_configs import ( ) +@pytest.fixture(autouse=True) +def _isolate_global_component_residency_manager(monkeypatch): + # Worker paths exercised here create the process-global residency manager + # for their fake pipeline; leaving it behind hands its modules to later + # tests (the worker prefers the global manager's placement modules). + monkeypatch.setattr( + component_manager_module, "_GLOBAL_COMPONENT_RESIDENCY_MANAGER", None + ) + + def _perf_record(memory_snapshots: dict[str, dict]) -> RequestPerfRecord: return RequestPerfRecord( request_id="request", @@ -35,6 +50,27 @@ def _perf_record(memory_snapshots: dict[str, dict]) -> RequestPerfRecord: ) +def test_request_metrics_attributes_steps_and_iterations_to_active_stage(): + metrics = RequestMetrics("request") + metrics.active_stage_name = "ShapeStage" + + metrics.record_step(0.1) + metrics.record_stage_iterations(4, 50) + metrics.active_stage_name = "PaintStage" + metrics.record_step(0.2) + metrics.record_stage_iterations(4, 30) + + assert metrics.steps == [100.0, 200.0] + assert metrics.steps_by_stage == { + "ShapeStage": [100.0], + "PaintStage": [200.0], + } + assert metrics.stage_iterations == { + "ShapeStage": (4, 50), + "PaintStage": (4, 30), + } + + def test_performance_summary_separates_load_and_runtime_peaks(): summary = PerformanceSummary.from_req_perf_record( _perf_record( @@ -62,13 +98,14 @@ def test_worker_records_replica_load_and_runtime_peaks(): worker.is_output_rank = True worker._load_peak_reserved_mb = 4096.0 worker._runtime_peak_reserved_mb = 0.0 + worker._warmup_peak_reserved_mb = 0.0 worker._load_peak_allocated_mb = 3000.0 worker._runtime_peak_allocated_mb = 0.0 output = OutputBatch() metrics = RequestMetrics("request") replica_group = Mock() replica_group.all_reduce.return_value = torch.tensor( - [5120.0, 3584.0, 3500.0, 2560.0], dtype=torch.float64 + [5120.0, 3584.0, 6144.0, 3500.0, 2560.0], dtype=torch.float64 ) snapshots = [ MemorySnapshot(0.0, 0.0, 2048.0, 3072.0), @@ -97,6 +134,88 @@ def test_worker_records_replica_load_and_runtime_peaks(): assert metrics.memory_snapshots["runtime_peak"].peak_reserved_mb == 3584.0 assert metrics.memory_snapshots["load_peak"].peak_allocated_mb == 3500.0 assert metrics.memory_snapshots["runtime_peak"].peak_allocated_mb == 2560.0 + assert metrics.memory_snapshots["warmup_peak"].peak_reserved_mb == 6144.0 + + +def test_server_warmup_preserves_peak_after_managed_stage_timeline(): + worker = GPUWorker.__new__(GPUWorker) + worker._auto_residency_warmup_records = [] + residency_manager = Mock() + residency_manager.take_warmup_phase_peaks.return_value = { + "0:denoise:use:transformer": WarmupPhasePeak(("transformer",), 8) + } + residency_manager.current_device_components.return_value = ("transformer",) + device_module = Mock() + device_module.max_memory_allocated.return_value = 9 + device_module.max_memory_reserved.return_value = 12 + req = SimpleNamespace( + width=64, + height=64, + num_frames=1, + num_inference_steps=1, + metrics=None, + ) + + with ( + patch.object(torch, "get_device_module", return_value=device_module), + patch.object( + gpu_worker_module, + "peek_global_component_residency_manager", + return_value=residency_manager, + ), + ): + worker._record_server_warmup_memory( + req=req, + workload=(128, 96, 9, 2), + baseline_allocated_bytes=3, + succeeded=True, + ) + + record = worker._auto_residency_warmup_records[0] + assert record.peak_allocated_bytes == 9 + assert record.peak_reserved_bytes == 12 + assert (record.width, record.height, record.num_frames) == (128, 96, 9) + assert record.num_inference_steps == 2 + assert record.phase_peak_allocated_bytes["request:untracked"] == 9 + assert record.phase_active_components["request:untracked"] == ("transformer",) + + +def test_server_warmup_does_not_treat_allocator_cache_as_untracked_live_memory(): + worker = GPUWorker.__new__(GPUWorker) + worker._auto_residency_warmup_records = [] + residency_manager = Mock() + residency_manager.take_warmup_phase_peaks.return_value = { + "0:denoise:use:transformer": WarmupPhasePeak(("transformer",), 8) + } + residency_manager.current_device_components.return_value = ("transformer",) + device_module = Mock() + device_module.max_memory_allocated.return_value = 8 + device_module.max_memory_reserved.return_value = 12 + req = SimpleNamespace( + width=64, + height=64, + num_frames=1, + num_inference_steps=1, + metrics=None, + ) + + with ( + patch.object(torch, "get_device_module", return_value=device_module), + patch.object( + gpu_worker_module, + "peek_global_component_residency_manager", + return_value=residency_manager, + ), + ): + worker._record_server_warmup_memory( + req=req, + workload=(64, 64, 1, 1), + baseline_allocated_bytes=3, + succeeded=True, + ) + + record = worker._auto_residency_warmup_records[0] + assert "request:untracked" not in record.phase_peak_allocated_bytes def test_baseline_config_loads_per_scenario_peak_vram(tmp_path): @@ -249,6 +368,16 @@ def test_peak_vram_validation_enforces_allocated_when_baselined(): with pytest.raises(AssertionError, match="Runtime Peak VRAM"): validator.validate_peak_vram(reserved_drift, 10_000.0, 10_000.0) + reserved_drift.warmup_peak_vram_mb = 12_000.0 + with pytest.raises(AssertionError, match="Warmup Peak VRAM"): + validator.validate_peak_vram( + reserved_drift, + 10_000.0, + 10_000.0, + expected_warmup_peak_vram_mb=10_000.0, + expected_runtime_peak_allocated_mb=8_000.0, + ) + @pytest.mark.parametrize( ("load_peak_vram_mb", "runtime_peak_vram_mb", "message"), @@ -308,3 +437,138 @@ def test_results_json_merges_retry_sessions(tmp_path): results = json.loads(path.read_text(encoding="utf-8")) assert {item["test_name"] for item in results} == {"first", "second"} + + +def _warmup_batch_for_iterations(steps: int, target_steps: int): + from sglang.multimodal_gen.runtime.utils.perf_logger import RequestMetrics + + metrics = RequestMetrics(request_id="probe") + metrics.suppress_stage_breakdown = False + metrics.active_stage_name = "DenoisingStage" + return SimpleNamespace( + metrics=metrics, + num_inference_steps=steps, + extra={"warmup_target_num_inference_steps": target_steps}, + is_warmup=True, + ) + + +def test_stage_formula_records_probe_and_default_iterations(): + from sglang.multimodal_gen.runtime.pipelines_core.stages.base import ( + record_default_workload_iterations, + ) + + batch = _warmup_batch_for_iterations(steps=4, target_steps=50) + stage = SimpleNamespace(default_workload_iterations=lambda batch, steps: steps - 1) + record_default_workload_iterations(stage, batch) + assert batch.metrics.stage_iterations == {"DenoisingStage": (3, 49)} + + +def test_fixed_schedule_formula_records_the_same_count_twice(): + from sglang.multimodal_gen.runtime.pipelines_core.stages.base import ( + record_default_workload_iterations, + ) + + batch = _warmup_batch_for_iterations(steps=4, target_steps=50) + stage = SimpleNamespace(default_workload_iterations=lambda batch, steps: 8) + record_default_workload_iterations(stage, batch) + assert batch.metrics.stage_iterations == {"DenoisingStage": (8, 8)} + + +def test_explicit_loop_record_wins_over_the_formula(): + from sglang.multimodal_gen.runtime.pipelines_core.stages.base import ( + record_default_workload_iterations, + ) + + batch = _warmup_batch_for_iterations(steps=4, target_steps=50) + batch.metrics.record_stage_iterations(12, 12) + stage = SimpleNamespace(default_workload_iterations=lambda batch, steps: steps) + record_default_workload_iterations(stage, batch) + assert batch.metrics.stage_iterations == {"DenoisingStage": (12, 12)} + + +def test_stage_without_a_formula_records_nothing(): + from sglang.multimodal_gen.runtime.pipelines_core.stages.base import ( + record_default_workload_iterations, + ) + + batch = _warmup_batch_for_iterations(steps=4, target_steps=50) + stage = SimpleNamespace(default_workload_iterations=lambda batch, steps: None) + record_default_workload_iterations(stage, batch) + assert batch.metrics.stage_iterations == {} + + +def test_warmup_peak_stays_out_of_the_runtime_peak(): + from unittest import mock + + worker = gpu_worker_module.GPUWorker.__new__(gpu_worker_module.GPUWorker) + worker._runtime_peak_reserved_mb = 0.0 + worker._runtime_peak_allocated_mb = 0.0 + worker._warmup_peak_reserved_mb = 0.0 + worker.is_output_rank = False + with ( + mock.patch.object( + gpu_worker_module.current_platform, "is_cpu", return_value=False + ), + mock.patch.object( + gpu_worker_module, + "capture_memory_snapshot", + side_effect=[ + MemorySnapshot(0.0, 0.0, 35000.0, 40000.0), + MemorySnapshot(0.0, 0.0, 15000.0, 20000.0), + ], + ), + ): + worker._record_output_peak_memory(SimpleNamespace(), is_warmup=True) + worker._record_output_peak_memory(SimpleNamespace(), is_warmup=False) + assert worker._warmup_peak_reserved_mb == 40000.0 + assert worker._runtime_peak_reserved_mb == 20000.0 + assert worker._runtime_peak_allocated_mb == 15000.0 + + +def _warmup_req(*, probe: bool = False) -> SimpleNamespace: + extra = {"auto_residency_full_shape_probe": True} if probe else {} + return SimpleNamespace(is_warmup=True, extra=extra) + + +def test_worker_releases_the_probe_pool_before_the_next_request(monkeypatch): + import sglang.multimodal_gen.runtime.distributed.device_communicators.ipc_a2a as ipc_a2a_module + + calls = [] + fake_device = SimpleNamespace(empty_cache=lambda: calls.append("empty_cache")) + monkeypatch.setattr( + gpu_worker_module.torch, "get_device_module", lambda: fake_device + ) + monkeypatch.setattr(type(current_platform), "is_cpu", lambda self: False) + monkeypatch.setattr(type(current_platform), "is_mps", lambda self: False) + monkeypatch.setattr( + ipc_a2a_module.IPC_A2A, "drop_staging", lambda: calls.append("drop_staging") + ) + worker = GPUWorker.__new__(GPUWorker) + worker._release_warmup_pool_before_serving = False + + worker._release_warmup_pool(_warmup_req()) + worker._release_warmup_pool(_warmup_req(probe=True)) + assert calls == [] + + # the bounded re-warm after the probe regrows the pool from empty, and the + # IPC staging buffers sized for the probe's messages go with it + worker._release_warmup_pool(_warmup_req()) + assert calls == ["drop_staging", "empty_cache"] + + worker._release_warmup_pool(SimpleNamespace(is_warmup=False, extra={})) + assert calls == ["drop_staging", "empty_cache"] + + +def test_worker_keeps_the_pool_when_no_probe_ran(monkeypatch): + calls = [] + fake_device = SimpleNamespace(empty_cache=lambda: calls.append("empty_cache")) + monkeypatch.setattr( + gpu_worker_module.torch, "get_device_module", lambda: fake_device + ) + worker = GPUWorker.__new__(GPUWorker) + worker._release_warmup_pool_before_serving = False + + worker._release_warmup_pool(_warmup_req()) + worker._release_warmup_pool(SimpleNamespace(is_warmup=False, extra={})) + assert calls == [] diff --git a/python/sglang/multimodal_gen/test/unit/test_warmup_probe_degradation.py b/python/sglang/multimodal_gen/test/unit/test_warmup_probe_degradation.py new file mode 100644 index 000000000..bcbb0ba67 --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/test_warmup_probe_degradation.py @@ -0,0 +1,179 @@ +"""A warmup probe that does not fit is retried smaller instead of abandoned.""" + +from types import SimpleNamespace + +from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams +from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req +from sglang.multimodal_gen.runtime.warmup_request_builder import lighten_warmup_req + + +def _server_args(temporal_compression_ratio: int = 4) -> SimpleNamespace: + arch_config = SimpleNamespace( + temporal_compression_ratio=temporal_compression_ratio, + vae_scale_factor=8, + spatial_compression_ratio=8, + ) + return SimpleNamespace( + pipeline_class_name=None, + pipeline_config=SimpleNamespace( + vae_config=SimpleNamespace(arch_config=arch_config), + vae_scale_factor=8, + ), + ) + + +def _req(width: int, height: int, num_frames: int) -> Req: + return Req( + sampling_params=SamplingParams( + width=width, height=height, num_frames=num_frames + ) + ) + + +class TestLightenWarmupReq: + def test_video_probe_halves_latent_frames_first(self): + lighter = lighten_warmup_req(_server_args(), _req(832, 480, 17)) + assert (lighter.width, lighter.height) == (832, 480) + assert lighter.num_frames == 9 + + def test_frames_step_down_to_a_single_frame(self): + server_args = _server_args() + frames = [] + req = _req(832, 480, 17) + for _ in range(4): + req = lighten_warmup_req(server_args, req) + if req is None: + break + frames.append(req.num_frames) + assert frames[:3] == [9, 5, 1] + + def test_image_probe_halves_the_area(self): + lighter = lighten_warmup_req(_server_args(), _req(1024, 1024, 1)) + assert lighter.num_frames == 1 + assert lighter.width * lighter.height <= 1024 * 1024 // 2 + assert lighter.width % 16 == 0 and lighter.height % 16 == 0 + + def test_the_original_request_is_left_alone(self): + req = _req(832, 480, 17) + lighten_warmup_req(_server_args(), req) + assert req.num_frames == 17 + + def test_a_probe_at_the_floor_cannot_shrink(self): + assert lighten_warmup_req(_server_args(), _req(16, 16, 1)) is None + + def test_frames_follow_the_model_frame_contract(self): + # LongLive2-style contract: latent frames come in causal blocks of 8, + # so with a temporal ratio of 4 only 29, 61, 93, ... frames are valid. + server_args = _server_args() + + def adjust_num_frames(num_frames: int) -> int: + latent = (num_frames - 1) // 4 + 1 + if latent % 8 == 0: + return num_frames + return (max(8, latent // 8 * 8) - 1) * 4 + 1 + + server_args.pipeline_config.adjust_num_frames = adjust_num_frames + + lighter = lighten_warmup_req(server_args, _req(960, 928, 61)) + assert lighter.num_frames == 29 + assert (lighter.width, lighter.height) == (960, 928) + + # At the smallest valid frame count the probe shrinks the area instead. + floor = lighten_warmup_req(server_args, lighter) + assert floor.num_frames == 29 + assert floor.width * floor.height <= 960 * 928 // 2 + + +def _record(width: int, height: int, num_frames: int, *, peak_gib: float): + from sglang.multimodal_gen.runtime.managers.memory_managers.auto_residency import ( + WarmupMemoryRecord, + ) + + return WarmupMemoryRecord( + width=width, + height=height, + num_frames=num_frames, + baseline_allocated_bytes=2 << 30, + peak_allocated_bytes=int(peak_gib * (1 << 30)), + succeeded=True, + ) + + +class TestFitAutoResidencyProbe: + def test_probe_shrinks_until_its_extrapolated_peak_fits(self): + from sglang.multimodal_gen.runtime.managers.gpu_worker import ( + fit_auto_residency_probe, + ) + + fitted, estimate, steps = fit_auto_residency_probe( + _req(1280, 720, 81), + records=[_record(832, 480, 81, peak_gib=20.0)], + free_bytes=40 << 30, + total_bytes=80 << 30, + server_args=_server_args(), + ) + assert steps >= 1 + assert (fitted.width, fitted.height) == (1280, 720) + assert fitted.num_frames < 81 + assert estimate is not None and estimate <= 40 << 30 + + def test_probe_that_fits_runs_at_full_shape(self): + from sglang.multimodal_gen.runtime.managers.gpu_worker import ( + fit_auto_residency_probe, + ) + + fitted, _, steps = fit_auto_residency_probe( + _req(1280, 720, 81), + records=[_record(832, 480, 81, peak_gib=20.0)], + free_bytes=79 << 30, + total_bytes=80 << 30, + server_args=_server_args(), + ) + assert steps == 0 + assert fitted.num_frames == 81 + + def test_probe_never_shrinks_below_the_bounded_warmup_shape(self): + from sglang.multimodal_gen.runtime.managers.gpu_worker import ( + fit_auto_residency_probe, + ) + + # Nothing fits the extrapolation, but the bounded 832x480x17f warmup + # already ran, so the ladder (81 -> 41 -> 21 -> 9 frames) stops at the + # first shape at or below it instead of reaching a 16x16x1f probe. + fitted, _, steps = fit_auto_residency_probe( + _req(832, 480, 81), + records=[_record(832, 480, 17, peak_gib=30.0)], + free_bytes=8 << 30, + total_bytes=80 << 30, + server_args=_server_args(), + ) + assert steps >= 1 + assert (fitted.width, fitted.height) == (832, 480) + assert fitted.num_frames == 9 + + def test_without_a_trusted_estimate_the_probe_runs_as_requested(self): + from sglang.multimodal_gen.runtime.managers.gpu_worker import ( + fit_auto_residency_probe, + ) + + fitted, estimate, steps = fit_auto_residency_probe( + _req(1280, 720, 81), + records=[], + free_bytes=1 << 30, + total_bytes=80 << 30, + server_args=_server_args(), + ) + assert (steps, estimate) == (0, None) + assert fitted.num_frames == 81 + + +class TestOutOfMemoryClassification: + def test_allocation_failures_from_libraries_count_as_out_of_memory(self): + from sglang.multimodal_gen.runtime.server_warmup import _is_out_of_memory + + assert _is_out_of_memory("CUDA error: out of memory") + assert _is_out_of_memory("cuBLAS error: CUBLAS_STATUS_ALLOC_FAILED") + assert _is_out_of_memory( + "RuntimeError: cudaErrorMemoryAllocation: out of memory" + ) + assert not _is_out_of_memory("shape mismatch in attention")