mirror of
https://github.com/storytold/FineTrainers-Conditioning.git
synced 2026-10-09 00:09:45 +00:00
LTX uses a default frame rate of 24 FPS
We need to modify the output validation framerate to match that value. Add Framerate args. Add Update video output and inference frame rate
This commit is contained in:
@@ -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",
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user