diff --git a/finetrainers/args.py b/finetrainers/args.py index cb4bc82..2948963 100644 --- a/finetrainers/args.py +++ b/finetrainers/args.py @@ -892,6 +892,12 @@ def _add_validation_arguments(parser: argparse.ArgumentParser) -> None: default=None, help="Run validation every X training steps. Validation consists of running the validation prompt `args.num_validation_videos` times.", ) + parser.add_argument( + "--validation_frame_rate", + type=int, + default=None, + help="Run inference validation at a specified frame rate every X training steps with the same output frame rate. `args.validation_frame_rate`", + ) parser.add_argument( "--enable_model_cpu_offload", action="store_true", diff --git a/finetrainers/ltx_video/ltx_video_lora.py b/finetrainers/ltx_video/ltx_video_lora.py index 0e1af9b..c5c1df2 100644 --- a/finetrainers/ltx_video/ltx_video_lora.py +++ b/finetrainers/ltx_video/ltx_video_lora.py @@ -225,7 +225,7 @@ def validation( height: Optional[int] = None, width: Optional[int] = None, num_frames: Optional[int] = None, - frame_rate: int = 25, + frame_rate: int = 24, num_videos_per_prompt: int = 1, generator: Optional[torch.Generator] = None, **kwargs, diff --git a/finetrainers/trainer.py b/finetrainers/trainer.py index d37eb81..67c1b1e 100644 --- a/finetrainers/trainer.py +++ b/finetrainers/trainer.py @@ -11,7 +11,6 @@ import diffusers import torch import torch.backends import transformers -import wandb from accelerate import Accelerator, DistributedType from accelerate.logging import get_logger from accelerate.utils import ( @@ -31,6 +30,8 @@ from huggingface_hub import create_repo, upload_folder from peft import LoraConfig, get_peft_model_state_dict, set_peft_model_state_dict from tqdm import tqdm +import wandb + from .args import _INVERSE_DTYPE_MAP, Args, validate_args from .constants import ( FINETRAINERS_LOG_LEVEL, @@ -926,7 +927,7 @@ class Trainer: height = self.args.validation_heights[i] width = self.args.validation_widths[i] num_frames = self.args.validation_num_frames[i] - + frame_rate = self.args.validation_frame_rate if image is not None: image = load_image(image) if video is not None: @@ -944,6 +945,7 @@ class Trainer: 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 @@ -983,7 +985,7 @@ class Trainer: elif artifact_type == "video": logger.debug(f"Saving video to {filename}") # TODO: this should be configurable here as well as in validation runs where we call the pipeline that has `fps`. - export_to_video(artifact_value, filename, fps=15) + export_to_video(artifact_value, filename, fps=frame_rate) artifact_value = wandb.Video(filename, caption=prompt) all_processes_artifacts.append(artifact_value)