diff --git a/extensions_built_in/diffusion_models/ltx2/ltx2.py b/extensions_built_in/diffusion_models/ltx2/ltx2.py index 346afa6e..a6d1e6dc 100644 --- a/extensions_built_in/diffusion_models/ltx2/ltx2.py +++ b/extensions_built_in/diffusion_models/ltx2/ltx2.py @@ -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 diff --git a/toolkit/data_loader.py b/toolkit/data_loader.py index 95075a61..51605c98 100644 --- a/toolkit/data_loader.py +++ b/toolkit/data_loader.py @@ -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: diff --git a/toolkit/data_transfer_object/data_loader.py b/toolkit/data_transfer_object/data_loader.py index 7af8de01..ebd53a7a 100644 --- a/toolkit/data_transfer_object/data_loader.py +++ b/toolkit/data_transfer_object/data_loader.py @@ -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 diff --git a/toolkit/dataloader_mixins.py b/toolkit/dataloader_mixins.py index e4fb95bb..20140eb8 100644 --- a/toolkit/dataloader_mixins.py +++ b/toolkit/dataloader_mixins.py @@ -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) diff --git a/toolkit/models/base_model.py b/toolkit/models/base_model.py index 1a9f23e7..ae117c85 100644 --- a/toolkit/models/base_model.py +++ b/toolkit/models/base_model.py @@ -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 diff --git a/toolkit/prompt_utils.py b/toolkit/prompt_utils.py index 0bcbe876..0f3c10a9 100644 --- a/toolkit/prompt_utils.py +++ b/toolkit/prompt_utils.py @@ -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) diff --git a/toolkit/stable_diffusion_model.py b/toolkit/stable_diffusion_model.py index f397c3d7..e235da99 100644 --- a/toolkit/stable_diffusion_model.py +++ b/toolkit/stable_diffusion_model.py @@ -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):