From 8ec35d58fb7422723fd6345a25dbc316db0f7bef Mon Sep 17 00:00:00 2001 From: CrossProduct Date: Tue, 28 Jan 2025 22:36:09 +0000 Subject: [PATCH] trying training again but broke something so going back. --- .../conditioning/conditioned_pipeline.py | 102 ++++++++++-------- .../ltx_video/full_finetune_condition.py | 11 +- finetrainers/trainer.py | 55 +++++++++- 3 files changed, 110 insertions(+), 58 deletions(-) diff --git a/finetrainers/conditioning/conditioned_pipeline.py b/finetrainers/conditioning/conditioned_pipeline.py index 087c28f..3f6a8dc 100644 --- a/finetrainers/conditioning/conditioned_pipeline.py +++ b/finetrainers/conditioning/conditioned_pipeline.py @@ -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 diff --git a/finetrainers/ltx_video/full_finetune_condition.py b/finetrainers/ltx_video/full_finetune_condition.py index a6fbd9d..3105fe1 100644 --- a/finetrainers/ltx_video/full_finetune_condition.py +++ b/finetrainers/ltx_video/full_finetune_condition.py @@ -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] diff --git a/finetrainers/trainer.py b/finetrainers/trainer.py index e64346f..1fba463 100644 --- a/finetrainers/trainer.py +++ b/finetrainers/trainer.py @@ -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