diff --git a/finetrainers/conditioning/LTXVideoConditionedTransformer3DModel.py b/finetrainers/conditioning/LTXVideoConditionedTransformer3DModel.py index ed4c32d..9c6d47d 100644 --- a/finetrainers/conditioning/LTXVideoConditionedTransformer3DModel.py +++ b/finetrainers/conditioning/LTXVideoConditionedTransformer3DModel.py @@ -48,58 +48,6 @@ class LTXVideoConditionedTransformer3DModel(LTXVideoTransformer3DModel): attention_out_bias=attention_out_bias ) - # adapter.down_proj.weight - self.adapter = ConditionedResidualAdapterBottleneck( - input_dim=adapter_in_dim, - output_dim=128, - bottleneck_dim=64, - adapter_dropout=0.1, - adapter_init_scale=1e-3 - ) - - @classmethod # rewrite this to ensure that if it doesn't have adapter weights it will load this way but if it has adapter weights it will not. - def from_pretrained(cls, pretrained_model_name_or_path,save_directory, **kwargs): - # First create an empty model with the desired architecture - print("Pretrain Loading") - model = cls() - model_dict = model.state_dict() - # # Then load the pretrained weights into it - pretrained = LTXVideoTransformer3DModel.from_pretrained(pretrained_model_name_or_path, **kwargs) - - # Copy over the pretrained weights for the shared components - pretrained_dict = pretrained.state_dict() - - # Filter out adapter weights from the pretrained dict - filtered_dict = {} - for k, v in pretrained_dict.items(): - if k in model_dict: - - filtered_dict[k] = v - else: - print(k) - # Update model with pretrained weights - model_dict.update(filtered_dict) - - adapter_weights_path = os.path.join(save_directory, "adapter_weights.pth") - if os.path.exists(adapter_weights_path): - adapter_weights = torch.load(adapter_weights_path) - model.adapter.load_state_dict(adapter_weights) - print("Adapter weights loaded successfully.") - else: - print("No adapter weights file found. Using default adapter weights.") - - model.load_state_dict(model_dict) - return model - - def save_pretrained(self,save_directory): - model_state_dict = self.state_dict() - torch.save(model_state_dict, save_directory) - - adapter_state_dict = self.adapter.state_dict() - save_directory = os.path.dirname(save_directory) - path = os.path.join(save_directory, "adapter_weights.pth") - torch.save(adapter_state_dict, path) - print("Saving Model with Adapter") def forward( self, @@ -125,9 +73,6 @@ class LTXVideoConditionedTransformer3DModel(LTXVideoTransformer3DModel): batch_size = hidden_states.size(0) - # inject the condition and the residual then project it into the pretrained proj_in - hidden_states = self.adapter(residual_x=residual_x, - conditioned_x=hidden_states) # whats this value ? hidden_states = self.proj_in(hidden_states) diff --git a/finetrainers/conditioning/conditioned_pipeline.py b/finetrainers/conditioning/conditioned_pipeline.py index 81d2ca2..b17ab46 100644 --- a/finetrainers/conditioning/conditioned_pipeline.py +++ b/finetrainers/conditioning/conditioned_pipeline.py @@ -213,31 +213,23 @@ class LTXConditionedPipeline(LTXPipeline): # residual x latent # add pose information to both channels - 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. - condition_latent = torch.cat([pose_img_ref_latents,noisy_latents], dim=1) + noisy_latents = noise_latent + pose_latent + img_ref_latent noise_latent_tokens = post_conditioned_latent_patchify(latents=noise_latent, num_frames=num_frames, height=height, width=width, patch_size = 1, patch_size_t = 1)["latents"] + # # 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) - condition_tokens = post_conditioned_latent_patchify(latents=condition_latent, + condition_tokens = post_conditioned_latent_patchify(latents=noisy_latents, num_frames=num_frames, height=height, width=width, 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"] + # Need to change the latents to # 5. Prepare timesteps diff --git a/finetrainers/trainer.py b/finetrainers/trainer.py index eaffea8..547b963 100644 --- a/finetrainers/trainer.py +++ b/finetrainers/trainer.py @@ -813,28 +813,18 @@ 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_img_ref_latents = img_refs_latents["latents"] + pose_video_latents["latents"] + pose_noisy_latents = noisy_latents + pose_video_latents["latents"] + img_refs_latents["latents"] - # expand channel information # B x 2C latent will be projected to adapter to scale it back to 128d using adapter. - condition_latents = torch.cat([pose_img_ref_latents,pose_noisy_latents], dim=1) - # project this turn into patches - + # concat tensor - condition_tokens = post_conditioned_latent_patchify(latents=condition_latents, + condition_tokens = post_conditioned_latent_patchify(latents=pose_noisy_latents, num_frames=latent_conditions["num_frames"], height=latent_conditions["height"], width=latent_conditions["width"], patch_size = 1, patch_size_t = 1) - # TODO REMOVE target video as a input residual might not need. - noisy_residual_tokens = post_conditioned_latent_patchify(latents=noisy_latents, - num_frames=latent_conditions["num_frames"], - height=latent_conditions["height"], - width=latent_conditions["width"], - patch_size = 1, - patch_size_t = 1) + # Target Latent Patchified concat tensor # pose template noisey input [cat] img_ref + pose video # That dict says latents but actually tokens.