diff --git a/finetrainers/conditioning/conditioned_pipeline.py b/finetrainers/conditioning/conditioned_pipeline.py index cff5cb3..087c28f 100644 --- a/finetrainers/conditioning/conditioned_pipeline.py +++ b/finetrainers/conditioning/conditioned_pipeline.py @@ -18,6 +18,8 @@ from torch import torch from transformers import T5EncoderModel, T5TokenizerFast from typing import Any, Callable, Dict, List, Optional, Union +from finetrainers.conditioning import condition_latents_prepare,post_conditioned_latent_patchify +from finetrainers.utils.torch_utils import randn_tensor if is_torch_xla_available(): import torch_xla.core.xla_model as xm @@ -36,6 +38,36 @@ class LTXConditionedPipeline(LTXPipeline): ): super().__init__(scheduler, vae, text_encoder, tokenizer, transformer) + def noise_condition_latent_prepare(self, + batch_size: int = 1, + num_channels_latents: int = 128, + height: int = 512, + width: int = 704, + num_frames: int = 161, + dtype: Optional[torch.dtype] = None, + device: Optional[torch.device] = None, + generator: Optional[torch.Generator] = None, + latents: Optional[torch.Tensor] = None, + ) -> torch.Tensor: + if latents is not None: + return latents.to(device=device, dtype=dtype) + + height = height // self.vae_spatial_compression_ratio + width = width // self.vae_spatial_compression_ratio + num_frames = (num_frames - 1) // self.vae_temporal_compression_ratio + 1 + + shape = (batch_size, num_channels_latents, num_frames, height, width) + + if isinstance(generator, list) and len(generator) != batch_size: + raise ValueError( + f"You have passed a list of generators of length {len(generator)}, but requested an effective batch" + f" size of {batch_size}. Make sure the batch size matches the length of the generators." + ) + + latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype) + + return latents + def __call__( self, prompt: Union[str, List[str]] = None, @@ -126,6 +158,9 @@ class LTXConditionedPipeline(LTXPipeline): # it needs to be the size of the image num_channels_latents = self.transformer.config.in_channels + + # TODO this is the noise latent patchified. + # we have to use the image size by default. latents = self.prepare_latents( batch_size * num_videos_per_prompt, num_channels_latents, @@ -138,11 +173,68 @@ class LTXConditionedPipeline(LTXPipeline): latents, ) # creates noise tensor the size suggested by the user - pose_latents = self.prepare_latents( + noise_latent = self.noise_condition_latent_prepare( + batch_size * num_videos_per_prompt, + num_channels_latents, + height, + width, + num_frames, + torch.float32, + device, + generator, ) + # device=device, + # dtype=weight_dtype, + + + # 4. create conditioning latents from the video. + pose_latent = condition_latents_prepare( + vae=self.vae, + image_or_video=pose_video, + patch_size=self.transformer_config.patch_size, + patch_size_t=self.transformer_config.patch_size_t, + device=device, + dtype=None, + generator=generator, + )["latents"] + + img_ref_latent = condition_latents_prepare( + vae=self.vae, + image_or_video=pose_video, + patch_size=self.transformer_config.patch_size, + patch_size_t=self.transformer_config.patch_size_t, + device=device, + dtype=None, + generator=generator, + )["latents"] + + + # add pose information to both channels + pose_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. + pose_img_ref_cond_latents = torch.cat([pose_img_ref_latents,pose_noisy_latents], dim=1) + + + pose_img_ref_latents = post_conditioned_latent_patchify(latents=pose_img_ref_latents, + num_frames=num_frames, + height=height, + width=width, + patch_size = 1, + patch_size_t = 1) + # target video as a input residual + noisy_latents_residual = post_conditioned_latent_patchify(latents=noise_latent, + num_frames=num_frames, + height=height, + width=width, + patch_size = 1, + patch_size_t = 1) + + # Need to change the latents to # 5. Prepare timesteps latent_num_frames = (num_frames - 1) // self.vae_temporal_compression_ratio + 1 latent_height = height // self.vae_spatial_compression_ratio