[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): def execute_serve_cmd(args: argparse.Namespace, unknown_args: list[str] | None = None):
"""The entry point for the serve command.""" """The entry point for the serve command."""
server_args = ServerArgs.from_cli_args(args, unknown_args) 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) dispatch_launch(server_args)
@@ -156,6 +156,7 @@ class Scheduler(SchedulerDisaggMixin):
# warmup progress tracking # warmup progress tracking
self._warmup_total = 0 self._warmup_total = 0
self._warmup_processed = 0 self._warmup_processed = 0
self._logged_server_ready_after_warmup = False
self.prepare_server_warmup_reqs() self.prepare_server_warmup_reqs()
@@ -296,6 +297,11 @@ class Scheduler(SchedulerDisaggMixin):
f"Warmup req processed in {GREEN}%.2f{RESET} seconds", f"Warmup req processed in {GREEN}%.2f{RESET} seconds",
total_duration_s, 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: else:
if self._warmup_total > 0: if self._warmup_total > 0:
logger.info( logger.info(
@@ -9,6 +9,7 @@ This module contains implementations of prompt encoding stages for diffusion pip
import inspect import inspect
from dataclasses import dataclass from dataclasses import dataclass
from functools import lru_cache
from typing import Any from typing import Any
import torch import torch
@@ -34,6 +35,18 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__) 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) @dataclass(frozen=True)
class TextEncodingFingerprint: class TextEncodingFingerprint:
prompt: Any prompt: Any
@@ -104,69 +117,195 @@ class TextEncodingStage(PipelineStage):
def get_or_compute_negative_text_embedding( def get_or_compute_negative_text_embedding(
self, batch: Req, server_args: ServerArgs, all_indices: list[int] 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( negative_cache_key = self._build_negative_text_cache_key(
batch, server_args, all_indices batch, server_args, all_indices
) )
use_negative_cache = not batch.is_warmup cached_negative = self._get_cached_negative_text_embedding(negative_cache_key)
cached_negative = None if cached_negative is not None:
if use_negative_cache: return cached_negative
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,
)
if use_negative_cache: negative_text_outputs = self.encode_text(
self._negative_text_cache_key = negative_cache_key batch.negative_prompt,
self._negative_text_cache_value = ( server_args,
tuple(neg_embeds_list), encoder_index=all_indices,
tuple(neg_masks_list), return_attention_mask=True,
tuple(neg_pooler_embeds_list), )
tuple(neg_embeds_masks_list), self._maybe_cache_negative_text_embedding(
tuple(neg_seq_lens_list), negative_cache_key, negative_text_outputs
) )
else: return negative_text_outputs
(
neg_embeds_list, def _should_cache_negative_text_embedding(
neg_masks_list, self, batch: Req, server_args: ServerArgs
neg_pooler_embeds_list, ) -> bool:
neg_embeds_masks_list, if not batch.is_warmup:
neg_seq_lens_list, return True
) = cached_negative return self._uses_model_default_negative_prompt(batch, server_args)
return (
neg_embeds_list, def _get_cached_negative_text_embedding(self, negative_cache_key):
neg_masks_list, if negative_cache_key is None:
neg_pooler_embeds_list, return None
neg_embeds_masks_list, if self._negative_text_cache_key == negative_cache_key:
neg_seq_lens_list, 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( def _build_negative_text_cache_key(
self, batch: Req, server_args: ServerArgs, encoder_indices: list[int] 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, # Negative text encoding changes when the template or max length changes,
# even if the visible negative prompt string is the same. # even if the visible negative prompt string is the same.
return ( return (
server_args.pipeline_class_name,
tuple(encoder_indices), tuple(encoder_indices),
self.freeze_for_dedup(batch.negative_prompt), self.freeze_for_dedup(batch.negative_prompt),
self.freeze_for_dedup(batch.prompt_template), self.freeze_for_dedup(batch.prompt_template),
batch.max_sequence_length, 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() @torch.no_grad()
def forward( def forward(
self, self,
@@ -204,25 +343,6 @@ class TextEncodingStage(PipelineStage):
max_length=max_seq_length, 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: if batch.do_classifier_free_guidance:
assert isinstance(batch.negative_prompt, str) assert isinstance(batch.negative_prompt, str)
( (
@@ -235,72 +355,26 @@ class TextEncodingStage(PipelineStage):
batch, server_args, all_indices 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. # Encode negative prompt if CFG is enabled
target_batch_sizes = [pe.shape[0] for pe in prompt_embeds_list] if batch.do_classifier_free_guidance:
self._append_negative_text_outputs(
def align_negative_batch_dim( batch,
tensor: torch.Tensor, target_batch: int, name: str prompt_embeds_list,
) -> torch.Tensor: neg_embeds_list,
if tensor.shape[0] == target_batch: neg_masks_list,
return tensor neg_pooler_embeds_list,
if tensor.shape[0] == 1 and target_batch > 1: neg_embeds_masks_list,
return tensor.expand(target_batch, *tensor.shape[1:]) neg_seq_lens_list,
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"
)
)
return batch return batch
@@ -388,7 +388,6 @@ if not current_platform.is_hip():
"hunyuan3d_shape_gen", "hunyuan3d_shape_gen",
DiffusionServerArgs( DiffusionServerArgs(
model_path="tencent/Hunyuan3D-2", model_path="tencent/Hunyuan3D-2",
enable_warmup=False,
), ),
HUNYUAN3D_SHAPE_sampling_params, HUNYUAN3D_SHAPE_sampling_params,
run_consistency_check=False, run_consistency_check=False,
@@ -128,9 +128,6 @@ def diffusion_server(case: DiffusionTestCase) -> ServerContext:
if server_args.lora_path: if server_args.lora_path:
extra_args += f" --lora-path {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 # 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). # picking another one (which causes the test client to connect to the wrong server).
extra_args += " --strict-ports" extra_args += " --strict-ports"
@@ -186,7 +186,6 @@ class DiffusionServerArgs:
dit_offload_prefetch_size: int | float | None = None dit_offload_prefetch_size: int | float | None = None
enable_cache_dit: bool = False enable_cache_dit: bool = False
text_encoder_cpu_offload: bool = False text_encoder_cpu_offload: bool = False
enable_warmup: bool = True
extras: list[str] = field(default_factory=lambda: []) extras: list[str] = field(default_factory=lambda: [])
env_vars: dict[str, str] = field(default_factory=dict) env_vars: dict[str, str] = field(default_factory=dict)
@@ -473,7 +472,6 @@ def _make_modelopt_ci_case(
DiffusionServerArgs( DiffusionServerArgs(
model_path=model_path, model_path=model_path,
modality=modality, modality=modality,
enable_warmup=False,
extras=extras, extras=extras,
env_vars=env_vars or {}, env_vars=env_vars or {},
), ),
@@ -21,14 +21,18 @@ class DummyTextEncodingStage(TextEncodingStage):
def encode_text(self, *args, **kwargs): def encode_text(self, *args, **kwargs):
self.calls += 1 self.calls += 1
embeds = torch.full((1, 1, 1), float(self.calls)) text = args[0]
mask = torch.ones((1, 1), dtype=torch.int64) batch_size = len(text) if isinstance(text, list) else 1
return [embeds], [mask], [], [mask], [[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): def make_req(**kwargs):
defaults = { defaults = {
"prompt": "hello",
"negative_prompt": "bad quality", "negative_prompt": "bad quality",
"do_classifier_free_guidance": True,
"prompt_template": {"template": "{}"}, "prompt_template": {"template": "{}"},
"max_sequence_length": 1024, "max_sequence_length": 1024,
"is_warmup": False, "is_warmup": False,
@@ -37,12 +41,30 @@ def make_req(**kwargs):
return SimpleNamespace(**defaults) 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(): def test_negative_text_cache_key_tracks_encode_options():
stage = DummyTextEncodingStage() 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]) get_negative_embedding_twice(stage, server_args, make_req())
stage.get_or_compute_negative_text_embedding(make_req(), server_args, [0])
assert stage.calls == 1 assert stage.calls == 1
stage.get_or_compute_negative_text_embedding( 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(): def test_negative_text_cache_skips_warmup():
stage = DummyTextEncodingStage() stage = DummyTextEncodingStage()
server_args = SimpleNamespace(pipeline_class_name="LTX2TwoStagePipeline") server_args = make_server_args()
stage.get_or_compute_negative_text_embedding( with patch.object(
make_req(is_warmup=True), server_args, [0] stage, "_get_model_default_negative_prompt", return_value="default negative"
) ):
stage.get_or_compute_negative_text_embedding(make_req(), server_args, [0]) get_negative_embedding_twice(stage, server_args, make_req(is_warmup=True))
assert stage.calls == 2 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