Fixed an issue with ltx 2.3 i2v training

This commit is contained in:
Jaret Burkett 2026-03-23 12:41:18 -06:00
parent 330059d8a1
commit 561e6f201c
1 changed files with 49 additions and 26 deletions

View File

@ -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.")