From 561e6f201c55dcd0b455493f4311b803aca4ce2c Mon Sep 17 00:00:00 2001 From: Jaret Burkett Date: Mon, 23 Mar 2026 12:41:18 -0600 Subject: [PATCH] Fixed an issue with ltx 2.3 i2v training --- .../diffusion_models/ltx2/ltx2.py | 75 ++++++++++++------- 1 file changed, 49 insertions(+), 26 deletions(-) diff --git a/extensions_built_in/diffusion_models/ltx2/ltx2.py b/extensions_built_in/diffusion_models/ltx2/ltx2.py index 3f64fa75..4359116b 100644 --- a/extensions_built_in/diffusion_models/ltx2/ltx2.py +++ b/extensions_built_in/diffusion_models/ltx2/ltx2.py @@ -77,7 +77,7 @@ dit_prefix = "model.diffusion_model." vae_prefix = "vae." audio_vae_prefix = "audio_vae." vocoder_prefix = "vocoder." -base_te_path = "Lightricks/gemma-3-12b-it-qat-q4_0-unquantized" +base_te_path = "google/gemma-3-12b-it-qat-q4_0-unquantized" HF_TOKEN = os.getenv("HF_TOKEN", None) @@ -240,7 +240,7 @@ class LTX2Model(BaseModel): combined_state_dict = None self.print_and_status_update("Loading transformer") - + if not os.path.exists(model_path) and model_path.endswith(".safetensors"): # download the model from the Hugging Face Hub if it is not a local path splits = model_path.split("/") @@ -254,7 +254,7 @@ class LTX2Model(BaseModel): filename=splits[2], token=HF_TOKEN, ) - + # if we have a safetensors file it is a mono checkpoint if os.path.exists(model_path) and model_path.endswith(".safetensors"): combined_state_dict = load_file(model_path) @@ -264,7 +264,9 @@ class LTX2Model(BaseModel): original_dit_ckpt = get_model_state_dict_from_combined_ckpt( combined_state_dict, dit_prefix ) - transformer = convert_ltx2_transformer(original_dit_ckpt, version=self.ltx_version) + transformer = convert_ltx2_transformer( + original_dit_ckpt, version=self.ltx_version + ) transformer = transformer.to(dtype) else: transformer_path = model_path @@ -440,22 +442,30 @@ class LTX2Model(BaseModel): original_vae_ckpt = get_model_state_dict_from_combined_ckpt( combined_state_dict, vae_prefix ) - vae = convert_ltx2_video_vae(original_vae_ckpt, version=self.ltx_version).to(dtype) + vae = convert_ltx2_video_vae( + original_vae_ckpt, version=self.ltx_version + ).to(dtype) del original_vae_ckpt original_audio_vae_ckpt = get_model_state_dict_from_combined_ckpt( combined_state_dict, audio_vae_prefix ) - audio_vae = convert_ltx2_audio_vae(original_audio_vae_ckpt, version=self.ltx_version).to(dtype) + audio_vae = convert_ltx2_audio_vae( + original_audio_vae_ckpt, version=self.ltx_version + ).to(dtype) del original_audio_vae_ckpt original_connectors_ckpt = get_model_state_dict_from_combined_ckpt( combined_state_dict, dit_prefix ) - connectors = convert_ltx2_connectors(original_connectors_ckpt, version=self.ltx_version).to(dtype) + connectors = convert_ltx2_connectors( + original_connectors_ckpt, version=self.ltx_version + ).to(dtype) del original_connectors_ckpt original_vocoder_ckpt = get_model_state_dict_from_combined_ckpt( combined_state_dict, vocoder_prefix ) - vocoder = convert_ltx2_vocoder(original_vocoder_ckpt, version=self.ltx_version).to(dtype) + vocoder = convert_ltx2_vocoder( + original_vocoder_ckpt, version=self.ltx_version + ).to(dtype) del original_vocoder_ckpt del combined_state_dict flush() @@ -665,17 +675,19 @@ class LTX2Model(BaseModel): # 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) - + if self.ltx_version == "2.3": - extra['stg_scale'] = 1.0 - extra['modality_scale'] = 3.0 - extra['guidance_rescale'] = 0.7 - extra['audio_guidance_scale'] = 7.0 - extra['audio_stg_scale'] = 1.0 - extra['audio_modality_scale'] = 3.0 - extra['audio_guidance_rescale'] = 0.7 - extra['spatio_temporal_guidance_blocks'] = [28] - extra['use_cross_timestep'] = True # they dont set this in some examples in diffusers, but I believe it should always be true for 2.3 + extra["stg_scale"] = 1.0 + extra["modality_scale"] = 3.0 + extra["guidance_rescale"] = 0.7 + extra["audio_guidance_scale"] = 7.0 + extra["audio_stg_scale"] = 1.0 + extra["audio_modality_scale"] = 3.0 + extra["audio_guidance_rescale"] = 0.7 + extra["spatio_temporal_guidance_blocks"] = [28] + extra["use_cross_timestep"] = ( + True # they dont set this in some examples in diffusers, but I believe it should always be true for 2.3 + ) video, audio = pipeline( prompt_embeds=conditional_embeds.text_embeds.to( @@ -865,12 +877,18 @@ class LTX2Model(BaseModel): # use conditioning mask to replace latents latent_model_input = ( - latent_model_input * (1 - conditioning_mask) - + init_latents * conditioning_mask + init_latents * conditioning_mask + + latent_model_input * (1 - conditioning_mask) + ) + + packed_conditioning_mask = self.pipeline._pack_latents( + conditioning_mask, + patch_size=self.pipeline.transformer_spatial_patch_size, + patch_size_t=self.pipeline.transformer_temporal_patch_size, ) # set video timestep - video_timestep = timestep.unsqueeze(-1) * (1 - conditioning_mask) + video_timestep = timestep.unsqueeze(-1) * (1 - packed_conditioning_mask) # todo get this somehow frame_rate = 24 @@ -911,7 +929,9 @@ class LTX2Model(BaseModel): ) duration_s = batch.dataset_config.num_frames / frame_rate audio_latents_per_second = ( - self.pipeline.audio_sampling_rate / self.pipeline.audio_hop_length / float(self.pipeline.audio_vae_temporal_compression_ratio) + self.pipeline.audio_sampling_rate + / self.pipeline.audio_hop_length + / float(self.pipeline.audio_vae_temporal_compression_ratio) ) audio_num_frames = round(duration_s * audio_latents_per_second) audio_latents = self.pipeline.prepare_audio_latents( @@ -929,7 +949,8 @@ class LTX2Model(BaseModel): if self.pipeline.connectors.device != self.transformer.device: self.pipeline.connectors.to(self.transformer.device) - tokenizer_padding_side = "left" # Padding side for default Gemma3-12B text encoder + # Padding side for default Gemma3-12B text encoder + tokenizer_padding_side = "left" if getattr(self, "tokenizer", None) is not None: tokenizer_padding_side = getattr(self.tokenizer, "padding_side", "left") ( @@ -954,7 +975,7 @@ class LTX2Model(BaseModel): audio_coords = self.transformer.audio_rope.prepare_audio_coords( audio_latents.shape[0], audio_num_frames, audio_latents.device ) - + # use_cross_timestep - Whether to use the cross modality (audio is the cross modality of video, and vice versa) sigma when # calculating the cross attention modulation parameters. `True` is the newer (e.g. LTX-2.3) behavior; # `False` is the legacy LTX-2.0 behavior. @@ -966,7 +987,7 @@ class LTX2Model(BaseModel): encoder_hidden_states=connector_prompt_embeds, audio_encoder_hidden_states=connector_audio_prompt_embeds, timestep=video_timestep, - sigma=video_timestep, # Used by LTX-2.3 + sigma=timestep, # Used by LTX-2.3 audio_timestep=timestep, encoder_attention_mask=connector_attention_mask, audio_encoder_attention_mask=connector_attention_mask, @@ -1088,7 +1109,9 @@ class LTX2Model(BaseModel): return new_sd def convert_lora_weights_before_load(self, state_dict): - state_dict = convert_lora_original_to_diffusers(state_dict, version=self.ltx_version) + state_dict = convert_lora_original_to_diffusers( + state_dict, version=self.ltx_version + ) new_sd = {} for key, value in state_dict.items(): new_key = key.replace("diffusion_model.", "transformer.")