mirror of
https://github.com/storytold/FineTrainers-Conditioning.git
synced 2026-10-09 00:09:45 +00:00
Passed data in for validation.
This commit is contained in:
@@ -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
@@ -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 = {
|
||||
|
||||
Reference in New Issue
Block a user