trying training again but broke something so going back.

This commit is contained in:
CrossProduct
2025-01-28 22:36:09 +00:00
parent dfa03f1ed1
commit 8ec35d58fb
3 changed files with 110 additions and 58 deletions
@@ -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
View File
@@ -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