[diffusion] feat: cache encoder results for default negative prompt (#24304)
This commit is contained in:
@@ -41,6 +41,16 @@ class TextEncodingFingerprint:
|
|||||||
max_sequence_length: int | None
|
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):
|
class TextEncodingStage(PipelineStage):
|
||||||
"""
|
"""
|
||||||
Stage for encoding text prompts into embeddings for diffusion models.
|
Stage for encoding text prompts into embeddings for diffusion models.
|
||||||
@@ -73,6 +83,8 @@ class TextEncodingStage(PipelineStage):
|
|||||||
super().__init__()
|
super().__init__()
|
||||||
self.tokenizers = tokenizers
|
self.tokenizers = tokenizers
|
||||||
self.text_encoders = text_encoders
|
self.text_encoders = text_encoders
|
||||||
|
self._negative_text_cache_key = None
|
||||||
|
self._negative_text_cache_value = None
|
||||||
|
|
||||||
def component_uses(
|
def component_uses(
|
||||||
self, server_args: ServerArgs, stage_name: str | None = None
|
self, server_args: ServerArgs, stage_name: str | None = None
|
||||||
@@ -87,6 +99,72 @@ class TextEncodingStage(PipelineStage):
|
|||||||
for i in range(len(self.text_encoders))
|
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()
|
@torch.no_grad()
|
||||||
def forward(
|
def forward(
|
||||||
self,
|
self,
|
||||||
@@ -147,11 +225,8 @@ class TextEncodingStage(PipelineStage):
|
|||||||
neg_pooler_embeds_list,
|
neg_pooler_embeds_list,
|
||||||
neg_embeds_masks_list,
|
neg_embeds_masks_list,
|
||||||
neg_seq_lens_list,
|
neg_seq_lens_list,
|
||||||
) = self.encode_text(
|
) = self.get_or_compute_negative_text_embedding(
|
||||||
batch.negative_prompt,
|
batch, server_args, all_indices
|
||||||
server_args,
|
|
||||||
encoder_index=all_indices,
|
|
||||||
return_attention_mask=True,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
assert batch.negative_prompt_embeds is not None
|
assert batch.negative_prompt_embeds is not None
|
||||||
@@ -319,8 +394,8 @@ class TextEncodingStage(PipelineStage):
|
|||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Depending on return_type and return_attention_mask:
|
Depending on return_type and return_attention_mask:
|
||||||
- list: List[Tensor] or
|
- list: (embeds, pooler_outputs) or
|
||||||
(embeds, attention_masks, pooled_embeds, embeds_masks, seq_lens)
|
(embeds, attention_masks, pooler_outputs, embeds_masks, seq_lens)
|
||||||
- dict: Dict[str, Tensor] or (Dict[str, Tensor], Dict[str, Tensor])
|
- dict: Dict[str, Tensor] or (Dict[str, Tensor], Dict[str, Tensor])
|
||||||
- stack: Tensor of shape [num_encoders, ...] or a tuple with stacked
|
- stack: Tensor of shape [num_encoders, ...] or a tuple with stacked
|
||||||
attention masks
|
attention masks
|
||||||
@@ -335,14 +410,14 @@ class TextEncodingStage(PipelineStage):
|
|||||||
)
|
)
|
||||||
|
|
||||||
# Resolve selection into indices
|
# Resolve selection into indices
|
||||||
encoder_cfgs = server_args.pipeline_config.text_encoder_configs
|
|
||||||
if encoder_index is None:
|
if encoder_index is None:
|
||||||
indices: list[int] = [0]
|
indices: list[int] = [0]
|
||||||
elif isinstance(encoder_index, int):
|
elif isinstance(encoder_index, int):
|
||||||
indices = [encoder_index]
|
indices = [encoder_index]
|
||||||
else:
|
else:
|
||||||
indices = list(encoder_index)
|
indices = list(encoder_index)
|
||||||
# validate range
|
|
||||||
|
# Validate indices are within range
|
||||||
num_encoders = len(self.text_encoders)
|
num_encoders = len(self.text_encoders)
|
||||||
for idx in indices:
|
for idx in indices:
|
||||||
if idx < 0 or idx >= num_encoders:
|
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}]"
|
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]
|
# Normalize input to list[str]
|
||||||
assert isinstance(text, str | list)
|
assert isinstance(text, str | list)
|
||||||
if isinstance(text, str):
|
if isinstance(text, str):
|
||||||
@@ -545,14 +617,7 @@ class TextEncodingStage(PipelineStage):
|
|||||||
return embeds_dict
|
return embeds_dict
|
||||||
|
|
||||||
# return_type == "stack"
|
# return_type == "stack"
|
||||||
# Validate shapes are compatible
|
stacked_embeds = stack_tensors("embeddings", embeds_list)
|
||||||
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)
|
|
||||||
if return_attention_mask:
|
if return_attention_mask:
|
||||||
stackable_masks = [
|
stackable_masks = [
|
||||||
(
|
(
|
||||||
@@ -564,13 +629,7 @@ class TextEncodingStage(PipelineStage):
|
|||||||
)
|
)
|
||||||
for embed, mask in zip(embeds_list, attn_masks_list, strict=True)
|
for embed, mask in zip(embeds_list, attn_masks_list, strict=True)
|
||||||
]
|
]
|
||||||
base_mask_shape = list(stackable_masks[0].shape)
|
stacked_masks = stack_tensors("attention masks", stackable_masks)
|
||||||
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)
|
|
||||||
return stacked_embeds, stacked_masks
|
return stacked_embeds, stacked_masks
|
||||||
return stacked_embeds
|
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
|
||||||
Reference in New Issue
Block a user