Shrink text embeds to max token length for LTX-2. Drastically reduces cached text embedding sizes
This commit is contained in:
parent
ea912d2d7b
commit
1ce2428722
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Reference in New Issue