diff --git a/python/sglang/multimodal_gen/runtime/managers/scheduler.py b/python/sglang/multimodal_gen/runtime/managers/scheduler.py index c11c2c850..1c26655f9 100644 --- a/python/sglang/multimodal_gen/runtime/managers/scheduler.py +++ b/python/sglang/multimodal_gen/runtime/managers/scheduler.py @@ -4,12 +4,13 @@ import asyncio import os import pickle +import tempfile from collections import deque from typing import Any, List import zmq -from sglang.multimodal_gen.configs.pipeline_configs.base import ModelTaskType +from sglang.multimodal_gen.runtime.distributed import get_world_group from sglang.multimodal_gen.runtime.entrypoints.openai.utils import ( _parse_size, save_image_to_path, @@ -198,32 +199,24 @@ class Scheduler: # insert warmup reqs constructed with each warmup-resolution self._warmup_total = len(self.server_args.warmup_resolutions) self._warmup_processed = 0 + task_type = self.server_args.pipeline_config.task_type + + requires_warmup_image = task_type.accepts_image_input() + warmup_input_path = None + if requires_warmup_image: + warmup_input_path = self._prepare_shared_warmup_image_path() for resolution in self.server_args.warmup_resolutions: width, height = _parse_size(resolution) - task_type = self.server_args.pipeline_config.task_type - if task_type in ( - ModelTaskType.I2I, - ModelTaskType.TI2I, - ModelTaskType.I2V, - ModelTaskType.TI2V, - ): - uploads_dir = os.path.join("outputs", "uploads") - os.makedirs(uploads_dir, exist_ok=True) - input_path = asyncio.run( - save_image_to_path( - MINIMUM_PICTURE_BASE64_FOR_WARMUP, - os.path.join(uploads_dir, "warmup_image.jpg"), - ) - ) + if requires_warmup_image: req = Req( data_type=task_type.data_type(), width=width, height=height, prompt="", negative_prompt="", - image_path=[input_path], + image_path=[warmup_input_path], ) else: req = Req( @@ -237,6 +230,53 @@ class Scheduler: # if server is warmed-up, set this flag to avoid req-based warmup self.warmed_up = True + def _prepare_shared_warmup_image_path(self) -> str: + world_group = get_world_group() + src_rank = world_group.ranks[0] + + warmup_sync: dict[str, str | None] + if world_group.rank == src_rank: + try: + if self.server_args.input_save_path is not None: + uploads_dir = self.server_args.input_save_path + os.makedirs(uploads_dir, exist_ok=True) + else: + uploads_dir = tempfile.mkdtemp(prefix="sglang_input_") + warmup_image_base = os.path.join(uploads_dir, "warmup_image") + input_path = asyncio.run( + save_image_to_path( + MINIMUM_PICTURE_BASE64_FOR_WARMUP, + warmup_image_base, + ) + ) + warmup_sync = {"input_path": input_path, "error": None} + except Exception as e: + warmup_sync = {"input_path": None, "error": str(e)} + else: + warmup_sync = {} + + # Sync rank 0's warmup-image write result (path or error) to all ranks. + warmup_sync = broadcast_pyobj( + warmup_sync, + world_group.rank, + world_group.cpu_group, + src=src_rank, + ) + if not isinstance(warmup_sync, dict): + raise RuntimeError("Invalid warmup sync payload received across ranks") + + error = warmup_sync.get("error") + if error is not None: + raise RuntimeError( + f"Warmup image preparation failed on rank {src_rank}: {error}" + ) + + input_path = warmup_sync.get("input_path") + if not isinstance(input_path, str) or not input_path: + raise RuntimeError("Warmup image preparation returned empty input path") + + return input_path + def process_received_reqs_with_req_based_warmup( self, recv_reqs: List[tuple[bytes, Any]] ) -> List[tuple[bytes, Any]]: @@ -406,7 +446,7 @@ class Scheduler: f"Warmup req ({self._warmup_processed}/{self._warmup_total}) processing failed" ) else: - logger.info(f"Warmup req processing failed") + logger.info("Warmup req processing failed") # TODO: Support sending back to multiple identities if batched self.return_result(output_batch, identities[0], is_warmup=is_warmup)