From 315431be0a784f572b6e854321ed5b52911f0255 Mon Sep 17 00:00:00 2001 From: CrossProduct Date: Thu, 16 Jan 2025 05:03:44 +0000 Subject: [PATCH] Recommit the conditioning code. --- finetrainers/args.py | 10 +++ finetrainers/dataset.py | 69 ++++++++++++++++---- finetrainers/ltx_video/lora.py | 1 + finetrainers/trainer.py | 113 +++++++++++++++++++++++++++------ 4 files changed, 160 insertions(+), 33 deletions(-) diff --git a/finetrainers/args.py b/finetrainers/args.py index 1f2cfd4..855173b 100644 --- a/finetrainers/args.py +++ b/finetrainers/args.py @@ -249,6 +249,8 @@ class Args: dataset_file: Optional[str] = None video_column: str = None caption_column: str = None + pose_column:str = None + id_token: Optional[str] = None image_resolution_buckets: List[Tuple[int, int]] = None video_resolution_buckets: List[Tuple[int, int, int]] = None @@ -353,6 +355,7 @@ class Args: "dataset_file": self.dataset_file, "video_column": self.video_column, "caption_column": self.caption_column, + "pose_column": self.pose_column, "id_token": self.id_token, "image_resolution_buckets": self.image_resolution_buckets, "video_resolution_buckets": self.video_resolution_buckets, @@ -550,6 +553,12 @@ def _add_dataset_arguments(parser: argparse.ArgumentParser) -> None: default="text", help="The column of the dataset containing the instance prompt for each video. Or, the name of the file in `--data_root` folder containing the line-separated instance prompts.", ) + parser.add_argument( + "--pose_column", + type=str, + default="text", + help="The column of the dataset containing the instance prompt for each video. Or, the name of the file in `--data_root` folder containing the line-separated instance prompts.", + ) parser.add_argument( "--id_token", type=str, @@ -1006,6 +1015,7 @@ def _map_to_args_type(args: Dict[str, Any]) -> Args: result_args.dataset_file = args.dataset_file result_args.video_column = args.video_column result_args.caption_column = args.caption_column + result_args.pose_column = args.pose_column result_args.id_token = args.id_token result_args.image_resolution_buckets = args.image_resolution_buckets or DEFAULT_IMAGE_RESOLUTION_BUCKETS result_args.video_resolution_buckets = args.video_resolution_buckets or DEFAULT_VIDEO_RESOLUTION_BUCKETS diff --git a/finetrainers/dataset.py b/finetrainers/dataset.py index 19ebb69..7b07b97 100644 --- a/finetrainers/dataset.py +++ b/finetrainers/dataset.py @@ -46,6 +46,7 @@ class ImageOrVideoDataset(Dataset): data_root: str, caption_column: str, video_column: str, + pose_column: str, resolution_buckets: List[Tuple[int, int, int]], dataset_file: Optional[str] = None, id_token: Optional[str] = None, @@ -55,8 +56,11 @@ class ImageOrVideoDataset(Dataset): self.data_root = Path(data_root) self.dataset_file = dataset_file + self.caption_column = caption_column self.video_column = video_column + self.pose_column = pose_column + self.id_token = f"{id_token.strip()} " if id_token else "" self.resolution_buckets = resolution_buckets @@ -72,6 +76,7 @@ class ImageOrVideoDataset(Dataset): ( self.prompts, self.video_paths, + self.pose_pathes, ) = self._load_dataset_from_local_path() elif dataset_file.endswith(".csv"): ( @@ -141,15 +146,28 @@ class ImageOrVideoDataset(Dataset): else: video = self._preprocess_video(video_path) - return { - "prompt": prompt, - "video": video, - "video_metadata": { - "num_frames": video.shape[0], - "height": video.shape[2], - "width": video.shape[3], - }, - } + if self.pose_condition_column != None: + pose_video = self._preprocess_video(self.pose_paths[index]) + return { + "prompt": prompt, + "video": video, + "pose_video": pose_video, + "video_metadata": { + "num_frames": video.shape[0], + "height": video.shape[2], + "width": video.shape[3], + }, + } + else: + return { + "prompt": prompt, + "video": video, + "video_metadata": { + "num_frames": video.shape[0], + "height": video.shape[2], + "width": video.shape[3], + }, + } def _load_dataset_from_local_path(self) -> Tuple[List[str], List[str]]: if not self.data_root.exists(): @@ -157,7 +175,8 @@ class ImageOrVideoDataset(Dataset): prompt_path = self.data_root.joinpath(self.caption_column) video_path = self.data_root.joinpath(self.video_column) - + pose_path = self.data_root.joinpath(self.pose_column) + if not prompt_path.exists() or not prompt_path.is_file(): raise ValueError( "Expected `--caption_column` to be path to a file in `--data_root` containing line-separated text prompts." @@ -166,18 +185,23 @@ class ImageOrVideoDataset(Dataset): raise ValueError( "Expected `--video_column` to be path to a file in `--data_root` containing line-separated paths to video data in the same directory." ) - + if not pose_path.exists() or not pose_path.is_file(): + raise ValueError( + "Expected `--pose_column` to be path to a file in `--data_root` containing line-separated paths to video data in the same directory." + ) with open(prompt_path, "r", encoding="utf-8") as file: prompts = [line.strip() for line in file.readlines() if len(line.strip()) > 0] with open(video_path, "r", encoding="utf-8") as file: video_paths = [self.data_root.joinpath(line.strip()) for line in file.readlines() if len(line.strip()) > 0] + with open(pose_path, "r", encoding="utf-8") as file: + pose_paths = [self.data_root.joinpath(line.strip()) for line in file.readlines() if len(line.strip()) > 0] if any(not path.is_file() for path in video_paths): raise ValueError( f"Expected `{self.video_column=}` to be a path to a file in `{self.data_root=}` containing line-separated paths to video data but found atleast one path that is not a valid file." ) - return prompts, video_paths + return prompts, video_paths, pose_paths def _load_dataset_from_csv(self) -> Tuple[List[str], List[str]]: df = pd.read_csv(self.dataset_file) @@ -243,7 +267,28 @@ class ImageOrVideoDataset(Dataset): frames = frames.permute(0, 3, 1, 2).contiguous() 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: + video_reader = decord.VideoReader(uri=path.as_posix()) + video_num_frames = len(video_reader) + nearest_frame_bucket = min( + [bucket for bucket in self.resolution_buckets if bucket[0] <= video_num_frames], + key=lambda x: abs(x[0] - min(video_num_frames, self.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([resize(frame, nearest_res) for frame in frames], dim=0) + frames = torch.stack([self.video_transforms(frame) for frame in frames_resized], dim=0) + + return frames class ImageOrVideoDatasetWithResizing(ImageOrVideoDataset): def __init__(self, *args, **kwargs) -> None: diff --git a/finetrainers/ltx_video/lora.py b/finetrainers/ltx_video/lora.py index a73abff..5e44c47 100644 --- a/finetrainers/ltx_video/lora.py +++ b/finetrainers/ltx_video/lora.py @@ -191,6 +191,7 @@ 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]]) } diff --git a/finetrainers/trainer.py b/finetrainers/trainer.py index 12e05d4..6de1128 100644 --- a/finetrainers/trainer.py +++ b/finetrainers/trainer.py @@ -100,19 +100,34 @@ class Trainer: self.state.model_name = self.args.model_name self.model_config = get_config_from_model_name(self.args.model_name, self.args.training_type) + if self.args.pose_column != None: + self.pose_condition = True + else: + self.pose_condition = False + def prepare_dataset(self) -> None: # TODO(aryan): Make a background process for fetching logger.info("Initializing dataset and dataloader") - - self.dataset = ImageOrVideoDatasetWithResizing( - data_root=self.args.data_root, - caption_column=self.args.caption_column, - video_column=self.args.video_column, - resolution_buckets=self.args.video_resolution_buckets, - dataset_file=self.args.dataset_file, - id_token=self.args.id_token, - remove_llm_prefixes=self.args.remove_common_llm_caption_prefixes, - ) + if self.pose_condition == True: + self.dataset = ImageOrVideoDatasetWithResizing( + data_root=self.args.data_root, + caption_column=self.args.caption_column, + video_column=self.args.video_column, + pose_column=self.args.pose_column, + resolution_buckets=self.args.video_resolution_buckets, + dataset_file=self.args.dataset_file, + id_token=self.args.id_token, + ) + else: + self.dataset = ImageOrVideoDatasetWithResizing( + data_root=self.args.data_root, + caption_column=self.args.caption_column, + video_column=self.args.video_column, + resolution_buckets=self.args.video_resolution_buckets, + dataset_file=self.args.dataset_file, + id_token=self.args.id_token, + remove_llm_prefixes=self.args.remove_common_llm_caption_prefixes, + ) self.dataloader = torch.utils.data.DataLoader( self.dataset, batch_size=1, @@ -647,6 +662,12 @@ class Trainer: if not self.args.precompute_conditions: videos = batch["videos"] prompts = batch["prompts"] + + if self.pose_condition == True: + poses = batch["poses"] + # first frame video ? cond video. + first_frame_videos = batch["first_frame_videos"] + batch_size = len(prompts) if self.args.caption_dropout_technique == "empty": @@ -691,6 +712,32 @@ class Trainer: batch_size = latent_conditions["latents"].shape[0] latent_conditions = make_contiguous(latent_conditions) + + if self.pose_conditioning: + pose_video_latents = self.model_config["prepare_latents"]( + vae=self.vae, + image_or_video=poses, + patch_size=self.transformer_config.patch_size, + patch_size_t=self.transformer_config.patch_size_t, + device=accelerator.device, + dtype=weight_dtype, + generator=generator, + ) + + pose_video_latents = make_contiguous(pose_video_latents) + + single_frame_video_latents = self.model_config["prepare_latents"]( + vae=self.vae, + image_or_video=first_frame_videos, + patch_size=self.transformer_config.patch_size, + patch_size_t=self.transformer_config.patch_size_t, + device=accelerator.device, + dtype=weight_dtype, + generator=generator, + ) + + single_frame_video_latents = make_contiguous(single_frame_video_latents) + text_conditions = make_contiguous(text_conditions) if self.args.caption_dropout_technique == "zero": @@ -734,10 +781,20 @@ class Trainer: timesteps=timesteps, ) 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) - + 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 = noisy_latents + pose_video_latents["latents"] + single_frame_video_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}) @@ -749,13 +806,27 @@ class Trainer: ) weights = expand_tensor_dims(weights, noise.ndim) - pred = self.model_config["forward_pass"]( - transformer=self.transformer, - scheduler=self.scheduler, - timesteps=timesteps, - **latent_conditions, - **text_conditions, - ) + if self.pose_conditioning: + pred = self.model_config["forward_pass"]( + transformer=self.transformer, + scheduler=self.scheduler, + timesteps=timesteps, + **single_frame_video_latents, + **text_conditions, + ) + + + else: + pred = self.model_config["forward_pass"]( + transformer=self.transformer, + scheduler=self.scheduler, + timesteps=timesteps, + **latent_conditions, + **text_conditions, + ) + + + target = prepare_target( scheduler=self.scheduler, noise=noise, latents=latent_conditions["latents"] )