[diffusion] feat: cache encoder results for default negative prompt (#24304)

This commit is contained in:
Mick
2026-05-05 11:56:01 +08:00
committed by GitHub
parent e483e60b72
commit 04926e1d9f
2 changed files with 154 additions and 27 deletions
@@ -41,6 +41,16 @@ class TextEncodingFingerprint:
max_sequence_length: int | None
def stack_tensors(name: str, tensors: list[torch.Tensor]) -> torch.Tensor:
base_shape = list(tensors[0].shape)
for tensor in tensors[1:]:
if list(tensor.shape) != base_shape:
raise ValueError(
f"Cannot stack {name} with differing shapes: {[list(t.shape) for t in tensors]}"
)
return torch.stack(tensors, dim=0)
class TextEncodingStage(PipelineStage):
"""
Stage for encoding text prompts into embeddings for diffusion models.
@@ -73,6 +83,8 @@ class TextEncodingStage(PipelineStage):
super().__init__()
self.tokenizers = tokenizers
self.text_encoders = text_encoders
self._negative_text_cache_key = None
self._negative_text_cache_value = None
def component_uses(
self, server_args: ServerArgs, stage_name: str | None = None
@@ -87,6 +99,72 @@ class TextEncodingStage(PipelineStage):
for i in range(len(self.text_encoders))
]
def get_or_compute_negative_text_embedding(
self, batch: Req, server_args: ServerArgs, all_indices: list[int]
):
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,
)
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,
)
def _build_negative_text_cache_key(
self, batch: Req, server_args: ServerArgs, encoder_indices: list[int]
):
# 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,
)
@torch.no_grad()
def forward(
self,
@@ -147,11 +225,8 @@ class TextEncodingStage(PipelineStage):
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,
) = self.get_or_compute_negative_text_embedding(
batch, server_args, all_indices
)
assert batch.negative_prompt_embeds is not None
@@ -319,8 +394,8 @@ class TextEncodingStage(PipelineStage):
Returns:
Depending on return_type and return_attention_mask:
- list: List[Tensor] or
(embeds, attention_masks, pooled_embeds, embeds_masks, seq_lens)
- list: (embeds, pooler_outputs) or
(embeds, attention_masks, pooler_outputs, embeds_masks, seq_lens)
- dict: Dict[str, Tensor] or (Dict[str, Tensor], Dict[str, Tensor])
- stack: Tensor of shape [num_encoders, ...] or a tuple with stacked
attention masks
@@ -335,14 +410,14 @@ class TextEncodingStage(PipelineStage):
)
# Resolve selection into indices
encoder_cfgs = server_args.pipeline_config.text_encoder_configs
if encoder_index is None:
indices: list[int] = [0]
elif isinstance(encoder_index, int):
indices = [encoder_index]
else:
indices = list(encoder_index)
# validate range
# Validate indices are within range
num_encoders = len(self.text_encoders)
for idx in indices:
if idx < 0 or idx >= num_encoders:
@@ -350,9 +425,6 @@ class TextEncodingStage(PipelineStage):
f"encoder index {idx} out of range [0, {num_encoders - 1}]"
)
# Validate indices are within range
num_encoders = len(self.text_encoders)
# Normalize input to list[str]
assert isinstance(text, str | list)
if isinstance(text, str):
@@ -545,14 +617,7 @@ class TextEncodingStage(PipelineStage):
return embeds_dict
# return_type == "stack"
# Validate shapes are compatible
base_shape = list(embeds_list[0].shape)
for t in embeds_list[1:]:
if list(t.shape) != base_shape:
raise ValueError(
f"Cannot stack embeddings with differing shapes: {[list(t.shape) for t in embeds_list]}"
)
stacked_embeds = torch.stack(embeds_list, dim=0)
stacked_embeds = stack_tensors("embeddings", embeds_list)
if return_attention_mask:
stackable_masks = [
(
@@ -564,13 +629,7 @@ class TextEncodingStage(PipelineStage):
)
for embed, mask in zip(embeds_list, attn_masks_list, strict=True)
]
base_mask_shape = list(stackable_masks[0].shape)
for m in stackable_masks[1:]:
if list(m.shape) != base_mask_shape:
raise ValueError(
f"Cannot stack attention masks with differing shapes: {[list(m.shape) for m in stackable_masks]}"
)
stacked_masks = torch.stack(stackable_masks, dim=0)
stacked_masks = stack_tensors("attention masks", stackable_masks)
return stacked_embeds, stacked_masks
return stacked_embeds
@@ -0,0 +1,68 @@
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import torch
from sglang.multimodal_gen.runtime.pipelines_core.stages.text_encoding import (
TextEncodingStage,
)
_GLOBAL_ARGS_PATCH = (
"sglang.multimodal_gen.runtime.pipelines_core.stages.base.get_global_server_args"
)
class DummyTextEncodingStage(TextEncodingStage):
def __init__(self):
with patch(_GLOBAL_ARGS_PATCH) as mock_global_args:
mock_global_args.return_value = MagicMock()
super().__init__(text_encoders=[], tokenizers=[])
self.calls = 0
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]]
def make_req(**kwargs):
defaults = {
"negative_prompt": "bad quality",
"prompt_template": {"template": "{}"},
"max_sequence_length": 1024,
"is_warmup": False,
}
defaults.update(kwargs)
return SimpleNamespace(**defaults)
def test_negative_text_cache_key_tracks_encode_options():
stage = DummyTextEncodingStage()
server_args = SimpleNamespace(pipeline_class_name="LTX2TwoStagePipeline")
stage.get_or_compute_negative_text_embedding(make_req(), server_args, [0])
stage.get_or_compute_negative_text_embedding(make_req(), server_args, [0])
assert stage.calls == 1
stage.get_or_compute_negative_text_embedding(
make_req(max_sequence_length=512), server_args, [0]
)
assert stage.calls == 2
stage.get_or_compute_negative_text_embedding(
make_req(prompt_template={"template": "negative: {}"}), server_args, [0]
)
assert stage.calls == 3
def test_negative_text_cache_skips_warmup():
stage = DummyTextEncodingStage()
server_args = SimpleNamespace(pipeline_class_name="LTX2TwoStagePipeline")
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])
assert stage.calls == 2