mirror of
https://github.com/storytold/FineTrainers-Conditioning.git
synced 2026-10-09 00:09:45 +00:00
trying training again but broke something so going back.
This commit is contained in:
@@ -19,7 +19,7 @@ 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
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
|
||||
if is_torch_xla_available():
|
||||
import torch_xla.core.xla_model as xm
|
||||
@@ -95,12 +95,12 @@ class LTXConditionedPipeline(LTXPipeline):
|
||||
callback_on_step_end_tensor_inputs: List[str] = ["latents"],
|
||||
max_sequence_length: int = 128,
|
||||
pose_video=None,
|
||||
image_ref_video=None,
|
||||
img_ref_video=None,
|
||||
):
|
||||
# These are the raw videos
|
||||
if pose_video is None:
|
||||
raise ValueError("pose_video cannot be None.")
|
||||
if image_ref_video is None:
|
||||
if img_ref_video is None:
|
||||
raise ValueError("image_ref_video cannot be None.")
|
||||
|
||||
if isinstance(callback_on_step_end, (PipelineCallback, MultiPipelineCallbacks)):
|
||||
@@ -159,20 +159,20 @@ class LTXConditionedPipeline(LTXPipeline):
|
||||
|
||||
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,
|
||||
height,
|
||||
width,
|
||||
num_frames,
|
||||
torch.float32,
|
||||
device,
|
||||
generator,
|
||||
latents,
|
||||
) # creates noise tensor the size suggested by the user
|
||||
|
||||
# # 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,
|
||||
# height,
|
||||
# width,
|
||||
# num_frames,
|
||||
# torch.float32,
|
||||
# device,
|
||||
# generator,
|
||||
# latents,
|
||||
# dtype=self.text_encoder.dtype
|
||||
# ) # creates noise tensor the size suggested by the user
|
||||
|
||||
noise_latent = self.noise_condition_latent_prepare(
|
||||
batch_size * num_videos_per_prompt,
|
||||
@@ -185,54 +185,59 @@ class LTXConditionedPipeline(LTXPipeline):
|
||||
generator,
|
||||
)
|
||||
|
||||
|
||||
# device=device,
|
||||
# dtype=weight_dtype,
|
||||
|
||||
|
||||
# 4. create conditioning latents from the video.
|
||||
pose_latent = condition_latents_prepare(
|
||||
pose_latent = condition_latents_prepare.prepare_latents_for_conditioning(
|
||||
vae=self.vae,
|
||||
image_or_video=pose_video,
|
||||
patch_size=self.transformer_config.patch_size,
|
||||
patch_size_t=self.transformer_config.patch_size_t,
|
||||
patch_size=self.transformer.config.patch_size,
|
||||
patch_size_t=self.transformer.config.patch_size_t,
|
||||
device=device,
|
||||
dtype=None,
|
||||
dtype=torch.float32,
|
||||
generator=generator,
|
||||
)["latents"]
|
||||
|
||||
img_ref_latent = condition_latents_prepare(
|
||||
img_ref_latent = condition_latents_prepare.prepare_latents_for_conditioning(
|
||||
vae=self.vae,
|
||||
image_or_video=pose_video,
|
||||
patch_size=self.transformer_config.patch_size,
|
||||
patch_size_t=self.transformer_config.patch_size_t,
|
||||
patch_size=self.transformer.config.patch_size,
|
||||
patch_size_t=self.transformer.config.patch_size_t,
|
||||
device=device,
|
||||
dtype=None,
|
||||
dtype=torch.float32,
|
||||
generator=generator,
|
||||
)["latents"]
|
||||
|
||||
# pose template noisey input [cat] img_ref + pose video
|
||||
# pose template ref video latent + patchify video latent.
|
||||
# img ref patchify video latent +
|
||||
# add them together as input
|
||||
# residual x latent
|
||||
|
||||
# add pose information to both channels
|
||||
pose_noisy_latents = noise_latent + pose_latent
|
||||
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)
|
||||
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,
|
||||
height=height,
|
||||
width=width,
|
||||
patch_size = 1,
|
||||
patch_size_t = 1)["latents"]
|
||||
|
||||
pose_img_ref_latents = post_conditioned_latent_patchify(latents=pose_img_ref_latents,
|
||||
condition_tokens = post_conditioned_latent_patchify(latents=condition_latent,
|
||||
num_frames=num_frames,
|
||||
height=height,
|
||||
width=width,
|
||||
patch_size = 1,
|
||||
patch_size_t = 1)
|
||||
patch_size_t = 1)["latents"]
|
||||
# target video as a input residual
|
||||
noisy_latents_residual = post_conditioned_latent_patchify(latents=noise_latent,
|
||||
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)
|
||||
patch_size_t = 1)["latents"]
|
||||
|
||||
# Need to change the latents to
|
||||
# 5. Prepare timesteps
|
||||
@@ -272,15 +277,18 @@ class LTXConditionedPipeline(LTXPipeline):
|
||||
for i, t in enumerate(timesteps):
|
||||
if self.interrupt:
|
||||
continue
|
||||
|
||||
latent_model_input = torch.cat([latents] * 2) if self.do_classifier_free_guidance else latents
|
||||
# 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 = noisy_latent_tokens
|
||||
latent_model_input = latent_model_input.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])
|
||||
|
||||
|
||||
noise_pred = self.transformer(
|
||||
hidden_states=latent_model_input,
|
||||
hidden_states=condition_tokens,
|
||||
encoder_hidden_states=prompt_embeds,
|
||||
timestep=timestep,
|
||||
encoder_attention_mask=prompt_attention_mask,
|
||||
@@ -290,7 +298,9 @@ class LTXConditionedPipeline(LTXPipeline):
|
||||
rope_interpolation_scale=rope_interpolation_scale,
|
||||
attention_kwargs=attention_kwargs,
|
||||
return_dict=False,
|
||||
residual_x=noisy_residual_tokens,
|
||||
)[0]
|
||||
|
||||
noise_pred = noise_pred.float()
|
||||
|
||||
if self.do_classifier_free_guidance:
|
||||
@@ -298,7 +308,7 @@ class LTXConditionedPipeline(LTXPipeline):
|
||||
noise_pred = noise_pred_uncond + self.guidance_scale * (noise_pred_text - noise_pred_uncond)
|
||||
|
||||
# compute the previous noisy sample x_t -> x_t-1
|
||||
latents = self.scheduler.step(noise_pred, t, latents, return_dict=False)[0]
|
||||
noisy_latents = self.scheduler.step(noise_pred, t, noisy_latents, return_dict=False)[0]
|
||||
|
||||
if callback_on_step_end is not None:
|
||||
callback_kwargs = {}
|
||||
@@ -306,7 +316,7 @@ class LTXConditionedPipeline(LTXPipeline):
|
||||
callback_kwargs[k] = locals()[k]
|
||||
callback_outputs = callback_on_step_end(self, i, t, callback_kwargs)
|
||||
|
||||
latents = callback_outputs.pop("latents", latents)
|
||||
noisy_latents = callback_outputs.pop("latents", noisy_latents)
|
||||
prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds)
|
||||
|
||||
# call the callback, if provided
|
||||
@@ -319,18 +329,18 @@ class LTXConditionedPipeline(LTXPipeline):
|
||||
if output_type == "latent":
|
||||
video = latents
|
||||
else:
|
||||
latents = self._unpack_latents(
|
||||
latents,
|
||||
noisy_latents = self._unpack_latents(
|
||||
noisy_latents,
|
||||
latent_num_frames,
|
||||
latent_height,
|
||||
latent_width,
|
||||
self.transformer_spatial_patch_size,
|
||||
self.transformer_temporal_patch_size,
|
||||
)
|
||||
latents = self._denormalize_latents(
|
||||
latents, self.vae.latents_mean, self.vae.latents_std, self.vae.config.scaling_factor
|
||||
noisy_latents = self._denormalize_latents(
|
||||
noisy_latents, self.vae.latents_mean, self.vae.latents_std, self.vae.config.scaling_factor
|
||||
)
|
||||
latents = latents.to(prompt_embeds.dtype)
|
||||
noisy_latents = noisy_latents.to(prompt_embeds.dtype)
|
||||
|
||||
if not self.vae.config.timestep_conditioning:
|
||||
timestep = None
|
||||
|
||||
@@ -137,7 +137,7 @@ def conditional_validation(
|
||||
image: Optional[Image.Image] = None,
|
||||
video: Optional[List[Image.Image]] = None,
|
||||
pose_video: Optional[List[Image.Image]] = None,
|
||||
image_ref_video: Optional[List[Image.Image]] = None,
|
||||
img_ref_video: Optional[List[Image.Image]] = None,
|
||||
height: Optional[int] = None,
|
||||
width: Optional[int] = None,
|
||||
num_frames: Optional[int] = None,
|
||||
@@ -157,15 +157,10 @@ def conditional_validation(
|
||||
"return_dict": True,
|
||||
"output_type": "pil",
|
||||
"pose_video": pose_video,
|
||||
"image_ref_video": image_ref_video,
|
||||
"img_ref_video": img_ref_video,
|
||||
}
|
||||
|
||||
# pose template noisey input [cat] img_ref + pose video
|
||||
|
||||
# pose template ref video latent + patchify video latent.
|
||||
# img ref patchify video latent +
|
||||
# add them together as input
|
||||
# residual x latent
|
||||
|
||||
|
||||
generation_kwargs = {k: v for k, v in generation_kwargs.items() if v is not None}
|
||||
video = pipeline(**generation_kwargs).frames[0]
|
||||
|
||||
+51
-4
@@ -6,7 +6,7 @@ import random
|
||||
from datetime import datetime, timedelta
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, List
|
||||
|
||||
from finetrainers.dataset import ImageOrVideoDataset
|
||||
import diffusers
|
||||
import torch
|
||||
import torch.backends
|
||||
@@ -21,6 +21,10 @@ from accelerate.utils import (
|
||||
gather_object,
|
||||
set_seed,
|
||||
)
|
||||
import decord
|
||||
from typing import Tuple,Optional
|
||||
|
||||
|
||||
from diffusers import DiffusionPipeline
|
||||
from diffusers.configuration_utils import FrozenDict
|
||||
from diffusers.models.autoencoders.vae import DiagonalGaussianDistribution
|
||||
@@ -816,6 +820,7 @@ class Trainer:
|
||||
pose_img_ref_latents = torch.cat([pose_img_ref_latents,pose_noisy_latents], dim=1)
|
||||
# project this turn into patches
|
||||
|
||||
# concat tensor
|
||||
pose_img_ref_latents = post_conditioned_latent_patchify(latents=pose_img_ref_latents,
|
||||
num_frames=latent_conditions["num_frames"],
|
||||
height=latent_conditions["height"],
|
||||
@@ -829,7 +834,7 @@ class Trainer:
|
||||
width=latent_conditions["width"],
|
||||
patch_size = 1,
|
||||
patch_size_t = 1)
|
||||
# Target Latent Patchified
|
||||
# Target Latent Patchified concat tensor
|
||||
# pose template noisey input [cat] img_ref + pose video
|
||||
latent_conditions.update({"noisy_latents": pose_img_ref_latents["latents"]})
|
||||
# input video noise at level residual information to adapter
|
||||
@@ -1027,9 +1032,9 @@ class Trainer:
|
||||
|
||||
if self.pose_condition:
|
||||
if pose_video is not None:
|
||||
pose_video = load_video(pose_video)
|
||||
pose_video = self.preprocess_condition_video(pose_video,max_num_frames=self.dataset.max_num_frames)
|
||||
if img_ref_video is not None:
|
||||
img_ref_video = load_video(img_ref_video)
|
||||
img_ref_video = self.preprocess_condition_video_image_reference_video(path=img_ref_video,resolution_buckets=self.dataset.resolution_buckets,max_num_frames=self.dataset.max_num_frames)
|
||||
|
||||
logger.debug(
|
||||
f"Validating sample {i + 1}/{num_validation_samples} on process {accelerator.process_index}. Prompt: {prompt}",
|
||||
@@ -1053,6 +1058,7 @@ class Trainer:
|
||||
self.args.seed if self.args.seed is not None else 0
|
||||
),
|
||||
)
|
||||
|
||||
else:
|
||||
validation_artifacts = self.model_config["validation"](
|
||||
pipeline=pipeline,
|
||||
@@ -1335,3 +1341,44 @@ class Trainer:
|
||||
for component in components:
|
||||
if component is not None:
|
||||
component.requires_grad_(True)
|
||||
|
||||
|
||||
# Two conditioning functions
|
||||
@staticmethod
|
||||
def preprocess_condition_video(path: Path,max_num_frames:int) -> Tuple[torch.Tensor, Optional[torch.Tensor]]:
|
||||
r"""
|
||||
Loads a single video, or latent and prompt embedding, based on initialization parameters.
|
||||
|
||||
Returns a [F, C, H, W] video tensor.
|
||||
"""
|
||||
video_reader = decord.VideoReader(uri=Path(path).as_posix())
|
||||
video_num_frames = len(video_reader)
|
||||
|
||||
indices = list(range(0, video_num_frames, video_num_frames // max_num_frames))
|
||||
frames = video_reader.get_batch(indices)
|
||||
frames = frames[: max_num_frames].float()
|
||||
frames = frames.permute(0, 3, 1, 2).contiguous()
|
||||
frames = torch.stack([frame for frame in frames], dim=0)
|
||||
return frames
|
||||
|
||||
@staticmethod
|
||||
def preprocess_condition_video_image_reference_video(path: Path,resolution_buckets,max_num_frames) -> torch.Tensor:
|
||||
video_reader = decord.VideoReader(uri=Path(path).as_posix())
|
||||
|
||||
video_num_frames = len(video_reader)
|
||||
nearest_frame_bucket = min(
|
||||
[bucket for bucket in resolution_buckets if bucket[0] <= video_num_frames],
|
||||
key=lambda x: abs(x[0] - min(video_num_frames, max_num_frames)),
|
||||
default=1,
|
||||
)[0]
|
||||
|
||||
frame_indices = [0 for _ in range(video_num_frames)]
|
||||
frames = video_reader.get_batch(frame_indices)
|
||||
frames = frames[:nearest_frame_bucket].float()
|
||||
|
||||
frames = frames.permute(0, 3, 1, 2).contiguous()
|
||||
|
||||
# nearest_res = self._find_nearest_resolution(frames.shape[2], frames.shape[3])
|
||||
# frames_resized = torch.stack([frame for frame in frames], dim=0)
|
||||
frames = torch.stack([frame for frame in frames], dim=0)
|
||||
return frames
|
||||
|
||||
Reference in New Issue
Block a user