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:
CrossProduct
2025-01-12 04:39:05 +00:00
parent a78fedd108
commit 4354fb3164
3 changed files with 12 additions and 4 deletions
+6
View File
@@ -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",
+1 -1
View File
@@ -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,
+5 -3
View File
@@ -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)