From cd66e6e58e857d5b9a954eabeff43703b8f97176 Mon Sep 17 00:00:00 2001 From: CrossProduct Date: Fri, 17 Jan 2025 01:46:35 +0000 Subject: [PATCH] fixed the training pipeline to address new changes from upstream --- finetrainers/dataset.py | 13 ++++++++----- finetrainers/ltx_video/lora.py | 9 ++------- finetrainers/trainer.py | 24 ++++++++++-------------- 3 files changed, 20 insertions(+), 26 deletions(-) diff --git a/finetrainers/dataset.py b/finetrainers/dataset.py index 7b07b97..721730d 100644 --- a/finetrainers/dataset.py +++ b/finetrainers/dataset.py @@ -76,7 +76,7 @@ class ImageOrVideoDataset(Dataset): ( self.prompts, self.video_paths, - self.pose_pathes, + self.pose_paths, ) = self._load_dataset_from_local_path() elif dataset_file.endswith(".csv"): ( @@ -146,12 +146,15 @@ class ImageOrVideoDataset(Dataset): else: video = self._preprocess_video(video_path) - if self.pose_condition_column != None: - pose_video = self._preprocess_video(self.pose_paths[index]) + if self.pose_column != None: + pose = self._preprocess_video(self.pose_paths[index]) + img_ref = self._preprocess_video_image_reference_video(video_path) + return { "prompt": prompt, "video": video, - "pose_video": pose_video, + "pose": pose, + "img_ref": img_ref, "video_metadata": { "num_frames": video.shape[0], "height": video.shape[2], @@ -268,7 +271,7 @@ class ImageOrVideoDataset(Dataset): frames = torch.stack([self.video_transforms(frame) for frame in frames], dim=0) return frames - def _preprocess_video_single_frame(self, path: Path) -> torch.Tensor: + def _preprocess_video_image_reference_video(self, path: Path) -> torch.Tensor: video_reader = decord.VideoReader(uri=path.as_posix()) video_num_frames = len(video_reader) diff --git a/finetrainers/ltx_video/lora.py b/finetrainers/ltx_video/lora.py index 5e44c47..e7c04e6 100644 --- a/finetrainers/ltx_video/lora.py +++ b/finetrainers/ltx_video/lora.py @@ -139,15 +139,9 @@ def prepare_latents( if not precompute: latents = vae.encode(image_or_video).latent_dist.sample(generator=generator) - - latents = latents.to(dtype=dtype) _, _, num_frames, height, width = latents.shape latents = _normalize_latents(latents, vae.latents_mean, vae.latents_std) - - # expand the channel and pack - latents = torch.cat([latents,latents],dim=1) - latents = _pack_latents(latents, patch_size, patch_size_t) return {"latents": latents, "num_frames": num_frames, "height": height, "width": width} else: @@ -191,7 +185,8 @@ def collate_fn_t2v(batch: List[List[Dict[str, torch.Tensor]]]) -> Dict[str, torc return { "prompts": [x["prompt"] for x in batch[0]], "videos": torch.stack([x["video"] for x in batch[0]]), - "poses": torch.stack([x["poses"] for x in batch[0]]) + "poses": torch.stack([x["pose"] for x in batch[0]]), + "img_refs": torch.stack([x["img_ref"] for x in batch[0]]) } diff --git a/finetrainers/trainer.py b/finetrainers/trainer.py index 6de1128..a7a919d 100644 --- a/finetrainers/trainer.py +++ b/finetrainers/trainer.py @@ -665,8 +665,8 @@ class Trainer: if self.pose_condition == True: poses = batch["poses"] - # first frame video ? cond video. - first_frame_videos = batch["first_frame_videos"] + # generate video across all frames + img_refs = batch["img_refs"] batch_size = len(prompts) @@ -726,9 +726,9 @@ class Trainer: pose_video_latents = make_contiguous(pose_video_latents) - single_frame_video_latents = self.model_config["prepare_latents"]( + img_refs_latents = self.model_config["prepare_latents"]( vae=self.vae, - image_or_video=first_frame_videos, + image_or_video=img_refs, patch_size=self.transformer_config.patch_size, patch_size_t=self.transformer_config.patch_size_t, device=accelerator.device, @@ -736,7 +736,7 @@ class Trainer: generator=generator, ) - single_frame_video_latents = make_contiguous(single_frame_video_latents) + img_refs_latents = make_contiguous(img_refs_latents) text_conditions = make_contiguous(text_conditions) @@ -782,18 +782,15 @@ class Trainer: ) else: if self.pose_conditioning: - # Use the training single frame - sample compare it to target - # todo fix it so the first frame isn't noised? - noisy_latents = (1.0 - sigmas) * single_frame_video_latents["latents"] + sigmas * noise - # add to noisy latent with the single frame repeating the conditioning to the pose latent + + noisy_latents = (1.0 - sigmas) * img_refs_latents["latents"] + sigmas * noise + noisy_latents = noisy_latents + pose_video_latents["latents"] - single_frame_video_latents.update({"noisy_latents": noisy_latents}) + img_refs_latents.update({"noisy_latents": noisy_latents}) else: # Default to flow-matching noise addition noisy_latents = (1.0 - sigmas) * latent_conditions["latents"] + sigmas * noise noisy_latents = noisy_latents.to(latent_conditions["latents"].dtype) - - # this is used for whatever reason to pass into the latent_conditions.update({"noisy_latents": noisy_latents}) @@ -811,11 +808,10 @@ class Trainer: transformer=self.transformer, scheduler=self.scheduler, timesteps=timesteps, - **single_frame_video_latents, + **img_refs_latents, **text_conditions, ) - else: pred = self.model_config["forward_pass"]( transformer=self.transformer,