[diffusion] feat: enable warmup for sglang serve by default (#25988)

This commit is contained in:
Mick
2026-05-22 08:54:48 +08:00
committed by GitHub
parent f6d98a17ba
commit 16b3edc84f
7 changed files with 257 additions and 146 deletions
@@ -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)
@@ -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(
@@ -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
@@ -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,
@@ -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"
@@ -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 {},
),
@@ -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