[diffusion] feat: enable warmup for sglang serve by default (#25988)
This commit is contained in:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user