diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py b/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py index c778b19b8..d91804181 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py @@ -31,6 +31,8 @@ from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBa from sglang.multimodal_gen.runtime.scheduler_client import sync_scheduler_client from sglang.multimodal_gen.runtime.server_args import PortArgs, ServerArgs from sglang.multimodal_gen.runtime.utils.logging_utils import ( + GREEN, + RESET, init_logger, log_batch_completion, log_generation_timer, @@ -266,6 +268,13 @@ class DiffGenerator: log_batch_completion(logger, len(results), total_gen_time) if results: + if self.server_args.enable_warmup: + total_duration_ms = results[0]["timings"]["total_duration_ms"] + logger.info( + f"Warmed-up request processed in {GREEN}%.2f{RESET} seconds (with warmup excluded)", + total_duration_ms / 1000.0, + ) + peak_memories = [r.get("peak_memory_mb", 0) for r in results] if peak_memories: max_peak_memory = max(peak_memories) @@ -293,14 +302,10 @@ class DiffGenerator: # LoRA def _send_lora_request(self, req: Any, success_msg: str, failure_msg: str): response = sync_scheduler_client.forward(req) - if isinstance(response, dict) and response.get("status") == "ok": + if response.error is None: logger.info(success_msg) else: - error_msg = ( - response.get("message", "Unknown error") - if isinstance(response, dict) - else "Unknown response format" - ) + error_msg = response.error raise RuntimeError(f"{failure_msg}: {error_msg}") def set_lora( 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 2a9dc2892..8f7372ad5 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/openai/common_api.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/openai/common_api.py @@ -7,6 +7,7 @@ from sglang.multimodal_gen.runtime.entrypoints.openai.utils import ( SetLoraReq, UnmergeLoraWeightsReq, ) +from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch 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 @@ -17,15 +18,11 @@ logger = init_logger(__name__) async def _handle_lora_request(req: Any, success_msg: str, failure_msg: str): try: - response = await async_scheduler_client.forward(req) - if isinstance(response, dict) and response.get("status") == "ok": + output: OutputBatch = await async_scheduler_client.forward(req) + if output.error is None: return {"status": "ok", "message": success_msg} else: - error_msg = ( - response.get("message", "Unknown error") - if isinstance(response, dict) - else "Unknown response format" - ) + error_msg = output.error raise HTTPException(status_code=500, detail=f"{failure_msg}: {error_msg}") except Exception as e: if isinstance(e, HTTPException): diff --git a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py index 961d305d9..c024757e3 100644 --- a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py +++ b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py @@ -20,6 +20,7 @@ from sglang.multimodal_gen.runtime.distributed.parallel_state import ( ) from sglang.multimodal_gen.runtime.pipelines_core import ( ComposedPipelineBase, + LoRAPipeline, Req, build_pipeline, ) @@ -184,7 +185,7 @@ class GPUWorker: lora_path: str | None = None, target: str = "all", strength: float = 1.0, - ) -> None: + ) -> OutputBatch: """ Set the LoRA adapter for the pipeline. @@ -194,10 +195,14 @@ class GPUWorker: target: Which transformer(s) to apply the LoRA to. strength: LoRA strength for merge, default 1.0. """ - assert self.pipeline is not None + if not isinstance(self.pipeline, LoRAPipeline): + return OutputBatch(error="Lora is not enabled") self.pipeline.set_lora(lora_nickname, lora_path, target, strength) + return OutputBatch() - def merge_lora_weights(self, target: str = "all", strength: float = 1.0) -> None: + def merge_lora_weights( + self, target: str = "all", strength: float = 1.0 + ) -> OutputBatch: """ Merge LoRA weights. @@ -205,18 +210,22 @@ class GPUWorker: target: Which transformer(s) to merge. strength: LoRA strength for merge, default 1.0. """ - assert self.pipeline is not None + if not isinstance(self.pipeline, LoRAPipeline): + return OutputBatch(error="Lora is not enabled") self.pipeline.merge_lora_weights(target, strength) + return OutputBatch() - def unmerge_lora_weights(self, target: str = "all") -> None: + def unmerge_lora_weights(self, target: str = "all") -> OutputBatch: """ Unmerge LoRA weights. Args: target: Which transformer(s) to unmerge. """ - assert self.pipeline is not None + if not isinstance(self.pipeline, LoRAPipeline): + return OutputBatch(error="Lora is not enabled") self.pipeline.unmerge_lora_weights(target) + return OutputBatch() def run_scheduler_process( diff --git a/python/sglang/multimodal_gen/runtime/managers/scheduler.py b/python/sglang/multimodal_gen/runtime/managers/scheduler.py index cc931968f..0a01b806e 100644 --- a/python/sglang/multimodal_gen/runtime/managers/scheduler.py +++ b/python/sglang/multimodal_gen/runtime/managers/scheduler.py @@ -3,6 +3,7 @@ # SPDX-License-Identifier: Apache-2.0 import pickle from collections import deque +from copy import deepcopy from typing import Any, List import zmq @@ -22,7 +23,7 @@ from sglang.multimodal_gen.runtime.server_args import ( ) from sglang.multimodal_gen.runtime.utils.common import get_zmq_socket from sglang.multimodal_gen.runtime.utils.distributed import broadcast_pyobj -from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger +from sglang.multimodal_gen.runtime.utils.logging_utils import GREEN, RESET, init_logger logger = init_logger(__name__) @@ -82,30 +83,37 @@ class Scheduler: # FIFO, new reqs are appended self.waiting_queue: deque[tuple[bytes, Req]] = deque() - def _handle_set_lora(self, reqs: List[Any]): + self.warmed_up = False + + def _handle_set_lora(self, reqs: List[Any]) -> OutputBatch: # TODO: return set status + # TODO: return with SetLoRAResponse or something more appropriate req = reqs[0] - self.worker.set_lora(req.lora_nickname, req.lora_path, req.target, req.strength) - return {"status": "ok"} + return self.worker.set_lora( + req.lora_nickname, req.lora_path, req.target, req.strength + ) def _handle_merge_lora(self, reqs: List[Any]): req = reqs[0] - self.worker.merge_lora_weights(req.target, req.strength) - return {"status": "ok"} + return self.worker.merge_lora_weights(req.target, req.strength) - def _handle_unmerge_lora(self, reqs: List[Any]): + def _handle_unmerge_lora(self, reqs: List[Any]) -> OutputBatch: req = reqs[0] - self.worker.unmerge_lora_weights(req.target) - return {"status": "ok"} + return self.worker.unmerge_lora_weights(req.target) def _handle_generation(self, reqs: List[Req]): return self.worker.execute_forward(reqs) - def return_result(self, output_batch: OutputBatch, identity: bytes | None = None): + def return_result( + self, + output_batch: OutputBatch, + identity: bytes | None = None, + is_warmup: bool = False, + ): """ replies to client, only on rank 0 """ - if self.receiver is not None and identity is not None: + if not is_warmup and self.receiver is not None and identity is not None: self.receiver.send_multipart([identity, b"", pickle.dumps(output_batch)]) def get_next_batch_to_run(self) -> list[tuple[bytes, Req]] | None: @@ -124,18 +132,22 @@ class Scheduler: """ if self.receiver is not None: try: - identity, _, payload = self.receiver.recv_multipart() - recv_reqs = pickle.loads(payload) + try: + identity, _, payload = self.receiver.recv_multipart(zmq.NOBLOCK) + recv_reqs = pickle.loads(payload) + except zmq.Again: + recv_reqs = [] except zmq.ZMQError: # re-raise or handle appropriately to let the outer loop continue raise - # Ensure recv_reqs is a list - if not isinstance(recv_reqs, list): - recv_reqs = [recv_reqs] + if recv_reqs: + # Ensure recv_reqs is a list + if not isinstance(recv_reqs, list): + recv_reqs = [recv_reqs] - # Pack with identity for rank 0 - recv_reqs = [(identity, req) for req in recv_reqs] + # Pack with identity for rank 0 + recv_reqs = [(identity, req) for req in recv_reqs] else: recv_reqs = None @@ -166,6 +178,22 @@ class Scheduler: assert recv_reqs is not None + # handle server warmup by inserting an identical req to the beginning of the waiting queue + # only the very first req through server's lifetime will be warmup + if ( + not self.warmed_up + and len(recv_reqs) == 1 + and self.server_args.enable_warmup + ): + identity, req = recv_reqs[0] + if isinstance(req, Req): + warmup_req = deepcopy(req) + warmup_req.is_warmup = True + warmup_req.num_inference_steps = 1 + recv_reqs.insert(0, (identity, warmup_req)) + self.warmed_up = True + logger.info("Server warming up....") + return recv_reqs def event_loop(self) -> None: @@ -192,43 +220,51 @@ class Scheduler: continue # 2: execute, make sure a reply is always sent - while self.waiting_queue: - items = self.get_next_batch_to_run() - if not items: - break + items = self.get_next_batch_to_run() + if not items: + continue - identities = [item[0] for item in items] - reqs = [item[1] for item in items] + identities = [item[0] for item in items] + reqs = [item[1] for item in items] - try: - first_req = reqs[0] - handler = self.request_handlers.get(type(first_req)) - if handler: - output_batch = handler(reqs) - else: - output_batch = { - "status": "error", - "message": f"Unknown request type: {type(first_req)}", - } - except Exception as e: - logger.error( - f"Error executing request in scheduler event loop: {e}", - exc_info=True, + try: + processed_req = reqs[0] + handler = self.request_handlers.get(type(processed_req)) + if handler: + output_batch = handler(reqs) + else: + output_batch = OutputBatch( + error=f"Unknown request type: {type(processed_req)}" ) - # Determine appropriate error response format - output_batch = ( - OutputBatch(error=str(e)) - if reqs and isinstance(reqs[0], Req) - else {"status": "error", "message": str(e)} + except Exception as e: + logger.error( + f"Error executing request in scheduler event loop: {e}", + exc_info=True, + ) + # Determine appropriate error response format + output_batch = ( + OutputBatch(error=str(e)) + if reqs and isinstance(reqs[0], Req) + else OutputBatch(error=str(e)) + ) + + # 3. return results + try: + # TODO: Support sending back to multiple identities if batched + is_warmup = ( + processed_req.is_warmup if isinstance(processed_req, Req) else False + ) + if is_warmup: + logger.info( + f"Server warmup done in {GREEN}%.2f{RESET} seconds", + output_batch.timings.total_duration_s, ) - try: - # TODO: Support sending back to multiple identities if batched - self.return_result(output_batch, identities[0]) - except zmq.ZMQError as e: - # Reply failed; log and keep loop alive to accept future requests - logger.error(f"ZMQ error sending reply: {e}") - continue + self.return_result(output_batch, identities[0], is_warmup=is_warmup) + except zmq.ZMQError as e: + # Reply failed; log and keep loop alive to accept future requests + logger.error(f"ZMQ error sending reply: {e}") + continue logger.info("Scheduler event loop terminated.") if self.receiver is not None: @@ -243,16 +279,6 @@ class Scheduler: for pipe in self.task_pipes_to_slaves: pipe.send(task) - def _execute_on_rank0(self, payload: dict[str, Any]) -> dict[str, Any]: - """Execute task locally on the rank 0 worker.""" - method = payload["method"] - kwargs = {k: v for k, v in payload.items() if k != "method"} - handler = getattr(self.worker, method, None) - if handler: - result = handler(**kwargs) - return {"status": "ok", "result": result} - return {"status": "error", "error": f"Unknown method: {method}"} - def _collect_slave_results(self) -> List[dict[str, Any]]: """Collect results from all slave worker processes.""" results = [] diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/executors/pipeline_executor.py b/python/sglang/multimodal_gen/runtime/pipelines_core/executors/pipeline_executor.py index 8e87cc181..f77e06087 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/executors/pipeline_executor.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/executors/pipeline_executor.py @@ -81,8 +81,7 @@ class PipelineExecutor(ABC): """ Context manager for profiling execution. """ - do_profile = batch.profile - + do_profile = batch.profile and not batch.is_warmup if not do_profile: # fast forward yield 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 ceabdd476..56108b290 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/schedule_batch.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/schedule_batch.py @@ -159,6 +159,7 @@ class Req: # Misc save_output: bool = True return_frames: bool = False + is_warmup: bool = False # TeaCache parameters enable_teacache: bool = False @@ -234,6 +235,8 @@ class Req: return pprint.pformat(asdict(self), indent=2, width=120) def log(self, server_args: ServerArgs): + if self.is_warmup: + return # TODO: in some cases (e.g., TI2I), height and weight might be undecided at this moment if self.height: target_height = align_to(self.height, 16) diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py old mode 100755 new mode 100644 index e126af0c2..6ca92bd8b --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py @@ -93,9 +93,8 @@ class DenoisingStage(PipelineStage): # torch compile if self.server_args.enable_torch_compile: - self.torch_compile_module(self.transformer) - if transformer_2 is not None: - self.torch_compile_module(self.transformer_2) + for transformer in filter(None, [self.transformer, self.transformer_2]): + self.compile_module_with_torch_compile(transformer) self.scheduler = scheduler self.vae = vae @@ -115,8 +114,9 @@ class DenoisingStage(PipelineStage): # cache-dit state (for delayed mounting and idempotent control) self._cache_dit_enabled = False self._cached_num_steps = None + self._is_warmed_up = False - def torch_compile_module(self, module): + def compile_module_with_torch_compile(self, module): """ Compile a module's forward with torch.compile, and enable inductor overlap tweak if available. No-op if torch compile is disabled or the object has no forward. @@ -488,18 +488,14 @@ class DenoisingStage(PipelineStage): assert self.transformer is not None pipeline = self.pipeline() if self.pipeline else None if not server_args.model_loaded["transformer"]: + # FIXME: reuse more code loader = TransformerLoader() self.transformer = loader.load( - server_args.model_paths["transformer"], server_args + server_args.model_paths["transformer"], server_args, "transformer" ) - # enable cache-dit before torch.compile (delayed mounting) self._maybe_enable_cache_dit(batch.num_inference_steps) - - if self.server_args.enable_torch_compile: - self.transformer = torch.compile( - self.transformer, mode="max-autotune", fullgraph=True - ) + self.compile_module_with_torch_compile(self.transformer) if pipeline: pipeline.add_module("transformer", self.transformer) server_args.model_loaded["transformer"] = True @@ -666,6 +662,7 @@ class DenoisingStage(PipelineStage): trajectory_latents: list, trajectory_timesteps: list, server_args: ServerArgs, + is_warmup: bool = False, ): # Gather results if using sequence parallelism if trajectory_latents: @@ -702,14 +699,15 @@ class DenoisingStage(PipelineStage): # Save STA mask search results if needed if ( - self.attn_backend.get_enum() == AttentionBackendEnum.SLIDING_TILE_ATTN + not is_warmup + and self.attn_backend.get_enum() == AttentionBackendEnum.SLIDING_TILE_ATTN and server_args.STA_mode == STA_Mode.STA_SEARCHING ): self.save_sta_search_results(batch) # deallocate transformer if on mps pipeline = self.pipeline() if self.pipeline else None - if torch.backends.mps.is_available(): + if torch.backends.mps.is_available() and not is_warmup: logger.info( "Memory before deallocating transformer: %s", torch.mps.current_allocated_memory(), @@ -954,7 +952,7 @@ class DenoisingStage(PipelineStage): denoising_start_time = time.time() # to avoid device-sync caused by timestep comparison - + is_warmup = batch.is_warmup self.scheduler.set_begin_index(0) timesteps_cpu = timesteps.cpu() num_timesteps = timesteps_cpu.shape[0] @@ -1051,11 +1049,12 @@ class DenoisingStage(PipelineStage): ): progress_bar.update() - self.step_profile() + if not is_warmup: + self.step_profile() denoising_end_time = time.time() - if num_timesteps > 0: + if num_timesteps > 0 and not is_warmup: self.log_info( "average time per step: %.4f seconds", (denoising_end_time - denoising_start_time) / len(timesteps), @@ -1067,6 +1066,7 @@ class DenoisingStage(PipelineStage): trajectory_latents=trajectory_latents, trajectory_timesteps=trajectory_timesteps, server_args=server_args, + is_warmup=is_warmup, ) return batch diff --git a/python/sglang/multimodal_gen/runtime/server_args.py b/python/sglang/multimodal_gen/runtime/server_args.py index e89690fb7..de823e3bd 100644 --- a/python/sglang/multimodal_gen/runtime/server_args.py +++ b/python/sglang/multimodal_gen/runtime/server_args.py @@ -206,6 +206,7 @@ class ServerArgs: # Compilation enable_torch_compile: bool = False + enable_warmup: bool = False disable_autocast: bool | None = None @@ -456,6 +457,14 @@ class ServerArgs: help="Use torch.compile to speed up DiT inference." + "However, will likely cause precision drifts. See (https://github.com/pytorch/pytorch/issues/145213)", ) + parser.add_argument( + "--enable-warmup", + action=StoreBoolean, + default=ServerArgs.enable_warmup, + help="Perform a 1-step end-to-end warmup request before the actual request. " + "Recommended to enable when benchmarking to ensure fair comparison and best performance." + "When enabled, look for the line ending with `with warmup excluded` for actual processing time.", + ) parser.add_argument( "--dit-cpu-offload", action=StoreBoolean, diff --git a/python/sglang/multimodal_gen/runtime/utils/logging_utils.py b/python/sglang/multimodal_gen/runtime/utils/logging_utils.py index e9e1c3de9..4be477cc5 100644 --- a/python/sglang/multimodal_gen/runtime/utils/logging_utils.py +++ b/python/sglang/multimodal_gen/runtime/utils/logging_utils.py @@ -466,7 +466,10 @@ def log_generation_timer( yield timer timer.end_time = time.perf_counter() timer.duration = timer.end_time - timer.start_time - logger.info("Pixel data generated successfully in %.2f seconds", timer.duration) + logger.info( + f"Pixel data generated successfully in {GREEN}%.2f{RESET} seconds", + timer.duration, + ) except Exception as e: if request_idx is not None: logger.error( @@ -487,7 +490,7 @@ def log_batch_completion( logger: logging.Logger, num_outputs: int, total_time: float ) -> None: logger.info( - "Completed batch processing. Generated %d outputs in %.2f seconds.", + f"Completed batch processing. Generated %d outputs in {GREEN}%.2f{RESET} seconds", num_outputs, total_time, ) diff --git a/python/sglang/multimodal_gen/runtime/utils/perf_logger.py b/python/sglang/multimodal_gen/runtime/utils/perf_logger.py index d05fef2f5..6b7da272d 100644 --- a/python/sglang/multimodal_gen/runtime/utils/perf_logger.py +++ b/python/sglang/multimodal_gen/runtime/utils/perf_logger.py @@ -33,6 +33,10 @@ class RequestTimings: self.steps: list[float] = [] self.total_duration_ms: float = 0.0 + @property + def total_duration_s(self) -> float: + return self.total_duration_ms / 1000.0 + def record_stage(self, stage_name: str, duration_s: float): """Records the duration of a pipeline stage""" self.stages[stage_name] = duration_s * 1000 # Store as milliseconds diff --git a/python/sglang/multimodal_gen/test/cli/test_generate_common.py b/python/sglang/multimodal_gen/test/cli/test_generate_common.py index 88534f205..5303bed1e 100644 --- a/python/sglang/multimodal_gen/test/cli/test_generate_common.py +++ b/python/sglang/multimodal_gen/test/cli/test_generate_common.py @@ -24,19 +24,13 @@ logger = init_logger(__name__) class TestResult: name: str key: str - duration: Optional[float] succeed: bool - @property - def duration_str(self): - return f"{self.duration:.4f}" if self.duration else "NA" - def run_command(command) -> Optional[float]: """Runs a command and returns the execution time and status.""" print(f"Running command: {shlex.join(command)}") - duration = None with subprocess.Popen( command, stdout=subprocess.PIPE, @@ -46,15 +40,11 @@ def run_command(command) -> Optional[float]: ) as process: for line in process.stdout: sys.stdout.write(line) - if "Pixel data generated" in line: - words = line.split(" ") - duration = float(words[-2]) - - if process.returncode == 0: - return duration - else: + process.wait() + if process.returncode == 0: + return True print(f"Command failed with exit code {process.returncode}") - return None + return False class CLIBase(unittest.TestCase): @@ -80,13 +70,7 @@ class CLIBase(unittest.TestCase): f"--output-path={self.output_path}", ] - results = [] - - @classmethod - def setUpClass(cls): - cls.results = [] - - def _run_command(self, name: str, model_path: str, test_key: str = "", args=[]): + def _run_command(self, name: str, model_path: str, args=[]): command = ( self.get_base_command() + [f"--model-path={model_path}"] @@ -94,28 +78,21 @@ class CLIBase(unittest.TestCase): + ["--output-file-name", f"{name}"] + self.extra_args ) - duration = run_command(command) - status = "Success" if duration else "Failed" - succeed = duration is not None + succeed = run_command(command) + status = "Success" if succeed else "Failed" - duration = float(duration) if succeed else None - self.results.append(TestResult(name, test_key, duration, succeed)) - - return name, duration, status + return name, status def _run_test(self, name: str, args, model_path: str, test_key: str): - name, duration, status = self._run_command( - name, args=args, model_path=model_path, test_key=test_key - ) - self.verify(status, name, duration) + name, status = self._run_command(name, args=args, model_path=model_path) + self.verify(status, name) - def verify(self, status, name, duration): + def verify(self, status, name): print("-" * 80) print("\n" * 3) # test task status self.assertEqual(status, "Success", f"{name} command failed") - self.assertIsNotNone(duration, f"Could not parse duration for {name}") # test output file path = os.path.join( @@ -125,7 +102,6 @@ class CLIBase(unittest.TestCase): if self.data_type == DataType.IMAGE: with Image.open(path) as image: check_image_size(self, image, self.width, self.height) - logger.info(f"{name} passed in {duration:.4f}s") def model_name(self): return self.model_path.split("/")[-1] diff --git a/python/sglang/multimodal_gen/test/run_suite.py b/python/sglang/multimodal_gen/test/run_suite.py index fefa71a96..da9eb2667 100644 --- a/python/sglang/multimodal_gen/test/run_suite.py +++ b/python/sglang/multimodal_gen/test/run_suite.py @@ -198,7 +198,9 @@ def run_pytest(files, filter_expr=None): and "AssertionError" in full_output ) - is_flaky_ci_assertion = "SafetensorError" in full_output + is_flaky_ci_assertion = ( + "SafetensorError" in full_output or "FileNotFoundError" in full_output + ) is_oom_error = ( "out of memory" in full_output.lower() diff --git a/python/sglang/multimodal_gen/test/server/perf_baselines.json b/python/sglang/multimodal_gen/test/server/perf_baselines.json index f680c5eb1..c7948bac7 100644 --- a/python/sglang/multimodal_gen/test/server/perf_baselines.json +++ b/python/sglang/multimodal_gen/test/server/perf_baselines.json @@ -529,6 +529,31 @@ "expected_avg_denoise_ms": 94.15, "expected_median_denoise_ms": 102.03 }, + "zimage_image_t2i_warmup": { + "stages_ms": { + "InputValidationStage": 0.02, + "TextEncodingStage": 100.65, + "ConditioningStage": 0.01, + "TimestepPreparationStage": 0.98, + "LatentPreparationStage": 0.06, + "DenoisingStage": 889.42, + "DecodingStage": 37.81 + }, + "denoise_step_ms": { + "0": 16.49, + "1": 94.63, + "2": 109.65, + "3": 110.05, + "4": 109.39, + "5": 110.58, + "6": 109.52, + "7": 110.54, + "8": 115.24 + }, + "expected_e2e_ms": 1029.96, + "expected_avg_denoise_ms": 98.46, + "expected_median_denoise_ms": 109.65 + }, "qwen_image_edit_ti2i": { "stages_ms": { "InputValidationStage": 38.62, 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 969998aa1..60eba4758 100644 --- a/python/sglang/multimodal_gen/test/server/test_server_common.py +++ b/python/sglang/multimodal_gen/test/server/test_server_common.py @@ -83,6 +83,9 @@ def diffusion_server(case: DiffusionTestCase) -> ServerContext: if server_args.lora_path: extra_args += f" --lora-path {server_args.lora_path}" + if server_args.enable_warmup: + extra_args += f" --enable-warmup" + # Build custom environment variables env_vars = {} if server_args.enable_cache_dit: diff --git a/python/sglang/multimodal_gen/test/server/testcase_configs.py b/python/sglang/multimodal_gen/test/server/testcase_configs.py index c3a3763b4..0d4ba93bb 100644 --- a/python/sglang/multimodal_gen/test/server/testcase_configs.py +++ b/python/sglang/multimodal_gen/test/server/testcase_configs.py @@ -151,6 +151,8 @@ class DiffusionServerArgs: ring_degree: int | None = None # LoRA lora_path: str | None = None # LoRA adapter path (HF repo or local path) + # misc + enable_warmup: bool = False dit_layerwise_offload: bool = False enable_cache_dit: bool = False @@ -359,6 +361,13 @@ ONE_GPU_CASES_A: list[DiffusionTestCase] = [ ), T2I_sampling_params, ), + DiffusionTestCase( + "zimage_image_t2i_warmup", + DiffusionServerArgs( + model_path="Tongyi-MAI/Z-Image-Turbo", modality="image", enable_warmup=True + ), + T2I_sampling_params, + ), # === Text and Image to Image (TI2I) === DiffusionTestCase( "qwen_image_edit_ti2i",