Fix before trying again.

This commit is contained in:
CrossProduct
2025-01-27 08:30:45 +00:00
parent 80f653f336
commit dfa03f1ed1
@@ -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