Conditioned Adapter make residual optional.

This commit is contained in:
CrossProduct
2025-01-29 23:49:22 +00:00
parent 41eec16c67
commit 8fe0ed225a
2 changed files with 27 additions and 26 deletions
@@ -219,7 +219,7 @@ class LTXConditionedPipeline(LTXPipeline):
# expand channel information # B x 2C latent will be projected to adapter to scale it back to 128d using adapter.
condition_latent = torch.cat([pose_img_ref_latents,noisy_latents], dim=1)
noisy_latent_tokens = post_conditioned_latent_patchify(latents=noisy_latents, num_frames=num_frames,
noise_latent_tokens = post_conditioned_latent_patchify(latents=noise_latent, num_frames=num_frames,
height=height,
width=width,
patch_size = 1,
@@ -232,12 +232,12 @@ class LTXConditionedPipeline(LTXPipeline):
patch_size = 1,
patch_size_t = 1)["latents"]
# target video as a input residual
noisy_residual_tokens = post_conditioned_latent_patchify(latents=noisy_latents,
num_frames=num_frames,
height=height,
width=width,
patch_size = 1,
patch_size_t = 1)["latents"]
# noisy_residual_tokens = post_conditioned_latent_patchify(latents=noisy_latents,
# num_frames=num_frames,
# height=height,
# width=width,
# patch_size = 1,
# patch_size_t = 1)["latents"]
# Need to change the latents to
# 5. Prepare timesteps
@@ -280,11 +280,9 @@ class LTXConditionedPipeline(LTXPipeline):
# 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 = torch.cat([noisy_latent_tokens] * 2) if self.do_classifier_free_guidance else noisy_latent_tokens
latent_model_input = torch.cat([condition_tokens] * 2) if self.do_classifier_free_guidance else condition_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,7 +298,7 @@ class LTXConditionedPipeline(LTXPipeline):
rope_interpolation_scale=rope_interpolation_scale,
attention_kwargs=attention_kwargs,
return_dict=False,
residual_x=latent_model_input,
residual_x=None,
)[0]
# patchified
noise_pred = noise_pred.float()
@@ -311,7 +309,7 @@ class LTXConditionedPipeline(LTXPipeline):
# compute the previous noisy sample x_t -> x_t-1
# Noisy latents needs to be
noisy_residual_tokens = self.scheduler.step(noise_pred, t, noisy_residual_tokens, return_dict=False)[0]
noise_latent_tokens = self.scheduler.step(noise_pred, t, noise_latent_tokens, return_dict=False)[0]
if callback_on_step_end is not None:
callback_kwargs = {}
@@ -319,7 +317,7 @@ class LTXConditionedPipeline(LTXPipeline):
callback_kwargs[k] = locals()[k]
callback_outputs = callback_on_step_end(self, i, t, callback_kwargs)
noisy_residual_tokens = callback_outputs.pop("latents", noisy_residual_tokens)
noise_latent_tokens = callback_outputs.pop("latents", noise_latent_tokens)
prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds)
# call the callback, if provided
@@ -330,25 +328,25 @@ class LTXConditionedPipeline(LTXPipeline):
xm.mark_step()
if output_type == "latent":
video = noisy_residual_tokens
video = noise_latent_tokens
else:
noisy_residual_tokens = self._unpack_latents(
noisy_residual_tokens,
noise_latent_tokens = self._unpack_latents(
noise_latent_tokens,
latent_num_frames,
latent_height,
latent_width,
self.transformer_spatial_patch_size,
self.transformer_temporal_patch_size,
)
noisy_residual_tokens = self._denormalize_latents(
noisy_residual_tokens, self.vae.latents_mean, self.vae.latents_std, self.vae.config.scaling_factor
noise_latent_tokens = self._denormalize_latents(
noise_latent_tokens, self.vae.latents_mean, self.vae.latents_std, self.vae.config.scaling_factor
)
noisy_residual_tokens = noisy_residual_tokens.to(prompt_embeds.dtype)
noise_latent_tokens = noise_latent_tokens.to(prompt_embeds.dtype)
if not self.vae.config.timestep_conditioning:
timestep = None
else:
noise = torch.randn(noisy_residual_tokens.shape, generator=generator, device=device, dtype=noisy_residual_tokens.dtype)
noise = torch.randn(noise_latent_tokens.shape, generator=generator, device=device, dtype=noise_latent_tokens.dtype)
if not isinstance(decode_timestep, list):
decode_timestep = [decode_timestep] * batch_size
if decode_noise_scale is None:
@@ -356,13 +354,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=noisy_residual_tokens.dtype)
decode_noise_scale = torch.tensor(decode_noise_scale, device=device, dtype=noisy_residual_tokens.dtype)[
timestep = torch.tensor(decode_timestep, device=device, dtype=noise_latent_tokens.dtype)
decode_noise_scale = torch.tensor(decode_noise_scale, device=device, dtype=noise_latent_tokens.dtype)[
:, None, None, None, None
]
noisy_residual_tokens = (1 - decode_noise_scale) * noisy_residual_tokens + decode_noise_scale * noise
noise_latent_tokens = (1 - decode_noise_scale) * noise_latent_tokens + decode_noise_scale * noise
video = self.vae.decode(noisy_residual_tokens, timestep, return_dict=False)[0]
video = self.vae.decode(noise_latent_tokens, timestep, return_dict=False)[0]
video = self.video_processor.postprocess_video(video, output_type=output_type)
# Offload all models
@@ -52,7 +52,10 @@ class ConditionedResidualAdapterBottleneck(nn.Module):
# Up projection to new dimension
hidden_states = self.up_proj(hidden_states)
hidden_states = hidden_states + residual_x
if residual_x is None:
return hidden_states
else:
hidden_states = hidden_states + residual_x
return hidden_states