From 8fe0ed225ae2e4d1c5cc7102de05ae81730145fb Mon Sep 17 00:00:00 2001 From: CrossProduct Date: Wed, 29 Jan 2025 23:49:22 +0000 Subject: [PATCH] Conditioned Adapter make residual optional. --- .../conditioning/conditioned_pipeline.py | 48 +++++++++---------- ...conditioned_residual_adapter_bottleneck.py | 5 +- 2 files changed, 27 insertions(+), 26 deletions(-) diff --git a/finetrainers/conditioning/conditioned_pipeline.py b/finetrainers/conditioning/conditioned_pipeline.py index bebf1ef..81d2ca2 100644 --- a/finetrainers/conditioning/conditioned_pipeline.py +++ b/finetrainers/conditioning/conditioned_pipeline.py @@ -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 diff --git a/finetrainers/conditioning/conditioned_residual_adapter_bottleneck.py b/finetrainers/conditioning/conditioned_residual_adapter_bottleneck.py index e8151c9..26d0938 100644 --- a/finetrainers/conditioning/conditioned_residual_adapter_bottleneck.py +++ b/finetrainers/conditioning/conditioned_residual_adapter_bottleneck.py @@ -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