Passed data in for validation.

This commit is contained in:
CrossProduct
2025-01-25 03:07:12 +00:00
parent e99d70ff46
commit 80f653f336
3 changed files with 71 additions and 19 deletions
+18
View File
@@ -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
@@ -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
+45 -19
View File
@@ -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 = {