mirror of
https://github.com/storytold/FineTrainers-Conditioning.git
synced 2026-10-09 00:09:45 +00:00
Fix before trying again.
This commit is contained in:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user