Shrink text embeds to max token length for LTX-2. Drastically reduces cached text embedding sizes

This commit is contained in:
Jaret Burkett 2026-01-28 12:54:49 -07:00
parent ea912d2d7b
commit 1ce2428722
7 changed files with 130 additions and 27 deletions

View File

@ -213,6 +213,9 @@ class LTX2Model(BaseModel):
# use the new format on this new model by default
self.use_old_lokr_format = False
self.audio_processor = None
# gemma needs left side padding
self.te_padding_side = "left"
# static method to get the noise scheduler
@staticmethod
@ -627,6 +630,10 @@ class LTX2Model(BaseModel):
tile_sample_stride_width=224,
tile_sample_stride_num_frames=4,
)
# We only encode and store the minimum prompt tokens, but need them padded to 1024 for LTX2
conditional_embeds = self.pad_embeds(conditional_embeds)
unconditional_embeds = self.pad_embeds(unconditional_embeds)
video, audio = pipeline(
prompt_embeds=conditional_embeds.text_embeds.to(
@ -731,6 +738,29 @@ class LTX2Model(BaseModel):
latents_std = self.pipeline.audio_vae.latents_std
output_tensor = (output_tensor - latents_mean) / latents_std
return output_tensor
def pad_embeds(self, embeds: PromptEmbeds):
# ltx-2 connector requires 1024 tokens for good results. Any smaller and it degrades.
target_length = 1024
current_length = embeds.text_embeds.shape[1]
if current_length < target_length:
pad_length = target_length - current_length
pad_tensor = torch.zeros(
(embeds.text_embeds.shape[0], pad_length, embeds.text_embeds.shape[2]),
device=embeds.text_embeds.device,
dtype=embeds.text_embeds.dtype,
)
embeds.text_embeds = torch.cat([pad_tensor, embeds.text_embeds], dim=1)
if embeds.attention_mask is not None:
pad_mask = torch.zeros(
(embeds.attention_mask.shape[0], pad_length),
device=embeds.attention_mask.device,
dtype=embeds.attention_mask.dtype,
)
embeds.attention_mask = torch.cat(
[pad_mask, embeds.attention_mask], dim=1
)
return embeds
def get_noise_prediction(
self,
@ -743,6 +773,9 @@ class LTX2Model(BaseModel):
with torch.no_grad():
if self.model.device == torch.device("cpu"):
self.model.to(self.device_torch)
# We only encode and store the minimum prompt tokens, but need them padded to 1024 for LTX2
text_embeddings = self.pad_embeds(text_embeddings)
batch_size, C, latent_num_frames, latent_height, latent_width = (
latent_model_input.shape
@ -916,11 +949,58 @@ class LTX2Model(BaseModel):
if self.pipeline.text_encoder.device != self.device_torch:
self.pipeline.text_encoder.to(self.device_torch)
prompt_embeds, prompt_attention_mask, _, _ = self.pipeline.encode_prompt(
device = self.device_torch
scale_factor = 8
batch_size = len(prompt)
# Gemma expects left padding for chat-style prompts
self.tokenizer[0].padding_side = "left"
if self.tokenizer[0].pad_token is None:
self.tokenizer[0].pad_token = self.tokenizer[0].eos_token
prompt = [p.strip() for p in prompt]
text_inputs = self.tokenizer[0](
prompt,
do_classifier_free_guidance=False,
device=self.device_torch,
# padding="max_length",
padding="longest",
max_length=1024,
truncation=True,
add_special_tokens=True,
return_tensors="pt",
)
text_input_ids = text_inputs.input_ids
prompt_attention_mask = text_inputs.attention_mask
text_input_ids = text_input_ids.to(device)
prompt_attention_mask = prompt_attention_mask.to(device)
text_encoder_outputs = self.text_encoder[0](
input_ids=text_input_ids,
attention_mask=prompt_attention_mask,
output_hidden_states=True,
)
text_encoder_hidden_states = text_encoder_outputs.hidden_states
text_encoder_hidden_states = torch.stack(text_encoder_hidden_states, dim=-1)
sequence_lengths = prompt_attention_mask.sum(dim=-1)
prompt_embeds = self.pipeline._pack_text_embeds(
text_encoder_hidden_states,
sequence_lengths,
device=device,
padding_side=self.tokenizer[0].padding_side,
scale_factor=scale_factor,
)
prompt_embeds = prompt_embeds.to(dtype=self.torch_dtype)
# duplicate text embeddings for each generation per prompt, using mps friendly method
_, seq_len, _ = prompt_embeds.shape
prompt_embeds = prompt_embeds.repeat(1, 1, 1)
prompt_embeds = prompt_embeds.view(
batch_size * 1, seq_len, -1
)
prompt_attention_mask = prompt_attention_mask.view(batch_size, -1)
prompt_attention_mask = prompt_attention_mask.repeat(1, 1)
pe = PromptEmbeds([prompt_embeds, None])
pe.attention_mask = prompt_attention_mask
return pe

View File

@ -487,6 +487,23 @@ class AiToolkitDataset(LatentCachingMixin, ControlCachingMixin, CLIPCachingMixin
self.size_database["__version__"] = dataloader_version
# set latent space version
latent_space_version = "sd1"
if self.sd is not None and self.sd.model_config.latent_space_version is not None:
latent_space_version = self.sd.model_config.latent_space_version
elif self.sd.is_xl:
latent_space_version = 'sdxl'
elif self.sd.is_v3:
latent_space_version = 'sd3'
elif self.sd.is_auraflow:
latent_space_version = 'sdxl'
elif self.sd.is_flux:
latent_space_version = 'flux1'
elif self.sd.model_config.is_pixart_sigma:
latent_space_version = 'sdxl'
else:
latent_space_version = self.sd.model_config.arch if self.sd is not None else "sd1"
bad_count = 0
for file in tqdm(file_list):
try:
@ -498,6 +515,9 @@ class AiToolkitDataset(LatentCachingMixin, ControlCachingMixin, CLIPCachingMixin
size_database=self.size_database,
dataset_root=dataset_folder,
encode_control_in_text_embeddings=self.sd.encode_control_in_text_embeddings if self.sd else False,
text_embedding_space_version=self.sd.model_config.arch if self.sd else "sd1",
te_padding_side=self.sd.te_padding_side if self.sd else "right",
latent_space_version=latent_space_version,
)
self.file_list.append(file_item)
except Exception as e:

View File

@ -60,6 +60,9 @@ class FileItemDTO(
self.encode_control_in_text_embeddings = kwargs.get(
"encode_control_in_text_embeddings", False
)
self.te_padding_side = kwargs.get("te_padding_side", "right")
self.latent_space_version = kwargs.get("latent_space_version", "sd1")
self.text_embedding_space_version = kwargs.get("text_embedding_space_version", "sd1")
if dataset_root is not None:
# remove dataset root from path
file_key = self.path.replace(dataset_root, "")
@ -399,7 +402,9 @@ class DataLoaderBatchDTO:
if not isinstance(y.text_embeds, list):
y.text_embeds = [y.text_embeds]
prompt_embeds_list.append(y)
self.prompt_embeds = concat_prompt_embeds(prompt_embeds_list)
padding_side = self.file_items[0].te_padding_side
self.prompt_embeds = concat_prompt_embeds(prompt_embeds_list, padding_side=padding_side)
if any([x.audio_tensor is not None for x in self.file_items]):
# find one to use as a base

View File

@ -1721,8 +1721,6 @@ class LatentCachingFileItemDTOMixin:
self.is_caching_to_disk = False
self.is_caching_to_memory = False
self.latent_load_device = 'cpu'
# sd1 or sdxl or others
self.latent_space_version = 'sd1'
# todo, increment this if we change the latent format to invalidate cache
self.latent_version = 1
@ -1829,21 +1827,6 @@ class LatentCachingMixin:
# use tqdm to show progress
i = 0
for file_item in tqdm(self.file_list, desc=f'Caching latents{" to disk" if to_disk else ""}'):
# set latent space version
if self.sd.model_config.latent_space_version is not None:
file_item.latent_space_version = self.sd.model_config.latent_space_version
elif self.sd.is_xl:
file_item.latent_space_version = 'sdxl'
elif self.sd.is_v3:
file_item.latent_space_version = 'sd3'
elif self.sd.is_auraflow:
file_item.latent_space_version = 'sdxl'
elif self.sd.is_flux:
file_item.latent_space_version = 'flux1'
elif self.sd.model_config.is_pixart_sigma:
file_item.latent_space_version = 'sdxl'
else:
file_item.latent_space_version = self.sd.model_config.arch
file_item.is_caching_to_disk = to_disk
file_item.is_caching_to_memory = to_memory
file_item.latent_load_device = self.sd.device
@ -1933,7 +1916,6 @@ class TextEmbeddingFileItemDTOMixin:
self._text_embedding_path: Union[str, None] = None
self.is_text_embedding_cached = False
self.text_embedding_load_device = 'cpu'
self.text_embedding_space_version = 'sd1'
self.text_embedding_version = 1
def get_text_embedding_info_dict(self: 'FileItemDTO'):
@ -1997,7 +1979,6 @@ class TextEmbeddingCachingMixin:
# use tqdm to show progress
i = 0
for file_item in tqdm(self.file_list, desc='Caching text embeddings to disk'):
file_item.text_embedding_space_version = self.sd.model_config.arch
file_item.latent_load_device = self.sd.device
text_embedding_path = file_item.get_text_embedding_path(recalculate=True)

View File

@ -190,6 +190,10 @@ class BaseModel:
# use new lokr format (default false for old models for backwards compatibility)
self.use_old_lokr_format = True
# when padding to make batch size work, which side padding to use, right or left
# some llms need left side padding, others need right side
self.te_padding_side = "right"
# properties for old arch for backwards compatibility
@property

View File

@ -244,7 +244,7 @@ class EncodedPromptPair:
return self
def concat_prompt_embeds(prompt_embeds: list["PromptEmbeds"]):
def concat_prompt_embeds(prompt_embeds: list["PromptEmbeds"], padding_side: str = "right") -> PromptEmbeds:
# --- pad text_embeds ---
if isinstance(prompt_embeds[0].text_embeds, (list, tuple)):
embed_list = []
@ -259,7 +259,10 @@ def concat_prompt_embeds(prompt_embeds: list["PromptEmbeds"]):
dtype=t.dtype,
device=t.device,
)
t = torch.cat([t, pad], dim=1)
if padding_side == "right":
t = torch.cat([t, pad], dim=1)
else:
t = torch.cat([pad, t], dim=1)
padded.append(t)
embed_list.append(torch.cat(padded, dim=0))
text_embeds = embed_list
@ -274,7 +277,10 @@ def concat_prompt_embeds(prompt_embeds: list["PromptEmbeds"]):
dtype=t.dtype,
device=t.device,
)
t = torch.cat([t, pad], dim=1)
if padding_side == "right":
t = torch.cat([t, pad], dim=1)
else:
t = torch.cat([pad, t], dim=1)
padded.append(t)
text_embeds = torch.cat(padded, dim=0)
@ -296,7 +302,10 @@ def concat_prompt_embeds(prompt_embeds: list["PromptEmbeds"]):
dtype=m.dtype,
device=m.device,
)
m = torch.cat([m, pad], dim=1)
if padding_side == "right":
m = torch.cat([m, pad], dim=1)
else:
m = torch.cat([pad, m], dim=1)
padded.append(m)
attention_mask = torch.cat(padded, dim=0)

View File

@ -229,6 +229,10 @@ class StableDiffusion:
# use new lokr format (default false for old models for backwards compatibility)
self.use_old_lokr_format = True
# when padding to make batch size work, which side padding to use, right or left
# some llms need left side padding, others need right side
self.te_padding_side = "right"
# properties for old arch for backwards compatibility
@property
def is_xl(self):