Automated PR - 2026-01-29

This commit is contained in:
sync-bot
2026-01-29 18:42:17 +00:00
parent 727c43e998
commit ca1623ad2a
31 changed files with 1723 additions and 663 deletions
@@ -16,13 +16,14 @@ from pathlib import Path
import typer
from decode_latents import LatentsDecoder
from process_captions import compute_captions_embeddings
from process_videos import compute_latents, parse_resolution_buckets
from process_videos import compute_latents, compute_scaled_resolution_buckets, parse_resolution_buckets
from rich.console import Console
from ltx_trainer import logger
from ltx_trainer.gpu_utils import free_gpu_memory_context
console = Console()
app = typer.Typer(
pretty_exceptions_enable=False,
no_args_is_help=True,
@@ -46,6 +47,7 @@ def preprocess_dataset( # noqa: PLR0913
device: str,
remove_llm_prefixes: bool = False,
reference_column: str | None = None,
reference_downscale_factor: int = 1,
with_audio: bool = False,
load_text_encoder_in_8bit: bool = False,
) -> None:
@@ -99,14 +101,33 @@ def preprocess_dataset( # noqa: PLR0913
# Process reference videos if reference_column is provided
if reference_column:
logger.info("Processing reference videos for IC-LoRA training...")
# Validate: scaled references with multiple buckets can cause ambiguous bucket matching
if reference_downscale_factor > 1 and len(resolution_buckets) > 1:
raise ValueError(
"When using --reference-downscale-factor > 1, only a single resolution bucket is supported. "
"Using multiple buckets with scaled references can cause ambiguous bucket matching "
"(e.g., a 512x256 reference could match either the scaled-down 1024x512 bucket or the 512x256 "
"bucket). Please use a single resolution bucket or set --reference-downscale-factor to 1."
)
# Calculate and validate scaled resolution buckets for reference videos
reference_buckets = compute_scaled_resolution_buckets(resolution_buckets, reference_downscale_factor)
if reference_downscale_factor > 1:
logger.info(
f"Processing reference videos for IC-LoRA training at 1/{reference_downscale_factor} resolution..."
)
logger.info(f"Reference resolution buckets: {reference_buckets}")
else:
logger.info("Processing reference videos for IC-LoRA training...")
reference_latents_dir = output_base / "reference_latents"
compute_latents(
dataset_file=dataset_file,
main_media_column=video_column,
video_column=reference_column,
resolution_buckets=resolution_buckets,
resolution_buckets=reference_buckets,
output_dir=str(reference_latents_dir),
model_path=model_path,
batch_size=batch_size,
@@ -226,6 +247,11 @@ def main( # noqa: PLR0913
default=False,
help="Load the Gemma text encoder in 8-bit precision to save GPU memory (requires bitsandbytes)",
),
reference_downscale_factor: int = typer.Option(
default=1,
help="Downscale factor for reference video resolution. When > 1, reference videos are processed at "
"1/n resolution (e.g., 2 means half resolution). Used for efficient IC-LoRA training.",
),
) -> None:
"""Preprocess a video dataset by computing and saving latents and text embeddings.
The dataset must be a CSV, JSON, or JSONL file with columns for captions and video paths.
@@ -242,6 +268,10 @@ def main( # noqa: PLR0913
python scripts/process_dataset.py dataset.json --resolution-buckets 768x768x25 \\
--model-path /path/to/ltx2.safetensors --text-encoder-path /path/to/gemma \\
--reference-column "reference_path"
# Process dataset with scaled reference videos (half resolution) for efficient IC-LoRA
python scripts/process_dataset.py dataset.json --resolution-buckets 768x768x25 \\
--model-path /path/to/ltx2.safetensors --text-encoder-path /path/to/gemma \\
--reference-column "reference_path" --reference-downscale-factor 2
# Process dataset with audio for audio-video training
python scripts/process_dataset.py dataset.json --resolution-buckets 768x512x97 \\
--model-path /path/to/ltx2.safetensors --text-encoder-path /path/to/gemma \\
@@ -255,6 +285,13 @@ def main( # noqa: PLR0913
"When training with multiple resolution buckets, you must use a batch size of 1."
)
# Validate reference_downscale_factor
if reference_downscale_factor < 1:
raise typer.BadParameter("--reference-downscale-factor must be >= 1")
if reference_downscale_factor > 1 and not reference_column:
logger.warning("--reference-downscale-factor specified but no --reference-column provided. Ignoring.")
preprocess_dataset(
dataset_file=dataset_path,
caption_column=caption_column,
@@ -270,6 +307,7 @@ def main( # noqa: PLR0913
device=device,
remove_llm_prefixes=remove_llm_prefixes,
reference_column=reference_column,
reference_downscale_factor=reference_downscale_factor,
with_audio=with_audio,
load_text_encoder_in_8bit=load_text_encoder_in_8bit,
)