Compare commits
14
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
a8dddfaa16 | ||
|
|
7210c68f1b | ||
|
|
5602dc1bad | ||
|
|
ac4bc4ab84 | ||
|
|
bee27f9f74 | ||
|
|
ad58f802f3 | ||
|
|
04fa356ee3 | ||
|
|
f9c076fe2b | ||
|
|
dff0ea401a | ||
|
|
becd379f58 | ||
|
|
1ed7d7e1b0 | ||
|
|
d9fabcc5ef | ||
|
|
9db48498de | ||
|
|
1cd7038315 |
@@ -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,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.**
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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()
|
||||
@@ -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):
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
@@ -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)
|
||||
|
||||
@@ -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
@@ -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(),
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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__)
|
||||
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -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__)
|
||||
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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))
|
||||
Reference in New Issue
Block a user