diff --git a/finetrainers/args.py b/finetrainers/args.py index fb79fe5..e0b28b3 100644 --- a/finetrainers/args.py +++ b/finetrainers/args.py @@ -316,7 +316,11 @@ class Args: validation_prompts: List[str] = None validation_images: List[str] = None validation_videos: List[str] = None + + # Condition validation_pose_videos: List[str] = None + validation_img_ref_videos: List[str] = None + validation_heights: List[int] = None validation_widths: List[int] = None validation_num_frames: List[int] = None @@ -888,6 +892,13 @@ def _add_validation_arguments(parser: argparse.ArgumentParser) -> None: help="One or more pose_videos path(s)/URLs that is used during validation to verify that the model is learning. Multiple validation paths should be separated by the '--validation_prompt_seperator' string. These should correspond to the order of the validation prompts.", ) + parser.add_argument( + "--validation_img_ref_videos", + type=str, + default=None, + help="One or more img_ref video or img will just take the first frame and create a video from it path(s)/URLs that is used during validation to verify that the model is learning. Multiple validation paths should be separated by the '--validation_prompt_seperator' string. These should correspond to the order of the validation prompts.", + ) + parser.add_argument( "--validation_videos", type=str, @@ -1089,7 +1100,10 @@ def _map_to_args_type(args: Dict[str, Any]) -> Args: validation_images = args.validation_images.split(args.validation_separator) if args.validation_images else None validation_videos = args.validation_videos.split(args.validation_separator) if args.validation_videos else None + # extend with pose conditioning ... validation_pose_videos = args.validation_pose_videos.split(args.validation_separator) if args.validation_pose_videos else None + validation_img_ref_videos = args.validation_img_ref_videos.split(args.validation_separator) if args.validation_img_ref_videos else None + stripped_validation_prompts = [] validation_heights = [] @@ -1118,7 +1132,11 @@ def _map_to_args_type(args: Dict[str, Any]) -> Args: result_args.validation_num_frames = validation_num_frames result_args.validation_images = validation_images result_args.validation_videos = validation_videos + + # extend with pose conditioning ... result_args.validation_pose_videos = validation_pose_videos + result_args.validation_img_ref_videos = validation_img_ref_videos + result_args.num_validation_videos_per_prompt = args.num_validation_videos result_args.validation_every_n_epochs = args.validation_epochs result_args.validation_every_n_steps = args.validation_steps diff --git a/finetrainers/conditioning/conditioned_pipeline.py b/finetrainers/conditioning/conditioned_pipeline.py index c13dd39..cff5cb3 100644 --- a/finetrainers/conditioning/conditioned_pipeline.py +++ b/finetrainers/conditioning/conditioned_pipeline.py @@ -62,7 +62,15 @@ class LTXConditionedPipeline(LTXPipeline): callback_on_step_end: Optional[Callable[[int, int, Dict], None]] = None, callback_on_step_end_tensor_inputs: List[str] = ["latents"], max_sequence_length: int = 128, + pose_video=None, + image_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: + raise ValueError("image_ref_video cannot be None.") + if isinstance(callback_on_step_end, (PipelineCallback, MultiPipelineCallbacks)): callback_on_step_end_tensor_inputs = callback_on_step_end.tensor_inputs diff --git a/finetrainers/trainer.py b/finetrainers/trainer.py index bb2d2a0..e64346f 100644 --- a/finetrainers/trainer.py +++ b/finetrainers/trainer.py @@ -1010,7 +1010,12 @@ class Trainer: prompt = self.args.validation_prompts[i] image = self.args.validation_images[i] video = self.args.validation_videos[i] - pose_video = self.args.validation_pose_videos[i] + # Condition extension + + if self.pose_condition: + pose_video = self.args.validation_pose_videos[i] + img_ref_video = self.args.validation_img_ref_videos[i] + height = self.args.validation_heights[i] width = self.args.validation_widths[i] num_frames = self.args.validation_num_frames[i] @@ -1019,30 +1024,51 @@ class Trainer: image = load_image(image) if video is not None: video = load_video(video) - if pose_video is not None: - pose_video = load_video(pose_video) + if self.pose_condition: + if pose_video is not None: + pose_video = load_video(pose_video) + if img_ref_video is not None: + img_ref_video = load_video(img_ref_video) + logger.debug( f"Validating sample {i + 1}/{num_validation_samples} on process {accelerator.process_index}. Prompt: {prompt}", main_process_only=False, ) - validation_artifacts = self.model_config["validation"]( - pipeline=pipeline, - prompt=prompt, - image=image, - video=video, - pose_video=pose_video, - height=height, - width=width, - num_frames=num_frames, - frame_rate=frame_rate, - num_videos_per_prompt=self.args.num_validation_videos_per_prompt, - generator=torch.Generator(device=accelerator.device).manual_seed( - self.args.seed if self.args.seed is not None else 0 - ), - # todo support passing `fps` for supported pipelines. - ) + if self.pose_condition: + validation_artifacts = self.model_config["validation"]( + pipeline=pipeline, + prompt=prompt, + image=image, + video=video, + height=height, + width=width, + num_frames=num_frames, + frame_rate=frame_rate, + pose_video=pose_video, + img_ref_video=img_ref_video, + num_videos_per_prompt=self.args.num_validation_videos_per_prompt, + generator=torch.Generator(device=accelerator.device).manual_seed( + self.args.seed if self.args.seed is not None else 0 + ), + ) + else: + validation_artifacts = self.model_config["validation"]( + pipeline=pipeline, + prompt=prompt, + image=image, + video=video, + height=height, + width=width, + num_frames=num_frames, + frame_rate=frame_rate, + num_videos_per_prompt=self.args.num_validation_videos_per_prompt, + generator=torch.Generator(device=accelerator.device).manual_seed( + self.args.seed if self.args.seed is not None else 0 + ), + # todo support passing `fps` for supported pipelines. + ) prompt_filename = string_to_filename(prompt)[:25] artifacts = {