From 16b3edc84fbef060dd0e37e08a7d066654296a57 Mon Sep 17 00:00:00 2001 From: Mick Date: Fri, 22 May 2026 08:54:48 +0800 Subject: [PATCH] [diffusion] feat: enable warmup for sglang serve by default (#25988) --- .../runtime/entrypoints/cli/serve.py | 3 + .../runtime/managers/scheduler.py | 6 + .../pipelines_core/stages/text_encoding.py | 332 +++++++++++------- .../multimodal_gen/test/server/gpu_cases.py | 1 - .../test/server/test_server_common.py | 3 - .../test/server/testcase_configs.py | 2 - .../test/unit/test_text_encoding_cache.py | 56 ++- 7 files changed, 257 insertions(+), 146 deletions(-) diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/cli/serve.py b/python/sglang/multimodal_gen/runtime/entrypoints/cli/serve.py index a5171d8e5..f5b2bbc35 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/cli/serve.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/cli/serve.py @@ -33,6 +33,9 @@ def add_multimodal_gen_serve_args(parser: argparse.ArgumentParser): def execute_serve_cmd(args: argparse.Namespace, unknown_args: list[str] | None = None): """The entry point for the serve command.""" server_args = ServerArgs.from_cli_args(args, unknown_args) + if not server_args.is_arg_explicitly_set("warmup"): + server_args.warmup = True + logger.info("Warmup is enabled by default for sglang serve.") dispatch_launch(server_args) diff --git a/python/sglang/multimodal_gen/runtime/managers/scheduler.py b/python/sglang/multimodal_gen/runtime/managers/scheduler.py index a51de4f8a..02b51718e 100644 --- a/python/sglang/multimodal_gen/runtime/managers/scheduler.py +++ b/python/sglang/multimodal_gen/runtime/managers/scheduler.py @@ -156,6 +156,7 @@ class Scheduler(SchedulerDisaggMixin): # warmup progress tracking self._warmup_total = 0 self._warmup_processed = 0 + self._logged_server_ready_after_warmup = False self.prepare_server_warmup_reqs() @@ -296,6 +297,11 @@ class Scheduler(SchedulerDisaggMixin): f"Warmup req processed in {GREEN}%.2f{RESET} seconds", total_duration_s, ) + if not self._logged_server_ready_after_warmup and ( + self._warmup_total <= 0 or self._warmup_processed >= self._warmup_total + ): + logger.info("The server is fired up and ready to roll!") + self._logged_server_ready_after_warmup = True else: if self._warmup_total > 0: logger.info( diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/text_encoding.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/text_encoding.py index 09f60d2b7..acb9387d0 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/text_encoding.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/text_encoding.py @@ -9,6 +9,7 @@ This module contains implementations of prompt encoding stages for diffusion pip import inspect from dataclasses import dataclass +from functools import lru_cache from typing import Any import torch @@ -34,6 +35,18 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger logger = init_logger(__name__) +@lru_cache(maxsize=1) +def get_model_default_negative_prompt( + model_path: str, backend: Any, model_id: str | None +): + from sglang.multimodal_gen.registry import get_model_info + + model_info = get_model_info(model_path, backend=backend, model_id=model_id) + if model_info is None: + return None + return model_info.sampling_param_cls().negative_prompt + + @dataclass(frozen=True) class TextEncodingFingerprint: prompt: Any @@ -104,69 +117,195 @@ class TextEncodingStage(PipelineStage): def get_or_compute_negative_text_embedding( self, batch: Req, server_args: ServerArgs, all_indices: list[int] ): + """Get the cached text embedding result or compute + + this is a one-slot cache for the model-default negative prompt: + most requests don't override the negative prompt, the cache hit rate is considerably high + """ negative_cache_key = self._build_negative_text_cache_key( batch, server_args, all_indices ) - use_negative_cache = not batch.is_warmup - cached_negative = None - if use_negative_cache: - cached_negative = ( - self._negative_text_cache_value - if self._negative_text_cache_key == negative_cache_key - else None - ) - if cached_negative is None: - ( - neg_embeds_list, - neg_masks_list, - neg_pooler_embeds_list, - neg_embeds_masks_list, - neg_seq_lens_list, - ) = self.encode_text( - batch.negative_prompt, - server_args, - encoder_index=all_indices, - return_attention_mask=True, - ) + cached_negative = self._get_cached_negative_text_embedding(negative_cache_key) + if cached_negative is not None: + return cached_negative - if use_negative_cache: - self._negative_text_cache_key = negative_cache_key - self._negative_text_cache_value = ( - tuple(neg_embeds_list), - tuple(neg_masks_list), - tuple(neg_pooler_embeds_list), - tuple(neg_embeds_masks_list), - tuple(neg_seq_lens_list), - ) - else: - ( - neg_embeds_list, - neg_masks_list, - neg_pooler_embeds_list, - neg_embeds_masks_list, - neg_seq_lens_list, - ) = cached_negative - return ( - neg_embeds_list, - neg_masks_list, - neg_pooler_embeds_list, - neg_embeds_masks_list, - neg_seq_lens_list, + negative_text_outputs = self.encode_text( + batch.negative_prompt, + server_args, + encoder_index=all_indices, + return_attention_mask=True, + ) + self._maybe_cache_negative_text_embedding( + negative_cache_key, negative_text_outputs + ) + return negative_text_outputs + + def _should_cache_negative_text_embedding( + self, batch: Req, server_args: ServerArgs + ) -> bool: + if not batch.is_warmup: + return True + return self._uses_model_default_negative_prompt(batch, server_args) + + def _get_cached_negative_text_embedding(self, negative_cache_key): + if negative_cache_key is None: + return None + if self._negative_text_cache_key == negative_cache_key: + return self._negative_text_cache_value + return None + + def _maybe_cache_negative_text_embedding( + self, + negative_cache_key, + negative_text_outputs, + ) -> None: + + # skip caching if None + if negative_cache_key is None: + return + self._negative_text_cache_key = negative_cache_key + self._negative_text_cache_value = tuple( + tuple(value) for value in negative_text_outputs ) def _build_negative_text_cache_key( self, batch: Req, server_args: ServerArgs, encoder_indices: list[int] ): + """if the current req doesn't worth caching, returns None""" + # skip if we don't cache for current req + if not self._should_cache_negative_text_embedding(batch, server_args): + return None + # Negative text encoding changes when the template or max length changes, # even if the visible negative prompt string is the same. return ( - server_args.pipeline_class_name, tuple(encoder_indices), self.freeze_for_dedup(batch.negative_prompt), self.freeze_for_dedup(batch.prompt_template), batch.max_sequence_length, ) + def _uses_model_default_negative_prompt( + self, batch: Req, server_args: ServerArgs + ) -> bool: + default_negative_prompt = self._get_model_default_negative_prompt(server_args) + if default_negative_prompt is None: + return False + return self._normalize_negative_prompt_for_default_match( + batch.negative_prompt + ) == self._normalize_negative_prompt_for_default_match(default_negative_prompt) + + def _get_model_default_negative_prompt(self, server_args: ServerArgs) -> str | None: + return get_model_default_negative_prompt( + server_args.model_path, + server_args.backend, + server_args.model_id, + ) + + @staticmethod + def _normalize_negative_prompt_for_default_match(value): + if isinstance(value, str) and not value.isspace(): + return value.strip() + return value + + def _append_positive_text_outputs( + self, + batch: Req, + prompt_embeds_list, + prompt_masks_list, + pooler_embeds_list, + prompt_embeds_masks_list, + prompt_seq_lens_list, + ) -> None: + for pe in prompt_embeds_list: + batch.prompt_embeds.append(pe) + + for pe in pooler_embeds_list: + batch.pooled_embeds.append(pe) + + if batch.prompt_attention_mask is None: + batch.prompt_attention_mask = [] + for am in prompt_masks_list: + batch.prompt_attention_mask.append(am) + + batch.prompt_embeds_mask = [] + batch.prompt_seq_lens = [] + for mask in prompt_embeds_masks_list: + batch.prompt_embeds_mask.append(mask) + for seq_lens in prompt_seq_lens_list: + batch.prompt_seq_lens.append(seq_lens) + + def _append_negative_text_outputs( + self, + batch: Req, + prompt_embeds_list, + neg_embeds_list, + neg_masks_list, + neg_pooler_embeds_list, + neg_embeds_masks_list, + neg_seq_lens_list, + ) -> None: + assert batch.negative_prompt_embeds is not None + + # a single negative prompt can be shared across positive prompts + target_batch_sizes = [pe.shape[0] for pe in prompt_embeds_list] + + def align_negative_batch_dim( + tensor: torch.Tensor, target_batch: int, name: str + ) -> torch.Tensor: + if tensor.shape[0] == target_batch: + return tensor + if tensor.shape[0] == 1 and target_batch > 1: + return tensor.expand(target_batch, *tensor.shape[1:]) + raise ValueError( + f"{name} batch dimension mismatch: got {tensor.shape[0]}, expected 1 or {target_batch}" + ) + + def align_negative_seq_lens( + seq_lens: list[int], target_batch: int, name: str + ) -> list[int]: + if len(seq_lens) == target_batch: + return [int(x) for x in seq_lens] + if len(seq_lens) == 1 and target_batch > 1: + return [int(seq_lens[0])] * target_batch + raise ValueError( + f"{name} batch dimension mismatch: got {len(seq_lens)}, expected 1 or {target_batch}" + ) + + for idx, ne in enumerate(neg_embeds_list): + target_batch = target_batch_sizes[min(idx, len(target_batch_sizes) - 1)] + ne = align_negative_batch_dim(ne, target_batch, "negative_prompt_embeds") + batch.negative_prompt_embeds.append(ne) + + for idx, pe in enumerate(neg_pooler_embeds_list): + target_batch = target_batch_sizes[min(idx, len(target_batch_sizes) - 1)] + pe = align_negative_batch_dim(pe, target_batch, "negative_pooled_embeds") + batch.neg_pooled_embeds.append(pe) + if batch.negative_attention_mask is None: + batch.negative_attention_mask = [] + for idx, nm in enumerate(neg_masks_list): + target_batch = target_batch_sizes[min(idx, len(target_batch_sizes) - 1)] + nm = align_negative_batch_dim( + nm, target_batch, "negative_attention_mask" + ) + batch.negative_attention_mask.append(nm) + + batch.negative_prompt_embeds_mask = [] + batch.negative_prompt_seq_lens = [] + for idx, nm in enumerate(neg_embeds_masks_list): + target_batch = target_batch_sizes[min(idx, len(target_batch_sizes) - 1)] + nm = align_negative_batch_dim( + nm, target_batch, "negative_prompt_embeds_mask" + ) + batch.negative_prompt_embeds_mask.append(nm) + for idx, seq_lens in enumerate(neg_seq_lens_list): + target_batch = target_batch_sizes[min(idx, len(target_batch_sizes) - 1)] + batch.negative_prompt_seq_lens.append( + align_negative_seq_lens( + seq_lens, target_batch, "negative_prompt_seq_lens" + ) + ) + @torch.no_grad() def forward( self, @@ -204,25 +343,6 @@ class TextEncodingStage(PipelineStage): max_length=max_seq_length, ) - for pe in prompt_embeds_list: - batch.prompt_embeds.append(pe) - - for pe in pooler_embeds_list: - batch.pooled_embeds.append(pe) - - if batch.prompt_attention_mask is None: - batch.prompt_attention_mask = [] - for am in prompt_masks_list: - batch.prompt_attention_mask.append(am) - - batch.prompt_embeds_mask = [] - batch.prompt_seq_lens = [] - for mask in prompt_embeds_masks_list: - batch.prompt_embeds_mask.append(mask) - for seq_lens in prompt_seq_lens_list: - batch.prompt_seq_lens.append(seq_lens) - - # Encode negative prompt if CFG is enabled if batch.do_classifier_free_guidance: assert isinstance(batch.negative_prompt, str) ( @@ -235,72 +355,26 @@ class TextEncodingStage(PipelineStage): batch, server_args, all_indices ) - assert batch.negative_prompt_embeds is not None + self._append_positive_text_outputs( + batch, + prompt_embeds_list, + prompt_masks_list, + pooler_embeds_list, + prompt_embeds_masks_list, + prompt_seq_lens_list, + ) - # A single negative prompt can be shared across positive prompts. - target_batch_sizes = [pe.shape[0] for pe in prompt_embeds_list] - - def align_negative_batch_dim( - tensor: torch.Tensor, target_batch: int, name: str - ) -> torch.Tensor: - if tensor.shape[0] == target_batch: - return tensor - if tensor.shape[0] == 1 and target_batch > 1: - return tensor.expand(target_batch, *tensor.shape[1:]) - raise ValueError( - f"{name} batch dimension mismatch: got {tensor.shape[0]}, expected 1 or {target_batch}" - ) - - def align_negative_seq_lens( - seq_lens: list[int], target_batch: int, name: str - ) -> list[int]: - if len(seq_lens) == target_batch: - return [int(x) for x in seq_lens] - if len(seq_lens) == 1 and target_batch > 1: - return [int(seq_lens[0])] * target_batch - raise ValueError( - f"{name} batch dimension mismatch: got {len(seq_lens)}, expected 1 or {target_batch}" - ) - - for idx, ne in enumerate(neg_embeds_list): - target_batch = target_batch_sizes[min(idx, len(target_batch_sizes) - 1)] - ne = align_negative_batch_dim( - ne, target_batch, "negative_prompt_embeds" - ) - batch.negative_prompt_embeds.append(ne) - - for idx, pe in enumerate(neg_pooler_embeds_list): - target_batch = target_batch_sizes[min(idx, len(target_batch_sizes) - 1)] - pe = align_negative_batch_dim( - pe, target_batch, "negative_pooled_embeds" - ) - batch.neg_pooled_embeds.append(pe) - if batch.negative_attention_mask is None: - batch.negative_attention_mask = [] - for idx, nm in enumerate(neg_masks_list): - target_batch = target_batch_sizes[ - min(idx, len(target_batch_sizes) - 1) - ] - nm = align_negative_batch_dim( - nm, target_batch, "negative_attention_mask" - ) - batch.negative_attention_mask.append(nm) - - batch.negative_prompt_embeds_mask = [] - batch.negative_prompt_seq_lens = [] - for idx, nm in enumerate(neg_embeds_masks_list): - target_batch = target_batch_sizes[min(idx, len(target_batch_sizes) - 1)] - nm = align_negative_batch_dim( - nm, target_batch, "negative_prompt_embeds_mask" - ) - batch.negative_prompt_embeds_mask.append(nm) - for idx, seq_lens in enumerate(neg_seq_lens_list): - target_batch = target_batch_sizes[min(idx, len(target_batch_sizes) - 1)] - batch.negative_prompt_seq_lens.append( - align_negative_seq_lens( - seq_lens, target_batch, "negative_prompt_seq_lens" - ) - ) + # Encode negative prompt if CFG is enabled + if batch.do_classifier_free_guidance: + self._append_negative_text_outputs( + batch, + prompt_embeds_list, + neg_embeds_list, + neg_masks_list, + neg_pooler_embeds_list, + neg_embeds_masks_list, + neg_seq_lens_list, + ) return batch diff --git a/python/sglang/multimodal_gen/test/server/gpu_cases.py b/python/sglang/multimodal_gen/test/server/gpu_cases.py index 0d3a8e44a..dab086cfc 100644 --- a/python/sglang/multimodal_gen/test/server/gpu_cases.py +++ b/python/sglang/multimodal_gen/test/server/gpu_cases.py @@ -388,7 +388,6 @@ if not current_platform.is_hip(): "hunyuan3d_shape_gen", DiffusionServerArgs( model_path="tencent/Hunyuan3D-2", - enable_warmup=False, ), HUNYUAN3D_SHAPE_sampling_params, run_consistency_check=False, 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 3635799f6..eb1329aad 100644 --- a/python/sglang/multimodal_gen/test/server/test_server_common.py +++ b/python/sglang/multimodal_gen/test/server/test_server_common.py @@ -128,9 +128,6 @@ 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 += " --warmup" - # Strict ports: fail immediately if port is occupied instead of silently # picking another one (which causes the test client to connect to the wrong server). extra_args += " --strict-ports" diff --git a/python/sglang/multimodal_gen/test/server/testcase_configs.py b/python/sglang/multimodal_gen/test/server/testcase_configs.py index a8ff76c3e..5d0bb5617 100644 --- a/python/sglang/multimodal_gen/test/server/testcase_configs.py +++ b/python/sglang/multimodal_gen/test/server/testcase_configs.py @@ -186,7 +186,6 @@ class DiffusionServerArgs: dit_offload_prefetch_size: int | float | None = None enable_cache_dit: bool = False text_encoder_cpu_offload: bool = False - enable_warmup: bool = True extras: list[str] = field(default_factory=lambda: []) env_vars: dict[str, str] = field(default_factory=dict) @@ -473,7 +472,6 @@ def _make_modelopt_ci_case( DiffusionServerArgs( model_path=model_path, modality=modality, - enable_warmup=False, extras=extras, env_vars=env_vars or {}, ), diff --git a/python/sglang/multimodal_gen/test/unit/test_text_encoding_cache.py b/python/sglang/multimodal_gen/test/unit/test_text_encoding_cache.py index 0b99d3efe..abcb84462 100644 --- a/python/sglang/multimodal_gen/test/unit/test_text_encoding_cache.py +++ b/python/sglang/multimodal_gen/test/unit/test_text_encoding_cache.py @@ -21,14 +21,18 @@ class DummyTextEncodingStage(TextEncodingStage): def encode_text(self, *args, **kwargs): self.calls += 1 - embeds = torch.full((1, 1, 1), float(self.calls)) - mask = torch.ones((1, 1), dtype=torch.int64) - return [embeds], [mask], [], [mask], [[1]] + text = args[0] + batch_size = len(text) if isinstance(text, list) else 1 + embeds = torch.full((batch_size, 1, 1), float(self.calls)) + mask = torch.ones((batch_size, 1), dtype=torch.int64) + return [embeds], [mask], [], [mask], [[1] * batch_size] def make_req(**kwargs): defaults = { + "prompt": "hello", "negative_prompt": "bad quality", + "do_classifier_free_guidance": True, "prompt_template": {"template": "{}"}, "max_sequence_length": 1024, "is_warmup": False, @@ -37,12 +41,30 @@ def make_req(**kwargs): return SimpleNamespace(**defaults) +def make_server_args(**kwargs): + defaults = { + "pipeline_class_name": "LTX2TwoStagePipeline", + "model_path": "dummy-model", + "backend": "auto", + "model_id": None, + "pipeline_config": SimpleNamespace(text_encoder_configs=[]), + } + defaults.update(kwargs) + return SimpleNamespace(**defaults) + + +def get_negative_embedding_twice(stage, server_args, first_req, second_req=None): + stage.get_or_compute_negative_text_embedding(first_req, server_args, [0]) + stage.get_or_compute_negative_text_embedding( + second_req if second_req is not None else make_req(), server_args, [0] + ) + + def test_negative_text_cache_key_tracks_encode_options(): stage = DummyTextEncodingStage() - server_args = SimpleNamespace(pipeline_class_name="LTX2TwoStagePipeline") + server_args = make_server_args() - stage.get_or_compute_negative_text_embedding(make_req(), server_args, [0]) - stage.get_or_compute_negative_text_embedding(make_req(), server_args, [0]) + get_negative_embedding_twice(stage, server_args, make_req()) assert stage.calls == 1 stage.get_or_compute_negative_text_embedding( @@ -58,11 +80,23 @@ def test_negative_text_cache_key_tracks_encode_options(): def test_negative_text_cache_skips_warmup(): stage = DummyTextEncodingStage() - server_args = SimpleNamespace(pipeline_class_name="LTX2TwoStagePipeline") + server_args = make_server_args() - stage.get_or_compute_negative_text_embedding( - make_req(is_warmup=True), server_args, [0] - ) - stage.get_or_compute_negative_text_embedding(make_req(), server_args, [0]) + with patch.object( + stage, "_get_model_default_negative_prompt", return_value="default negative" + ): + get_negative_embedding_twice(stage, server_args, make_req(is_warmup=True)) assert stage.calls == 2 + + +def test_negative_text_cache_keeps_default_warmup(): + stage = DummyTextEncodingStage() + server_args = make_server_args() + + with patch.object( + stage, "_get_model_default_negative_prompt", return_value="bad quality" + ): + get_negative_embedding_twice(stage, server_args, make_req(is_warmup=True)) + + assert stage.calls == 1