Noted place that requires intervention.

This commit is contained in:
CrossProduct
2025-01-29 05:30:17 +00:00
parent 707e001f7b
commit a16285141e
2 changed files with 6 additions and 3 deletions
@@ -279,16 +279,16 @@ class LTXConditionedPipeline(LTXPipeline):
continue
# change this
# latent_model_input = torch.cat([latents] * 2) if self.do_classifier_free_guidance else latents
# latent_model_input = latent_model_input.to(prompt_embeds.dtype)
#latent_model_input = latent_model_input.to(prompt_embeds.dtype)
# might be the bug ..
latent_model_input = noisy_latent_tokens
latent_model_input = torch.cat([noisy_latent_tokens] * 2) if self.do_classifier_free_guidance else noisy_latent_tokens
latent_model_input = latent_model_input.to(prompt_embeds.dtype)
# broadcast to batch dimension in a way that's compatible with ONNX/Core ML
timestep = t.expand(latent_model_input.shape[0])
#encoder hidden states are different...
noise_pred = self.transformer(
hidden_states=condition_tokens,
encoder_hidden_states=prompt_embeds,
@@ -116,6 +116,9 @@ def conditioned_forward_pass(
) -> torch.Tensor:
rope_interpolation_scale = [1 / 25, 32, 32]
# encoder_hidden_states=prompt_embeds,
# timestep=timesteps,
# encoder_attention_mask=prompt_attention_mask,
denoised_latents = transformer(
hidden_states=noisy_latents,
encoder_hidden_states=prompt_embeds,