Compare commits

...
Author SHA1 Message Date
SolitaryThinker a8dddfaa16 missing file 2026-02-10 08:33:00 +00:00
Will Lin 7210c68f1b lint 2026-02-10 00:30:03 -08:00
SolitaryThinker 5602dc1bad revert 2026-02-10 08:22:05 +00:00
SolitaryThinker ac4bc4ab84 uipdate 2026-02-10 07:58:43 +00:00
Matthew Noto bee27f9f74 Merge branch 'main' into ltx-base 2026-02-09 17:51:37 -08:00
ad58f802f3 [Feat] Port LTX2 trainer (#1074)
Co-authored-by: Davids048 <jundasu@ucsd.edu>
Co-authored-by: Matthew Noto <99706358+RandNMR73@users.noreply.github.com>
2026-02-09 17:32:57 -08:00
Wei Zhou 04fa356ee3 [Misc] [Training] Fixed a bunch of bugs in current training pipeline (#1084) 2026-02-09 16:01:05 -08:00
Matthew Notoandgemini-code-assist[bot] f9c076fe2b [misc] add AGENTS.md file (#1085)
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
2026-02-09 01:03:49 -08:00
Will Lin dff0ea401a update 2026-02-08 00:22:21 -08:00
Davids048 becd379f58 Split LTX2 mappings and add registry coverage.
Detail:

- Merged LTX2 sampling behavior into the global registry and removed the ambiguous local converted-path default.
- Replaced single LTX2 sampling mapping with explicit model-ID mappings:
    - Lightricks/LTX-2 -> LTX2BaseSamplingParam
    - FastVideo/LTX2-base -> LTX2BaseSamplingParam
    - FastVideo/LTX2-Distilled-Diffusers -> LTX2DistilledSamplingParam
- Kept LTX2T2VConfig as the pipeline config for all explicitly mapped LTX2 IDs.
- Removed implicit mapping for converted/ltx2_diffusers to avoid guessing base vs distilled for user-local
  conversions.
- Added focused local tests at tests/local_tests/test_ltx2_registry.py for:
    - exact base/distilled sampling resolution,
    - pipeline config resolution,
    - no fallback behavior for ambiguous local converted paths.

Assumptions:

- Canonical rename/ID intent:
    - “base” names map to LTX2BaseSamplingParam.
    - “Distilled” names map to LTX2DistilledSamplingParam.
- converted/ltx2_diffusers is intentionally ambiguous across users and must not be auto-assigned.
- Unknown/non-canonical names containing “LTX”/“LTX2” (but not matching explicit registered IDs) should not auto-
  resolve to base or distilled.
    - Sampling resolver returns None (caller falls back to generic defaults/user overrides).
    - Pipeline config lookup raises a “No match found” error.

Notes:

- This change prioritizes explicitness over convenience: only predetermined, canonical model IDs get LTX2-specific
  defaults; everything else requires user intent.
2026-02-05 20:28:46 -08:00
Davids048 1ed7d7e1b0 Add gemma tokenizer to LTX2 conversion script.
- Also clean up for PR.
2026-02-05 18:53:22 -08:00
Davids048 d9fabcc5ef Add some annotations. 2026-02-05 18:53:22 -08:00
Davids048 9db48498de Update quality test script. 2026-02-05 18:53:22 -08:00
Davids048 1cd7038315 Add LTX2 base model. 2026-02-05 18:53:12 -08:00
46 changed files with 2377 additions and 195 deletions
+42
View File
@@ -0,0 +1,42 @@
# Repository Guidelines
## Project Structure & Module Organization
- Core Python package: `fastvideo/` (models, pipelines, training, distributed runtime, CLI entrypoints).
- CUDA/custom kernels: `fastvideo-kernel/` (separate build/test flow).
- Tests:
- `fastvideo/tests/` for package-level tests (dataset, encoders, inference, training, SSIM, workflow).
- `tests/local_tests/` for additional local/component checks.
- Docs and guides: `docs/` (MkDocs source), with contributor docs in `docs/contributing/`.
- Runnable examples and scripts: `examples/` and `scripts/`.
- Static assets: `assets/`, `images/`, `videos/`, and `comfyui/assets/`.
## Build, Test, and Development Commands
- `uv pip install -e .[dev]`: editable install with lint/test extras.
- `pre-commit install --hook-type pre-commit --hook-type commit-msg`: enable local hooks.
- `pre-commit run --all-files`: run formatter/lint/type/spelling checks.
- `pytest tests/`: run top-level test suite.
- `pytest fastvideo/tests/ -v`: run package tests.
- `pytest fastvideo/tests/ssim/ -vs`: run SSIM regression tests (GPU-heavy).
- `cd fastvideo-kernel && ./build.sh`: build kernel extensions.
## Coding Style & Naming Conventions
- Python 3.10+; 4-space indentation; keep code and imports readable and explicit.
- Style tools are configured in `pyproject.toml` and `.pre-commit-config.yaml`:
- `yapf` (format), `ruff` (lint, auto-fix), `mypy` (typing), `codespell`.
- Target line length is 80.
- Naming: `snake_case` for functions/files, `PascalCase` for classes, `UPPER_SNAKE_CASE` for constants.
## Testing Guidelines
- Use `pytest` and place tests near relevant domains (e.g., `fastvideo/tests/encoders/`).
- Prefer descriptive names like `test_<feature>_<expected_behavior>.py`.
- For new pipelines/backends, include at least one regression-oriented test; add SSIM coverage when output quality must be preserved.
- Document GPU assumptions in tests that require specific hardware.
## Commit & Pull Request Guidelines
- Follow existing commit style: short subject with optional tag prefix, e.g. `[bugfix]: ...`, `[feat]: ...`, `[misc]: ...`, and include PR reference like `(#1234)` when applicable.
- Keep commits focused by concern (feature, refactor, fix).
- PRs should include:
- clear problem/solution summary,
- test evidence (`pytest`/SSIM outputs or rationale if skipped),
- linked issue/PR context,
- screenshots or sample outputs for UI/demo/docs changes.
+1
View File
@@ -0,0 +1 @@
@AGENTS.md
+6 -1
View File
@@ -1,5 +1,10 @@
<div align="center">
<img src=assets/logos/logo.svg width="30%"/>
</div>
| **[Documentation](https://hao-ai-lab.github.io/FastVideo)** | **[Quick Start](https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start/)** | **[Weekly Dev Meeting](https://github.com/hao-ai-lab/FastVideo/discussions/982)** | 🟣💬 **[Slack](https://join.slack.com/t/fastvideo/shared_invite/zt-3f4lao1uq-u~Ipx6Lt4J27AlD2y~IdLQ)** |
<p align="center">
| <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start/"><b> Quick Start</b></a> | <a href="https://github.com/hao-ai-lab/FastVideo/discussions/982" target="_blank"><b>Weekly Dev Meeting</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-3f4lao1uq-u~Ipx6Lt4J27AlD2y~IdLQ" target="_blank"> <b>Slack</b> </a> | 🟣💬 <a href="https://ibb.co/sv3MMKyv" target="_blank"> <b> WeChat </b> </a> |
</p>
**FastVideo is a unified post-training and inference framework for accelerated video generation.**
+7 -2
View File
@@ -1,5 +1,6 @@
from fastvideo import VideoGenerator
PROMPT = (
"A warm sunny backyard. The camera starts in a tight cinematic close-up "
"of a woman and a man in their 30s, facing each other with serious "
@@ -16,16 +17,20 @@ PROMPT = (
def main() -> None:
# Uses FastVideo default sampling settings for LTX2 base.
generator = VideoGenerator.from_pretrained(
"FastVideo/LTX2-Distilled-Diffusers",
"Davids048/LTX2-Base-Diffusers",
num_gpus=1,
)
output_path = "outputs_video/ltx2_basic/output_ltx2_distilled_t2v.mp4"
output_path = "outputs_video/ltx2_basic/output_ltx2_base_t2v_1088_1920_1.1.mp4"
generator.generate_video(
prompt=PROMPT,
output_path=output_path,
save_video=True,
num_frames=121,
height=1088,
width=1920,
)
generator.shutdown()
@@ -0,0 +1,34 @@
from fastvideo import VideoGenerator
PROMPT = (
"A warm sunny backyard. The camera starts in a tight cinematic close-up "
"of a woman and a man in their 30s, facing each other with serious "
"expressions. The woman, emotional and dramatic, says softly, \"That's "
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
"annoyed: \"Stop being so dramatic, Jess.\" A beat. He glances aside, "
"then mutters defensively, \"He's just having fun.\" The camera slowly "
"pans right, revealing the grandfather in the garden wearing enormous "
"butterfly wings, waving his arms in the air like he's trying to take "
"off. He shouts, \"Wheeeew!\" as he flaps his wings with full commitment. "
"The woman covers her face, on the verge of tears. The tone is deadpan, "
"absurd, and quietly tragic."
)
def main() -> None:
generator = VideoGenerator.from_pretrained(
"FastVideo/LTX2-Distilled-Diffusers",
num_gpus=1,
)
output_path = "outputs_video/ltx2_basic/output_ltx2_distilled_t2v.mp4"
generator.generate_video(
prompt=PROMPT,
output_path=output_path,
save_video=True,
)
generator.shutdown()
if __name__ == "__main__":
main()
@@ -4,7 +4,7 @@ export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export TOKENIZERS_PARALLELISM=false
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
DATA_DIR="data/crush-smol_processed_t2v_1_3b_ode_init/"
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
NUM_GPUS=1
@@ -14,7 +14,6 @@ NUM_GPUS=1
training_args=(
--tracker_project_name "wan_ode_init"
--output_dir "wan_ode_init_crush_smol"
--override_transformer_cls_name "CausalWanTransformer3DModel"
--wandb_run_name "wan_ode_init_crush_smol"
--max_train_steps 6000
--train_batch_size 1
@@ -34,7 +33,7 @@ parallel_args=(
--sp_size 1
--tp_size 1
--hsdp_replicate_dim 1
--hsdp_shard_dim 1
--hsdp_shard_dim $NUM_GPUS
)
# Model arguments
@@ -51,20 +50,17 @@ dataset_args=(
# Validation arguments
validation_args=(
--log_validation
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 50
--validation_sampling_steps "50"
--validation_guidance_scale "6.0"
--log-visualization
--visualization-steps 100
)
# Optimizer arguments
optimizer_args=(
--learning_rate 6e-6
--learning_rate 1e-5
--mixed_precision "bf16"
--weight_only_checkpointing_steps 1000
--training_state_checkpointing_steps 1000
--weight_decay 1e-4
--weight_decay 0.01
--max_grad_norm 1.0
)
@@ -0,0 +1,23 @@
# LTX-2 Crush-Smol Example
# TODO: Update this doc.
These are e2e example scripts for finetuning LTX-2 on the crush-smol dataset.
## Execute the following commands from `FastVideo/` to run training:
### Download crush-smol dataset:
`bash examples/training/finetune/ltx2/overfit/download_dataset.sh`
### Preprocess the videos and captions into latents:
`bash examples/training/finetune/ltx2/overfit/preprocess_ltx2_data_t2v_new.sh`
### Edit the following file and run finetuning:
`bash examples/training/finetune/ltx2/overfit/finetune_t2v.sh`
Notes:
- Update `DATASET_PATH` in the preprocess script to point to your merged dataset root (`videos/` + `videos2caption.json`).
- `MODEL_PATH` should point to a local LTX-2 diffusers-style directory that contains `model_index.json` and `text_encoder/gemma`.
@@ -0,0 +1,3 @@
# #!/bin/bash
#
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
@@ -0,0 +1,95 @@
#!/bin/bash
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export TOKENIZERS_PARALLELISM=false
MODEL_PATH="Davids048/LTX2-Base-Diffusers"
# Also can use simple 1 video for overfitting experiments.
# DATA_DIR="/home/hal-jundas/codes/FastVideo/data/crush-smol"
DATA_DIR="<PATH_TO_PROCESSED_DATASET>"
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
echo VALIDATION_DATASET_FILE: $VALIDATION_DATASET_FILE
NUM_GPUS=4
OVERFIT_HEIGHT=480
OVERFIT_WIDTH=832
OVERFIT_FRAMES=73
training_args=(
--tracker_project_name "ltx2_t2v_finetune"
--output_dir "checkpoints/ltx2_t2v_finetune"
--max_train_steps 5000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 10
--num_height $OVERFIT_HEIGHT
--num_width $OVERFIT_WIDTH
--num_frames $OVERFIT_FRAMES
--ltx2-first-frame-conditioning-p 0.1
--enable_gradient_checkpointing_type "full"
--mode "finetuning"
)
parallel_args=(
--num_gpus $NUM_GPUS
--sp_size $NUM_GPUS
--tp_size 1
--hsdp_replicate_dim 1
--hsdp_shard_dim $NUM_GPUS
)
model_args=(
--model_path $MODEL_PATH
--pretrained_model_name_or_path $MODEL_PATH
)
dataset_args=(
--data_path $DATA_DIR
--dataloader_num_workers 1
)
validation_args=(
--log_validation
--validation_dataset_file $VALIDATION_DATASET_FILE
--validation_steps 50
--validation_sampling_steps "50"
--validation_guidance_scale "3.0"
)
optimizer_args=(
--learning_rate 1e-5
--mixed_precision "bf16"
--weight_only_checkpointing_steps 1000
--training_state_checkpointing_steps 1000
--weight_decay 1e-4
--max_grad_norm 1.0
--lr_scheduler "linear"
)
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
--dit_precision "fp32"
--dit_cpu_offload False
--dit_layerwise_offload False
--text_encoder_cpu_offload False
--image_encoder_cpu_offload False
--vae_cpu_offload False
)
# NOTE: Setting this environment variable to TORCH_SDPA to avoid the issue of stacking that failed in flash attn.
export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
torchrun \
--nnodes 1 \
--nproc_per_node $NUM_GPUS \
fastvideo/training/ltx2_training_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}"
@@ -0,0 +1,80 @@
#!/bin/bash
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export TOKENIZERS_PARALLELISM=false
MODEL_PATH="/path/to/LTX-2"
DATA_DIR="data/crush-smol"
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
NUM_GPUS=1
training_args=(
--tracker_project_name "ltx2_t2v_lora_finetune"
--output_dir "checkpoints/ltx2_t2v_lora_finetune"
--max_train_steps 2000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 8
--num_latent_t 10
--num_height 480
--num_width 832
--num_frames 77
--ltx2-first-frame-conditioning-p 0.1
--enable_gradient_checkpointing_type "full"
)
parallel_args=(
--num_gpus $NUM_GPUS
--sp_size $NUM_GPUS
--tp_size 1
--hsdp_replicate_dim 1
--hsdp_shard_dim $NUM_GPUS
)
model_args=(
--model_path $MODEL_PATH
--pretrained_model_name_or_path $MODEL_PATH
)
dataset_args=(
--data_path $DATA_DIR
--dataloader_num_workers 1
)
validation_args=(
--log_validation
--validation_dataset_file $VALIDATION_DATASET_FILE
--validation_steps 200
--validation_sampling_steps "50"
--validation_guidance_scale "3.0"
)
optimizer_args=(
--learning_rate 2e-4
--mixed_precision "bf16"
--weight_only_checkpointing_steps 1000
--training_state_checkpointing_steps 1000
--weight_decay 1e-4
--max_grad_norm 1.0
--lora_training True
--lora_rank 16
--lora_alpha 16
)
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
)
torchrun \
--nnodes 1 \
--nproc_per_node $NUM_GPUS \
fastvideo/training/ltx2_training_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}"
@@ -0,0 +1,35 @@
#!/bin/bash
GPU_NUM=1
MODEL_PATH="Davids048/LTX2-Base-Diffusers"
# DATASET_PATH="data/overfit"
DATASET_PATH="data/crush-smol"
OUTPUT_DIR="$DATASET_PATH"
WITH_AUDIO=true
# Convert one-file overfit metadata into merged format if needed.
if [ ! -f "$DATASET_PATH/videos2caption.json" ] && [ -f "$DATASET_PATH/overfit.json" ]; then
python scripts/dataset_preparation/convert_to_merged_dataset.py \
--items-json "$DATASET_PATH/overfit.json" \
--output-dir "$DATASET_PATH"
fi
torchrun --nproc_per_node=$GPU_NUM \
--master_port=29513 \
-m fastvideo.pipelines.preprocess.v1_preprocessing_new \
--model_path $MODEL_PATH \
--mode preprocess \
--workload_type t2v \
--preprocess.video_loader_type torchvision \
--preprocess.dataset_type merged \
--preprocess.dataset_path $DATASET_PATH \
--preprocess.dataset_output_dir $OUTPUT_DIR \
--preprocess.with_audio $WITH_AUDIO \
--preprocess.preprocess_video_batch_size 1 \
--preprocess.dataloader_num_workers 0 \
--preprocess.max_height 480 \
--preprocess.max_width 832 \
--preprocess.num_frames 73 \
--preprocess.train_fps 16 \
--preprocess.video_length_tolerance_range 5
@@ -0,0 +1,13 @@
{
"data": [
{
"caption": "The camera opens in a calm, sunlit frog yoga studio. Warm morning light washes over the wooden floor as incense smoke drifts lazily in the air. The senior frog instructor sits cross-legged at the center, eyes closed, voice deep and calm. “We are one with the pond.” All the frogs answer softly: “Ommm...” “We are one with the mud.” “Ommm...” He smiles faintly. “We are one with the flies.” A quiet pause. The camera slowly pans to the side — one frog twitches, eyes darting. Suddenly — *thwip!* — its tongue snaps out, catching a fly mid-air and pulling it into its mouth. The master exhales slowly, still serene.",
"image_path": null,
"video_path": null,
"num_inference_steps": 50,
"height": 1088,
"width": 1920,
"num_frames": 121
}
]
}
@@ -0,0 +1,31 @@
{
"data": [
{
"caption": "A large metal cylinder is seen pressing down on a pile of Oreo cookies, flattening them as if they were under a hydraulic press.",
"image_path": null,
"video_path": null,
"num_inference_steps": 50,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.",
"image_path": null,
"video_path": null,
"num_inference_steps": 50,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "A large metal cylinder is seen pressing down on a pile of colorful candies, flattening them as if they were under a hydraulic press. The candies are crushed and broken into small pieces, creating a mess on the table.",
"image_path": null,
"video_path": null,
"num_inference_steps": 50,
"height": 480,
"width": 832,
"num_frames": 77
}
]
}
+5
View File
@@ -86,6 +86,7 @@ class PreprocessConfig:
# Model configuration
training_cfg_rate: float = 0.0
with_audio: bool = False
# framework configuration
seed: int = 42
@@ -190,6 +191,10 @@ class PreprocessConfig:
type=float,
default=PreprocessConfig.training_cfg_rate,
help="Training CFG rate")
preprocess_args.add_argument(f"--{prefix_with_dot}with-audio",
action=StoreBoolean,
default=PreprocessConfig.with_audio,
help="Whether to extract and encode audio")
preprocess_args.add_argument(f"--{prefix_with_dot}seed",
type=int,
default=PreprocessConfig.seed,
+4 -2
View File
@@ -7,10 +7,12 @@ from dataclasses import dataclass, field
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
import re
def is_ltx2_blocks(name: str, _module) -> bool:
"""FSDP shard condition for LTX-2 transformer blocks."""
return "transformer_blocks" in name
res = re.search(r"(?:^|\.)transformer_blocks\.\d+$", name) is not None
return res
@dataclass
+40 -2
View File
@@ -5,10 +5,44 @@ from fastvideo.configs.sample.base import SamplingParam
@dataclass
class LTX2SamplingParam(SamplingParam):
"""Default sampling parameters for LTX-2 distilled T2V.
class LTX2BaseSamplingParam(SamplingParam):
"""Default sampling parameters for LTX-2 base one-stage T2V.
Values follow the official LTX-2 one-stage defaults.
"""
seed: int = 10
num_frames: int = 121
height: int = 512
width: int = 768
fps: int = 24
num_inference_steps: int = 40
guidance_scale: float = 3.0
# Copied/following official LTX-2 DEFAULT_NEGATIVE_PROMPT.
negative_prompt: str = (
"blurry, out of focus, overexposed, underexposed, low contrast, "
"washed out colors, excessive noise, grainy texture, poor lighting, "
"flickering, motion blur, distorted proportions, unnatural skin "
"tones, deformed facial features, asymmetrical face, missing facial "
"features, extra limbs, disfigured hands, wrong hand count, "
"artifacts around text, inconsistent perspective, camera shake, "
"incorrect depth of field, background too sharp, background clutter, "
"distracting reflections, harsh shadows, inconsistent lighting "
"direction, color banding, cartoonish rendering, 3D CGI look, "
"unrealistic materials, uncanny valley effect, incorrect ethnicity, "
"wrong gender, exaggerated expressions, wrong gaze direction, "
"mismatched lip sync, silent or muted audio, distorted voice, "
"robotic voice, echo, background noise, off-sync audio, incorrect "
"dialogue, added dialogue, repetitive speech, jittery movement, "
"awkward pauses, incorrect timing, unnatural transitions, "
"inconsistent framing, tilted camera, flat lighting, inconsistent "
"tone, cinematic oversaturation, stylized filters, or AI artifacts.")
@dataclass
class LTX2DistilledSamplingParam(SamplingParam):
"""Default sampling parameters for LTX-2 distilled one-stage T2V."""
seed: int = 10
num_frames: int = 121
height: int = 1024
@@ -18,3 +52,7 @@ class LTX2SamplingParam(SamplingParam):
guidance_scale: float = 1.0
# No default negative_prompt for distilled models
negative_prompt: str = ""
# Backward compatibility alias.
LTX2SamplingParam = LTX2DistilledSamplingParam
+8 -2
View File
@@ -4,6 +4,8 @@ from torchvision.transforms import Lambda
from fastvideo.dataset.parquet_dataset_map_style import (
build_parquet_map_style_dataloader)
from fastvideo.dataset.ltx2_precomputed_dataset import (
build_ltx2_precomputed_dataloader, LTX2PrecomputedDataset)
from fastvideo.dataset.preprocessing_datasets import VideoCaptionMergedDataset, TextDataset
from fastvideo.dataset.transform import (CenterCropResizeVideo, Normalize255,
TemporalRandomCrop)
@@ -46,6 +48,10 @@ def gettextdataset(args) -> TextDataset:
__all__ = [
"build_parquet_map_style_dataloader", "ValidationDataset",
"VideoCaptionMergedDataset", "TextDataset"
"build_parquet_map_style_dataloader",
"build_ltx2_precomputed_dataloader",
"LTX2PrecomputedDataset",
"ValidationDataset",
"VideoCaptionMergedDataset",
"TextDataset",
]
@@ -0,0 +1,210 @@
# SPDX-License-Identifier: Apache-2.0
# Dataset utilities for loading LTX2 precomputed training artifacts.
#
# Usage:
# - Input root can be either `<data_root>/` or `<data_root>/.precomputed/`.
# - Required sources are `latents/` and `conditions/` with matching `.pt` files.
# - Optional source `audio_latents/` is loaded when provided in `data_sources`.
# - `build_ltx2_precomputed_dataloader(...)` is the intended entrypoint used by
# `fastvideo/training/ltx2_training_pipeline.py`.
from __future__ import annotations
from pathlib import Path
from typing import Any
import torch
from einops import rearrange
from torch.utils.data import Dataset
from torchdata.stateful_dataloader import StatefulDataLoader
from fastvideo.dataset.parquet_dataset_map_style import DP_SP_BatchSampler
from fastvideo.distributed import get_sp_world_size, get_world_rank, get_world_size
from fastvideo.logger import init_logger
logger = init_logger(__name__)
PRECOMPUTED_DIR_NAME = ".precomputed"
class LTX2PrecomputedDataset(Dataset):
"""Dataset for LTX-2 precomputed latents and conditions.
Expected directory structure (data_root):
.precomputed/
latents/*.pt
conditions/*.pt
audio_latents/*.pt (optional)
"""
def __init__(
self,
data_root: str,
data_sources: dict[str, str] | list[str] | None = None,
) -> None:
super().__init__()
self.data_root = self._setup_data_root(data_root)
self.data_sources = self._normalize_data_sources(data_sources)
self.source_paths = self._setup_source_paths()
self.sample_files = self._discover_samples()
self._validate_setup()
@staticmethod
def _setup_data_root(data_root: str) -> Path:
data_root_path = Path(data_root).expanduser().resolve()
if not data_root_path.exists():
raise FileNotFoundError(
f"Data root directory does not exist: {data_root_path}")
if (data_root_path / PRECOMPUTED_DIR_NAME).exists():
data_root_path = data_root_path / PRECOMPUTED_DIR_NAME
return data_root_path
@staticmethod
def _normalize_data_sources(
data_sources: dict[str, str] | list[str] | None,
) -> dict[str, str]:
if data_sources is None:
return {"latents": "latents", "conditions": "conditions"}
if isinstance(data_sources, list):
return {source: source for source in data_sources}
if isinstance(data_sources, dict):
return data_sources.copy()
raise TypeError(
f"data_sources must be dict, list, or None, got {type(data_sources)}")
def _setup_source_paths(self) -> dict[str, Path]:
source_paths: dict[str, Path] = {}
for dir_name in self.data_sources:
source_path = self.data_root / dir_name
if not source_path.exists():
raise FileNotFoundError(
f"Required {dir_name} directory does not exist: {source_path}")
source_paths[dir_name] = source_path
return source_paths
def _discover_samples(self) -> dict[str, list[Path]]:
data_key = ("latents"
if "latents" in self.data_sources else next(iter(
self.data_sources.keys())))
data_path = self.source_paths[data_key]
data_files = list(data_path.glob("**/*.pt"))
if not data_files:
raise ValueError(f"No data files found in {data_path}")
sample_files = {output_key: [] for output_key in self.data_sources.values()}
for data_file in data_files:
rel_path = data_file.relative_to(data_path)
if self._all_source_files_exist(data_file, rel_path):
self._fill_sample_data_files(data_file, rel_path, sample_files)
return sample_files
def _all_source_files_exist(self, data_file: Path, rel_path: Path) -> bool:
for dir_name in self.data_sources:
expected_path = self._get_expected_file_path(dir_name, data_file,
rel_path)
if not expected_path.exists():
logger.warning(
"No matching %s file found for: %s (expected in: %s)",
dir_name,
data_file.name,
expected_path,
)
return False
return True
def _get_expected_file_path(self, dir_name: str, data_file: Path,
rel_path: Path) -> Path:
source_path = self.source_paths[dir_name]
if dir_name == "conditions" and data_file.name.startswith("latent_"):
return source_path / f"condition_{data_file.stem[7:]}.pt"
return source_path / rel_path
def _fill_sample_data_files(self, data_file: Path, rel_path: Path,
sample_files: dict[str, list[Path]]) -> None:
for dir_name, output_key in self.data_sources.items():
expected_path = self._get_expected_file_path(dir_name, data_file,
rel_path)
sample_files[output_key].append(
expected_path.relative_to(self.source_paths[dir_name]))
def _validate_setup(self) -> None:
if not self.sample_files:
raise ValueError(
"No valid samples found - all data sources must have matching files"
)
sample_counts = {
key: len(files)
for key, files in self.sample_files.items()
}
if len(set(sample_counts.values())) > 1:
raise ValueError(
f"Mismatched sample counts across sources: {sample_counts}")
def __len__(self) -> int:
first_key = next(iter(self.sample_files.keys()))
return len(self.sample_files[first_key])
def __getitem__(self, index: int) -> dict[str, torch.Tensor]:
result: dict[str, Any] = {}
for dir_name, output_key in self.data_sources.items():
source_path = self.source_paths[dir_name]
file_rel_path = self.sample_files[output_key][index]
file_path = source_path / file_rel_path
try:
data = torch.load(file_path, map_location="cpu", weights_only=True)
if "latent" in dir_name.lower():
data = self._normalize_video_latents(data)
result[output_key] = data
except Exception as e:
raise RuntimeError(
f"Failed to load {output_key} from {file_path}: {e}") from e
result["idx"] = index
return result
@staticmethod
def _normalize_video_latents(data: dict) -> dict:
latents = data["latents"]
if latents.dim() == 2:
num_frames = data["num_frames"]
height = data["height"]
width = data["width"]
latents = rearrange(
latents,
"(f h w) c -> c f h w",
f=num_frames,
h=height,
w=width,
)
data = data.copy()
data["latents"] = latents
return data
def build_ltx2_precomputed_dataloader(
path: str,
batch_size: int,
num_data_workers: int,
data_sources: dict[str, str] | list[str] | None = None,
drop_last: bool = True,
seed: int = 42,
) -> tuple[LTX2PrecomputedDataset, StatefulDataLoader]:
dataset = LTX2PrecomputedDataset(path, data_sources=data_sources)
sampler = DP_SP_BatchSampler(
batch_size=batch_size,
dataset_size=len(dataset),
num_sp_groups=get_world_size() // get_sp_world_size(),
sp_world_size=get_sp_world_size(),
global_rank=get_world_rank(),
drop_last=drop_last,
drop_first_row=False,
seed=seed,
)
loader = StatefulDataLoader(
dataset,
batch_sampler=sampler,
collate_fn=None,
num_workers=num_data_workers,
pin_memory=True,
persistent_workers=num_data_workers > 0,
)
return dataset, loader
+12
View File
@@ -903,6 +903,7 @@ class TrainingArgs(FastVideoArgs):
lora_rank: int | None = None
lora_alpha: int | None = None
lora_training: bool = False
ltx2_first_frame_conditioning_p: float = 0.1
# distillation args
generator_update_interval: int = 5
@@ -916,6 +917,7 @@ class TrainingArgs(FastVideoArgs):
training_state_checkpointing_steps: int = 0 # for resuming training
weight_only_checkpointing_steps: int = 0 # for inference
log_visualization: bool = False
visualization_steps: int = 0
# simulate generator forward to match inference
simulate_generator_forward: bool = False
warp_denoising_step: bool = False
@@ -1079,6 +1081,9 @@ class TrainingArgs(FastVideoArgs):
parser.add_argument("--log-validation",
action=StoreBoolean,
help="Whether to log validation results")
parser.add_argument("--visualization-steps",
type=int,
help="Number of visualization steps")
parser.add_argument("--tracker-project-name",
type=str,
help="Project name for tracking")
@@ -1253,6 +1258,13 @@ class TrainingArgs(FastVideoArgs):
help="Whether to use LoRA training")
parser.add_argument("--lora-rank", type=int, help="LoRA rank")
parser.add_argument("--lora-alpha", type=int, help="LoRA alpha")
parser.add_argument(
"--ltx2-first-frame-conditioning-p",
type=float,
default=TrainingArgs.ltx2_first_frame_conditioning_p,
help=
"Probability of conditioning on the first frame during LTX-2 training",
)
# V-MoBA parameters
parser.add_argument(
@@ -0,0 +1,61 @@
# SPDX-License-Identifier: Apache-2.0
"""Audio preprocessing helpers for LTX-2 training."""
from __future__ import annotations
import torch
import torchaudio
from torch import nn
class AudioProcessor(nn.Module):
"""Converts audio waveforms to log-mel spectrograms with resampling."""
def __init__(
self,
sample_rate: int,
mel_bins: int,
mel_hop_length: int,
n_fft: int,
) -> None:
super().__init__()
self.sample_rate = sample_rate
self.mel_transform = torchaudio.transforms.MelSpectrogram(
sample_rate=sample_rate,
n_fft=n_fft,
win_length=n_fft,
hop_length=mel_hop_length,
f_min=0.0,
f_max=sample_rate / 2.0,
n_mels=mel_bins,
window_fn=torch.hann_window,
center=True,
pad_mode="reflect",
power=1.0,
mel_scale="slaney",
norm="slaney",
)
def resample_waveform(
self,
waveform: torch.Tensor,
source_rate: int,
target_rate: int,
) -> torch.Tensor:
if source_rate == target_rate:
return waveform
resampled = torchaudio.functional.resample(
waveform, source_rate, target_rate)
return resampled.to(device=waveform.device, dtype=waveform.dtype)
def waveform_to_mel(
self,
waveform: torch.Tensor,
waveform_sample_rate: int,
) -> torch.Tensor:
waveform = self.resample_waveform(
waveform, waveform_sample_rate, self.sample_rate)
mel = self.mel_transform(waveform)
mel = torch.log(torch.clamp(mel, min=1e-5))
mel = mel.to(device=waveform.device, dtype=waveform.dtype)
return mel.permute(0, 1, 3, 2).contiguous()
+3 -2
View File
@@ -33,7 +33,7 @@ from fastvideo.layers.visual_embedding import (PatchEmbed)
from fastvideo.logger import init_logger
from fastvideo.models.dits.base import BaseDiT
from fastvideo.models.dits.wanvideo import WanT2VCrossAttention, WanTimeTextImageEmbedding
from fastvideo.platforms import AttentionBackendEnum
from fastvideo.platforms import AttentionBackendEnum, current_platform
logger = init_logger(__name__)
class CausalWanSelfAttention(nn.Module):
@@ -286,6 +286,8 @@ class CausalWanTransformerBlock(nn.Module):
null_shift = null_scale = torch.tensor([0], device=hidden_states.device)
norm_hidden_states, hidden_states = self.self_attn_residual_norm(
hidden_states, attn_output, gate_msa, null_shift, null_scale)
norm_hidden_states, hidden_states = norm_hidden_states.to(
orig_dtype), hidden_states.to(orig_dtype)
# 2. Cross-attention
attn_output = self.attn2(norm_hidden_states,
@@ -452,7 +454,6 @@ class CausalWanTransformer3DModel(BaseDiT):
This function will be run for num_frame times.
Process the latent frames one by one (1560 tokens each)
"""
from fastvideo.platforms import current_platform
orig_dtype = hidden_states.dtype
if not isinstance(encoder_hidden_states, torch.Tensor):
+27
View File
@@ -814,6 +814,8 @@ class TransformerArgsPreprocessor:
batch_size = x.shape[0]
if context.device != x.device:
context = context.to(x.device)
if context.dtype != x.dtype:
context = context.to(x.dtype)
if attention_mask is not None and attention_mask.device != x.device:
attention_mask = attention_mask.to(x.device)
context = self.caption_projection(context)
@@ -1476,6 +1478,26 @@ class BasicAVTransformerBlock(torch.nn.Module):
self.norm_eps = norm_eps
def _register_fsdp_backward_hooks_on_output(self, vx, ax):
"""Register backward hooks on output tensors to trigger FSDP2 unshard.
FSDP2's module-level backward hooks don't fire when the module returns
dataclass outputs. We must register hooks directly on the output tensors.
"""
if not hasattr(self, 'unshard'):
return # Not wrapped by FSDP2
def make_unshard_hook():
def hook(grad):
self.unshard()
return grad
return hook
if vx is not None and vx.requires_grad:
vx.register_hook(make_unshard_hook())
if ax is not None and ax.requires_grad:
ax.register_hook(make_unshard_hook())
def get_ada_values(
self, scale_shift_table: torch.Tensor, batch_size: int, timestep: torch.Tensor, indices: slice
) -> tuple[torch.Tensor, ...]:
@@ -1704,6 +1726,10 @@ class BasicAVTransformerBlock(torch.nn.Module):
f"audio_sum={audio_sum:.6f}"
)
# Register FSDP2 backward hooks on output tensors (module-level hooks don't
# fire for dataclass outputs, so we must hook the tensors directly)
self._register_fsdp_backward_hooks_on_output(vx, ax)
return (
replace(video, x=vx) if video is not None else None,
replace(audio, x=ax) if audio is not None else None,
@@ -2075,6 +2101,7 @@ class LTX2Transformer3DModel(CachableDiT):
param_names_mapping = LTX2VideoConfig().param_names_mapping
reverse_param_names_mapping = LTX2VideoConfig().reverse_param_names_mapping
lora_param_names_mapping = LTX2VideoConfig().lora_param_names_mapping
_fsdp_shard_conditions = LTX2VideoConfig()._fsdp_shard_conditions
def __init__(self, config: LTX2VideoConfig, hf_config: dict[str, Any]):
super().__init__(config=config, hf_config=hf_config)
+56 -2
View File
@@ -3,11 +3,11 @@ from __future__ import annotations
from dataclasses import dataclass
import os
from typing import Iterable
from typing import Any, Iterable
import torch
from torch import nn
from transformers import Gemma3ForConditionalGeneration
from transformers import AutoTokenizer, Gemma3ForConditionalGeneration
from fastvideo.configs.models.encoders import BaseEncoderOutput, TextEncoderConfig
from fastvideo.models.encoders.base import TextEncoder
@@ -447,6 +447,60 @@ class LTX2GemmaTextEncoderModel(TextEncoder):
return encoded, encoded_for_audio, attention_mask.squeeze(-1)
@torch.no_grad()
def preprocess_text_embeddings(
self,
prompts: str | list[str],
tokenizer: AutoTokenizer,
tokenizer_kwargs: dict[str, Any] | None = None,
padding_side: str | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Compute pre-connector text embeddings for LTX-2 training preprocessing."""
if isinstance(prompts, str):
prompts = [prompts]
model = self.gemma_model
kwargs: dict[str, Any] = {
"padding": "max_length",
"truncation": True,
"return_tensors": "pt",
}
if tokenizer_kwargs is not None:
kwargs.update(tokenizer_kwargs)
if "max_length" not in kwargs:
kwargs["max_length"] = self.config.arch_config.text_len
original_padding_side = tokenizer.padding_side
target_padding_side = padding_side or self.padding_side
tokenizer.padding_side = target_padding_side
try:
text_inputs = tokenizer(prompts, **kwargs)
finally:
tokenizer.padding_side = original_padding_side
input_ids = text_inputs["input_ids"].to(device=model.device)
attention_mask = text_inputs["attention_mask"].to(device=model.device)
outputs = model(
input_ids=input_ids,
attention_mask=attention_mask,
output_hidden_states=True,
return_dict=True,
)
prompt_embeds = self._run_feature_extractor(
outputs.hidden_states,
attention_mask,
padding_side=target_padding_side,
)
return prompt_embeds, attention_mask
def run_connectors(
self,
encoded_input: torch.Tensor,
attention_mask: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""Apply embedding connectors to precomputed Gemma features."""
return self._run_connectors(encoded_input, attention_mask)
def forward(
self,
input_ids: torch.Tensor | None,
+1 -1
View File
@@ -277,7 +277,7 @@ def load_video(
if convert_method is not None:
pil_images = convert_method(pil_images)
return pil_images, original_fps if return_fps else pil_images
return (pil_images, original_fps) if return_fps else pil_images
def get_default_height_width(
@@ -230,6 +230,14 @@ class TrainingBatch:
noise_latents: torch.Tensor | None = None
encoder_hidden_states: torch.Tensor | None = None
encoder_attention_mask: torch.Tensor | None = None
# LTX related audio inputs
audio_latents: torch.Tensor | None = None
audio_noisy_model_input: torch.Tensor | None = None
audio_timesteps: torch.Tensor | None = None
audio_noise: torch.Tensor | None = None
audio_encoder_hidden_states: torch.Tensor | None = None
audio_encoder_attention_mask: torch.Tensor | None = None
conditioning_mask: torch.Tensor | None = None
# i2v
preprocessed_image: torch.Tensor | None = None
image_embeds: torch.Tensor | None = None
@@ -0,0 +1,290 @@
# SPDX-License-Identifier: Apache-2.0
"""LTX-2 preprocessing pipeline for native FastVideo training data generation.
This module defines the LTX-2 preprocess pipeline used by FastVideo workflows
to build precomputed training artifacts from raw text/video datasets.
Usage:
- Entry is through preprocess workflows that register `PreprocessPipelineT2V`.
- Input samples should provide prompt text plus video metadata/loader fields
consumed by `TextTransformStage` and `VideoTransformStage`.
- Output artifacts are written by the shared preprocessing workflow into
`.precomputed/` (latents, conditions, and optional audio_latents).
Optional audio path:
- When audio preprocessing is enabled, this pipeline loads the native
LTX-2 audio encoder and stores per-sample audio latents in
`batch.extra["ltx2_audio_latents"]`.
"""
from __future__ import annotations
import glob
import os
from typing import Any, cast
import torch
import torchaudio
from safetensors.torch import load_file as safetensors_load_file
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.models.audio.ltx2_audio_processing import AudioProcessor
from fastvideo.models.audio.ltx2_audio_vae import LTX2AudioEncoder
from fastvideo.models.hf_transformer_utils import get_diffusers_config
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch, PreprocessBatch
from fastvideo.pipelines.preprocess.preprocess_stages import (
TextTransformStage, VideoTransformStage)
from fastvideo.pipelines.stages import EncodingStage, PipelineStage
logger = init_logger(__name__)
class LTX2TextPrecomputeStage(PipelineStage):
"""Compute pre-connector Gemma embeddings for LTX-2 training."""
def __init__(
self,
text_encoder: torch.nn.Module,
tokenizer: Any,
preprocess_text_fn,
tokenizer_kwargs: dict[str, Any],
padding_side: str,
) -> None:
self.text_encoder = text_encoder
self.tokenizer = tokenizer
self.preprocess_text_fn = preprocess_text_fn
self.tokenizer_kwargs = tokenizer_kwargs
self.padding_side = padding_side
@torch.no_grad()
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
batch = cast(PreprocessBatch, batch)
assert isinstance(batch.prompt, list)
prompts = []
for prompt in batch.prompt:
if not isinstance(prompt, str):
prompt = str(prompt)
processed_prompt = self.preprocess_text_fn(prompt)
prompts.append(
processed_prompt if processed_prompt is not None else "")
prompt_embeds, prompt_attention_mask = (
self.text_encoder.preprocess_text_embeddings(
prompts=prompts,
tokenizer=self.tokenizer,
tokenizer_kwargs=self.tokenizer_kwargs,
padding_side=self.padding_side,
))
batch.prompt_embeds = [prompt_embeds]
batch.prompt_attention_mask = [prompt_attention_mask]
return batch
class LTX2AudioEncodingStage(PipelineStage):
"""Extract audio from input videos and encode into LTX-2 audio latents."""
def __init__(
self,
audio_encoder: torch.nn.Module,
audio_processor: AudioProcessor,
fallback_fps: int,
) -> None:
self.audio_encoder = audio_encoder.eval()
self.audio_processor = audio_processor
self.fallback_fps = fallback_fps
self.audio_dtype = next(audio_encoder.parameters()).dtype
self.audio_device = next(audio_encoder.parameters()).device
@staticmethod
def _extract_audio(
video_path: str,
target_duration: float,
) -> tuple[torch.Tensor, int] | None:
try:
waveform, sample_rate = torchaudio.load(video_path)
except Exception as e:
logger.error("Failed to load audio from %s: %s", video_path, e)
raise e
target_samples = int(target_duration * sample_rate)
if target_samples <= 0:
return None
current_samples = waveform.shape[-1]
if current_samples > target_samples:
waveform = waveform[..., :target_samples]
elif current_samples < target_samples:
padding = target_samples - current_samples
waveform = torch.nn.functional.pad(waveform, (0, padding))
return waveform, sample_rate
@torch.no_grad()
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
batch = cast(PreprocessBatch, batch)
assert isinstance(batch.video_loader, list)
assert isinstance(batch.num_frames, list)
assert isinstance(batch.fps, list)
audio_latents: list[torch.Tensor | None] = []
for idx, video_input in enumerate(batch.video_loader):
if not isinstance(video_input, str):
logger.warning(
"Skipping audio for sample %s: video loader is not a path string",
idx,
)
audio_latents.append(None)
continue
fps = float(batch.fps[idx]) if batch.fps[idx] else float(
self.fallback_fps)
if fps <= 0:
fps = float(self.fallback_fps)
target_duration = float(batch.num_frames[idx]) / fps
audio_data = self._extract_audio(video_input, target_duration)
if audio_data is None:
audio_latents.append(None)
continue
waveform, sample_rate = audio_data
waveform = waveform.unsqueeze(0).to(device=self.audio_device,
dtype=self.audio_dtype)
mel = self.audio_processor.waveform_to_mel(
waveform,
waveform_sample_rate=sample_rate).to(device=self.audio_device,
dtype=self.audio_dtype)
latents = self.audio_encoder(mel).squeeze(0).detach().cpu()
audio_latents.append(latents)
batch.extra["ltx2_audio_latents"] = audio_latents
return batch
class PreprocessPipelineT2V(ComposedPipelineBase):
"""Native LTX-2 preprocessing pipeline (text/video with optional audio)."""
_required_config_modules = ["text_encoder", "tokenizer", "vae"]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
tokenizer = self.get_module("tokenizer")
if tokenizer is not None:
tokenizer.padding_side = "left"
if tokenizer.pad_token is None and tokenizer.eos_token is not None:
tokenizer.pad_token = tokenizer.eos_token
def _load_ltx2_audio_encoder(
self) -> tuple[torch.nn.Module, AudioProcessor]:
audio_vae_path = os.path.join(self.model_path, "audio_vae")
if not os.path.isdir(audio_vae_path):
raise FileNotFoundError(
f"Expected audio_vae directory for LTX-2 audio preprocessing: {audio_vae_path}"
)
config = get_diffusers_config(model=audio_vae_path)
audio_encoder = LTX2AudioEncoder(config).to(
device=get_local_torch_device(),
dtype=torch.float32,
)
safetensors_list = glob.glob(
os.path.join(audio_vae_path, "*.safetensors"))
loaded: dict[str, torch.Tensor] = {}
for sf_file in safetensors_list:
loaded.update(safetensors_load_file(sf_file))
encoder_state = {}
for name, tensor in loaded.items():
if name.startswith("encoder."):
encoder_state[name.replace("encoder.", "")] = tensor
elif name.startswith("per_channel_statistics."):
encoder_state[name] = tensor
target_module = getattr(audio_encoder, "model", audio_encoder)
missing, unexpected = target_module.load_state_dict(encoder_state,
strict=False)
if missing:
logger.warning("Missing LTX-2 audio encoder keys: %s", missing[:8])
if unexpected:
logger.warning("Unexpected LTX-2 audio encoder keys: %s",
unexpected[:8])
target_module.eval()
audio_processor = AudioProcessor(
sample_rate=target_module.sample_rate,
mel_bins=target_module.mel_bins,
mel_hop_length=target_module.mel_hop_length,
n_fft=target_module.n_fft,
).to(next(target_module.parameters()).device)
return target_module, audio_processor
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
assert fastvideo_args.preprocess_config is not None
preprocess_cfg = fastvideo_args.preprocess_config
self.add_stage(
stage_name="text_transform_stage",
stage=TextTransformStage(
cfg_uncondition_drop_rate=preprocess_cfg.training_cfg_rate,
seed=preprocess_cfg.seed,
),
)
text_encoder = self.get_module("text_encoder")
tokenizer = self.get_module("tokenizer")
encoder_config = fastvideo_args.pipeline_config.text_encoder_configs[0]
tokenizer_kwargs = dict(encoder_config.tokenizer_kwargs)
if "max_length" not in tokenizer_kwargs:
tokenizer_kwargs["max_length"] = encoder_config.arch_config.text_len
self.add_stage(
stage_name="prompt_precompute_stage",
stage=LTX2TextPrecomputeStage(
text_encoder=text_encoder,
tokenizer=tokenizer,
preprocess_text_fn=fastvideo_args.pipeline_config.
preprocess_text_funcs[0],
tokenizer_kwargs=tokenizer_kwargs,
padding_side=encoder_config.arch_config.padding_side,
),
)
self.add_stage(
stage_name="video_transform_stage",
stage=VideoTransformStage(
train_fps=preprocess_cfg.train_fps,
num_frames=preprocess_cfg.num_frames,
max_height=preprocess_cfg.max_height,
max_width=preprocess_cfg.max_width,
do_temporal_sample=preprocess_cfg.do_temporal_sample,
),
)
if preprocess_cfg.with_audio:
audio_encoder, audio_processor = self._load_ltx2_audio_encoder()
self.add_stage(
stage_name="audio_encoding_stage",
stage=LTX2AudioEncodingStage(
audio_encoder=audio_encoder,
audio_processor=audio_processor,
fallback_fps=preprocess_cfg.train_fps,
),
)
self.add_stage(
stage_name="video_encoding_stage",
stage=EncodingStage(vae=self.get_module("vae")),
)
EntryClass = PreprocessPipelineT2V
+1
View File
@@ -314,6 +314,7 @@ class DenoisingStage(PipelineStage):
t_expand = timestep.repeat(latent_model_input.shape[0], 1)
else:
t_expand = t.repeat(latent_model_input.shape[0])
t_expand = t_expand.to(get_local_torch_device())
use_meanflow = getattr(self.transformer.config, "use_meanflow",
False)
+4 -6
View File
@@ -23,7 +23,6 @@ from fastvideo.models.dits.ltx2 import (
DEFAULT_LTX2_AUDIO_DOWNSAMPLE, DEFAULT_LTX2_AUDIO_HOP_LENGTH,
DEFAULT_LTX2_AUDIO_MEL_BINS, DEFAULT_LTX2_AUDIO_SAMPLE_RATE,
VideoLatentShape)
from fastvideo.utils import PRECISION_TO_TYPE
BASE_SHIFT_ANCHOR = 1024
MAX_SHIFT_ANCHOR = 4096
@@ -46,6 +45,7 @@ def _ltx2_sigmas(
stretch: bool = True,
terminal: float = 0.1,
) -> torch.Tensor:
# Copied/following official LTX-2 scheduler (LTX2Scheduler.execute).
tokens = math.prod(
latent.shape[2:]) if latent is not None else MAX_SHIFT_ANCHOR
sigmas = torch.linspace(1.0,
@@ -114,12 +114,9 @@ class LTX2DenoisingStage(PipelineStage):
if neg_prompt_mask is not None and neg_prompt_mask.device != latents.device:
neg_prompt_mask = neg_prompt_mask.to(latents.device)
target_dtype = PRECISION_TO_TYPE[
fastvideo_args.pipeline_config.dit_precision]
disable_autocast = os.getenv("LTX2_DISABLE_AUTOCAST", "1") == "1"
target_dtype = torch.bfloat16
autocast_enabled = (target_dtype != torch.float32
) and not fastvideo_args.disable_autocast and (
not disable_autocast)
) and not fastvideo_args.disable_autocast
# Use official distilled sigma schedule for 8 steps (distilled models)
use_distilled_sigmas = os.getenv("LTX2_USE_DISTILLED_SIGMAS",
@@ -137,6 +134,7 @@ class LTX2DenoisingStage(PipelineStage):
latent=None,
device=latents.device,
)
logger.info("[LTX2] Using computed sigma schedule")
if hasattr(self.transformer, "patchifier"):
video_shape = VideoLatentShape.from_torch_shape(latents.shape)
token_count = self.transformer.patchifier.get_token_count(
+18 -5
View File
@@ -57,7 +57,8 @@ from fastvideo.configs.sample.hunyuan15 import (
Hunyuan15_720P_SamplingParam, Hunyuan15_720P_Distilled_I2V_SamplingParam,
Hunyuan15_SR_1080P_SamplingParam)
from fastvideo.configs.sample.hyworld import HYWorld_SamplingParam
from fastvideo.configs.sample.ltx2 import LTX2SamplingParam
from fastvideo.configs.sample.ltx2 import (LTX2BaseSamplingParam,
LTX2DistilledSamplingParam)
from fastvideo.configs.sample.stepvideo import StepVideoT2VSamplingParam
from fastvideo.configs.sample.turbodiffusion import (
TurboDiffusionI2V_A14B_SamplingParam,
@@ -239,17 +240,29 @@ def _get_config_info(
def _register_configs() -> None:
# LTX-2
# LTX-2 (base)
register_configs(
sampling_param_cls=LTX2SamplingParam,
sampling_param_cls=LTX2BaseSamplingParam,
pipeline_config_cls=LTX2T2VConfig,
hf_model_paths=[
"Lightricks/LTX-2",
"converted/ltx2_diffusers",
"FastVideo/LTX2-base",
],
model_detectors=[
lambda path: ("ltx2" in path.lower() or "ltx-2" in path.lower()) and
"distilled" not in path.lower(),
],
)
# LTX-2 (distilled)
register_configs(
sampling_param_cls=LTX2DistilledSamplingParam,
pipeline_config_cls=LTX2T2VConfig,
hf_model_paths=[
"FastVideo/LTX2-Distilled-Diffusers",
],
model_detectors=[
lambda path: "ltx2" in path.lower() or "ltx-2" in path.lower(),
lambda path: ("ltx2" in path.lower() or "ltx-2" in path.lower()) and
"distilled" in path.lower(),
],
)
+7 -1
View File
@@ -1,5 +1,11 @@
from .distillation_pipeline import DistillationPipeline
from .training_pipeline import TrainingPipeline
from .wan_training_pipeline import WanTrainingPipeline
from .ltx2_training_pipeline import LTX2TrainingPipeline
__all__ = ["TrainingPipeline", "WanTrainingPipeline", "DistillationPipeline"]
__all__ = [
"TrainingPipeline",
"WanTrainingPipeline",
"LTX2TrainingPipeline",
"DistillationPipeline",
]
+4 -1
View File
@@ -43,7 +43,10 @@ from fastvideo.training.training_utils import (
from fastvideo.utils import (is_vsa_available, maybe_download_model,
set_random_seed, verify_model_config_and_directory)
vsa_available = is_vsa_available()
try:
vsa_available = is_vsa_available()
except Exception:
vsa_available = False
logger = init_logger(__name__)
@@ -0,0 +1,477 @@
# SPDX-License-Identifier: Apache-2.0
import os
import sys
from copy import deepcopy
from pathlib import Path
import torch
import torch.distributed as dist
from fastvideo.dataset import build_ltx2_precomputed_dataloader
from fastvideo.distributed import get_local_torch_device, get_sp_group, get_world_group
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
from fastvideo.models.dits.ltx2 import VideoLatentShape
from fastvideo.pipelines.basic.ltx2.ltx2_pipeline import LTX2Pipeline
from fastvideo.pipelines.pipeline_batch_info import TrainingBatch
from fastvideo.training.training_pipeline import TrainingPipeline
from fastvideo.training.trackers import (DummyTracker, TrackerType, Trackers,
initialize_trackers)
from fastvideo.training.training_utils import (
clip_grad_norm_while_handling_failing_dtensor_cases, get_scheduler)
from fastvideo.utils import set_random_seed
logger = init_logger(__name__)
class LTX2TrainingPipeline(TrainingPipeline):
"""Training pipeline for LTX-2 text-to-video (optional audio)."""
_required_config_modules = [
"transformer",
"text_encoder",
"tokenizer",
"vae",
"audio_vae",
"vocoder",
]
text_encoder: torch.nn.Module
with_audio: bool = False
tracker: TrackerType
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
# TODO (David): Change to port LTX2 scheduler into self.modules["scheduler"]
if "scheduler" in self.modules:
del self.modules["scheduler"]
def initialize_training_pipeline(self, training_args: TrainingArgs):
logger.info("Initializing LTX-2 training pipeline...")
self.device = get_local_torch_device()
self.training_args = training_args
world_group = get_world_group()
self.world_size = world_group.world_size
self.global_rank = world_group.rank
self.sp_group = get_sp_group()
self.rank_in_sp_group = self.sp_group.rank_in_group
self.sp_world_size = self.sp_group.world_size
self.local_rank = world_group.local_rank
self.transformer = self.get_module("transformer")
self.transformer_2 = self.get_module("transformer_2", None)
self.text_encoder = self.get_module("text_encoder")
self.text_encoder.eval()
self.text_encoder.to(self.device)
self.seed = training_args.seed
assert self.seed is not None, "seed must be set"
set_random_seed(self.seed)
self.transformer.train()
if training_args.enable_gradient_checkpointing_type is not None:
from fastvideo.training.activation_checkpoint import (
apply_activation_checkpointing)
self.transformer = apply_activation_checkpointing(
self.transformer,
checkpointing_type=training_args.
enable_gradient_checkpointing_type)
self.set_trainable()
params_to_optimize = list(
filter(lambda p: p.requires_grad, self.transformer.parameters()))
betas_str = training_args.betas
betas = tuple(float(x.strip()) for x in betas_str.split(","))
self.optimizer = torch.optim.AdamW(
params_to_optimize,
lr=training_args.learning_rate,
betas=betas,
weight_decay=training_args.weight_decay,
eps=1e-8,
)
self.init_steps = 0
logger.info("optimizer: %s", self.optimizer)
self.lr_scheduler = get_scheduler(
training_args.lr_scheduler,
optimizer=self.optimizer,
num_warmup_steps=training_args.lr_warmup_steps,
num_training_steps=training_args.max_train_steps,
num_cycles=training_args.lr_num_cycles,
power=training_args.lr_power,
min_lr_ratio=training_args.min_lr_ratio,
last_epoch=self.init_steps - 1,
)
data_sources = self._get_ltx2_data_sources(training_args.data_path)
self.with_audio = "audio_latents" in data_sources
self.train_dataset, self.train_dataloader = (
build_ltx2_precomputed_dataloader(
training_args.data_path,
training_args.train_batch_size,
num_data_workers=training_args.dataloader_num_workers,
data_sources=data_sources,
drop_last=True,
seed=self.seed,
))
self.num_update_steps_per_epoch = max(
1,
len(self.train_dataloader) //
training_args.gradient_accumulation_steps)
self.num_train_epochs = max(
1, training_args.max_train_steps // self.num_update_steps_per_epoch)
self.current_epoch = 0
trackers = list(training_args.trackers)
if not trackers and training_args.tracker_project_name:
trackers.append(Trackers.WANDB.value)
if self.global_rank != 0:
trackers = []
tracker_log_dir = training_args.output_dir or os.getcwd()
if trackers:
tracker_log_dir = os.path.join(tracker_log_dir, "tracker")
tracker_config = training_args.__dict__ if trackers else None
tracker_run_name = training_args.wandb_run_name or None
project = training_args.tracker_project_name or "fastvideo"
self.tracker = initialize_trackers(
trackers,
experiment_name=project,
config=tracker_config,
log_dir=tracker_log_dir,
run_name=tracker_run_name,
) if trackers else DummyTracker()
def initialize_validation_pipeline(self, training_args: TrainingArgs):
logger.info("Initializing LTX-2 validation pipeline...")
args_copy = deepcopy(training_args)
args_copy.inference_mode = True
validation_pipeline = LTX2Pipeline.from_pretrained(
training_args.model_path,
args=args_copy,
inference_mode=True,
loaded_modules={
"transformer": self.get_module("transformer"),
"text_encoder": self.get_module("text_encoder"),
"tokenizer": self.get_module("tokenizer"),
"vae": self.get_module("vae"),
"audio_vae": self.get_module("audio_vae"),
"vocoder": self.get_module("vocoder"),
},
tp_size=training_args.tp_size,
sp_size=training_args.sp_size,
num_gpus=training_args.num_gpus,
pin_cpu_memory=training_args.pin_cpu_memory,
dit_cpu_offload=training_args.dit_cpu_offload,
)
self.validation_pipeline = validation_pipeline
def _get_ltx2_data_sources(self, data_path: str) -> dict[str, str]:
data_root = Path(data_path).expanduser().resolve()
if (data_root / ".precomputed").exists():
data_root = data_root / ".precomputed"
sources = {"latents": "latents", "conditions": "conditions"}
audio_dir = data_root / "audio_latents"
if audio_dir.exists() and any(audio_dir.rglob("*.pt")):
sources["audio_latents"] = "audio_latents"
return sources
def _get_next_batch(self, training_batch: TrainingBatch) -> TrainingBatch:
with self.tracker.timed("timing/get_next_batch"):
batch = next(self.train_loader_iter, None) # type: ignore
if batch is None:
self.current_epoch += 1
logger.info("Starting epoch %s", self.current_epoch)
self.train_loader_iter = iter(self.train_dataloader)
batch = next(self.train_loader_iter)
latents = batch["latents"]["latents"].to(get_local_torch_device(),
dtype=torch.bfloat16)
latents = latents[:, :, :self.training_args.num_latent_t]
conditions = batch["conditions"]
if ("video_prompt_embeds" in conditions
and "audio_prompt_embeds" in conditions
and "prompt_attention_mask" in conditions):
video_embeds = conditions["video_prompt_embeds"].to(
get_local_torch_device())
audio_embeds = conditions["audio_prompt_embeds"].to(
get_local_torch_device())
attention_mask = conditions["prompt_attention_mask"].to(
get_local_torch_device(), dtype=torch.int64)
else:
prompt_embeds = conditions["prompt_embeds"].to(
get_local_torch_device())
prompt_attention_mask = conditions["prompt_attention_mask"].to(
get_local_torch_device(), dtype=torch.int64)
video_embeds, audio_embeds, attention_mask = (
self.text_encoder.run_connectors(prompt_embeds,
prompt_attention_mask))
training_batch.latents = latents.to(get_local_torch_device(),
dtype=torch.bfloat16)
training_batch.encoder_hidden_states = video_embeds.to(
get_local_torch_device(), dtype=torch.bfloat16)
training_batch.encoder_attention_mask = attention_mask.to(
get_local_torch_device(), dtype=torch.bfloat16)
if self.with_audio and "audio_latents" in batch:
audio_latents = batch["audio_latents"]["latents"].to(
get_local_torch_device(), dtype=torch.bfloat16)
training_batch.audio_latents = audio_latents
training_batch.audio_encoder_hidden_states = audio_embeds.to(
get_local_torch_device(), dtype=torch.bfloat16)
training_batch.audio_encoder_attention_mask = attention_mask.to(
get_local_torch_device())
idxs = batch.get("idx")
if idxs is not None and torch.is_tensor(idxs):
training_batch.infos = [{"idx": int(i)} for i in idxs.tolist()]
else:
training_batch.infos = []
training_batch.raw_latent_shape = latents.shape
return training_batch
def _normalize_dit_input(self,
training_batch: TrainingBatch) -> TrainingBatch:
return training_batch
def _sample_sigmas(self, batch_size: int, seq_length: int,
device: torch.device,
dtype: torch.dtype) -> torch.Tensor:
min_tokens = 1024
max_tokens = 4096
min_shift = 0.95
max_shift = 2.05
slope = (max_shift - min_shift) / (max_tokens - min_tokens)
shift = slope * seq_length + (min_shift - slope * min_tokens)
normal_samples = torch.randn(
(batch_size, ),
generator=self.noise_random_generator,
device="cpu",
) + shift
return torch.sigmoid(normal_samples).to(device=device, dtype=dtype)
def _prepare_dit_inputs(self,
training_batch: TrainingBatch) -> TrainingBatch:
assert training_batch.latents is not None
latents = training_batch.latents
batch_size = latents.shape[0]
video_shape = VideoLatentShape.from_torch_shape(latents.shape)
if hasattr(self.transformer, "patchifier"):
token_count = self.transformer.patchifier.get_token_count(
video_shape)
else:
token_count = latents.shape[2] * latents.shape[3] * latents.shape[4]
sigmas = self._sample_sigmas(batch_size, token_count, latents.device,
latents.dtype)
noise = torch.randn(latents.shape,
generator=self.noise_gen_cuda,
device=latents.device,
dtype=latents.dtype)
sigmas_expanded = sigmas.view(-1, 1, 1, 1, 1)
noisy_model_input = (
1.0 - sigmas_expanded) * latents + sigmas_expanded * noise
conditioning_mask = None
first_frame_p = self.training_args.ltx2_first_frame_conditioning_p
if (first_frame_p > 0 and torch.rand(
1,
generator=self.noise_random_generator,
).item() < first_frame_p):
conditioning_mask = torch.zeros(
(batch_size, 1, latents.shape[2], latents.shape[3],
latents.shape[4]),
dtype=torch.bool,
device=latents.device,
)
conditioning_mask[:, :, 0:1] = True
noisy_model_input = torch.where(conditioning_mask, latents,
noisy_model_input)
if conditioning_mask is None:
mask_patch = torch.zeros(
(batch_size, token_count),
dtype=torch.bool,
device=latents.device,
)
elif hasattr(self.transformer, "patchifier"):
mask_patch = self.transformer.patchifier.patchify(
conditioning_mask.float()).sum(dim=-1) > 0
else:
mask_patch = conditioning_mask[:, 0].reshape(batch_size, -1)
timesteps = torch.where(
mask_patch,
torch.zeros_like(mask_patch, dtype=latents.dtype),
sigmas.view(-1, 1).expand_as(mask_patch).to(latents.dtype),
)
training_batch.noisy_model_input = noisy_model_input
training_batch.timesteps = timesteps
training_batch.sigmas = sigmas
training_batch.noise = noise
training_batch.conditioning_mask = conditioning_mask
if hasattr(
training_batch,
"audio_latents") and training_batch.audio_latents is not None:
audio_latents = training_batch.audio_latents
audio_noise = torch.randn(audio_latents.shape,
generator=self.noise_gen_cuda,
device=audio_latents.device,
dtype=audio_latents.dtype)
audio_sigmas_expanded = sigmas.view(-1, 1, 1, 1)
audio_noisy = (1.0 - audio_sigmas_expanded) * audio_latents + (
audio_sigmas_expanded * audio_noise)
audio_timesteps = sigmas.view(-1, 1).expand(batch_size,
audio_latents.shape[2])
training_batch.audio_noisy_model_input = audio_noisy
training_batch.audio_timesteps = audio_timesteps.to(
dtype=audio_latents.dtype)
training_batch.audio_noise = audio_noise
return training_batch
def _build_attention_metadata(
self, training_batch: TrainingBatch) -> TrainingBatch:
training_batch.attn_metadata = None
return training_batch
def _build_input_kwargs(self,
training_batch: TrainingBatch) -> TrainingBatch:
assert training_batch.noisy_model_input is not None
assert training_batch.encoder_hidden_states is not None
assert training_batch.timesteps is not None
training_batch.input_kwargs = {
"hidden_states":
training_batch.noisy_model_input,
"encoder_hidden_states":
training_batch.encoder_hidden_states,
"timestep":
training_batch.timesteps.to(get_local_torch_device(),
dtype=torch.bfloat16),
"encoder_attention_mask":
training_batch.encoder_attention_mask,
"return_dict":
False,
}
if training_batch.audio_noisy_model_input is not None:
training_batch.input_kwargs.update({
"audio_hidden_states":
training_batch.audio_noisy_model_input,
"audio_encoder_hidden_states":
training_batch.audio_encoder_hidden_states,
"audio_timestep":
training_batch.audio_timesteps.to(get_local_torch_device(),
dtype=torch.bfloat16),
"audio_encoder_attention_mask":
training_batch.audio_encoder_attention_mask,
})
return training_batch
def _masked_mse(self, pred: torch.Tensor, target: torch.Tensor,
conditioning_mask: torch.Tensor | None) -> torch.Tensor:
loss = (pred.float() - target.float())**2
if conditioning_mask is None:
return loss.mean()
loss_mask = (~conditioning_mask).float()
loss_mask = loss_mask.expand(-1, pred.shape[1], -1, -1, -1)
denom = loss_mask.mean()
if denom.item() == 0.0:
return loss.sum() * 0.0
loss = loss * loss_mask
return loss.mean() / denom
def _transformer_forward_and_compute_loss(
self, training_batch: TrainingBatch) -> TrainingBatch:
input_kwargs = training_batch.input_kwargs
assert input_kwargs is not None
assert training_batch.sigmas is not None
assert training_batch.latents is not None
assert training_batch.noisy_model_input is not None
assert training_batch.noise is not None
with self.tracker.timed("timing/forward_backward"), set_forward_context(
current_timestep=training_batch.current_timestep,
attn_metadata=training_batch.attn_metadata):
with torch.autocast("cuda", dtype=training_batch.latents.dtype
), torch.autograd.set_detect_anomaly(True):
outputs = self.transformer(**input_kwargs)
if isinstance(outputs, tuple):
video_denoised, audio_denoised = outputs
else:
video_denoised = outputs
audio_denoised = None
sigmas_expanded = training_batch.sigmas.view(-1, 1, 1, 1, 1)
video_pred_velocity = (training_batch.noisy_model_input -
video_denoised) / sigmas_expanded
video_target = training_batch.noise - training_batch.latents
loss = self._masked_mse(video_pred_velocity, video_target,
training_batch.conditioning_mask)
if audio_denoised is not None and training_batch.audio_latents is not None:
audio_sigmas = training_batch.sigmas.view(-1, 1, 1, 1)
audio_pred_velocity = (training_batch.audio_noisy_model_input -
audio_denoised) / audio_sigmas
audio_target = training_batch.audio_noise - training_batch.audio_latents
audio_loss = (audio_pred_velocity.float() -
audio_target.float())**2
loss = loss + audio_loss.mean()
logger.info("Audio loss: %s", audio_loss.mean())
else:
logger.warning("Audio denoised is None")
loss = loss / self.training_args.gradient_accumulation_steps
with torch.autograd.set_detect_anomaly(True):
loss.backward()
avg_loss = loss.detach().clone()
logger.info("finished backward")
with self.tracker.timed("timing/reduce_loss"):
world_group = get_world_group()
avg_loss = world_group.all_reduce(avg_loss, op=dist.ReduceOp.AVG)
training_batch.total_loss += avg_loss.item()
return training_batch
def _clip_grad_norm(self, training_batch: TrainingBatch) -> TrainingBatch:
max_grad_norm = self.training_args.max_grad_norm
if max_grad_norm is not None:
with self.tracker.timed("timing/clip_grad_norm"):
grad_norm = clip_grad_norm_while_handling_failing_dtensor_cases(
[p for p in self.transformer.parameters()],
max_grad_norm,
foreach=None,
)
grad_norm = grad_norm.item() if grad_norm is not None else 0.0
else:
grad_norm = 0.0
training_batch.grad_norm = grad_norm
return training_batch
def main(args) -> None:
logger.info("Starting LTX-2 training pipeline...")
pipeline = LTX2TrainingPipeline.from_pretrained(
args.pretrained_model_name_or_path, args=args)
args = pipeline.training_args
pipeline.train()
logger.info("Training pipeline done")
if __name__ == "__main__":
argv = sys.argv
from fastvideo.fastvideo_args import TrainingArgs
from fastvideo.utils import FlexibleArgumentParser
parser = FlexibleArgumentParser()
parser = TrainingArgs.add_cli_args(parser)
parser = FastVideoArgs.add_cli_args(parser)
args = parser.parse_args()
args.dit_cpu_offload = False
main(args)
@@ -18,7 +18,10 @@ from fastvideo.pipelines.pipeline_batch_info import ForwardBatch, TrainingBatch
from fastvideo.training.training_pipeline import TrainingPipeline
from fastvideo.utils import is_vsa_available, shallow_asdict
vsa_available = is_vsa_available()
try:
vsa_available = is_vsa_available()
except Exception:
vsa_available = False
logger = init_logger(__name__)
+13 -26
View File
@@ -17,6 +17,7 @@ from fastvideo.models.schedulers.scheduling_self_forcing_flow_match import (
SelfForcingFlowMatchScheduler)
from fastvideo.pipelines.basic.wan.wan_causal_dmd_pipeline import (
WanCausalDMDPipeline)
from fastvideo.pipelines.stages.decoding import DecodingStage
from fastvideo.pipelines.pipeline_batch_info import TrainingBatch
from fastvideo.training.training_pipeline import TrainingPipeline
from fastvideo.training.training_utils import (
@@ -57,15 +58,17 @@ class ODEInitTrainingPipeline(TrainingPipeline):
self.vae.requires_grad_(False)
self.timestep_shift = self.training_args.pipeline_config.flow_shift
assert self.timestep_shift == 5.0, "flow_shift must be 5.0"
self.noise_scheduler = SelfForcingFlowMatchScheduler(
shift=self.timestep_shift, sigma_min=0.0, extra_one_step=True)
self.noise_scheduler.set_timesteps(num_inference_steps=1000,
training=True)
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
logger.info("dmd_denoising_steps: %s",
self.training_args.pipeline_config.dmd_denoising_steps)
self.dmd_denoising_steps = torch.tensor([1000, 750, 500, 250],
self.dmd_denoising_steps = torch.tensor([1000, 750, 500, 250, 0],
dtype=torch.long,
device=get_local_torch_device())
if training_args.warp_denoising_step: # Warp the denoising step according to the scheduler time shift
@@ -161,27 +164,12 @@ class ODEInitTrainingPipeline(TrainingPipeline):
device, dtype=torch.bfloat16)
training_batch.infos = infos
return training_batch, trajectory_latents.to(
device, dtype=torch.bfloat16), trajectory_timesteps.to(device)
## TEMP Used for loading the sf .pt files directly
"""
self.manual_idx = self.manual_idx % 155
path = f"/mnt/weka/home/hao.zhang/wl/Self-Forcing/ode_pt_vidprom_1000/{self.manual_idx:05d}.pt"
logger.info("path: %s", path)
self.manual_idx += 1
# path = "/mnt/weka/home/hao.zhang/wl/Self-Forcing/ode_single_full/00000.pt"
b = torch.load(path)
training_batch.encoder_hidden_states = b["text_embedding"][0].unsqueeze(
0).to(device, dtype=torch.bfloat16)
trajectory_latents = b["ode_latent"].to(device, dtype=torch.bfloat16)
logger.info("trajectory_latents: %s", trajectory_latents.shape)
logger.info("encoder_hidden_states: %s",
training_batch.encoder_hidden_states.shape)
assert trajectory_latents.shape[1] <= 10, "trajectory_latents.shape[1] must be <= 10"
return training_batch, trajectory_latents.to(
device, dtype=torch.bfloat16), trajectory_timesteps.to(device)
"""
return training_batch, trajectory_latents[:, :, :self.training_args.
num_latent_t].to(
device,
dtype=torch.bfloat16
), trajectory_timesteps.to(
device)
def _get_timestep(self,
min_timestep: int,
@@ -225,7 +213,7 @@ class ODEInitTrainingPipeline(TrainingPipeline):
# Lazily cache nearest trajectory index per DMD step based on the (fixed) S timesteps
if self._cached_closest_idx_per_dmd is None:
self._cached_closest_idx_per_dmd = torch.tensor(
[0, 12, 24, 36], dtype=torch.long).cpu()
[0, 12, 24, 36, S - 1], dtype=torch.long).cpu()
# [0, 1, 2, 3], dtype=torch.long).cpu()
logger.info("self._cached_closest_idx_per_dmd: %s",
self._cached_closest_idx_per_dmd)
@@ -367,8 +355,7 @@ class ODEInitTrainingPipeline(TrainingPipeline):
assert latent_key in latents_vis_dict and latents_vis_dict[
latent_key] is not None
latent = latents_vis_dict[latent_key]
pixel_latent = self.validation_pipeline.decoding_stage.decode(
latent, training_args)
pixel_latent = self.decoding_stage.decode(latent, training_args)
video = pixel_latent.cpu().float()
video = video.permute(0, 2, 1, 3, 4)
@@ -30,7 +30,10 @@ from fastvideo.profiler import profile_region
logger = init_logger(__name__)
vsa_available = is_vsa_available()
try:
vsa_available = is_vsa_available()
except Exception:
vsa_available = False
class SelfForcingDistillationPipeline(DistillationPipeline):
+193 -74
View File
@@ -2,6 +2,8 @@
from dataclasses import asdict
import math
import os
import shutil
import tempfile
import time
from abc import ABC, abstractmethod
from collections import deque
@@ -13,16 +15,18 @@ import numpy as np
import torch
import torch.distributed as dist
import torchvision
from diffusers import FlowMatchEulerDiscreteScheduler
from einops import rearrange
from torch.utils.data import DataLoader
from torchdata.stateful_dataloader import StatefulDataLoader
from tqdm.auto import tqdm
import fastvideo.envs as envs
from fastvideo.attention.backends.video_sparse_attn import (
VideoSparseAttentionMetadataBuilder)
from fastvideo.attention.backends.vmoba import VideoMobaAttentionMetadataBuilder
try:
from fastvideo.attention.backends.video_sparse_attn import (
VideoSparseAttentionMetadataBuilder)
from fastvideo.attention.backends.vmoba import VideoMobaAttentionMetadataBuilder
except Exception:
pass
from fastvideo.configs.sample import SamplingParam
from fastvideo.dataset import build_parquet_map_style_dataloader
from fastvideo.dataset.dataloader.schema import pyarrow_schema_t2v
@@ -48,8 +52,12 @@ from fastvideo.training.training_utils import (
from fastvideo.utils import (is_vmoba_available, is_vsa_available,
set_random_seed, shallow_asdict)
vsa_available = is_vsa_available()
vmoba_available = is_vmoba_available()
try:
vsa_available = is_vsa_available()
vmoba_available = is_vmoba_available()
except Exception:
vsa_available = False
vmoba_available = False
logger = init_logger(__name__)
@@ -108,7 +116,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
# Set random seeds for deterministic training
assert self.seed is not None, "seed must be set"
set_random_seed(self.seed)
set_random_seed(self.seed + self.global_rank)
self.transformer.train()
if training_args.enable_gradient_checkpointing_type is not None:
self.transformer = apply_activation_checkpointing(
@@ -588,15 +596,15 @@ class TrainingPipeline(LoRAPipeline, ABC):
round(num_trainable_params / 1e9, 3))
# Set random seeds for deterministic training
self.noise_random_generator = torch.Generator(device="cpu").manual_seed(
self.seed)
self.noise_random_generator = torch.Generator(
device="cpu").manual_seed(self.seed + self.global_rank)
self.noise_gen_cuda = torch.Generator(
device=current_platform.device_name).manual_seed(self.seed)
device=current_platform.device_name).manual_seed(self.seed +
self.global_rank)
self.validation_random_generator = torch.Generator(
device="cpu").manual_seed(self.seed)
logger.info("Initialized random seeds with seed: %s", self.seed)
self.noise_scheduler = FlowMatchEulerDiscreteScheduler()
device="cpu").manual_seed(self.seed + self.global_rank)
logger.info("Initialized random seeds with seed: %s",
self.seed + self.global_rank)
if self.training_args.resume_from_checkpoint:
self._resume_from_checkpoint()
@@ -661,26 +669,31 @@ class TrainingPipeline(LoRAPipeline, ABC):
"grad_norm": grad_norm,
"vsa_sparsity": current_vsa_sparsity,
}
metrics["batch_size"] = int(training_batch.raw_latent_shape[0])
try:
metrics["batch_size"] = int(
training_batch.raw_latent_shape[0])
patch_t, patch_h, patch_w = self.training_args.pipeline_config.dit_config.patch_size
seq_len = (training_batch.raw_latent_shape[2] // patch_t) * (
training_batch.raw_latent_shape[3] //
patch_h) * (training_batch.raw_latent_shape[4] // patch_w)
if training_batch.encoder_hidden_states is not None:
context_len = int(
training_batch.encoder_hidden_states.shape[1])
else:
context_len = 0
patch_t, patch_h, patch_w = self.training_args.pipeline_config.dit_config.patch_size
seq_len = (
training_batch.raw_latent_shape[2] // patch_t) * (
training_batch.raw_latent_shape[3] // patch_h) * (
training_batch.raw_latent_shape[4] // patch_w)
if training_batch.encoder_hidden_states is not None:
context_len = int(
training_batch.encoder_hidden_states.shape[1])
else:
context_len = 0
metrics["dit_seq_len"] = int(seq_len)
metrics["context_len"] = context_len
metrics["dit_seq_len"] = int(seq_len)
metrics["context_len"] = context_len
arch_config = self.training_args.pipeline_config.dit_config.arch_config
arch_config = self.training_args.pipeline_config.dit_config.arch_config
metrics["hidden_dim"] = arch_config.hidden_size
metrics["num_layers"] = arch_config.num_layers
metrics["ffn_dim"] = arch_config.ffn_dim
metrics["hidden_dim"] = arch_config.hidden_size
metrics["num_layers"] = arch_config.num_layers
metrics["ffn_dim"] = arch_config.ffn_dim
except Exception:
pass
self.tracker.log(metrics, step)
if step % self.training_args.training_state_checkpointing_steps == 0:
@@ -693,12 +706,14 @@ class TrainingPipeline(LoRAPipeline, ABC):
self.noise_random_generator)
self.transformer.train()
self.sp_group.barrier()
if self.training_args.log_visualization and step % self.training_args.visualization_steps == 0:
self.visualize_intermediate_latents(training_batch,
self.training_args, step)
if self.training_args.log_validation and step % self.training_args.validation_steps == 0:
with self.profiler_controller.region(
"profiler_region_training_validation"):
if self.training_args.log_visualization:
self.visualize_intermediate_latents(
training_batch, self.training_args, step)
self._log_validation(self.transformer, self.training_args,
step)
gpu_memory_usage = current_platform.get_torch_device(
@@ -818,7 +833,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
validation_dataloader = DataLoader(validation_dataset,
batch_size=None,
num_workers=0)
# return
self.transformer.eval()
if getattr(self, "transformer_2", None) is not None:
self.transformer_2.eval()
@@ -873,49 +888,61 @@ class TrainingPipeline(LoRAPipeline, ABC):
# Only sp_group leaders (rank_in_sp_group == 0) need to send their
# results to global rank 0
if self.rank_in_sp_group == 0:
if self.global_rank == 0:
# Global rank 0 collects results from all sp_group leaders
all_videos = step_videos # Start with own results
all_captions = step_captions
if self.rank_in_sp_group == 0 and self.global_rank == 0:
# Global rank 0 collects results from all sp_group leaders
all_videos = step_videos # Start with own results
all_captions = step_captions
# Receive from other sp_group leaders
for sp_group_idx in range(1, num_sp_groups):
src_rank = sp_group_idx * self.sp_world_size # Global rank of other sp_group leaders
recv_videos = world_group.recv_object(src=src_rank)
recv_captions = world_group.recv_object(src=src_rank)
all_videos.extend(recv_videos)
all_captions.extend(recv_captions)
# Receive from other sp_group leaders
for sp_group_idx in range(1, num_sp_groups):
src_rank = sp_group_idx * self.sp_world_size # Global rank of other sp_group leaders
recv_videos = world_group.recv_object(src=src_rank)
recv_captions = world_group.recv_object(src=src_rank)
all_videos.extend(recv_videos)
all_captions.extend(recv_captions)
video_filenames = []
for i, (video, caption) in enumerate(
zip(all_videos, all_captions, strict=True)):
os.makedirs(training_args.output_dir, exist_ok=True)
filename = os.path.join(
training_args.output_dir,
f"validation_step_{global_step}_inference_steps_{num_inference_steps}_video_{i}.mp4"
)
imageio.mimsave(filename, video, fps=sampling_param.fps)
video_filenames.append(filename)
video_filenames = []
for i, (video, caption) in enumerate(
zip(all_videos, all_captions, strict=True)):
os.makedirs(training_args.output_dir, exist_ok=True)
filename = os.path.join(
training_args.output_dir,
f"validation_step_{global_step}_inference_steps_{num_inference_steps}_video_{i}.mp4"
)
imageio.mimsave(filename, video, fps=sampling_param.fps)
# Mux audio if available
audio = output_batch.extra.get("audio")
audio_sample_rate = output_batch.extra.get(
"audio_sample_rate")
if (audio is not None and audio_sample_rate is not None
and not self._mux_audio(
filename,
audio,
audio_sample_rate,
)):
logger.warning(
"Audio mux failed for validation video %s; saved video without audio.",
filename)
video_filenames.append(filename)
artifacts = []
for filename, caption in zip(video_filenames,
all_captions,
strict=True):
video_artifact = self.tracker.video(filename,
caption=caption)
if video_artifact is not None:
artifacts.append(video_artifact)
if artifacts:
logs = {
f"validation_videos_{num_inference_steps}_steps":
artifacts
}
self.tracker.log_artifacts(logs, global_step)
else:
# Other sp_group leaders send their results to global rank 0
world_group.send_object(step_videos, dst=0)
world_group.send_object(step_captions, dst=0)
artifacts = []
for filename, caption in zip(video_filenames,
all_captions,
strict=True):
video_artifact = self.tracker.video(filename,
caption=caption)
if video_artifact is not None:
artifacts.append(video_artifact)
if artifacts:
logs = {
f"validation_videos_{num_inference_steps}_steps":
artifacts
}
self.tracker.log_artifacts(logs, global_step)
elif self.rank_in_sp_group == 0:
# Other sp_group leaders send their results to global rank 0
world_group.send_object(step_videos, dst=0)
world_group.send_object(step_captions, dst=0)
# Re-enable gradients for training
training_args.inference_mode = False
@@ -923,6 +950,98 @@ class TrainingPipeline(LoRAPipeline, ABC):
if getattr(self, "transformer_2", None) is not None:
self.transformer_2.train()
@staticmethod
def _mux_audio(
video_path: str,
audio: torch.Tensor | np.ndarray,
sample_rate: int,
) -> bool:
"""Mux audio into video using PyAV."""
try:
import av
except ImportError:
logger.warning("PyAV not installed; cannot mux audio. "
"Install with: pip install av")
return False
if torch.is_tensor(audio):
audio_np = audio.detach().cpu().float().numpy()
else:
audio_np = np.asarray(audio, dtype=np.float32)
if audio_np.ndim == 1:
audio_np = audio_np[:, None]
elif audio_np.ndim == 2:
if audio_np.shape[0] <= 8 and audio_np.shape[1] > audio_np.shape[0]:
audio_np = audio_np.T
else:
logger.warning("Unexpected audio shape %s; skipping mux.",
audio_np.shape)
return False
audio_np = np.clip(audio_np, -1.0, 1.0)
audio_int16 = (audio_np * 32767.0).astype(np.int16)
num_channels = audio_int16.shape[1]
layout = "stereo" if num_channels == 2 else "mono"
try:
import wave
with tempfile.TemporaryDirectory() as tmpdir:
out_path = os.path.join(tmpdir, "muxed.mp4")
wav_path = os.path.join(tmpdir, "audio.wav")
# Write audio to WAV file
with wave.open(wav_path, "wb") as wav_file:
wav_file.setnchannels(num_channels)
wav_file.setsampwidth(2)
wav_file.setframerate(sample_rate)
wav_file.writeframes(audio_int16.tobytes())
# Open input video and audio
input_video = av.open(video_path)
input_audio = av.open(wav_path)
# Create output with both streams
output = av.open(out_path, mode="w")
# Add video stream (copy codec from input)
in_video_stream = input_video.streams.video[0]
out_video_stream = output.add_stream(
codec_name=in_video_stream.codec_context.name,
rate=in_video_stream.average_rate,
)
out_video_stream.width = in_video_stream.width
out_video_stream.height = in_video_stream.height
out_video_stream.pix_fmt = in_video_stream.pix_fmt
# Add audio stream (AAC)
out_audio_stream = output.add_stream("aac", rate=sample_rate)
out_audio_stream.layout = layout
# Remux video (decode and re-encode to be safe)
for frame in input_video.decode(video=0):
for packet in out_video_stream.encode(frame):
output.mux(packet)
for packet in out_video_stream.encode():
output.mux(packet)
# Encode audio
for frame in input_audio.decode(audio=0):
frame.pts = None # Let encoder assign PTS
for packet in out_audio_stream.encode(frame):
output.mux(packet)
for packet in out_audio_stream.encode():
output.mux(packet)
input_video.close()
input_audio.close()
output.close()
shutil.move(out_path, video_path)
return True
except Exception as e:
logger.warning("Audio mux failed: %s", e)
return False
def visualize_intermediate_latents(self, training_batch: TrainingBatch,
training_args: TrainingArgs, step: int):
"""Add visualization data to tracker logging and save frames to disk."""
+191 -51
View File
@@ -257,8 +257,6 @@ def save_distillation_checkpoint(
if generator_scheduler is not None:
generator_states["scheduler"] = SchedulerWrapper(
generator_scheduler)
if generator_ema is not None:
generator_states["ema"] = generator_ema.state_dict()
generator_dcp_dir = os.path.join(save_dir, "distributed_checkpoint",
"generator")
@@ -290,8 +288,6 @@ def save_distillation_checkpoint(
if generator_scheduler_2 is not None:
generator_2_states["scheduler"] = SchedulerWrapper(
generator_scheduler_2)
if generator_ema_2 is not None:
generator_2_states["ema"] = generator_ema_2.state_dict()
generator_2_dcp_dir = os.path.join(save_dir,
"distributed_checkpoint",
@@ -417,6 +413,67 @@ def save_distillation_checkpoint(
rank,
local_main_process_only=False)
# Persist EMA separately to avoid shape mismatches across ranks.
# Supports:
# - mode="rank0_full": save consolidated EMA only on rank 0
# - mode="local_shard": save per-rank EMA shard for each rank
try:
if generator_ema is not None and getattr(generator_ema, "mode",
None) == "rank0_full":
_save_rank0_full_ema_safetensors(generator_ema,
generator_transformer, rank,
save_dir, "generator_ema")
elif generator_ema is not None and getattr(generator_ema, "mode",
None) == "local_shard":
# Save per-rank shard
ema_dir_shard = os.path.join(save_dir, "ema_local_shard")
os.makedirs(ema_dir_shard, exist_ok=True)
ema_shard_path = os.path.join(ema_dir_shard,
f"generator_ema_rank{rank}.pt")
torch.save(generator_ema.state_dict(), ema_shard_path)
logger.info(
"rank: %s, saved generator EMA shard (local_shard) to %s",
rank,
ema_shard_path,
local_main_process_only=False)
# Also consolidate EMA to a single full-state file on rank 0 by applying EMA to the model and gathering
_consolidate_local_shard_ema_and_save_safetensors(
generator_ema, generator_transformer, rank, save_dir,
"generator_ema")
except Exception as e:
logger.warning("rank: %s, failed saving EMA separately: %s", rank,
str(e))
try:
if generator_ema_2 is not None and getattr(generator_ema_2, "mode",
None) == "rank0_full":
_save_rank0_full_ema_safetensors(generator_ema_2,
generator_transformer_2, rank,
save_dir, "generator_ema_2")
elif generator_ema_2 is not None and getattr(generator_ema_2, "mode",
None) == "local_shard":
# Save per-rank shard for EMA_2
ema_dir_shard_2 = os.path.join(save_dir, "ema_local_shard")
os.makedirs(ema_dir_shard_2, exist_ok=True)
ema2_shard_path = os.path.join(ema_dir_shard_2,
f"generator_ema_2_rank{rank}.pt")
torch.save(generator_ema_2.state_dict(), ema2_shard_path)
logger.info(
"rank: %s, saved generator_2 EMA shard (local_shard) to %s",
rank,
ema2_shard_path,
local_main_process_only=False)
# Also consolidate EMA_2 to a single full-state file on rank 0
_consolidate_local_shard_ema_and_save_safetensors(
generator_ema_2, generator_transformer_2, rank, save_dir,
"generator_ema_2")
except Exception as e:
logger.warning("rank: %s, failed saving EMA_2 separately: %s", rank,
str(e))
# Save generator model weights (consolidated) for inference
cpu_state = gather_state_dict_on_cpu_rank0(generator_transformer,
device=None)
@@ -454,46 +511,45 @@ def save_distillation_checkpoint(
logger.info("--> distillation checkpoint saved at step %s to %s", step,
weight_path)
# Save generator_2 model weights (consolidated) for inference (MoE support)
if generator_transformer_2 is not None:
inference_save_dir_2 = os.path.join(
save_dir, "generator_2_inference_transformer")
cpu_state_2 = gather_state_dict_on_cpu_rank0(
generator_transformer_2, device=None)
# Save generator_2 model weights (consolidated) for inference (MoE support)
if generator_transformer_2 is not None:
inference_save_dir_2 = os.path.join(
save_dir, "generator_2_inference_transformer")
cpu_state_2 = gather_state_dict_on_cpu_rank0(generator_transformer_2,
device=None)
if rank == 0:
os.makedirs(inference_save_dir_2, exist_ok=True)
weight_path_2 = os.path.join(
inference_save_dir_2, "diffusion_pytorch_model.safetensors")
logger.info(
"rank: %s, saving consolidated generator_2 inference checkpoint to %s",
rank,
weight_path_2,
local_main_process_only=False)
if rank == 0:
os.makedirs(inference_save_dir_2, exist_ok=True)
weight_path_2 = os.path.join(inference_save_dir_2,
"diffusion_pytorch_model.safetensors")
logger.info(
"rank: %s, saving consolidated generator_2 inference checkpoint to %s",
rank,
weight_path_2,
local_main_process_only=False)
# Convert training format to diffusers format and save
diffusers_state_dict_2 = custom_to_hf_state_dict(
cpu_state_2,
generator_transformer_2.reverse_param_names_mapping)
save_file(diffusers_state_dict_2, weight_path_2)
# Convert training format to diffusers format and save
diffusers_state_dict_2 = custom_to_hf_state_dict(
cpu_state_2,
generator_transformer_2.reverse_param_names_mapping)
save_file(diffusers_state_dict_2, weight_path_2)
logger.info(
"rank: %s, consolidated generator_2 inference checkpoint saved to %s",
rank,
weight_path_2,
local_main_process_only=False)
logger.info(
"rank: %s, consolidated generator_2 inference checkpoint saved to %s",
rank,
weight_path_2,
local_main_process_only=False)
# Save model config
config_dict_2 = generator_transformer_2.hf_config
if "dtype" in config_dict_2:
del config_dict_2["dtype"] # TODO
config_path_2 = os.path.join(inference_save_dir_2,
"config.json")
with open(config_path_2, "w") as f:
json.dump(config_dict_2, f, indent=4)
logger.info(
"--> generator_2 distillation checkpoint saved at step %s to %s",
step, weight_path_2)
# Save model config
config_dict_2 = generator_transformer_2.hf_config
if "dtype" in config_dict_2:
del config_dict_2["dtype"] # TODO
config_path_2 = os.path.join(inference_save_dir_2, "config.json")
with open(config_path_2, "w") as f:
json.dump(config_dict_2, f, indent=4)
logger.info(
"--> generator_2 distillation checkpoint saved at step %s to %s",
step, weight_path_2)
def load_checkpoint(transformer,
@@ -644,18 +700,37 @@ def load_distillation_checkpoint(
end_time - begin_time,
local_main_process_only=False)
# Load EMA state if available and generator_ema is provided
# Load EMA separately if saved in rank0_full mode
if generator_ema is not None:
try:
ema_state = generator_states.get("ema")
if ema_state is not None:
generator_ema.load_state_dict(ema_state)
logger.info("rank: %s, generator EMA state loaded successfully",
rank)
else:
logger.info("rank: %s, no EMA state found in checkpoint", rank)
if getattr(generator_ema, "mode", None) == "rank0_full":
ema_path = os.path.join(checkpoint_path, "ema",
"generator_ema.pt")
if rank == 0 and os.path.exists(ema_path):
ema_state = torch.load(ema_path, map_location="cpu")
generator_ema.load_state_dict(ema_state)
logger.info(
"rank: %s, generator EMA (rank0_full) loaded from %s",
rank, ema_path)
elif rank == 0:
logger.info(
"rank: %s, generator EMA file not found at %s; skipping",
rank, ema_path)
elif getattr(generator_ema, "mode", None) == "local_shard":
ema_path = os.path.join(checkpoint_path, "ema_local_shard",
f"generator_ema_rank{rank}.pt")
if os.path.exists(ema_path):
ema_state = torch.load(ema_path, map_location="cpu")
generator_ema.load_state_dict(ema_state)
logger.info(
"rank: %s, generator EMA shard (local_shard) loaded from %s",
rank, ema_path)
else:
logger.info(
"rank: %s, generator EMA shard file not found at %s; skipping",
rank, ema_path)
except Exception as e:
logger.warning("rank: %s, failed to load EMA state: %s", rank,
logger.warning("rank: %s, failed to load generator EMA: %s", rank,
str(e))
# Load generator_2 distributed checkpoint (MoE support)
@@ -850,7 +925,7 @@ def load_distillation_checkpoint(
def normalize_dit_input(model_type, latents, vae) -> torch.Tensor:
if model_type == "hunyuan_hf" or model_type == "hunyuan":
return latents * 0.476986
return latents * vae.config.scaling_factor
elif model_type == "wan":
latents_mean = torch.tensor(vae.latents_mean)
latents_std = 1.0 / torch.tensor(vae.latents_std)
@@ -1169,6 +1244,71 @@ def custom_to_hf_state_dict(
return new_state_dict
def _save_full_ema_safetensors_from_state(
state_dict: dict[str, Any],
reverse_param_names_mapping: dict[str, tuple[str, int, int]],
output_path: str,
) -> None:
"""
Convert a training-format state_dict to HF format and save as safetensors.
"""
diffusers_state_dict = custom_to_hf_state_dict(state_dict,
reverse_param_names_mapping)
save_file(diffusers_state_dict, output_path)
def _save_rank0_full_ema_safetensors(
ema: "EMA_FSDP",
module,
rank: int,
save_dir: str,
base_name: str,
) -> None:
if rank != 0:
return
ema_dir = os.path.join(save_dir, "ema")
os.makedirs(ema_dir, exist_ok=True)
output_path = os.path.join(ema_dir, f"{base_name}.safetensors")
ema_state = ema.state_dict()
_save_full_ema_safetensors_from_state(ema_state,
module.reverse_param_names_mapping,
output_path)
logger.info("rank: %s, saved %s as consolidated EMA safetensors to %s",
rank,
base_name,
output_path,
local_main_process_only=False)
def _consolidate_local_shard_ema_and_save_safetensors(
ema: "EMA_FSDP",
module,
rank: int,
save_dir: str,
base_name: str,
) -> None:
try:
# Temporarily apply EMA to the live (sharded) module and gather full CPU state on rank 0
with ema.apply_to_model(module):
cpu_state_full = gather_state_dict_on_cpu_rank0(module, device=None)
if rank == 0:
ema_dir = os.path.join(save_dir, "ema")
os.makedirs(ema_dir, exist_ok=True)
output_path = os.path.join(ema_dir, f"{base_name}.safetensors")
_save_full_ema_safetensors_from_state(
cpu_state_full, module.reverse_param_names_mapping, output_path)
logger.info(
"rank: %s, saved consolidated %s EMA (from local_shard) as safetensors to %s",
rank,
base_name,
output_path,
local_main_process_only=False)
except Exception as ce:
logger.warning(
"rank: %s, failed consolidating %s EMA (local_shard): %s", rank,
base_name, str(ce))
def shift_timestep(timestep: torch.Tensor, shift: float,
num_train_timestep: float) -> torch.Tensor:
if shift == 1:
@@ -1795,5 +1935,5 @@ class EMA_FSDP:
self.saved.clear()
return False
def apply_to_model(self, module):
def apply_to_model(self, module: torch.nn.Module) -> _ApplyEMACtx:
return EMA_FSDP._ApplyEMACtx(self, module)
@@ -10,7 +10,10 @@ from fastvideo.pipelines.basic.wan.wan_dmd_pipeline import WanDMDPipeline
from fastvideo.training.distillation_pipeline import DistillationPipeline
from fastvideo.utils import is_vsa_available
vsa_available = is_vsa_available()
try:
vsa_available = is_vsa_available()
except Exception:
vsa_available = False
logger = init_logger(__name__)
@@ -19,7 +19,10 @@ from fastvideo.pipelines.pipeline_batch_info import ForwardBatch, TrainingBatch
from fastvideo.training.distillation_pipeline import DistillationPipeline
from fastvideo.utils import is_vsa_available, shallow_asdict
vsa_available = is_vsa_available()
try:
vsa_available = is_vsa_available()
except Exception:
vsa_available = False
logger = init_logger(__name__)
@@ -18,7 +18,10 @@ from fastvideo.pipelines.pipeline_batch_info import ForwardBatch, TrainingBatch
from fastvideo.training.training_pipeline import TrainingPipeline
from fastvideo.utils import is_vsa_available, shallow_asdict
vsa_available = is_vsa_available()
try:
vsa_available = is_vsa_available()
except Exception:
vsa_available = False
logger = init_logger(__name__)
@@ -10,7 +10,10 @@ from fastvideo.training.self_forcing_distillation_pipeline import (
SelfForcingDistillationPipeline)
from fastvideo.utils import is_vsa_available
vsa_available = is_vsa_available()
try:
vsa_available = is_vsa_available()
except Exception:
vsa_available = False
logger = init_logger(__name__)
+4 -1
View File
@@ -10,7 +10,10 @@ from fastvideo.pipelines.basic.wan.wan_pipeline import WanPipeline
from fastvideo.training.training_pipeline import TrainingPipeline
from fastvideo.utils import is_vsa_available
vsa_available = is_vsa_available()
try:
vsa_available = is_vsa_available()
except Exception:
vsa_available = False
logger = init_logger(__name__)
@@ -119,6 +119,13 @@ class PreprocessWorkflow(WorkflowBase):
@classmethod
def get_workflow_cls(cls,
fastvideo_args: FastVideoArgs) -> "PreprocessWorkflow":
is_ltx2_t2v = (fastvideo_args.workload_type == WorkloadType.T2V
and fastvideo_args.pipeline_config.__class__.__name__
== "LTX2T2VConfig")
if is_ltx2_t2v:
from fastvideo.workflow.preprocess.preprocess_workflow_ltx2_t2v import (
PreprocessWorkflowLTX2T2V)
return cast(PreprocessWorkflow, PreprocessWorkflowLTX2T2V)
if fastvideo_args.workload_type == WorkloadType.T2V:
from fastvideo.workflow.preprocess.preprocess_workflow_t2v import (
PreprocessWorkflowT2V)
@@ -0,0 +1,185 @@
# SPDX-License-Identifier: Apache-2.0
"""LTX-2 preprocessing workflow writing native .precomputed training artifacts."""
from __future__ import annotations
from pathlib import Path
from typing import TYPE_CHECKING
import torch
from torch.utils.data import DataLoader
from tqdm import tqdm
from fastvideo.configs.configs import PreprocessConfig
from fastvideo.distributed.parallel_state import get_world_rank
from fastvideo.logger import init_logger
from fastvideo.pipelines.pipeline_batch_info import PreprocessBatch
from fastvideo.workflow.preprocess.components import (
PreprocessingDataValidator, VideoForwardBatchBuilder, build_dataset)
from fastvideo.workflow.preprocess.preprocess_workflow import PreprocessWorkflow
if TYPE_CHECKING:
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
logger = init_logger(__name__)
class LTX2PrecomputedSaver:
"""Save LTX-2 preprocessing outputs to .pt files under .precomputed/."""
def __init__(self, output_root: Path):
self.output_root = output_root
self.latents_dir = self.output_root / "latents"
self.conditions_dir = self.output_root / "conditions"
self.audio_latents_dir: Path | None = None
self.latents_dir.mkdir(parents=True, exist_ok=True)
self.conditions_dir.mkdir(parents=True, exist_ok=True)
def _to_rel_pt_path(self, video_name: str) -> Path:
return Path(video_name).with_suffix(".pt")
def save_batch(self, batch: PreprocessBatch) -> None:
assert isinstance(batch.latents, torch.Tensor)
assert isinstance(batch.prompt_embeds, list) and len(
batch.prompt_embeds) > 0
assert isinstance(batch.prompt_attention_mask, list) and len(
batch.prompt_attention_mask) > 0
assert isinstance(batch.video_file_name, list)
assert isinstance(batch.num_frames, list)
assert isinstance(batch.height, list)
assert isinstance(batch.width, list)
assert isinstance(batch.fps, list)
prompt_embeds = batch.prompt_embeds[0]
prompt_attention_mask = batch.prompt_attention_mask[0]
assert isinstance(prompt_embeds, torch.Tensor)
assert isinstance(prompt_attention_mask, torch.Tensor)
audio_latents: list[torch.Tensor | None] = batch.extra.get(
"ltx2_audio_latents", [])
for idx, video_name in enumerate(batch.video_file_name):
rel_path = self._to_rel_pt_path(video_name)
latent_output = self.latents_dir / rel_path
latent_output.parent.mkdir(parents=True, exist_ok=True)
latent = batch.latents[idx].detach().cpu().contiguous()
latent_payload = {
"latents": latent,
"num_frames": int(batch.num_frames[idx]),
"height": int(batch.height[idx]),
"width": int(batch.width[idx]),
"fps": float(batch.fps[idx]),
}
torch.save(latent_payload, latent_output)
condition_output = self.conditions_dir / rel_path
condition_output.parent.mkdir(parents=True, exist_ok=True)
condition_payload = {
"prompt_embeds":
prompt_embeds[idx].detach().cpu().contiguous(),
"prompt_attention_mask":
prompt_attention_mask[idx].detach().cpu().contiguous(),
}
torch.save(condition_payload, condition_output)
if idx >= len(audio_latents) or audio_latents[idx] is None:
continue
if self.audio_latents_dir is None:
self.audio_latents_dir = self.output_root / "audio_latents"
self.audio_latents_dir.mkdir(parents=True, exist_ok=True)
assert self.audio_latents_dir is not None
audio_latent = audio_latents[idx].detach().cpu().contiguous()
audio_output = self.audio_latents_dir / rel_path
audio_output.parent.mkdir(parents=True, exist_ok=True)
audio_payload = {
"latents": audio_latent,
"num_time_steps": int(audio_latent.shape[1]),
"frequency_bins": int(audio_latent.shape[2]),
"duration":
float(batch.num_frames[idx]) / float(batch.fps[idx]),
}
torch.save(audio_payload, audio_output)
class PreprocessWorkflowLTX2T2V(PreprocessWorkflow):
"""LTX-2 workflow for generating native precomputed training tensors."""
training_dataloader: DataLoader
preprocess_pipeline: ComposedPipelineBase
video_forward_batch_builder: VideoForwardBatchBuilder
precomputed_saver: LTX2PrecomputedSaver
@staticmethod
def _resolve_precomputed_output_dir(
preprocess_cfg: PreprocessConfig) -> Path:
output_dir = Path(
preprocess_cfg.dataset_output_dir).expanduser().resolve()
if output_dir.name != ".precomputed":
output_dir = output_dir / ".precomputed"
return output_dir
def register_components(self) -> None:
assert self.fastvideo_args.preprocess_config is not None
preprocess_cfg = self.fastvideo_args.preprocess_config
raw_data_validator = PreprocessingDataValidator(
max_height=preprocess_cfg.max_height,
max_width=preprocess_cfg.max_width,
num_frames=preprocess_cfg.num_frames,
train_fps=preprocess_cfg.train_fps,
speed_factor=preprocess_cfg.speed_factor,
video_length_tolerance_range=preprocess_cfg.
video_length_tolerance_range,
drop_short_ratio=preprocess_cfg.drop_short_ratio,
)
self.add_component("raw_data_validator", raw_data_validator)
training_dataset = build_dataset(preprocess_cfg,
split="train",
validator=raw_data_validator)
training_dataloader = DataLoader(
training_dataset,
batch_size=preprocess_cfg.preprocess_video_batch_size,
num_workers=preprocess_cfg.dataloader_num_workers,
collate_fn=lambda x: x,
)
self.add_component("training_dataloader", training_dataloader)
video_forward_batch_builder = VideoForwardBatchBuilder(
seed=preprocess_cfg.seed)
self.add_component("video_forward_batch_builder",
video_forward_batch_builder)
output_root = self._resolve_precomputed_output_dir(preprocess_cfg)
precomputed_saver = LTX2PrecomputedSaver(output_root)
self.add_component("precomputed_saver", precomputed_saver)
def prepare_system_environment(self) -> None:
assert self.fastvideo_args.preprocess_config is not None
preprocess_cfg = self.fastvideo_args.preprocess_config
output_root = self._resolve_precomputed_output_dir(preprocess_cfg)
output_root.mkdir(parents=True, exist_ok=True)
self.precomputed_output_dir = output_root
logger.info("LTX-2 precomputed output directory: %s",
self.precomputed_output_dir)
def run(self) -> None:
total_samples = 0
for batch in tqdm(self.training_dataloader,
desc="Preprocessing LTX-2 training dataset",
unit="batch"):
forward_batch: PreprocessBatch = self.video_forward_batch_builder(
batch)
forward_batch = self.preprocess_pipeline.forward(
forward_batch, self.fastvideo_args)
self.precomputed_saver.save_batch(forward_batch)
total_samples += len(forward_batch.video_file_name)
logger.info(
"Finished LTX-2 preprocessing on rank %s with %s samples written to %s",
get_world_rank(),
total_samples,
self.precomputed_output_dir,
)
@@ -1,6 +1,19 @@
# SPDX-License-Identifier: Apache-2.0
"""
Convert LTX-2 weights to FastVideo naming conventions and split by component.
LTX 2 conversion requires two huggingface models:
- LTX 2 model
- Gemma model
Example usage:
python scripts/checkpoint_conversion/convert_ltx2_weights.py \\
--source "<PATH_TO_LOCAL_REPO>/Lightricks/LTX-2/ltx-2-19b-dev.safetensors" \\
--output "converted_weights/ltx2-base" \\
--class-name "LTX2Transformer3DModel" \\
--pipeline-class-name "LTX2Pipeline" \\
--diffusers-version "0.33.0.dev0" \\
--gemma-path "<PATH_TO_LOCAL_REPO>/google/gemma-3-12b-it"
"""
from __future__ import annotations
@@ -336,6 +349,32 @@ def maybe_download(repo_id: str, target_dir: Path, token: str | None, allow_patt
return target_dir
def copy_gemma_tokenizer(gemma_src: Path, tokenizer_dest: Path) -> None:
tokenizer_dest.mkdir(parents=True, exist_ok=True)
tokenizer_file_names = [
"tokenizer.json",
"tokenizer.model",
"tokenizer_config.json",
"special_tokens_map.json",
"added_tokens.json",
"chat_template.json",
"chat_template.jinja",
"preprocessor_config.json",
"processor_config.json",
]
copied = 0
for file_name in tokenizer_file_names:
src_path = gemma_src / file_name
if src_path.is_file():
shutil.copy2(src_path, tokenizer_dest / file_name)
copied += 1
if copied == 0:
raise FileNotFoundError(
f"No tokenizer files found in {gemma_src}. Expected at least one tokenizer file."
)
print(f"Copied {copied} tokenizer files to {tokenizer_dest}")
def main() -> None:
parser = argparse.ArgumentParser(description="Convert LTX-2 weights to FastVideo format")
parser.add_argument("--source", type=str, help="Path to transformer weights directory")
@@ -421,6 +460,7 @@ def main() -> None:
shutil.rmtree(gemma_dest)
gemma_dest.parent.mkdir(parents=True, exist_ok=True)
shutil.copytree(gemma_src, gemma_dest)
copy_gemma_tokenizer(gemma_src, output_dir / "tokenizer")
gemma_model_path = "gemma"
convert_components(
+108
View File
@@ -0,0 +1,108 @@
# SPDX-License-Identifier: Apache-2.0
import json
import importlib
import sys
import types
from pathlib import Path
import pytest
pytest.importorskip("torchvision")
def _bootstrap_fastvideo_namespace() -> None:
"""Avoid importing fastvideo/__init__.py during local registry tests."""
if "fastvideo" in sys.modules:
return
repo_root = Path(__file__).resolve().parents[2]
package_dir = repo_root / "fastvideo"
fastvideo_pkg = types.ModuleType("fastvideo")
fastvideo_pkg.__path__ = [str(package_dir)] # type: ignore[attr-defined]
fastvideo_pkg.__file__ = str(package_dir / "__init__.py")
sys.modules["fastvideo"] = fastvideo_pkg
def _get_registry_test_symbols() -> tuple[type, type, type, object, object]:
_bootstrap_fastvideo_namespace()
pipeline_module = importlib.import_module("fastvideo.configs.pipelines.ltx2")
sample_module = importlib.import_module("fastvideo.configs.sample.ltx2")
registry_module = importlib.import_module("fastvideo.registry")
return (
pipeline_module.LTX2T2VConfig,
sample_module.LTX2BaseSamplingParam,
sample_module.LTX2DistilledSamplingParam,
registry_module.get_pipeline_config_cls_from_name,
registry_module.get_sampling_param_cls_for_name,
)
@pytest.mark.parametrize(
("model_id", "expected_variant"),
[
("Lightricks/LTX-2", "base"),
("FastVideo/LTX2-base", "base"),
("FastVideo/LTX2-Distilled-Diffusers", "distilled"),
],
)
def test_ltx2_sampling_registry_exact_ids(model_id: str,
expected_variant: str) -> None:
_, base_cls, distilled_cls, _, get_sampling_param_cls_for_name = _get_registry_test_symbols()
expected_cls = base_cls if expected_variant == "base" else distilled_cls
resolved_cls = get_sampling_param_cls_for_name(model_id)
assert resolved_cls is expected_cls
@pytest.mark.parametrize(
"model_id",
[
"Lightricks/LTX-2",
"FastVideo/LTX2-base",
"FastVideo/LTX2-Distilled-Diffusers",
],
)
def test_ltx2_pipeline_registry_exact_ids(model_id: str) -> None:
pipeline_cls, _, _, get_pipeline_config_cls_from_name, _ = _get_registry_test_symbols()
resolved_cls = get_pipeline_config_cls_from_name(model_id)
assert resolved_cls is pipeline_cls
def _write_minimal_diffusers_repo(model_dir: Path, class_name: str) -> None:
model_dir.mkdir(parents=True, exist_ok=True)
(model_dir / "transformer").mkdir(exist_ok=True)
(model_dir / "vae").mkdir(exist_ok=True)
with (model_dir / "model_index.json").open("w", encoding="utf-8") as f:
json.dump(
{
"_class_name": class_name,
"_diffusers_version": "0.33.0.dev0",
"transformer": ["diffusers", "LTX2Transformer3DModel"],
"vae": ["diffusers", "CausalVideoAutoencoder"],
},
f,
)
def test_ltx2_ambiguous_local_path_has_no_sampling_fallback(
tmp_path: Path) -> None:
_, _, _, _, get_sampling_param_cls_for_name = _get_registry_test_symbols()
# Simulate a user-local converted LTX2 path that might be either base or
# distilled. Registry must not assume a variant for local converted paths.
model_dir = tmp_path / "converted" / "ltx2_diffusers"
_write_minimal_diffusers_repo(model_dir, "LTX2Pipeline")
resolved_cls = get_sampling_param_cls_for_name(str(model_dir))
assert resolved_cls is None
def test_ltx2_ambiguous_local_path_has_no_pipeline_mapping(
tmp_path: Path) -> None:
_, _, _, get_pipeline_config_cls_from_name, _ = _get_registry_test_symbols()
model_dir = tmp_path / "converted" / "ltx2_diffusers"
_write_minimal_diffusers_repo(model_dir, "LTX2Pipeline")
with pytest.raises(ValueError,
match="No match found for pipeline .*check the pipeline name or path"):
get_pipeline_config_cls_from_name(str(model_dir))