Add back in the pose for both channels.

This commit is contained in:
CrossProduct
2025-01-29 22:40:18 +00:00
parent bca44f6d7f
commit e4a14af8d7
2 changed files with 13 additions and 13 deletions
@@ -67,7 +67,7 @@ class LTXConditionedPipeline(LTXPipeline):
latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
return latents
@torch.no_grad()
def __call__(
self,
prompt: Union[str, List[str]] = None,
@@ -213,7 +213,7 @@ class LTXConditionedPipeline(LTXPipeline):
# residual x latent
# add pose information to both channels
noisy_latents = noise_latent # + pose_latent
noisy_latents = noise_latent + pose_latent
pose_img_ref_latents = img_ref_latent + pose_latent
# expand channel information # B x 2C latent will be projected to adapter to scale it back to 128d using adapter.
@@ -284,7 +284,7 @@ class LTXConditionedPipeline(LTXPipeline):
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)
condition_tokens = condition_tokens.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])
@@ -300,11 +300,11 @@ class LTXConditionedPipeline(LTXPipeline):
rope_interpolation_scale=rope_interpolation_scale,
attention_kwargs=attention_kwargs,
return_dict=False,
residual_x=noisy_residual_tokens,
residual_x=latent_model_input,
)[0]
# patchified ....
# 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)
@@ -330,7 +330,7 @@ class LTXConditionedPipeline(LTXPipeline):
xm.mark_step()
if output_type == "latent":
video = latents
video = noisy_residual_tokens
else:
noisy_residual_tokens = self._unpack_latents(
noisy_residual_tokens,
@@ -348,7 +348,7 @@ class LTXConditionedPipeline(LTXPipeline):
if not self.vae.config.timestep_conditioning:
timestep = None
else:
noise = torch.randn(latents.shape, generator=generator, device=device, dtype=latents.dtype)
noise = torch.randn(noisy_residual_tokens.shape, generator=generator, device=device, dtype=noisy_residual_tokens.dtype)
if not isinstance(decode_timestep, list):
decode_timestep = [decode_timestep] * batch_size
if decode_noise_scale is None:
@@ -356,13 +356,13 @@ class LTXConditionedPipeline(LTXPipeline):
elif not isinstance(decode_noise_scale, list):
decode_noise_scale = [decode_noise_scale] * batch_size
timestep = torch.tensor(decode_timestep, device=device, dtype=latents.dtype)
decode_noise_scale = torch.tensor(decode_noise_scale, device=device, dtype=latents.dtype)[
timestep = torch.tensor(decode_timestep, device=device, dtype=noisy_residual_tokens.dtype)
decode_noise_scale = torch.tensor(decode_noise_scale, device=device, dtype=noisy_residual_tokens.dtype)[
:, None, None, None, None
]
latents = (1 - decode_noise_scale) * latents + decode_noise_scale * noise
noisy_residual_tokens = (1 - decode_noise_scale) * noisy_residual_tokens + decode_noise_scale * noise
video = self.vae.decode(latents, timestep, return_dict=False)[0]
video = self.vae.decode(noisy_residual_tokens, timestep, return_dict=False)[0]
video = self.video_processor.postprocess_video(video, output_type=output_type)
# Offload all models
+1 -1
View File
@@ -813,7 +813,7 @@ class Trainer:
noisy_latents = noisy_latents.to(target_video_latents.dtype)
# add pose information to both channels
pose_noisy_latents = noisy_latents # + pose_video_latents["latents"]
pose_noisy_latents = noisy_latents + pose_video_latents["latents"]
pose_img_ref_latents = img_refs_latents["latents"] + pose_video_latents["latents"]
# expand channel information # B x 2C latent will be projected to adapter to scale it back to 128d using adapter.