Figure out why classifier free guidence breaks the loop. I suspect it wasn't baked into pipeline should have been 0

This commit is contained in:
CrossProduct
2025-01-29 06:10:29 +00:00
parent a16285141e
commit 536dfc8f6a
@@ -78,7 +78,7 @@ class LTXConditionedPipeline(LTXPipeline):
frame_rate: int = 25,
num_inference_steps: int = 50,
timesteps: List[int] = None,
guidance_scale: float = 3,
guidance_scale: float = 0,
num_videos_per_prompt: Optional[int] = 1,
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
latents: Optional[torch.Tensor] = None,
@@ -302,15 +302,16 @@ class LTXConditionedPipeline(LTXPipeline):
return_dict=False,
residual_x=noisy_residual_tokens,
)[0]
# patchified ....
noise_pred = noise_pred.float()
# noisy latents isn't patchified ... fix that ...
if self.do_classifier_free_guidance:
noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
noise_pred = noise_pred_uncond + self.guidance_scale * (noise_pred_text - noise_pred_uncond)
# compute the previous noisy sample x_t -> x_t-1
noisy_latents = self.scheduler.step(noise_pred, t, noisy_latents, return_dict=False)[0]
# Noisy latents needs to be
noisy_residual_tokens = self.scheduler.step(noise_pred, t, noisy_residual_tokens, return_dict=False)[0]
if callback_on_step_end is not None:
callback_kwargs = {}
@@ -318,7 +319,7 @@ class LTXConditionedPipeline(LTXPipeline):
callback_kwargs[k] = locals()[k]
callback_outputs = callback_on_step_end(self, i, t, callback_kwargs)
noisy_latents = callback_outputs.pop("latents", noisy_latents)
noisy_residual_tokens = callback_outputs.pop("latents", noisy_residual_tokens)
prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds)
# call the callback, if provided
@@ -331,18 +332,18 @@ class LTXConditionedPipeline(LTXPipeline):
if output_type == "latent":
video = latents
else:
noisy_latents = self._unpack_latents(
noisy_latents,
noisy_residual_tokens = self._unpack_latents(
noisy_residual_tokens,
latent_num_frames,
latent_height,
latent_width,
self.transformer_spatial_patch_size,
self.transformer_temporal_patch_size,
)
noisy_latents = self._denormalize_latents(
noisy_latents, self.vae.latents_mean, self.vae.latents_std, self.vae.config.scaling_factor
noisy_residual_tokens = self._denormalize_latents(
noisy_residual_tokens, self.vae.latents_mean, self.vae.latents_std, self.vae.config.scaling_factor
)
noisy_latents = noisy_latents.to(prompt_embeds.dtype)
noisy_residual_tokens = noisy_residual_tokens.to(prompt_embeds.dtype)
if not self.vae.config.timestep_conditioning:
timestep = None