Compare commits

..
Author SHA1 Message Date
JerryZhou54 8cb2ae9d27 ckpt 2026-02-09 05:50:48 +00:00
222 changed files with 579913 additions and 14401 deletions
+1 -7
View File
@@ -45,12 +45,6 @@ jobs:
- name: Setup Pages
uses: actions/configure-pages@v4
- name: Generate docs examples
run: python docs/generate_examples.py
- name: Check docs links
run: python scripts/check_docs_links.py
- name: Build documentation
run: mkdocs build
@@ -69,4 +63,4 @@ jobs:
steps:
- name: Deploy to GitHub Pages
id: deployment
uses: actions/deploy-pages@v4
uses: actions/deploy-pages@v4
@@ -64,10 +64,7 @@ jobs:
# - torch-version: '2.7.1'
# cuda-version: '12.8.0'
# torch-cuda-short: 'cu128'
# - torch-version: '2.9.1'
# cuda-version: '12.8.0'
# torch-cuda-short: 'cu128'
- torch-version: '2.10.0'
- torch-version: '2.9.1'
cuda-version: '12.8.0'
torch-cuda-short: 'cu128'
-11
View File
@@ -18,7 +18,6 @@ venv/
.venv/
runs/
samples/
Miniconda3-latest-Linux-x86_64.sh
*validation/
data/
outputs/
@@ -34,11 +33,6 @@ env
*.log
weights/
# SSIM test outputs
fastvideo/tests/ssim/generated_videos/
**/.cache/**
# Distribution / packaging
build/
dist/
@@ -75,11 +69,6 @@ docs/distillation/examples/
!docs/assets/images/**/*.png
!comfyui/assets/**/*.png
!comfyui/assets/**/*.gif
!assets/images/**/*.png
!assets/images/**/*.jpg
!assets/images/**/*.jpeg
!assets/images/**/*.gif
!assets/videos/**/*.mp4
dmd_t2v_output/
preprocess_output_text/
+1 -1
View File
@@ -10,7 +10,7 @@ exclude: |
demo/.*|
predict\.py|
scripts/.*|
assets/prompts/.*|
prompts/.*|
fastvideo/data_preprocess/.*|
fastvideo/dataset/.*|
fastvideo/models/.*|
-42
View File
@@ -1,42 +0,0 @@
# 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/` (including `assets/images/`, `assets/videos/`, and `assets/prompts/`) and `comfyui/assets/`.
## Build, Test, and Development Commands
- `uv pip install -e .[dev]`: editable install with lint/test extras.
- `pre-commit install --hook-type pre-commit --hook-type commit-msg`: enable local hooks.
- `pre-commit run --all-files`: run formatter/lint/type/spelling checks.
- `pytest tests/`: run top-level test suite.
- `pytest fastvideo/tests/ -v`: run package tests.
- `pytest fastvideo/tests/ssim/ -vs`: run SSIM regression tests (GPU-heavy).
- `cd fastvideo-kernel && ./build.sh`: build kernel extensions.
## Coding Style & Naming Conventions
- Python 3.10+; 4-space indentation; keep code and imports readable and explicit.
- Style tools are configured in `pyproject.toml` and `.pre-commit-config.yaml`:
- `yapf` (format), `ruff` (lint, auto-fix), `mypy` (typing), `codespell`.
- Target line length is 80.
- Naming: `snake_case` for functions/files, `PascalCase` for classes, `UPPER_SNAKE_CASE` for constants.
## Testing Guidelines
- Use `pytest` and place tests near relevant domains (e.g., `fastvideo/tests/encoders/`).
- Prefer descriptive names like `test_<feature>_<expected_behavior>.py`.
- For new pipelines/backends, include at least one regression-oriented test; add SSIM coverage when output quality must be preserved.
- Document GPU assumptions in tests that require specific hardware.
## Commit & Pull Request Guidelines
- Follow existing commit style: short subject with optional tag prefix, e.g. `[bugfix]: ...`, `[feat]: ...`, `[misc]: ...`, and include PR reference like `(#1234)` when applicable.
- Keep commits focused by concern (feature, refactor, fix).
- PRs should include:
- clear problem/solution summary,
- test evidence (`pytest`/SSIM outputs or rationale if skipped),
- linked issue/PR context,
- screenshots or sample outputs for UI/demo/docs changes.
-1
View File
@@ -1 +0,0 @@
@AGENTS.md
+118 -72
View File
@@ -1,101 +1,147 @@
# Sliding Tile Attention (STA) Branch
This branch is a stash/testing branch for people who want to try
Sliding Tile Attention (STA). The top-level README is intentionally
STA-only.
| **[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)** |
## What is STA
**FastVideo is a unified post-training and inference framework for accelerated video generation.**
Sliding Tile Attention is an optimized attention backend for
window-based video generation.
## NEWS
- Blog: https://hao-ai-lab.github.io/blogs/sta/
- Paper: https://arxiv.org/abs/2502.04507
- In-repo STA docs: `docs/attention/sta/index.md`
- `2025/11/19`: Release [CausalWan2.2 I2V A14B Preview](https://huggingface.co/FastVideo/CausalWan2.2-I2V-A14B-Preview-Diffusers) models, [Blog](https://hao-ai-lab.github.io/blogs/fastvideo_causalwan_preview/) and [Inference Code!](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_self_forcing_causal_wan2_2_i2v.py)
- `2025/08/04`: Release [FastWan](https://hao-ai-lab.github.io/FastVideo/distillation/dmd) models and [Sparse-Distillation](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/).
## Setup
### More News
Install FastVideo from source:
- `2025/06/14`: Release finetuning and inference code for [VSA](https://arxiv.org/pdf/2505.13389)
- `2025/04/24`: [FastVideo V1](https://hao-ai-lab.github.io/blogs/fastvideo/) is released!
- `2025/02/18`: Release the inference code for [Sliding Tile Attention](https://hao-ai-lab.github.io/blogs/sta/).
## Key Features
FastVideo has the following features:
- End-to-end post-training support for bidirectional and autoregressive models:
- Support full finetuning and LoRA finetuning for state-of-the-art open video DiTs
- Data preprocessing pipeline for video, image, and text data
- Distribution Matching Distillation (DMD2) stepwise distillation.
- Sparse attention with [Video Sparse Attention](https://arxiv.org/pdf/2505.13389)
- [Sparse distillation](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/) to achieve >50x denoising speedup
- Scalable training with FSDP2, sequence parallelism, and selective activation checkpointing.
- Causal distillation through Self-Forcing
- See this [page](https://hao-ai-lab.github.io/FastVideo/training/overview/) for full list of supported models and recipes.
- State-of-the-art performance optimizations for inference
- Sequence Parallelism for distributed inference
- Multiple state-of-the-art attention backends
- User-friendly CLI and Python API
- See this [page](https://hao-ai-lab.github.io/FastVideo/inference/optimizations/) for full list of supported optimizations.
- Diverse hardware and OS support
- Support H100, A100, 4090
- Support Linux, Windows, MacOS
- See this [page](https://hao-ai-lab.github.io/FastVideo/inference/hardware_support/) for full list of supported hardware and OS.
## Getting Started
We recommend using an environment manager such as `Conda` to create a clean environment:
```bash
uv pip install -e .
# Create and activate a new conda environment
conda create -n fastvideo python=3.12
conda activate fastvideo
# Install FastVideo
pip install fastvideo
```
Build the STA kernel package:
Please see our [docs](https://hao-ai-lab.github.io/FastVideo/getting_started/installation/) for more detailed installation instructions.
## Sparse Distillation
For our sparse distillation techniques, please see our [distillation docs](https://hao-ai-lab.github.io/FastVideo/distillation/dmd/) and check out our [blog](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/).
See below for recipes and datasets:
| Model | Sparse Distillation | Dataset |
| ------------------------------------------------------------------------------------- | --------------------------------------------------------------------------------------------------------------- | -------------------------------------------------------------------------------------------------------- |
| [FastWan2.1-T2V-1.3B](https://huggingface.co/FastVideo/FastWan2.1-T2V-1.3B-Diffusers) | [Recipe](https://github.com/hao-ai-lab/FastVideo/tree/main/examples/distill/Wan2.1-T2V/Wan-Syn-Data-480P) | [FastVideo Synthetic Wan2.1 480P](https://huggingface.co/datasets/FastVideo/Wan-Syn_77x448x832_600k) |
| [FastWan2.2-TI2V-5B](https://huggingface.co/FastVideo/FastWan2.2-TI2V-5B-Diffusers) | [Recipe](https://github.com/hao-ai-lab/FastVideo/tree/main/examples/distill/Wan2.2-TI2V-5B-Diffusers/Data-free) | [FastVideo Synthetic Wan2.2 720P](https://huggingface.co/datasets/FastVideo/Wan2.2-Syn-121x704x1280_32k) |
## Inference
### Generating Your First Video
Here's a minimal example to generate a video using the default settings. Make sure VSA kernels are [installed](https://hao-ai-lab.github.io/FastVideo/video_sparse_attention/installation/). Create a file called `example.py` with the following code:
```python
import os
from fastvideo import VideoGenerator
def main():
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "VIDEO_SPARSE_ATTN"
# Create a video generator with a pre-trained model
generator = VideoGenerator.from_pretrained(
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers",
num_gpus=1, # Adjust based on your hardware
)
# Define a prompt for your video
prompt = "A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes wide with interest."
# Generate the video
video = generator.generate_video(
prompt,
return_frames=True, # Also return frames from this call (defaults to False)
output_path="my_videos/", # Controls where videos are saved
save_video=True
)
if __name__ == '__main__':
main()
```
Run the script with:
```bash
cd fastvideo-kernel
./build.sh
cd ..
python example.py
```
## Run STA Inference
For a more detailed guide, please see our [inference quick start](https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start/).
STA backend:
## More Guides
```bash
export FASTVIDEO_ATTENTION_BACKEND=SLIDING_TILE_ATTN
```
- [Design Overview](https://hao-ai-lab.github.io/FastVideo/design/overview/)
- [Distillation Guide](https://hao-ai-lab.github.io/FastVideo/distillation/dmd/)
- [Contribution Guide](https://hao-ai-lab.github.io/FastVideo/contributing/overview/)
Ready-to-run examples:
## Awesome work using FastVideo or our research projects
- HunyuanVideo: `scripts/inference/v1_inference_hunyuan_STA.sh`
- Wan2.1-T2V-14B: `scripts/inference/v1_inference_wan_STA.sh`
- [SGLang](https://github.com/sgl-project/sglang/tree/main/python/sglang/multimodal_gen): SGLang's diffusion inference functionality is based on a fork of FastVideo on Sept. 24, 2025.
- [DanceGRPO](https://github.com/XueZeyue/DanceGRPO): A unified framework to adapt Group Relative Policy Optimization (GRPO) to visual generation paradigms. Code based on FastVideo.
- [SRPO](https://github.com/Tencent-Hunyuan/SRPO): A method to directly align the full diffusion trajectory with fine-grained human preference. Code based on FastVideo.
- [DCM](https://github.com/Vchitect/DCM): Dual-expert consistency model for efficient and high-quality video generation. Code based on FastVideo.
- [Hunyuan Video 1.5](https://github.com/Tencent-Hunyuan/HunyuanVideo-1.5): A leading lightweight video generation model, where they proposed SSTA based on Sliding Tile Attention.
- [Kandinsky-5.0](https://github.com/kandinskylab/kandinsky-5): A family of diffusion models for video & image generation, where their NABLA attention includes a Sliding Tile Attention branch.
- [LongCat Video](https://github.com/meituan-longcat/LongCat-Video): A foundational video generation model with 13.6B parameters with block-sparse attention similar to Video Sparse Attention.
Run:
## 🤝 Contributing
```bash
bash scripts/inference/v1_inference_hunyuan_STA.sh
# or
bash scripts/inference/v1_inference_wan_STA.sh
```
We welcome all contributions. Please check out our guide [here](https://hao-ai-lab.github.io/FastVideo/contributing/overview/).
See details in [development roadmap](https://github.com/hao-ai-lab/FastVideo/issues/899).
Both scripts already set STA-related env vars:
## Acknowledgement
- `FASTVIDEO_ATTENTION_BACKEND=SLIDING_TILE_ATTN`
- `FASTVIDEO_ATTENTION_CONFIG` to an STA mask strategy JSON
## STA Mask Strategy Files
- HunyuanVideo config: `assets/mask_strategy_hunyuan.json`
- Wan config: `assets/mask_strategy_wan.json`
## STA Mask Search (Wan2.1-T2V-14B)
Run mask search + tuning from repo root:
```bash
bash examples/inference/sta_mask_search/inference_wan_sta.sh
```
What this script does:
- Runs `STA_searching` first.
- Runs `STA_tuning` next (`skip_time_steps=12` by default).
- Uses prompt shards from `assets/prompt_0.txt` to `assets/prompt_7.txt`.
Important notes:
- Default script is set to 8 GPUs (`num_gpu=8`). If needed, edit
`examples/inference/sta_mask_search/inference_wan_sta.sh`.
- STA searching/tuning currently supports `69x768x1280` (Wan setting).
Generated files:
- Search results: `output/mask_search_result_pos_1280x768/`
- Tuned strategy: `output/mask_search_strategy_1280x768/mask_strategy_s12.json`
Use the tuned mask for STA inference:
```bash
export FASTVIDEO_ATTENTION_BACKEND=SLIDING_TILE_ATTN
export FASTVIDEO_ATTENTION_CONFIG=output/mask_search_strategy_1280x768/mask_strategy_s12.json
python examples/inference/sta_mask_search/wan_example.py --STA_mode STA_inference --num_gpus 1
```
We learned the design and reused code from the following projects: [Wan-Video](https://github.com/Wan-Video), [ThunderKittens](https://github.com/HazyResearch/ThunderKittens), [DMD2](https://github.com/tianweiy/DMD2), [diffusers](https://github.com/huggingface/diffusers), [xDiT](https://github.com/xdit-project/xDiT), [vLLM](https://github.com/vllm-project/vllm), [SGLang](https://github.com/sgl-project/sglang). We thank [MBZUAI](https://ifm.mbzuai.ac.ae/), [Anyscale](https://www.anyscale.com/), and [GMI Cloud](https://www.gmicloud.ai/) for their support throughout this project.
## Citation
If you find FastVideo useful, please consider citing our research work:
```bibtex
@article{zhang2025vsa,
title={Vsa: Faster video diffusion with trainable sparse attention},
author={Zhang, Peiyuan and Chen, Yongqi and Huang, Haofeng and Lin, Will and Liu, Zhengzhong and Stoica, Ion and Xing, Eric and Zhang, Hao},
journal={arXiv preprint arXiv:2505.13389},
year={2025}
}
@article{zhang2025fast,
title={Fast video generation with sliding tile attention},
author={Zhang, Peiyuan and Chen, Yongqi and Su, Runlong and Ding, Hangliang and Stoica, Ion and Liu, Zhengzhong and Zhang, Hao},
File diff suppressed because it is too large Load Diff
-1
View File
@@ -1 +0,0 @@
Will Smith casually eats noodles, his relaxed demeanor contrasting with the energetic background of a bustling street food market. The scene captures a mix of humor and authenticity. Mid-shot framing, vibrant lighting.
-1
View File
@@ -1 +0,0 @@
A lone hiker stands atop a towering cliff, silhouetted against the vast horizon. The rugged landscape stretches endlessly beneath, its earthy tones blending into the soft blues of the sky. The scene captures the spirit of exploration and human resilience. High angle, dynamic framing, with soft natural lighting emphasizing the grandeur of nature.
-1
View File
@@ -1 +0,0 @@
A hand with delicate fingers picks up a bright yellow lemon from a wooden bowl filled with lemons and sprigs of mint against a peach-colored background. The hand gently tosses the lemon up and catches it, showcasing its smooth texture. A beige string bag sits beside the bowl, adding a rustic touch to the scene. Additional lemons, one halved, are scattered around the base of the bowl. The even lighting enhances the vibrant colors and creates a fresh, inviting atmosphere.
-1
View File
@@ -1 +0,0 @@
A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes wide with interest. The playful yet serene atmosphere is complemented by soft natural light filtering through the petals. Mid-shot, warm and cheerful tones.
-1
View File
@@ -1 +0,0 @@
A superintelligent humanoid robot waking up. The robot has a sleek metallic body with futuristic design features. Its glowing red eyes are the focal point, emanating a sharp, intense light as it powers on. The scene is set in a dimly lit, high-tech laboratory filled with glowing control panels, robotic arms, and holographic screens. The setting emphasizes advanced technology and an atmosphere of mystery. The ambiance is eerie and dramatic, highlighting the moment of awakening and the robots immense intelligence. Photorealistic style with a cinematic, dark sci-fi aesthetic. Aspect ratio: 16:9 --v 6.1
-1
View File
@@ -1 +0,0 @@
fox in the forest close-up quickly turned its head to the left
-1
View File
@@ -1 +0,0 @@
Man walking his dog in the woods on a hot sunny day
-1
View File
@@ -1 +0,0 @@
A majestic lion strides across the golden savanna, its powerful frame glistening under the warm afternoon sun. The tall grass ripples gently in the breeze, enhancing the lion's commanding presence. The tone is vibrant, embodying the raw energy of the wild. Low angle, steady tracking shot, cinematic.
+195
View File
@@ -0,0 +1,195 @@
import argparse
import os
import tempfile
import gradio as gr
import torch
from diffusers import FlowMatchEulerDiscreteScheduler
from diffusers.utils import export_to_video
from fastvideo.distill.solver import PCMFMScheduler
from fastvideo.models.mochi_hf.modeling_mochi import MochiTransformer3DModel
from fastvideo.models.mochi_hf.pipeline_mochi import MochiPipeline
def init_args():
parser = argparse.ArgumentParser()
parser.add_argument("--prompts", nargs="+", default=[])
parser.add_argument("--num_frames", type=int, default=25)
parser.add_argument("--height", type=int, default=480)
parser.add_argument("--width", type=int, default=848)
parser.add_argument("--num_inference_steps", type=int, default=8)
parser.add_argument("--guidance_scale", type=float, default=4.5)
parser.add_argument("--model_path", type=str, default="data/mochi")
parser.add_argument("--seed", type=int, default=12345)
parser.add_argument("--transformer_path", type=str, default=None)
parser.add_argument("--scheduler_type", type=str, default="pcm_linear_quadratic")
parser.add_argument("--lora_checkpoint_dir", type=str, default=None)
parser.add_argument("--shift", type=float, default=8.0)
parser.add_argument("--num_euler_timesteps", type=int, default=50)
parser.add_argument("--linear_threshold", type=float, default=0.1)
parser.add_argument("--linear_range", type=float, default=0.75)
parser.add_argument("--cpu_offload", action="store_true")
return parser.parse_args()
def load_model(args):
if args.scheduler_type == "euler":
scheduler = FlowMatchEulerDiscreteScheduler()
else:
linear_quadratic = True if "linear_quadratic" in args.scheduler_type else False
scheduler = PCMFMScheduler(
1000,
args.shift,
args.num_euler_timesteps,
linear_quadratic,
args.linear_threshold,
args.linear_range,
)
if args.transformer_path:
transformer = MochiTransformer3DModel.from_pretrained(args.transformer_path)
else:
transformer = MochiTransformer3DModel.from_pretrained(args.model_path, subfolder="transformer/")
pipe = MochiPipeline.from_pretrained(args.model_path, transformer=transformer, scheduler=scheduler)
pipe.enable_vae_tiling()
# pipe.to(device)
# if args.cpu_offload:
pipe.enable_sequential_cpu_offload()
return pipe
def generate_video(
prompt,
negative_prompt,
use_negative_prompt,
seed,
guidance_scale,
num_frames,
height,
width,
num_inference_steps,
randomize_seed=False,
):
if randomize_seed:
seed = torch.randint(0, 1000000, (1, )).item()
generator = torch.Generator(device="cuda").manual_seed(seed)
if not use_negative_prompt:
negative_prompt = None
with torch.autocast("cuda", dtype=torch.bfloat16):
output = pipe(
prompt=[prompt],
negative_prompt=negative_prompt,
height=height,
width=width,
num_frames=num_frames,
num_inference_steps=num_inference_steps,
guidance_scale=guidance_scale,
generator=generator,
).frames[0]
output_path = os.path.join(tempfile.mkdtemp(), "output.mp4")
export_to_video(output, output_path, fps=30)
return output_path, seed
examples = [
"A hand enters the frame, pulling a sheet of plastic wrap over three balls of dough placed on a wooden surface. The plastic wrap is stretched to cover the dough more securely. The hand adjusts the wrap, ensuring that it is tight and smooth over the dough. The scene focuses on the hand’s movements as it secures the edges of the plastic wrap. No new objects appear, and the camera remains stationary, focusing on the action of covering the dough.",
"A vintage train snakes through the mountains, its plume of white steam rising dramatically against the jagged peaks. The cars glint in the late afternoon sun, their deep crimson and gold accents lending a touch of elegance. The tracks carve a precarious path along the cliffside, revealing glimpses of a roaring river far below. Inside, passengers peer out the large windows, their faces lit with awe as the landscape unfolds.",
"A crowded rooftop bar buzzes with energy, the city skyline twinkling like a field of stars in the background. Strings of fairy lights hang above, casting a warm, golden glow over the scene. Groups of people gather around high tables, their laughter blending with the soft rhythm of live jazz. The aroma of freshly mixed cocktails and charred appetizers wafts through the air, mingling with the cool night breeze.",
]
args = init_args()
pipe = load_model(args)
print("load model successfully")
with gr.Blocks() as demo:
gr.Markdown("# Fastvideo Mochi Video Generation Demo")
with gr.Group():
with gr.Row():
prompt = gr.Text(
label="Prompt",
show_label=False,
max_lines=1,
placeholder="Enter your prompt",
container=False,
)
run_button = gr.Button("Run", scale=0)
result = gr.Video(label="Result", show_label=False)
with gr.Accordion("Advanced options", open=False):
with gr.Group():
with gr.Row():
height = gr.Slider(
label="Height",
minimum=256,
maximum=1024,
step=32,
value=args.height,
)
width = gr.Slider(label="Width", minimum=256, maximum=1024, step=32, value=args.width)
with gr.Row():
num_frames = gr.Slider(
label="Number of Frames",
minimum=21,
maximum=163,
value=args.num_frames,
)
guidance_scale = gr.Slider(
label="Guidance Scale",
minimum=1,
maximum=12,
value=args.guidance_scale,
)
num_inference_steps = gr.Slider(
label="Inference Steps",
minimum=4,
maximum=100,
value=args.num_inference_steps,
)
with gr.Row():
use_negative_prompt = gr.Checkbox(label="Use negative prompt", value=False)
negative_prompt = gr.Text(
label="Negative prompt",
max_lines=1,
placeholder="Enter a negative prompt",
visible=False,
)
seed = gr.Slider(label="Seed", minimum=0, maximum=1000000, step=1, value=args.seed)
randomize_seed = gr.Checkbox(label="Randomize seed", value=True)
seed_output = gr.Number(label="Used Seed")
gr.Examples(examples=examples, inputs=prompt)
use_negative_prompt.change(
fn=lambda x: gr.update(visible=x),
inputs=use_negative_prompt,
outputs=negative_prompt,
)
run_button.click(
fn=generate_video,
inputs=[
prompt,
negative_prompt,
use_negative_prompt,
seed,
guidance_scale,
num_frames,
height,
width,
num_inference_steps,
randomize_seed,
],
outputs=[result, seed_output],
)
if __name__ == "__main__":
demo.queue(max_size=20).launch(server_name="0.0.0.0", server_port=7860)
+15
View File
@@ -0,0 +1,15 @@
Fast-Hunyuan comparison with original Hunyuan, achieving an 8X diffusion speed boost with the FastVideo framework.
https://github.com/user-attachments/assets/064ac1d2-11ed-4a0c-955b-4d412a96ef30
Fast-Mochi comparison with original Mochi, achieving an 8X diffusion speed boost with the FastVideo framework.
https://github.com/user-attachments/assets/5fbc4596-56d6-43aa-98e0-da472cf8e26c
Comparison between OpenAI Sora, original Hunyuan and FastHunyuan
https://github.com/user-attachments/assets/d323b712-3f68-42b2-952b-94f6a49c4836
Comparison between original FastHunyuan, LLM-INT8 quantized FastHunyuan and NF4 quantized FastHunyuan
https://github.com/user-attachments/assets/cf89efb5-5f68-4949-a085-f41c1ef26c94
+3 -3
View File
@@ -6,7 +6,7 @@ All documented examples are autogenerated using [generate_examples.py](https://g
## Examples
- [Examples Distillation Index](../distillation/examples/examples_distillation_index.md)
- [Examples Training Index](../training/examples/examples_training_index.md)
- [Examples Inference Index](../inference/examples/examples_inference_index.md)
- [Examples Distillation Index](distillation/examples/examples_distillation_index.md)
- [Examples Training Index](training/examples/examples_training_index.md)
- [Examples Inference Index](inference/examples/examples_inference_index.md)
+5 -18
View File
@@ -2,7 +2,6 @@
# adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/docs/source/generate_examples.py
import itertools
import os
import re
from dataclasses import dataclass, field
from pathlib import Path
@@ -127,20 +126,8 @@ class Example:
Raises:
IndexError: If no Markdown files are found in the directory.
""" # noqa: E501
if self.path.is_file():
return self.path
markdown_files = sorted(self.path.glob("*.md"))
if not markdown_files:
raise IndexError(f"No Markdown files found in {self.path}")
readme_files = [
f for f in markdown_files if f.name.lower() == "readme.md"
]
if readme_files:
return readme_files[0]
return markdown_files[0]
return self.path if self.path.is_file() else list(
self.path.glob("*.md")).pop()
def determine_other_files(self) -> list[Path]:
"""
@@ -534,11 +521,11 @@ def generate_examples(generate_main_index: bool = False) -> None:
# Add to main index if it exists
if generate_main_index and examples_index:
main_index_dir = examples_index.path.parent
rel_path = os.path.relpath(category_index.path,
start=main_index_dir)
rel_path = category_index.path.relative_to(
main_index_dir.parent)
examples_index.documents.insert(
0,
str(rel_path).replace("\\", "/").replace(".md", ""))
str(rel_path).replace(".md", ""))
# Write the category index file
with open(category_index.path, "w+") as f:
-13
View File
@@ -18,25 +18,12 @@ conda activate fastvideo
pip install fastvideo
```
### Using uv
```bash
# Create and activate a new uv environment
uv venv --python 3.12 --seed
source .venv/bin/activate
uv pip install fastvideo
```
### From source
```bash
git clone https://github.com/hao-ai-lab/FastVideo.git
cd FastVideo
pip install -e .
# or if you are using uv
uv pip install -e .
```
Also optionally install flash-attn:
+1 -2
View File
@@ -79,6 +79,5 @@ if __name__ == '__main__':
- [Installation Guide](installation.md) - Detailed installation instructions
- [Configuration](../inference/configuration.md) - Learn about configuration options
- [Examples](../inference/examples/examples_inference_index.md) - Explore more
examples
- [Examples](../inference/examples/) - Explore more examples
- [Optimizations](../inference/optimizations.md) - Performance optimization tips
+4 -14
View File
@@ -101,17 +101,14 @@ Replace standard attention with FastVideo's optimized attention:
```python
# Local attention patterns
from fastvideo.attention import LocalAttention
from fastvideo.platforms.interface import AttentionBackendEnum
from fastvideo.attention.backends.abstract import _Backend
self.attn = LocalAttention(
num_heads=num_heads,
head_size=head_dim,
dropout_rate=0.0,
softmax_scale=None,
causal=False,
supported_attention_backends=(
AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA,
)
supported_attention_backends=(_Backend.FLASH_ATTN, _Backend.TORCH_SDPA)
)
# Distributed attention for long sequences
@@ -122,21 +119,14 @@ self.attn = DistributedAttention(
dropout_rate=0.0,
softmax_scale=None,
causal=False,
supported_attention_backends=(
AttentionBackendEnum.SLIDING_TILE_ATTN,
AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA,
)
supported_attention_backends=(_Backend.SLIDING_TILE_ATTN, _Backend.FLASH_ATTN, _Backend.TORCH_SDPA)
)
```
#### Define supported backend selection
```python
_supported_attention_backends = (
AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA,
)
_supported_attention_backends = (_Backend.FLASH_ATTN, _Backend.TORCH_SDPA)
```
### Registering Models
+94 -65
View File
@@ -1,85 +1,104 @@
# FastVideo CLI Inference
The FastVideo CLI exposes the same core inference controls as the Python API.
The FastVideo CLI provides a quick way to access the FastVideo inference pipeline for video generation. For more advanced usage,
see the Python interface [here](examples/basic.md).
## Basic Usage
Use either:
1. `--model-path` + `--prompt`
2. `--model-path` + `--prompt-txt` (batch prompts, one line per prompt)
3. `--config` (JSON/YAML)
The basic command to generate a video is:
```bash
fastvideo generate --model-path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
--prompt "A cat playing with a ball of yarn"
fastvideo generate --model-path {MODEL_PATH} --prompt {PROMPT}
```
```bash
fastvideo generate --model-path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
--prompt-txt prompts.txt
```
### Required Parameters
You cannot provide both `--prompt` and `--prompt-txt` in the same run.
- `--model-path {MODEL_PATH}`: Path to the model or model ID
- `--prompt {PROMPT}`: Text description for the video you want to generate
## View All Arguments
## Common Arguments
To see all the options, you can use the `--help` flag:
```bash
fastvideo generate --help
```
Arguments come from:
### Hardware Configuration
- FastVideo runtime args (`FastVideoArgs`)
- Sampling args (`SamplingParam`)
- Pipeline config args (`PipelineConfig`)
- `--num-gpus {NUM_GPUS}`: Number of GPUs to use
- `--tp-size {TP_SIZE}`: Tensor parallelism size (only for the encoder, should not be larger than 1 if text encoder offload is enabled, as layerwise offload + prefetch is faster)
- `--sp-size {SP_SIZE}`: Sequence parallelism size (Typically should match the number of GPUs)
## Common Arguments
#### Video Configuration
### Parallelism
- `--height {HEIGHT}`: Height of the generated video
- `--width {WIDTH}`: Width of the generated video
- `--num-frames {NUM_FRAMES}`: Number of frames to generate
- `--fps {FPS}`: Frames per second for the saved video
- `--num-gpus`
- `--sp-size`
- `--tp-size`
#### Generation Parameters
### Sampling
- `--num-inference-steps {STEPS}`: Number of denoising steps
- `--negative-prompt {PROMPT}`: Negative prompt to guide generation away from certain concepts
- `--seed {SEED}`: Random seed for reproducible generation
- `--num-frames`
- `--height` / `--width`
- `--num-inference-steps`
- `--guidance-scale`
- `--seed`
- `--negative-prompt`
#### Output Options
### Output
- `--output-path {PATH}`: Directory to save the generated video
- `--save-video`: Whether to save the video to disk
- `--return-frames`: Whether to return the raw frames
- `--output-path`
- `--save-video` / `--no-save-video`
- `--return-frames`
## Using Configuration Files
### Offloading and Performance
- `--dit-layerwise-offload`
- `--use-fsdp-inference`
- `--text-encoder-cpu-offload`
- `--image-encoder-cpu-offload`
- `--vae-cpu-offload`
- `--enable-torch-compile`
- `--torch-compile-kwargs`
## Using Config Files
Instead of specifying all parameters on the command line, you can use a configuration file:
```bash
fastvideo generate --config config.yaml
fastvideo generate --config {CONFIG_FILE_PATH}
```
Config files can be JSON or YAML. CLI flags override config-file values.
The config file should be in JSON or YAML format with the same parameter names as the CLI options. Command-line arguments will take precedence over settings in the configuration file, allowing you to override specific values while keeping the rest from the config file.
Example `config.yaml`:
Example configuration file (config.json):
```json
{
"model_path": "FastVideo/FastHunyuan-diffusers",
"prompt": "A beautiful woman in a red dress walking down a street",
"output_path": "outputs/",
"num_gpus": 2,
"sp_size": 2,
"tp_size": 1,
"num_frames": 45,
"height": 720,
"width": 1280,
"num_inference_steps": 6,
"seed": 1024,
"fps": 24,
"precision": "bf16",
"vae_precision": "fp16",
"vae_tiling": true,
"vae_sp": true,
"vae_config": {
"load_encoder": false,
"load_decoder": true,
"tile_sample_min_height": 256,
"tile_sample_min_width": 256
},
"text_encoder_precisions": [
"fp16",
"fp16"
],
"mask_strategy_file_path": null,
"enable_torch_compile": false
}
```
Or using YAML format (config.yaml):
```yaml
model_path: "FastVideo/FastHunyuan-diffusers"
prompt: "A capybara lounging in a hammock"
prompt: "A beautiful woman in a red dress walking down a street"
output_path: "outputs/"
num_gpus: 2
sp_size: 2
@@ -89,34 +108,44 @@ height: 720
width: 1280
num_inference_steps: 6
seed: 1024
dit_precision: "bf16"
fps: 24
precision: "bf16"
vae_precision: "fp16"
vae_tiling: true
vae_sp: true
vae_config:
load_encoder: false
load_decoder: true
tile_sample_min_height: 256
tile_sample_min_width: 256
text_encoder_precisions:
- "fp16"
- "fp16"
mask_strategy_file_path: null
enable_torch_compile: false
```
Notes:
- Use `dit_precision` / `vae_precision` (not `precision`).
- Nested config objects are supported, for example `vae_config` and
`dit_config`.
## Examples
Simple generation:
Generating a simple video:
```bash
fastvideo generate \
--model-path FastVideo/FastHunyuan-diffusers \
--prompt "A cat playing with a ball of yarn" \
--num-frames 45 --height 720 --width 1280 \
--num-inference-steps 6 --seed 1024 \
--output-path outputs/
fastvideo generate --model-path FastVideo/FastHunyuan-diffusers --prompt "A cat playing with a ball of yarn" --num-frames 45 --height 720 --width 1280 --num-inference-steps 6 --seed 1024 --output-path outputs/
```
Config + CLI override:
Using a negative prompt to avoid certain elements:
```bash
fastvideo generate --config config.yaml --prompt "A panda skiing at sunset"
fastvideo generate --model-path FastVideo/FastHunyuan-diffusers --prompt "A beautiful forest landscape" --negative-prompt "people, buildings, roads"
```
Combining command line arguments and a configuration file:
```bash
fastvideo generate --config config.json --prompt "A capybara lounging in a hammock"
```
## Troubleshooting
- If you encounter CUDA out-of-memory errors, try reducing the video dimensions or number of frames, or the number of inference steps.
- For reproducible results, set the same seed value between runs.
+3 -31
View File
@@ -1,3 +1,4 @@
# Configuration
## Multi-GPU Setup
@@ -17,8 +18,7 @@ generator = VideoGenerator.from_pretrained(
- `PipelineConfig`: Initialization time parameters
- `SamplingParam`: Generation time parameters
You can customize generation behavior using `PipelineConfig` and
`SamplingParam`:
You can customize various parameters when generating videos using the `PipelineConfig` and `SamplingParam` class:
```python
from fastvideo import VideoGenerator, SamplingParam, PipelineConfig
@@ -27,12 +27,12 @@ def main():
model_name = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
config = PipelineConfig.from_pretrained(model_name)
config.vae_precision = "fp16"
config.dit_cpu_offload = True
# Create the generator
generator = VideoGenerator.from_pretrained(
model_name,
num_gpus=1,
dit_layerwise_offload=True, # FastVideoArgs option
pipeline_config=config
)
@@ -72,34 +72,6 @@ if __name__ == '__main__':
main()
```
## JSON/YAML Config Files (CLI)
The CLI supports `--config` with JSON or YAML. Command-line arguments override
config file values.
```bash
fastvideo generate --config config.yaml
```
Use CLI argument names as keys (underscore or hyphen is accepted). Example:
```yaml
model_path: "FastVideo/FastHunyuan-diffusers"
prompt: "A capybara relaxing in a hammock"
num_gpus: 2
sp_size: 2
num_frames: 45
height: 720
width: 1280
num_inference_steps: 6
seed: 1024
dit_precision: "bf16"
vae_precision: "fp16"
vae_tiling: true
vae_sp: true
enable_torch_compile: false
```
## Performance Optimization
For configuring optimizations, please see our [optimizations guide](optimizations.md)
+4 -4
View File
@@ -61,8 +61,7 @@ python example.py
The generated video will be saved in the current directory under `my_videos/`
More inference scripts and recipes can be found in `examples/inference/` and
`scripts/inference/`.
More inference example scripts can be found in `scripts/inference/`
## Available Models
@@ -104,10 +103,11 @@ Common issues and their solutions:
If you encounter CUDA out of memory errors:
- Reduce `num_frames` or video resolution
- Enable FastVideo offloading options such as `dit_layerwise_offload=True`
(single GPU) or `use_fsdp_inference=True` (multi-GPU)
- Enable memory optimization with `enable_model_cpu_offload`
- Try a smaller model or use distilled versions
- Use `num_gpus` > 1 if multiple GPUs are available
- Try enabling FSDP inference with `use_fsdp_inference=True` (may slow down generation)
- Try enabling DiT layerwise offload with `dit_layerwise_offload=True` (now only a few models support this, but may introduce less overhead than FSDP)
### Slow Generation
+11 -13
View File
@@ -25,9 +25,6 @@ This page describes the various options for speeding up generation times in Fast
- Video Sparse Attention: `FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN`
- Sage Attention: `FASTVIDEO_ATTENTION_BACKEND=SAGE_ATTN`
- Sage Attention 3: `FASTVIDEO_ATTENTION_BACKEND=SAGE_ATTN_THREE`
- Video MoBA Attention: `FASTVIDEO_ATTENTION_BACKEND=VMOBA_ATTN`
- Sparse Linear Attention: `FASTVIDEO_ATTENTION_BACKEND=SLA_ATTN`
- SageSLA Attention: `FASTVIDEO_ATTENTION_BACKEND=SAGE_SLA_ATTN`
### Configuring Backends
@@ -73,15 +70,22 @@ python setup.py install
**`SLIDING_TILE_ATTN`**
Sliding Tile Attention is provided by `fastvideo-kernel`.
See [STA docs](../attention/sta/index.md) for installation details.
```bash
pip install st_attn==0.0.4
```
Please see [this page](../attention/sta/index.md) for more installation instructions.
### Video Sparse Attention
**`VIDEO_SPARSE_ATTN`**
Video Sparse Attention is provided by `fastvideo-kernel`.
See [VSA docs](../attention/vsa/index.md) for installation details.
```bash
git submodule update --init --recursive
python setup_vsa.py install
```
Please see [this page](../attention/vsa/index.md) for more installation instructions.
### Sage Attention
@@ -111,12 +115,6 @@ Note that Sage Attention 3 requires `python>=3.13`, `torch>=2.8.0`, `CUDA >=12.8
To use Sage Attention 3 in FastVideo, follow the `README.md` in the linked repository to install the package from source.
### V-MoBA / SLA / SageSLA
These backends are model-specific and require the corresponding kernels and
dependencies. Use the support matrix and model examples to confirm compatibility
before enabling them.
## Teacache
TeaCache is an optimization technique supported in FastVideo that can significantly speed up video generation by skipping redundant calculations across diffusion steps. This guide explains how to enable and configure TeaCache for optimal performance in FastVideo.
+6 -15
View File
@@ -1,10 +1,6 @@
# Compatibility Matrix
This page summarizes common model + optimization combinations.
For the canonical, code-level list of model IDs recognized by
`VideoGenerator.from_pretrained(...)`, see the registrations in
`fastvideo/registry.py` (`register_configs(...)` entries).
The table below shows every supported model and optimizations supported for them.
The symbols used have the following meanings:
@@ -14,9 +10,7 @@ The symbols used have the following meanings:
## Models x Optimization
The `HuggingFace Model ID` can be passed directly to
`from_pretrained()`. FastVideo then uses model-specific default settings for
pipeline initialization and sampling.
The `HuggingFace Model ID` can be directly pass to `from_pretrained()` methods and FastVideo will use the optimal default parameters when initializing and generating videos.
<style>
/* Target tables in this section */
@@ -59,9 +53,9 @@ pipeline initialization and sampling.
| Wan2.1 T2V 14B | `Wan-AI/Wan2.1-T2V-14B-Diffusers` | 480P, 720P | ✅ | ✅* | ✅ | ⭕ | ⭕ |
| Wan2.1 I2V 480P | `Wan-AI/Wan2.1-I2V-14B-480P-Diffusers` | 480P | ✅ | ✅* | ✅ | ⭕ | ⭕ |
| Wan2.1 I2V 720P | `Wan-AI/Wan2.1-I2V-14B-720P-Diffusers` | 720P | ✅ | ✅ | ✅ | ⭕ | ⭕ |
| StepVideo T2V | `FastVideo/stepvideo-t2v-diffusers` | 768px768px204f<br>544px992px204f<br>544px992px136f | ❌ | ❌ | ✅ | ⭕ | ⭕ |
| TurboWan2.1 T2V 1.3B | `loayrashid/TurboWan2.1-T2V-1.3B-Diffusers` | 480P | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
| TurboWan2.1 T2V 14B | `loayrashid/TurboWan2.1-T2V-14B-Diffusers` | 480P, 720P | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
| TurboWan2.2 I2V A14B | `loayrashid/TurboWan2.2-I2V-A14B-Diffusers` | 480P<br>720P | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
| LongCat T2V 13.6B | See note** | 480P<br>720P | ❌ | ❌ | ❌ | ⭕ | ✅ |
| Matrix Game 2.0 Base | `FastVideo/Matrix-Game-2.0-Base-Diffusers` | 352x640 | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
| Matrix Game 2.0 GTA | `FastVideo/Matrix-Game-2.0-GTA-Diffusers` | 352x640 | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
@@ -69,14 +63,11 @@ pipeline initialization and sampling.
**Note**: Wan2.2 TI2V 5B has some quality issues when performing I2V generation. We are working on fixing this issue.
## Canonical Supported IDs
The authoritative source for model-ID recognition is
`fastvideo/registry.py`. If a model ID is registered there, FastVideo can
resolve default pipeline and sampling configuration for it.
## Special requirements
### StepVideo T2V
- The self-attention in text-encoder (step_llm) only supports CUDA capabilities sm_80 sm_86 and sm_90
### Sliding Tile Attention
- Currently only Hopper GPUs (H100s) are supported.
-75
View File
@@ -1,75 +0,0 @@
# Debugging
This page collects practical debugging steps for FastVideo inference issues.
## Collect Environment Info
From the repository root, run:
```bash
python collect_env.py
```
Attach the output when filing a GitHub issue.
## Increase Logging
FastVideo logging level is controlled by environment variables:
```bash
FASTVIDEO_LOGGING_LEVEL=DEBUG \
FASTVIDEO_STAGE_LOGGING=1 \
python your_script.py
```
Useful variables:
- `FASTVIDEO_LOGGING_LEVEL`: `DEBUG`, `INFO`, `WARNING`, `ERROR`
- `FASTVIDEO_STAGE_LOGGING`: print per-stage timings during pipeline execution
- `FASTVIDEO_ATTENTION_BACKEND`: force an attention backend (for example
`TORCH_SDPA` or `FLASH_ATTN`)
## Common Failure Modes
### Out-of-memory
Try, in order:
1. Reduce `height`, `width`, `num_frames`, or `num_inference_steps`.
2. Enable offloading flags such as `dit_layerwise_offload` (single GPU) or
`use_fsdp_inference` (multi-GPU).
3. Enable `vae_cpu_offload`, `image_encoder_cpu_offload`, and
`text_encoder_cpu_offload`.
See [Inference Offloading](../inference/offloading.md) for recommended
combinations.
### Attention backend import errors
If forcing a backend fails, verify optional dependencies are installed:
- `FLASH_ATTN`: `flash-attn`
- `SLIDING_TILE_ATTN` and `VIDEO_SPARSE_ATTN`: `fastvideo-kernel`
- `SAGE_ATTN` / `SAGE_ATTN_THREE`: SageAttention packages
As a fallback, use:
```bash
export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
```
### Configuration parsing errors
When using `--config`, keep keys aligned with CLI argument names (underscores or
hyphens are both accepted). For nested config values, use nested objects
(`vae_config`, `dit_config`) rather than dotted keys.
## Issue Template
When opening an issue, include:
- exact command or Python snippet,
- model ID/path,
- full traceback,
- `collect_env.py` output,
- whether the problem reproduces with `FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA`.
@@ -20,7 +20,7 @@ def main():
sampling_param = SamplingParam.from_pretrained(model_path)
# image2world example from official repo
image_path = "assets/images/bus_terminal.jpg"
image_path = "images/bus_terminal.jpg"
prompt = (
"A nighttime city bus terminal gradually shifts from stillness to subtle movement. "
@@ -48,3 +48,4 @@ def main():
if __name__ == "__main__":
main()
@@ -20,7 +20,7 @@ def main():
sampling_param = SamplingParam.from_pretrained(model_path)
# video2world example from official repo
video_path = "assets/videos/robot_pouring.mp4"
video_path = "videos/robot_pouring.mp4"
prompt = (
"A robotic arm, primarily white with black joints and cables, is shown in a clean, modern indoor setting with a white tabletop. "
@@ -51,3 +51,4 @@ def main():
if __name__ == "__main__":
main()
-119
View File
@@ -1,119 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""
Basic inference script for HunyuanGameCraft video generation.
HunyuanGameCraft generates game-like videos with camera/action control.
It takes an optional image input and generates video with camera motion
based on simple action commands (forward, left, right, backward, rotations).
Available actions:
- forward (w): Move camera forward
- backward (s): Move camera backward
- left (a): Move camera left (strafe)
- right (d): Move camera right (strafe)
- left_rot: Rotate camera left (pan)
- right_rot: Rotate camera right (pan)
- up_rot: Rotate camera up (tilt)
- down_rot: Rotate camera down (tilt)
T2V vs I2V:
- Default: I2V (uses a default reference image). Set GAMECRAFT_I2V_IMAGE to a
URL or path to use a different image.
- T2V only (no reference image): run with GAMECRAFT_I2V_IMAGE= (empty).
"""
import os
import torch
from fastvideo import VideoGenerator
from fastvideo.models.camera import create_camera_trajectory
# Model configuration (use GAMECRAFT_MODEL_PATH for local weights)
MODEL_PATH = os.environ.get("GAMECRAFT_MODEL_PATH", "FastVideo/HunyuanGameCraft-Diffusers")
# Default prompts for demo
DEFAULT_PROMPTS = {
"village": "A charming medieval village with cobblestone streets, thatched-roof houses, and vibrant flower gardens under a bright blue sky.",
"temple": "A majestic ancient temple stands under a clear blue sky, its grandeur highlighted by towering Doric columns and intricate architectural details.",
"forest": "A lush green forest with tall trees, dappled sunlight filtering through the leaves, and a winding dirt path.",
"beach": "A tropical beach with crystal clear turquoise water, white sand, and palm trees swaying in the breeze.",
}
# I2V: default reference image (URL). Can override with a local path.
DEFAULT_I2V_IMAGE_URL = (
"https://huggingface.co/datasets/huggingface/documentation-images/"
"resolve/main/diffusers/astronaut.jpg"
)
DEFAULT_I2V_PROMPT = (
"An astronaut hatching from an egg, on the surface of the moon, "
"the darkness and depth of space realised in the background."
)
OUTPUT_PATH = "video_samples_gamecraft"
def main():
# Initialize generator
# FastVideo will automatically download weights from HuggingFace
generator = VideoGenerator.from_pretrained(
MODEL_PATH,
num_gpus=1,
use_fsdp_inference=True,
dit_cpu_offload=True,
vae_cpu_offload=True,
text_encoder_cpu_offload=True,
pin_cpu_memory=True,
)
# Video parameters
height = 704
width = 1280
num_frames = 33
action = "forward"
action_speed = 0.2
# Create camera trajectory (Plücker coordinates)
camera_states = create_camera_trajectory(
action=action,
height=height,
width=width,
num_frames=num_frames,
action_speed=action_speed,
dtype=torch.bfloat16,
)
print(f"Camera states shape: {camera_states.shape}")
# I2V vs T2V: unset GAMECRAFT_I2V_IMAGE -> I2V (default image). Set to "" -> T2V.
env_image = os.environ.get("GAMECRAFT_I2V_IMAGE")
if env_image is None:
image_path = DEFAULT_I2V_IMAGE_URL # default: I2V
elif env_image.strip() == "":
image_path = None # T2V
else:
image_path = env_image.strip() # I2V with given URL/path
is_i2v = image_path is not None
prompt = DEFAULT_I2V_PROMPT if is_i2v else DEFAULT_PROMPTS["temple"]
print(f"Mode: {'I2V' if is_i2v else 'T2V'}, prompt: {prompt[:60]}...")
gen_kw = dict(
prompt=prompt,
negative_prompt="",
camera_states=camera_states,
height=height,
width=width,
num_frames=num_frames,
num_inference_steps=50,
guidance_scale=6.0,
seed=42,
fps=24,
output_path=OUTPUT_PATH,
save_video=True,
)
if is_i2v:
gen_kw["image_path"] = image_path
generator.generate_video(**gen_kw)
if __name__ == "__main__":
main()
@@ -1,49 +0,0 @@
from fastvideo import VideoGenerator
from fastvideo.models.dits.lingbotworld.cam_utils import prepare_camera_embedding
# from fastvideo.configs.sample import SamplingParam
OUTPUT_PATH = "video_samples_lingbotworld"
def main():
# FastVideo will automatically use the optimal default arguments for the
# model.
# If a local path is provided, FastVideo will make a best effort
# attempt to identify the optimal arguments.
generator = VideoGenerator.from_pretrained(
"FastVideo/LingBot-World-Base-Cam-Diffusers",
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
dit_cpu_offload=True, # DiT need to be offloaded for MoE
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
pin_cpu_memory=True,
# image_encoder_cpu_offload=False,
)
num_frames = 81
prompt = "The video presents a soaring journey through a fantasy jungle. The wind whips past the rider's blue hands gripping the reins, causing the leather straps to vibrate. The ancient gothic castle approaches steadily, its stone details becoming clearer against the backdrop of floating islands and distant waterfalls."
image_path = "https://raw.githubusercontent.com/Robbyant/lingbot-world/main/examples/00/image.jpg"
action_path = "examples/inference/basic/lingbotworld_examples/00"
c2ws_plucker_emb, num_frames = prepare_camera_embedding(
action_path=action_path,
num_frames=num_frames,
height=480,
width=832,
spatial_scale=8,
)
generator.generate_video(
prompt,
image_path=image_path,
output_path=OUTPUT_PATH,
save_video=True,
num_frames=num_frames,
height=480,
width=832,
c2ws_plucker_emb=c2ws_plucker_emb,
)
if __name__ == "__main__":
main()
+3 -8
View File
@@ -1,6 +1,5 @@
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 "
@@ -17,23 +16,19 @@ PROMPT = (
def main() -> None:
# Uses FastVideo default sampling settings for LTX2 base.
generator = VideoGenerator.from_pretrained(
"Davids048/LTX2-Base-Diffusers",
"FastVideo/LTX2-Distilled-Diffusers",
num_gpus=1,
)
output_path = "outputs_video/ltx2_basic/output_ltx2_base_t2v_1088_1920_1.1.mp4"
output_path = "outputs_video/ltx2_basic/output_ltx2_distilled_t2v.mp4"
generator.generate_video(
prompt=PROMPT,
output_path=output_path,
save_video=True,
num_frames=121,
height=1088,
width=1920,
)
generator.shutdown()
if __name__ == "__main__":
main()
main()
@@ -1,35 +0,0 @@
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."
)
import os
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "FLASH_ATTN"
def main() -> None:
generator = VideoGenerator.from_pretrained(
"FastVideo/LTX2-Distilled-Diffusers",
num_gpus=4,
)
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()
-137
View File
@@ -1,137 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import argparse
import os
import re
from typing import List
DEFAULT_PROMPTS = [
"a photo of a cat",
"a cinematic photo of a red panda wearing a tiny backpack, standing on a rainy neon-lit street at night, shallow depth of field, sharp focus, 35mm, bokeh",
]
def _safe_filename(text: str, max_len: int = 100) -> str:
"""
Make a stable, filesystem-friendly filename base.
VideoGenerator uses prompt[:100].strip() internally, so we mirror that,
but also remove path separators and other problematic characters.
"""
s = text[:max_len].strip()
s = s.replace(os.sep, "_")
if os.altsep:
s = s.replace(os.altsep, "_")
s = re.sub(r"\s+", " ", s)
s = re.sub(r"[^A-Za-z0-9 .,_-]", "_", s)
s = s.strip(" .")
return s or "prompt"
def _remove_existing_outputs(out_dir: str, filename_base: str) -> None:
"""
Ensure deterministic naming by deleting any existing outputs that would
cause VideoGenerator to append suffixes like _1, _2, etc.
"""
if not os.path.isdir(out_dir):
return
pattern = re.compile(rf"^{re.escape(filename_base)}(_\d+)?\.(mp4|png)$")
for fn in os.listdir(out_dir):
if pattern.match(fn):
try:
os.remove(os.path.join(out_dir, fn))
except FileNotFoundError:
pass
def parse_args() -> argparse.Namespace:
p = argparse.ArgumentParser(description="Run SD3.5 Medium text-to-image with FastVideo VideoGenerator.")
p.add_argument("--model-path", default="stabilityai/stable-diffusion-3.5-medium", help="Path to local diffusers-format SD3.5 weights directory.")
p.add_argument(
"--out-dir",
"--outdir",
default="outputs/sd35/samples",
help="Output directory for generated mp4 files.",
)
p.add_argument(
"--prompt",
action="append",
default=None,
help="Prompt text. Repeat --prompt multiple times to generate multiple samples.",
)
p.add_argument("--negative", default="lowres, blurry, jpeg artifacts, watermark, text", help="Negative prompt.")
p.add_argument(
"--backend",
default=None,
help="Set FASTVIDEO_ATTENTION_BACKEND (e.g. TORCH_SDPA). If omitted, respects the existing env var.",
)
p.add_argument("--seed", type=int, default=42, help="Base seed. Each prompt uses seed + prompt_idx.")
p.add_argument("--height", type=int, default=768, help="Output height.")
p.add_argument("--width", type=int, default=768, help="Output width.")
p.add_argument("--steps", type=int, default=28, help="Number of inference steps.")
p.add_argument("--guidance", type=float, default=6.0, help="Guidance scale.")
p.add_argument("--num-gpus", type=int, default=1, help="Number of GPUs to use.")
return p.parse_args()
def main() -> None:
args = parse_args()
prompts: List[str] = args.prompt if args.prompt else DEFAULT_PROMPTS
if args.backend:
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = args.backend
from fastvideo import VideoGenerator
os.makedirs(args.out_dir, exist_ok=True)
init_kwargs = {
"num_gpus": args.num_gpus,
"workload_type": "t2i",
"sp_size": 1,
"tp_size": 1,
"dit_cpu_offload": False,
"dit_layerwise_offload": False,
"text_encoder_cpu_offload": False,
"vae_cpu_offload": False,
"image_encoder_cpu_offload": False,
"pin_cpu_memory": False,
"use_fsdp_inference": False,
}
generator = VideoGenerator.from_pretrained(model_path=args.model_path, **init_kwargs)
try:
for i, prompt in enumerate(prompts):
seed = args.seed + i
filename_base = f"sd35_{i:02d}_seed{seed}_{_safe_filename(prompt, max_len=80)}"
_remove_existing_outputs(args.out_dir, filename_base)
output_path = os.path.join(args.out_dir, f"{filename_base}.png")
print(f"[sd35] prompt_idx={i} seed={seed} output_path={output_path}")
generation_kwargs = {
"output_path": output_path,
"height": args.height,
"width": args.width,
"num_frames": 1,
"fps": 1,
"num_inference_steps": args.steps,
"guidance_scale": args.guidance,
"seed": seed,
"negative_prompt": args.negative,
"save_video": True,
}
generator.generate_video(prompt, **generation_kwargs)
print(f"[sd35] done. outputs written to: {args.out_dir}")
finally:
generator.shutdown()
if __name__ == "__main__":
main()
@@ -31,7 +31,7 @@ def main():
sampling_param.height = 480
sampling_param.seed = 1000
with open("assets/prompts/mixkit_i2v.jsonl", "r") as f:
with open("prompts/mixkit_i2v.jsonl", "r") as f:
prompt_image_pairs = json.load(f)
for prompt_image_pair in prompt_image_pairs:
@@ -188,9 +188,8 @@ def load_example_prompts():
prompt_to_image = {}
# Try to find the JSON file relative to project root
possible_json_paths = [
Path("assets/prompts/mixkit_i2v.jsonl"),
Path(__file__).resolve().parents[4] / "assets" / "prompts" /
"mixkit_i2v.jsonl",
Path("prompts/mixkit_i2v.jsonl"),
Path(__file__).parent.parent.parent.parent / "prompts" / "mixkit_i2v.jsonl",
]
json_path = None
for path in possible_json_paths:
@@ -202,8 +201,8 @@ def load_example_prompts():
try:
with open(json_path, "r", encoding='utf-8') as f:
data = json.load(f)
# Resolve paths relative to repository root.
project_root = Path(__file__).resolve().parents[4]
# Get the project root directory (parent of prompts directory)
project_root = json_path.parent.parent
for item in data:
prompt_text = item.get("prompt", "").strip()
image_path = item.get("image_path", "")
@@ -737,8 +736,8 @@ def main():
allowed_paths=[
os.path.abspath("outputs"),
os.path.abspath("fastvideo-logos"),
os.path.abspath("assets/prompts"),
os.path.abspath("assets/images"),
os.path.abspath("prompts"),
os.path.abspath("images"),
os.path.abspath(tempfile.gettempdir()),
os.path.abspath(os.path.join(tempfile.gettempdir(), "gradio")),
]
@@ -748,4 +747,4 @@ def main():
if __name__ == "__main__":
main()
main()
@@ -5,7 +5,7 @@ export FASTVIDEO_ATTENTION_BACKEND=SLIDING_TILE_ATTN
export MODEL_BASE=Wan-AI/Wan2.1-T2V-14B-Diffusers
base_port=29503
num_gpu=8
num_gpu=1
gpu_ids=$(seq 0 $((num_gpu-1)))
skip_time_steps=12
@@ -44,8 +44,8 @@ VALIDATION_DATASET_FILE=your_validation_dataset_file
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Core training arguments
core_training_args=(
# Training arguments
training_args=(
--tracker_project_name wan_i2v_VSA
--output_dir "checkpoints/wan_i2v_finetune_VSA"
--max_train_steps 4000
@@ -53,14 +53,10 @@ core_training_args=(
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 21
--enable_gradient_checkpointing_type "full"
)
# Validation generation shape arguments (used during training validation)
validation_generation_shape_args=(
--num_height 720
--num_width 1280
--num_frames 81
--enable_gradient_checkpointing_type "full"
)
# Parallel arguments
@@ -131,9 +127,8 @@ srun torchrun \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${core_training_args[@]}" \
"${validation_generation_shape_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${vsa_args[@]}"
"${vsa_args[@]}"
@@ -6,7 +6,9 @@ These are e2e example scripts for finetuning Wan2.1 T2V with VSA to accelerate i
## Make sure you have installed VSA
Go to fastvideo-kernel/README.md for instructions.
```bash
pip install vsa
```
### Download the synthetic dataset:
@@ -44,8 +44,8 @@ VALIDATION_DATASET_FILE=your_validation_dataset_file
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Core training arguments
core_training_args=(
# Training arguments
training_args=(
--tracker_project_name wan_t2v_VSA
--output_dir "checkpoints/wan_t2v_finetune_VSA"
--max_train_steps 4000
@@ -53,14 +53,10 @@ core_training_args=(
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 21
# --enable_gradient_checkpointing_type "full" # if OOM enable this
)
# Validation generation shape arguments (used during training validation)
validation_generation_shape_args=(
--num_height 480
--num_width 832
--num_frames 81
# --enable_gradient_checkpointing_type "full" # if OOM enable this
)
# Parallel arguments
@@ -131,9 +127,8 @@ srun torchrun \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${core_training_args[@]}" \
"${validation_generation_shape_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${vsa_args[@]}"
"${vsa_args[@]}"
@@ -1,114 +0,0 @@
# Basic info
export WANDB_MODE="online"
export NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
export TRITON_CACHE_DIR="/tmp/triton_cache_${USER}_$$"
export MASTER_ADDR="${MASTER_ADDR:-127.0.0.1}"
export MASTER_PORT="${MASTER_PORT:-29500}"
export NODE_RANK=0
export TOKENIZERS_PARALLELISM=false
export WANDB_BASE_URL="https://api.wandb.ai"
export FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN
export WANDB_API_KEY="your_wandb_api_key_here" # TODO: Replace with your actual key or load from a secure location
# Configs
NUM_GPUS=4
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATA_DIR=data/Wan-Syn_77x448x832_600k
VALIDATION_DATASET_FILE=examples/training/finetune/Wan2.1-VSA/Wan-Syn-Data/validation_4.json
# Core training arguments
core_training_args=(
--tracker_project_name wan_t2v_VSA
--output_dir "checkpoints/wan_t2v_finetune_VSA"
--max_train_steps 4000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 20
--enable_gradient_checkpointing_type "full" # if OOM enable this
)
# Validation generation shape arguments (used during training validation)
validation_generation_shape_args=(
--num_height 448
--num_width 832
--num_frames 77
)
# Parallel arguments
parallel_args=(
--num_gpus $NUM_GPUS
--sp_size 1
--tp_size 1
--hsdp_replicate_dim $NUM_GPUS
--hsdp_shard_dim 1
)
# Model arguments
model_args=(
--model_path "$MODEL_PATH"
--pretrained_model_name_or_path "$MODEL_PATH"
)
# Dataset arguments
dataset_args=(
--data_path "$DATA_DIR"
--dataloader_num_workers 4
)
# Validation arguments
validation_args=(
--log_validation
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 200
--validation_sampling_steps "50"
--validation_guidance_scale "5.0"
)
# Optimizer arguments
optimizer_args=(
--learning_rate 1e-6
--mixed_precision "bf16"
--weight_only_checkpointing_steps 1000
--training_state_checkpointing_steps 1000
--weight_decay 0.01
--max_grad_norm 1.0
)
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
--training_cfg_rate 0.1
--dit_precision "fp32"
--ema_start_step 0
--flow_shift 1
--seed 1000
)
# VSA arguments
vsa_args=(
--VSA_decay_rate 0.03
--VSA_decay_interval_steps 50
--VSA_sparsity 0.9
)
torchrun \
--nnodes 1 \
--nproc_per_node "$NUM_GPUS" \
--node_rank "$NODE_RANK" \
--rdzv_backend c10d \
--rdzv_endpoint "$MASTER_ADDR:$MASTER_PORT" \
fastvideo/training/wan_training_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${core_training_args[@]}" \
"${validation_generation_shape_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${vsa_args[@]}"
@@ -44,8 +44,8 @@ VALIDATION_DATASET_FILE=your_validation_dataset_file
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Core training arguments
core_training_args=(
# Training arguments
training_args=(
--tracker_project_name wan_t2v_VSA
--output_dir "checkpoints/wan_t2v_finetune_VSA"
--max_train_steps 4000
@@ -53,14 +53,10 @@ core_training_args=(
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 21
--enable_gradient_checkpointing_type "full"
)
# Validation generation shape arguments (used during training validation)
validation_generation_shape_args=(
--num_height 720
--num_width 1280
--num_frames 81
--enable_gradient_checkpointing_type "full"
)
# Parallel arguments
@@ -131,9 +127,8 @@ srun torchrun \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${core_training_args[@]}" \
"${validation_generation_shape_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${vsa_args[@]}"
"${vsa_args[@]}"
@@ -1,8 +1,7 @@
#!/bin/bash
# Note: for debugging, you do not need to download them all and just control c.
# 480P dataset
python scripts/huggingface/download_hf.py --repo_id "FastVideo/Wan-Syn_77x448x832_600k" --local_dir "data/Wan-Syn_77x448x832_600k" --repo_type "dataset"
python scripts/huggingface/download_hf.py --repo_id "FastVideo/Wan-Syn_77x448x832_600k" --local_dir "FastVideo/Wan-Syn_77x448x832_600k" --repo_type "dataset"
# 720P dataset
python scripts/huggingface/download_hf.py --repo_id "FastVideo/Wan-Syn_77x768x1280_250k" --local_dir "data/Wan-Syn_77x768x1280_250k" --repo_type "dataset"
python scripts/huggingface/download_hf.py --repo_id "FastVideo/Wan-Syn_77x768x1280_250k" --local_dir "FastVideo/Wan-Syn_77x768x1280_250k" --repo_type "dataset"
@@ -1,36 +0,0 @@
{
"data": [
{
"caption": "In the video, a woman is elegantly showcasing her earrings, bringing attention to their intricate design with a gentle touch of her fingers. She is bathed in ambient purple and pink lighting, which casts a soft glow on her delicate features and enhances the vivid tones of her lipstick and eye makeup. Her hair is styled to frame her face smoothly, emphasizing the contours of her jawline and cheekbones. The background features a blurred neon light, adding an artistic and modern touch to the overall aesthetic.",
"video_path": "Fashion/mixkit-face-of-an-elegant-and-captivating-woman-41914_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "In the video, a lone rider guides a majestic horse across an expansive, open field as the sun sets in the background. The rider, dressed in a classic blue shirt and wide-brimmed hat, sits confidently in the saddle, silhouetted against the warm glow of the evening sky. The horse moves gracefully, its mane and tail flowing with each step, creating a sense of harmony between horse and rider. Surrounding the pair, towering trees form a natural border, their leaves gently rustling in the breeze. The shadows lengthen on the ground, accentuating the serene and timeless feel of the scene. The distant hills and wooden fences frame the horizon, adding depth to the tranquil landscape. A few horses graze peacefully in the background, blending into the pastoral setting. The overall ambiance evokes a sense of calmness and quietude, capturing a perfect moment in the golden light of dusk.",
"video_path": "Man/mixkit-a-rancher-riding-a-horse-at-sunset-1143_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "In a dimly lit, eerie setting, a mysterious pink bottle labeled \"Authentic 100% organic POISON\" sits prominently in the foreground, casting a menacing aura. The bottle is accentuated by green fog, which swirls lightly around it, enhancing its sinister allure. Behind it, a shadowy golden bottle adorned with a spider emblem subtly emerges, adding an extra layer of mystery to the scene. Dim candles provide faint, flickering light, which complements the dark atmosphere, making the setting ideal for an illusion of hidden dangers.",
"video_path": "smoke/mixkit-poison-in-halloween-ritual-33879_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
},
{
"caption": "The video opens with a tranquil scene in the heart of a dense forest, emphasizing two large, textured tree trunks in the foreground framing the view. Sunlight filters through the canopy above, casting intricate patterns of light and shadow on the trees and the ground. Between the tree trunks, a clear view of a calm, muddy river unfolds, its surface shimmering under the gentle sunlight. The riverbank is decorated with a variety of small bushes and vibrant foliage, subtly transitioning into the deep greens of tall, leafy plants. In the background, the dense forest looms, filled with dark, towering trees, their branches intertwining to form an intricate canopy. The scene is bathed in the soft glow of the sun, creating a serene and picturesque setting. Occasional sunbeams pierce through the foliage, adding a magical aura to the landscape. The vibrant reds and oranges of the smaller plants add contrast, bringing warmth to the earthy tones of the scenery. Overall, this harmonious blend of natural elements creates a peaceful and idyllic forest setting.",
"video_path": "forest/mixkit-view-of-a-river-between-two-old-trees-560_clip_1.mp4",
"num_inference_steps": 3,
"height": 448,
"width": 832,
"num_frames": 61
}
]
}
@@ -1,3 +0,0 @@
*.err
*.out
*slurm_logs/*
-25
View File
@@ -1,25 +0,0 @@
# 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/download_dataset.sh`
#### Or use the scripts at scripts/dataset_preparation to download and prepare the dataset
### Preprocess the videos and captions into latents:
`bash examples/training/finetune/ltx2/preprocess_ltx2_data_t2v_new.sh`
### Edit the following file and run finetuning:
`bash examples/training/finetune/ltx2/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`.
@@ -1,3 +0,0 @@
# #!/bin/bash
#
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
@@ -1,93 +0,0 @@
#!/bin/bash
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export TOKENIZERS_PARALLELISM=false
MODEL_PATH="FastVideo/LTX2-Distilled-Diffusers"
DATA_DIR="data/crush-smol"
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
echo VALIDATION_DATASET_FILE: $VALIDATION_DATASET_FILE
NUM_GPUS=4
HEIGHT=1088
WIDTH=1920
FRAMES=121
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 16
--num_height $HEIGHT
--num_width $WIDTH
--num_frames $FRAMES
--enable_gradient_checkpointing_type "full"
--mode "finetuning"
)
parallel_args=(
--num_gpus $NUM_GPUS
--sp_size 2
--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 20
--validation_sampling_steps "8"
--validation_guidance_scale "1.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
--ltx2-first-frame-conditioning-p 0.0
)
# NOTE: Setting this environment variable to TORCH_SDPA to avoid the issue of stacking that failed in flash attn.
torchrun \
--nnodes 1 \
--master_port 29501 \
--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[@]}"
@@ -1,80 +0,0 @@
#!/bin/bash
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export TOKENIZERS_PARALLELISM=false
MODEL_PATH="FastVideo/LTX2-Distilled-Diffusers"
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[@]}"
@@ -1,27 +0,0 @@
#!/bin/bash
GPU_NUM=4
MODEL_PATH="FastVideo/LTX2-Distilled-Diffusers"
DATASET_PATH="data/crush-smol"
OUTPUT_DIR="$DATASET_PATH"
WITH_AUDIO=true
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 1088 \
--preprocess.max_width 1920 \
--preprocess.num_frames 121 \
--preprocess.train_fps 24 \
--preprocess.video_length_tolerance_range 5
@@ -1,16 +0,0 @@
{
"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."
},
{
"caption": "A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press."
},
{
"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."
},
{
"caption": "A massive steel piston descends onto a stack of chocolate chip cookies, crushing them into crumbs as though they are being compressed by a hydraulic press."
}
]
}
+1 -1
View File
@@ -9,7 +9,7 @@ build-backend = "scikit_build_core.build"
[project]
name = "fastvideo-kernel"
version = "0.2.6"
version = "0.2.5"
description = "Unified CUDA kernels for FastVideo"
readme = "README.md"
requires-python = ">=3.10"
@@ -287,8 +287,9 @@ def block_sparse_attn(
block_sparse_fwd, block_sparse_bwd = _get_sm90_ops()
if (not _force_triton()) and _is_sm90() and (block_sparse_fwd is not None) and (block_sparse_bwd is not None):
return block_sparse_attn_sm90(q, k, v, block_map, variable_block_sizes)
# Triton path: supports q_seq_len != kv_seq_len as long as both are padded
# to a multiple of the block size (64 tokens).
# Triton path: generally assumes q/k/v share the same padded length
if q.shape[2] != k.shape[2] or q.shape[2] != v.shape[2]:
raise RuntimeError("Triton fallback requires q/k/v to have the same padded length.")
return block_sparse_attn_triton(q, k, v, block_map, variable_block_sizes)
@@ -141,6 +141,12 @@ def video_sparse_attn(
# Use autograd-enabled wrapper so backward works (and still uses SM90 kernel when available)
out_s = block_sparse_attn(q, k, v, mask, variable_block_sizes)[0]
else:
if q_seq_len != kv_seq_len:
raise RuntimeError(
"q/k have different lengths, but the compiled CUDA kernel (block_sparse_fwd) "
"is not available. The Triton fallback currently requires q and k/v to have "
"the same padded length."
)
# Triton-only forward (kept for environments without the wrapper deps)
out_s, _ = triton_block_sparse_attn_forward(q, k, v, idx, num, variable_block_sizes)
@@ -29,7 +29,7 @@ configs = [
# ──────────────────────────── SPARSE ADDITION BEGIN ───────────────────────────
@triton.autotune(configs, key=["N_CTX_Q", "HEAD_DIM"])
@triton.autotune(configs, key=["N_CTX", "HEAD_DIM"])
@triton.jit
def _attn_fwd_sparse(
Q,
@@ -60,8 +60,7 @@ def _attn_fwd_sparse(
stride_on,
Z,
H,
N_CTX_Q, #
N_CTX_KV, #
N_CTX, #
HEAD_DIM: tl.constexpr, #
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
@@ -76,29 +75,24 @@ def _attn_fwd_sparse(
off_hz = tl.program_id(1) # fused (batch, head)
b = off_hz // H
h = off_hz % H
q_tiles = N_CTX_Q // BLOCK_M
q_tiles = N_CTX // BLOCK_M
meta_base = ((b * H + h) * q_tiles + q_blk)
kv_blocks = tl.load(q2k_num + meta_base) # int32
kv_ptr = q2k_index + meta_base * max_kv_blks # ptr to list
# ----- base pointers -----
# Note: when q and kv have different sequence lengths, their per-(batch,head)
# strides differ, so we must compute separate base offsets.
q_off = (b.to(tl.int64) * stride_qz + h.to(tl.int64) * stride_qh)
k_off = (b.to(tl.int64) * stride_kz + h.to(tl.int64) * stride_kh)
v_off = (b.to(tl.int64) * stride_vz + h.to(tl.int64) * stride_vh)
o_off = (b.to(tl.int64) * stride_oz + h.to(tl.int64) * stride_oh)
qvk_off = (b.to(tl.int64) * stride_qz + h.to(tl.int64) * stride_qh)
Q_ptr = tl.make_block_ptr(base=Q + q_off,
shape=(N_CTX_Q, HEAD_DIM),
Q_ptr = tl.make_block_ptr(base=Q + qvk_off,
shape=(N_CTX, HEAD_DIM),
strides=(stride_qm, stride_qk),
offsets=(q_blk * BLOCK_M, 0),
block_shape=(BLOCK_M, HEAD_DIM),
order=(1, 0))
K_base = tl.make_block_ptr(base=K + k_off,
shape=(HEAD_DIM, N_CTX_KV),
K_base = tl.make_block_ptr(base=K + qvk_off,
shape=(HEAD_DIM, N_CTX),
strides=(stride_kk, stride_kn),
offsets=(0, 0),
block_shape=(HEAD_DIM, BLOCK_N),
@@ -106,15 +100,15 @@ def _attn_fwd_sparse(
v_order: tl.constexpr = (0, 1) if V.dtype.element_ty == tl.float8e5 else (1,
0)
V_base = tl.make_block_ptr(base=V + v_off,
shape=(N_CTX_KV, HEAD_DIM),
V_base = tl.make_block_ptr(base=V + qvk_off,
shape=(N_CTX, HEAD_DIM),
strides=(stride_vk, stride_vn),
offsets=(0, 0),
block_shape=(BLOCK_N, HEAD_DIM),
order=v_order)
O_ptr = tl.make_block_ptr(base=Out + o_off,
shape=(N_CTX_Q, HEAD_DIM),
O_ptr = tl.make_block_ptr(base=Out + qvk_off,
shape=(N_CTX, HEAD_DIM),
strides=(stride_om, stride_on),
offsets=(q_blk * BLOCK_M, 0),
block_shape=(BLOCK_M, HEAD_DIM),
@@ -156,7 +150,7 @@ def _attn_fwd_sparse(
# ----- epilogue -----
m_i += tl.math.log2(l_i)
acc = acc / l_i[:, None]
tl.store(M + off_hz * N_CTX_Q + offs_m, m_i)
tl.store(M + off_hz * N_CTX + offs_m, m_i)
tl.store(O_ptr, acc.to(Out.type.element_ty))
@@ -207,7 +201,7 @@ def _attn_bwd_dkdv(
stride_tok,
stride_d, #
H,
N_CTX_KV,
N_CTX,
BLOCK_M1: tl.constexpr, #
BLOCK_N1: tl.constexpr, #
HEAD_DIM: tl.constexpr, #
@@ -227,8 +221,8 @@ def _attn_bwd_dkdv(
off_hz = tl.program_id(2) # fused (batch, head)
b = off_hz // H
h = off_hz % H
kv_tiles = N_CTX_KV // BLOCK_N1
meta_base = ((b * H + h) * kv_tiles + kv_blk)
q_tiles = N_CTX // BLOCK_N1
meta_base = ((b * H + h) * q_tiles + kv_blk)
q_blocks = tl.load(k2q_num + meta_base) # int32
q_ptr = k2q_index + meta_base * max_q_blks # ptr to list
@@ -308,21 +302,16 @@ def _attn_bwd_dq(
kv_blocks = tl.load(q2k_num + meta_base) # int32
kv_ptr = q2k_index + meta_base * max_kv_blks # ptr to list
block_size = tl.load(variable_block_sizes + q_blk)
for blk_idx in range(kv_blocks * 2):
kv_idx = tl.load(kv_ptr + blk_idx // 2).to(tl.int32)
# variable_block_sizes is defined per KV block (tile). Mask must therefore
# use kv_idx (not q_blk). Also, because we split each 64-token block into
# two 32-token halves, the mask must account for the half-block offset.
block_size = tl.load(variable_block_sizes + kv_idx).to(tl.int32)
half = (blk_idx % 2).to(tl.int32)
block_sparse_offset = (kv_idx * 2 + half) * step_n * stride_tok
block_sparse_offset = (tl.load(kv_ptr + blk_idx // 2).to(tl.int32) * 2 +
blk_idx % 2) * step_n * stride_tok
kT = tl.load(kT_ptrs + block_sparse_offset)
vT = tl.load(vT_ptrs + block_sparse_offset)
qk = tl.dot(q, kT)
p = tl.math.exp2(qk - m)
offs_in_block = half * step_n + tl.arange(0, BLOCK_N2)
mask = offs_in_block < block_size
mask = tl.arange(0, BLOCK_N2) < block_size.to(tl.int32)
p = tl.where(mask[None, :], p, 0.0)
# Compute dP and dS.
dp = tl.dot(do, vT).to(tl.float32)
@@ -478,235 +467,19 @@ def _attn_bwd(
tl.store(dq_ptrs, dq)
@triton.jit
def _attn_bwd_dkdv_kernel(
Q,
K,
V,
sm_scale, #
DO, #
DK,
DV, #
M,
D,
k2q_index,
k2q_num,
max_q_blks,
variable_block_sizes,
# shared token/dim strides (assumed contiguous along token and dim)
stride_tok,
stride_d, #
# batch/head strides (may differ between Q and KV)
stride_qz,
stride_qh,
stride_kz,
stride_kh,
stride_vz,
stride_vh,
stride_doz,
stride_doh,
stride_dkz,
stride_dkh,
stride_dvz,
stride_dvh,
H,
N_CTX_Q,
N_CTX_KV,
BLOCK_M1: tl.constexpr, #
BLOCK_N1: tl.constexpr, #
HEAD_DIM: tl.constexpr):
"""
Backward kernel that computes dK and dV for each KV block (64 tokens).
Grid:
pid0: kv_blk in [0, N_CTX_KV/BLOCK_N1)
pid2: fused (batch, head) in [0, B*H)
"""
bhid = tl.program_id(2)
b = bhid // H
h = bhid % H
kv_blk = tl.program_id(0)
q_adj = (b.to(tl.int64) * stride_qz + h.to(tl.int64) * stride_qh)
kv_adj_k = (b.to(tl.int64) * stride_kz + h.to(tl.int64) * stride_kh)
kv_adj_v = (b.to(tl.int64) * stride_vz + h.to(tl.int64) * stride_vh)
do_adj = (b.to(tl.int64) * stride_doz + h.to(tl.int64) * stride_doh)
dk_adj = (b.to(tl.int64) * stride_dkz + h.to(tl.int64) * stride_dkh)
dv_adj = (b.to(tl.int64) * stride_dvz + h.to(tl.int64) * stride_dvh)
Q = Q + q_adj
K = K + kv_adj_k
V = V + kv_adj_v
DO = DO + do_adj
DK = DK + dk_adj
DV = DV + dv_adj
# M and D (delta) are always sized by Q length.
M = M + (bhid * N_CTX_Q).to(tl.int64)
D = D + (bhid * N_CTX_Q).to(tl.int64)
offs_k = tl.arange(0, HEAD_DIM)
start_n = kv_blk * BLOCK_N1
offs_n = start_n + tl.arange(0, BLOCK_N1)
# load K and V: they stay in SRAM throughout the inner loop.
k = tl.load(K + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d)
v = tl.load(V + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d)
dv_acc = tl.zeros([BLOCK_N1, HEAD_DIM], dtype=tl.float32)
dk_acc = tl.zeros([BLOCK_N1, HEAD_DIM], dtype=tl.float32)
num_steps = N_CTX_Q // BLOCK_M1
dk_acc, dv_acc = _attn_bwd_dkdv(
dk_acc,
dv_acc,
Q,
k,
v,
sm_scale,
DO,
M,
D,
k2q_index,
k2q_num,
max_q_blks,
variable_block_sizes,
stride_tok,
stride_d,
H,
N_CTX_KV,
BLOCK_M1=BLOCK_M1,
BLOCK_N1=BLOCK_N1,
HEAD_DIM=HEAD_DIM,
start_n=start_n,
start_m=0,
num_steps=num_steps,
)
dv_ptrs = DV + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d
tl.store(dv_ptrs, dv_acc)
dk_acc *= sm_scale
dk_ptrs = DK + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d
tl.store(dk_ptrs, dk_acc)
@triton.jit
def _attn_bwd_dq_kernel(
Q,
K,
V,
DO, #
DQ,
M,
D,
q2k_index,
q2k_num,
max_kv_blks,
variable_block_sizes,
# shared token/dim strides (assumed contiguous along token and dim)
stride_tok,
stride_d, #
# batch/head strides (may differ between Q and KV)
stride_qz,
stride_qh,
stride_kz,
stride_kh,
stride_vz,
stride_vh,
stride_doz,
stride_doh,
stride_dqz,
stride_dqh,
H,
N_CTX_Q,
BLOCK_M2: tl.constexpr, #
BLOCK_N2: tl.constexpr, #
HEAD_DIM: tl.constexpr):
"""
Backward kernel that computes dQ for each Q block (64 tokens).
Grid:
pid0: q_blk in [0, N_CTX_Q/BLOCK_M2)
pid2: fused (batch, head) in [0, B*H)
"""
LN2 = 0.6931471824645996 # = ln(2)
bhid = tl.program_id(2)
b = bhid // H
h = bhid % H
q_blk = tl.program_id(0)
q_adj = (b.to(tl.int64) * stride_qz + h.to(tl.int64) * stride_qh)
kv_adj_k = (b.to(tl.int64) * stride_kz + h.to(tl.int64) * stride_kh)
kv_adj_v = (b.to(tl.int64) * stride_vz + h.to(tl.int64) * stride_vh)
do_adj = (b.to(tl.int64) * stride_doz + h.to(tl.int64) * stride_doh)
dq_adj = (b.to(tl.int64) * stride_dqz + h.to(tl.int64) * stride_dqh)
Q = Q + q_adj
K = K + kv_adj_k
V = V + kv_adj_v
DO = DO + do_adj
DQ = DQ + dq_adj
M = M + (bhid * N_CTX_Q).to(tl.int64)
D = D + (bhid * N_CTX_Q).to(tl.int64)
offs_k = tl.arange(0, HEAD_DIM)
start_m = q_blk * BLOCK_M2
offs_m = start_m + tl.arange(0, BLOCK_M2)
q = tl.load(Q + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d)
do = tl.load(DO + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d)
m = tl.load(M + offs_m)[:, None]
dq_acc = tl.zeros([BLOCK_M2, HEAD_DIM], dtype=tl.float32)
num_steps = 0 # unused in _attn_bwd_dq
dq_acc = _attn_bwd_dq(
dq_acc,
q,
K,
V,
do,
m,
D,
q2k_index,
q2k_num,
max_kv_blks,
variable_block_sizes,
stride_tok,
stride_d,
H,
N_CTX_Q,
BLOCK_M2=BLOCK_M2,
BLOCK_N2=BLOCK_N2,
HEAD_DIM=HEAD_DIM,
start_m=start_m,
start_n=0,
num_steps=num_steps,
)
dq_ptrs = DQ + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d
dq_acc *= LN2
tl.store(dq_ptrs, dq_acc)
# ──────────────────────────── SPARSE ADDITION BEGIN ───────────────────────────
def triton_block_sparse_attn_forward(q, k, v, q2k_index, q2k_num,
variable_block_sizes):
B, H, Tq, D = q.shape
Tkv = k.shape[2]
B, H, T, D = q.shape
sm_scale = 1.0 / math.sqrt(D)
max_kv_blks = q2k_index.shape[-1]
assert Tq % 64 == 0, f"q length must be a multiple of 64, but got {Tq}"
assert Tkv % 64 == 0, f"kv length must be a multiple of 64, but got {Tkv}"
assert T % 64 == 0, f"T must be a multiple of 64, but got {T}"
assert q2k_num.shape[
-1] == Tq // 64, f"shape mismatch, Tq // 64 = {Tq // 64}, q2k_num.shape[-2] = {q2k_num.shape[-2]}"
assert variable_block_sizes.numel() == Tkv // 64, (
f"shape mismatch, variable_block_sizes must have length {Tkv // 64}, "
f"got {variable_block_sizes.numel()}"
)
-1] == T // 64, f"shape mismatch, T // 64 = {T // 64}, q2k_num.shape[-2] = {q2k_num.shape[-2]}"
o = torch.empty_like(q)
M = torch.empty((B, H, Tq), dtype=torch.float32, device=q.device)
M = torch.empty((B, H, T), dtype=torch.float32, device=q.device)
grid = lambda _: (triton.cdiv(Tq, 64), B * H, 1)
grid = lambda _: (triton.cdiv(T, 64), B * H, 1)
_attn_fwd_sparse[grid](q,
k,
v,
@@ -735,8 +508,7 @@ def triton_block_sparse_attn_forward(q, k, v, q2k_index, q2k_num,
o.stride(3),
B,
H,
Tq,
Tkv,
T,
HEAD_DIM=D,
STAGE=3)
@@ -746,21 +518,21 @@ def triton_block_sparse_attn_forward(q, k, v, q2k_index, q2k_num,
def triton_block_sparse_attn_backward(do, q, k, v, o, M, q2k_index, q2k_num,
k2q_index, k2q_num, variable_block_sizes):
assert do.is_contiguous()
assert q.stride() == k.stride() == v.stride() == o.stride() == do.stride()
B, H, Tq, D = q.shape
Tkv = k.shape[2]
B, H, T, D = q.shape
sm_scale = 1.0 / math.sqrt(D)
dq = torch.empty_like(q)
dk = torch.empty_like(k)
dv = torch.empty_like(v)
BATCH, N_HEAD = q.shape[:2]
BATCH, N_HEAD, N_CTX = q.shape[:3]
BLOCK_M1, BLOCK_N1, BLOCK_M2, BLOCK_N2 = 32, 64, 64, 32
RCP_LN2 = 1.4426950408889634 # = 1.0 / ln(2)
arg_k = k
arg_k = arg_k * (sm_scale * RCP_LN2)
PRE_BLOCK = 64
assert Tq % PRE_BLOCK == 0
pre_grid = (Tq // PRE_BLOCK, BATCH * N_HEAD)
assert N_CTX % PRE_BLOCK == 0
pre_grid = (N_CTX // PRE_BLOCK, BATCH * N_HEAD)
delta = torch.empty_like(M)
_attn_bwd_preprocess[pre_grid](
o,
@@ -768,7 +540,7 @@ def triton_block_sparse_attn_backward(do, q, k, v, o, M, q2k_index, q2k_num,
delta, #
BATCH,
N_HEAD,
Tq, #
N_CTX, #
BLOCK_M=PRE_BLOCK,
HEAD_DIM=D #
)
@@ -776,75 +548,36 @@ def triton_block_sparse_attn_backward(do, q, k, v, o, M, q2k_index, q2k_num,
max_q_blks = k2q_index.shape[-1]
max_kv_blks = q2k_index.shape[-1]
# dK/dV kernel: grid over KV blocks
grid_kv = (Tkv // BLOCK_N1, 1, BATCH * N_HEAD)
_attn_bwd_dkdv_kernel[grid_kv](
grid = (N_CTX // BLOCK_N1, 1, BATCH * N_HEAD)
_attn_bwd[grid](
q,
arg_k,
v,
sm_scale,
do,
dq,
dk,
dv,
dv, #
M,
delta,
delta, #
q2k_index,
q2k_num,
max_kv_blks,
k2q_index,
k2q_num,
max_q_blks,
variable_block_sizes,
q.stride(2),
q.stride(3),
q.stride(0),
q.stride(1),
arg_k.stride(0),
arg_k.stride(1),
v.stride(0),
v.stride(1),
do.stride(0),
do.stride(1),
dk.stride(0),
dk.stride(1),
dv.stride(0),
dv.stride(1),
q.stride(2),
q.stride(3), #
N_HEAD,
Tq,
Tkv,
N_CTX, #
BLOCK_M1=BLOCK_M1,
BLOCK_N1=BLOCK_N1,
HEAD_DIM=D,
)
# dQ kernel: grid over Q blocks
grid_q = (Tq // BLOCK_M2, 1, BATCH * N_HEAD)
_attn_bwd_dq_kernel[grid_q](
q,
arg_k,
v,
do,
dq,
M,
delta,
q2k_index,
q2k_num,
max_kv_blks,
variable_block_sizes,
q.stride(2),
q.stride(3),
q.stride(0),
q.stride(1),
arg_k.stride(0),
arg_k.stride(1),
v.stride(0),
v.stride(1),
do.stride(0),
do.stride(1),
dq.stride(0),
dq.stride(1),
N_HEAD,
Tq,
BLOCK_N1=BLOCK_N1, #
BLOCK_M2=BLOCK_M2,
BLOCK_N2=BLOCK_N2,
HEAD_DIM=D,
BLOCK_N2=BLOCK_N2, #
HEAD_DIM=D #
)
return dq, dk, dv
@@ -1 +1 @@
__version__ = "0.2.6"
__version__ = "0.2.5"
@@ -350,7 +350,7 @@ def select_best_mask_strategy(
def save_mask_search_results(
mask_search_final_result: list[Any],
mask_search_final_result: list[dict[str, list[float]]],
prompt: str,
mask_strategies: list[str],
output_dir: str = 'output/mask_search_result/') -> str | None:
@@ -358,9 +358,8 @@ def save_mask_search_results(
print("No mask search results to save")
return None
# Create result dictionary with nested lists:
# [timesteps][layers][heads].
mask_search_dict: dict[str, dict[str, list[list[list[float]]]]] = {
# Create result dictionary with defaultdict for nested lists
mask_search_dict: dict[str, dict[str, list[list[float]]]] = {
"L2_loss": defaultdict(list),
"L1_loss": defaultdict(list)
}
@@ -372,79 +371,23 @@ def save_mask_search_results(
masks_list = [int(x) for x in mask.split(',')]
selected_masks.append(masks_list)
def _to_float_list(loss_values: Any, loss_name: str) -> list[float]:
if isinstance(loss_values, np.ndarray):
loss_values = loss_values.tolist()
if not isinstance(loss_values, list | tuple):
raise ValueError(
f"{loss_name} must be a sequence of numeric values")
float_values = []
for loss in loss_values:
try:
float_values.append(float(loss))
except (TypeError, ValueError) as exc:
raise ValueError(
f"Invalid {loss_name} value {loss!r}: expected a number"
) from exc
return float_values
def _extract_timestep_layer_losses(step_data: Any, loss_name: str,
strategy_idx: int) -> list[list[float]]:
if isinstance(step_data, dict):
layer_data_list = [step_data]
elif isinstance(step_data, list):
layer_data_list = step_data
else:
return []
timestep_layer_losses: list[list[float]] = []
for layer_data in layer_data_list:
if not isinstance(layer_data, dict) or loss_name not in layer_data:
continue
raw_losses = layer_data[loss_name]
if isinstance(raw_losses, np.ndarray):
raw_losses = raw_losses.tolist()
if not isinstance(raw_losses, list | tuple):
raise ValueError(f"{loss_name} must be a list or tuple")
raw_losses = list(raw_losses)
if raw_losses and isinstance(raw_losses[0], list | tuple
| np.ndarray):
if strategy_idx >= len(raw_losses):
raise ValueError(
f"Missing strategy index {strategy_idx} in {loss_name}")
strategy_losses = raw_losses[strategy_idx]
else:
if strategy_idx > 0:
continue
strategy_losses = raw_losses
timestep_layer_losses.append(
_to_float_list(strategy_losses, loss_name))
return timestep_layer_losses
# Process each mask strategy
for i, mask_strategy in enumerate(selected_masks):
mask_strategy_str = str(mask_strategy)
l2_step_results: list[list[list[float]]] = []
l1_step_results: list[list[list[float]]] = []
# Process L2 loss
step_results: list[list[float]] = []
for step_data in mask_search_final_result:
l2_layer_losses = _extract_timestep_layer_losses(
step_data, "L2_loss", i)
if l2_layer_losses:
l2_step_results.append(l2_layer_losses)
if isinstance(step_data, dict) and "L2_loss" in step_data:
layer_losses = [float(loss) for loss in step_data["L2_loss"]]
step_results.append(layer_losses)
mask_search_dict["L2_loss"][mask_strategy_str] = step_results
l1_layer_losses = _extract_timestep_layer_losses(
step_data, "L1_loss", i)
if l1_layer_losses:
l1_step_results.append(l1_layer_losses)
mask_search_dict["L2_loss"][mask_strategy_str] = l2_step_results
mask_search_dict["L1_loss"][mask_strategy_str] = l1_step_results
step_results = []
for step_data in mask_search_final_result:
if isinstance(step_data, dict) and "L1_loss" in step_data:
layer_losses = [float(loss) for loss in step_data["L1_loss"]]
step_results.append(layer_losses)
mask_search_dict["L1_loss"][mask_strategy_str] = step_results
# Create the output directory if it doesn't exist
os.makedirs(output_dir, exist_ok=True)
+1 -49
View File
@@ -92,59 +92,11 @@ class FlashAttentionImpl(AttentionImpl):
value: torch.Tensor,
attn_metadata: FlashAttnMetadata,
):
def _key_padding_mask_from_attn_mask(attn_mask: torch.Tensor,
key_len: int) -> torch.Tensor:
# Normalize attn_mask to [B, key_len] where True means valid token.
if attn_mask.dim() == 4:
attn_mask = attn_mask[:, 0, 0, :]
elif attn_mask.dim() == 3:
attn_mask = attn_mask[:, 0, :]
elif attn_mask.dim() != 2:
raise ValueError(
f"Unsupported attn_mask shape for FLASH_ATTN: {attn_mask.shape}"
)
if attn_mask.dtype == torch.bool:
key_padding_mask = attn_mask
else:
# SDPA additive mask convention: valid=0, masked=-inf/large negative.
key_padding_mask = attn_mask >= 0
if key_padding_mask.shape[-1] != key_len:
raise ValueError(
"Invalid key padding mask length for FLASH_ATTN: "
f"expected {key_len}, got {key_padding_mask.shape[-1]}")
return key_padding_mask
if attn_metadata is not None and hasattr(
attn_metadata,
"attn_mask") and attn_metadata.attn_mask is not None:
from fastvideo.attention.utils.flash_attn_no_pad import (
flash_attn_no_pad, flash_attn_varlen_qk_no_pad)
from fastvideo.attention.utils.flash_attn_no_pad import flash_attn_no_pad
attn_mask = attn_metadata.attn_mask
# flash_attn_no_pad packs q/k/v as one tensor and assumes equal q/k
# sequence lengths. Cross-attention can violate this.
if query.shape[1] != key.shape[1]:
query_padding_mask = torch.ones(
(query.shape[0], query.shape[1]),
dtype=torch.bool,
device=query.device,
)
key_padding_mask = _key_padding_mask_from_attn_mask(
attn_mask, key.shape[1]).to(device=key.device)
return flash_attn_varlen_qk_no_pad(
query,
key,
value,
query_padding_mask=query_padding_mask,
key_padding_mask=key_padding_mask,
causal=self.causal,
dropout_p=0.0,
softmax_scale=self.softmax_scale,
)
qkv = torch.stack([query, key, value], dim=2)
attn_mask = F.pad(attn_mask, (qkv.shape[1] - attn_mask.shape[1], 0),
@@ -97,60 +97,3 @@ def flash_attn_no_pad_v3(qkv,
"b s (h d) -> b s h d",
h=nheads)
return output
def flash_attn_varlen_qk_no_pad(
query,
key,
value,
query_padding_mask,
key_padding_mask,
causal=False,
dropout_p=0.0,
softmax_scale=None,
deterministic=False,
):
from flash_attn.bert_padding import pad_input, unpad_input
try:
from flash_attn_interface import flash_attn_varlen_func as flash_attn_varlen_func_impl
except ImportError:
from flash_attn import flash_attn_varlen_func as flash_attn_varlen_func_impl
if flash_attn_varlen_func_impl is None:
raise ImportError("FlashAttention varlen backend not available")
batch_size, q_seqlen, nheads, _ = query.shape
query_unpad, q_indices, cu_seqlens_q, max_seqlen_q, _ = unpad_input(
rearrange(query, "b s h d -> b s (h d)"), query_padding_mask)
key_unpad, _, cu_seqlens_k, max_seqlen_k, _ = unpad_input(
rearrange(key, "b s h d -> b s (h d)"), key_padding_mask)
value_unpad, _, _, _, _ = unpad_input(
rearrange(value, "b s h d -> b s (h d)"), key_padding_mask)
query_unpad = rearrange(query_unpad, "nnz (h d) -> nnz h d", h=nheads)
key_unpad = rearrange(key_unpad, "nnz (h d) -> nnz h d", h=nheads)
value_unpad = rearrange(value_unpad, "nnz (h d) -> nnz h d", h=nheads)
output_unpad = flash_attn_varlen_func_impl(
query_unpad,
key_unpad,
value_unpad,
cu_seqlens_q,
cu_seqlens_k,
max_seqlen_q,
max_seqlen_k,
dropout_p=dropout_p,
softmax_scale=softmax_scale,
causal=causal,
deterministic=deterministic,
)
output = rearrange(
pad_input(rearrange(output_unpad, "nnz h d -> nnz (h d)"), q_indices,
batch_size, q_seqlen),
"b s (h d) -> b s h d",
h=nheads,
)
return output
-5
View File
@@ -86,7 +86,6 @@ class PreprocessConfig:
# Model configuration
training_cfg_rate: float = 0.0
with_audio: bool = False
# framework configuration
seed: int = 42
@@ -191,10 +190,6 @@ 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,
+3 -3
View File
@@ -1,15 +1,15 @@
from fastvideo.configs.models.dits.cosmos import CosmosVideoConfig
from fastvideo.configs.models.dits.cosmos2_5 import Cosmos25VideoConfig
from fastvideo.configs.models.dits.hunyuangamecraft import HunyuanGameCraftConfig
from fastvideo.configs.models.dits.hunyuanvideo import HunyuanVideoConfig
from fastvideo.configs.models.dits.hunyuanvideo15 import HunyuanVideo15Config
from fastvideo.configs.models.dits.longcat import LongCatVideoConfig
from fastvideo.configs.models.dits.ltx2 import LTX2VideoConfig
from fastvideo.configs.models.dits.stepvideo import StepVideoConfig
from fastvideo.configs.models.dits.wanvideo import WanVideoConfig
from fastvideo.configs.models.dits.hyworld import HYWorldConfig
__all__ = [
"HunyuanVideoConfig", "HunyuanVideo15Config", "HunyuanGameCraftConfig",
"WanVideoConfig", "CosmosVideoConfig", "Cosmos25VideoConfig",
"HunyuanVideoConfig", "HunyuanVideo15Config", "WanVideoConfig",
"StepVideoConfig", "CosmosVideoConfig", "Cosmos25VideoConfig",
"LongCatVideoConfig", "LTX2VideoConfig", "HYWorldConfig"
]
@@ -1,163 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""
Configuration for HunyuanGameCraft transformer model.
HunyuanGameCraft extends HunyuanVideo with:
1. CameraNet for camera/action conditioning
2. 33 input channels (16 latent + 16 gt_latent + 1 mask)
3. Mask-based conditioning for autoregressive generation
"""
from dataclasses import dataclass, field
import torch
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
def is_double_block(n: str, m) -> bool:
return "double" in n and str.isdigit(n.split(".")[-1])
def is_single_block(n: str, m) -> bool:
return "single" in n and str.isdigit(n.split(".")[-1])
def is_refiner_block(n: str, m) -> bool:
return "refiner" in n and str.isdigit(n.split(".")[-1])
def is_txt_in(n: str, m) -> bool:
return n.split(".")[-1] == "txt_in"
def is_camera_net(n: str, m) -> bool:
return "camera_net" in n
@dataclass
class HunyuanGameCraftArchConfig(DiTArchConfig):
"""Architecture config for HunyuanGameCraft transformer."""
# Version field for compatibility with saved config.json
_fastvideo_version: str = "0.1.0"
# Camera net flag (for config.json compatibility)
camera_net: bool = True
_fsdp_shard_conditions: list = field(
default_factory=lambda:
[is_double_block, is_single_block, is_refiner_block, is_camera_net])
_compile_conditions: list = field(
default_factory=lambda: [is_double_block, is_single_block, is_txt_in])
# Parameter names mapping from official checkpoint to FastVideo naming
# GameCraft weights are already close to FastVideo format with minor adjustments
param_names_mapping: dict = field(
default_factory=lambda: {
# MLP naming: fc1 -> fc_in, fc2 -> fc_out
r"^(.*)\.img_mlp\.fc1\.(.*)$":
r"\1.img_mlp.fc_in.\2",
r"^(.*)\.img_mlp\.fc2\.(.*)$":
r"\1.img_mlp.fc_out.\2",
r"^(.*)\.txt_mlp\.fc1\.(.*)$":
r"\1.txt_mlp.fc_in.\2",
r"^(.*)\.txt_mlp\.fc2\.(.*)$":
r"\1.txt_mlp.fc_out.\2",
# Single block MLP naming
r"^single_blocks\.(\d+)\.mlp\.fc1\.(.*)$":
r"single_blocks.\1.mlp.fc_in.\2",
r"^single_blocks\.(\d+)\.mlp\.fc2\.(.*)$":
r"single_blocks.\1.mlp.fc_out.\2",
# Token refiner naming
r"^txt_in\.individual_token_refiner\.blocks\.(\d+)\.(.*)$":
r"txt_in.refiner_blocks.\1.\2",
# Vector in naming
r"^vector_in\.in_layer\.(.*)$":
r"vector_in.fc_in.\1",
r"^vector_in\.out_layer\.(.*)$":
r"vector_in.fc_out.\1",
# Time embedder naming
r"^time_in\.mlp\.0\.(.*)$":
r"time_in.mlp.fc_in.\1",
r"^time_in\.mlp\.2\.(.*)$":
r"time_in.mlp.fc_out.\1",
# Guidance embedder naming (if present)
r"^guidance_in\.mlp\.0\.(.*)$":
r"guidance_in.mlp.fc_in.\1",
r"^guidance_in\.mlp\.2\.(.*)$":
r"guidance_in.mlp.fc_out.\1",
# Final layer adaLN modulation
r"^final_layer\.adaLN_modulation\.1\.(.*)$":
r"final_layer.adaLN_modulation.linear.\1",
# Refiner block MLP naming
r"^txt_in\.refiner_blocks\.(\d+)\.mlp\.fc1\.(.*)$":
r"txt_in.refiner_blocks.\1.mlp.fc_in.\2",
r"^txt_in\.refiner_blocks\.(\d+)\.mlp\.fc2\.(.*)$":
r"txt_in.refiner_blocks.\1.mlp.fc_out.\2",
# Camera net weights are already correctly named
})
# Reverse mapping for saving checkpoints
reverse_param_names_mapping: dict = field(default_factory=lambda: {})
# Model architecture parameters
# patch_size can be int or tuple - if tuple, it's [T, H, W]
patch_size: int | tuple[int, int, int] = 2
patch_size_t: int = 1
in_channels: int = 33 # 16 latent + 16 gt_latent + 1 mask
out_channels: int = 16
num_attention_heads: int = 24
attention_head_dim: int = 128
mlp_ratio: float = 4.0
num_layers: int = 20 # Double stream blocks
num_single_layers: int = 40 # Single stream blocks
num_refiner_layers: int = 2
rope_axes_dim: tuple[int, int, int] = (16, 56, 56)
guidance_embeds: bool = False # GameCraft doesn't use guidance
dtype: torch.dtype | None = None
text_embed_dim: int = 4096 # LLaMA-3 hidden size
pooled_projection_dim: int = 768 # CLIP pooled output dim
rope_theta: int = 256
qk_norm: str = "rms_norm"
# Camera net parameters
camera_in_channels: int = 6 # Plücker coordinates
camera_downscale_coef: int = 8
camera_out_channels: int = 16
# Layers to exclude from LoRA
exclude_lora_layers: list[str] = field(
default_factory=lambda:
["img_in", "txt_in", "time_in", "vector_in", "camera_net"])
def __post_init__(self):
super().__post_init__()
self.hidden_size: int = self.attention_head_dim * self.num_attention_heads
self.num_channels_latents: int = 16 # Output is 16 channels
# Convert patch_size list to tuple if needed (from JSON deserialization)
if isinstance(self.patch_size, list):
self.patch_size = tuple(self.patch_size)
# Convert rope_axes_dim list to tuple if needed
if isinstance(self.rope_axes_dim, list):
self.rope_axes_dim = tuple(self.rope_axes_dim)
@dataclass
class HunyuanGameCraftConfig(DiTConfig):
"""Full config for HunyuanGameCraft model."""
arch_config: DiTArchConfig = field(
default_factory=HunyuanGameCraftArchConfig)
prefix: str = "HunyuanGameCraft"
@@ -1,110 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
def is_blocks(n: str, m) -> bool:
return "blocks" in n and str.isdigit(n.split(".")[-1])
@dataclass
class LingBotWorldArchConfig(DiTArchConfig):
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_blocks])
param_names_mapping: dict = field(
default_factory=lambda: {
r"^patch_embedding\.(.*)$": r"patch_embedding.proj.\1",
r"^patch_embedding_wancamctrl\.(.*)$":
r"patch_embedding_wancamctrl.proj.\1",
r"^c2ws_hidden_states_layer1\.(.*)$": r"c2ws_mlp.fc_in.\1",
r"^c2ws_hidden_states_layer2\.(.*)$": r"c2ws_mlp.fc_out.\1",
r"^text_embedding\.0\.(.*)$":
r"condition_embedder.text_embedder.fc_in.\1",
r"^text_embedding\.2\.(.*)$":
r"condition_embedder.text_embedder.fc_out.\1",
r"^time_embedding\.0\.(.*)$":
r"condition_embedder.time_embedder.mlp.fc_in.\1",
r"^time_embedding\.2\.(.*)$":
r"condition_embedder.time_embedder.mlp.fc_out.\1",
r"^time_projection\.1\.(.*)$":
r"condition_embedder.time_modulation.linear.\1",
r"^blocks\.(\d+)\.modulation$": r"blocks.\1.scale_shift_table",
r"^blocks\.(\d+)\.self_attn\.q\.(.*)$": r"blocks.\1.to_q.\2",
r"^blocks\.(\d+)\.self_attn\.k\.(.*)$": r"blocks.\1.to_k.\2",
r"^blocks\.(\d+)\.self_attn\.v\.(.*)$": r"blocks.\1.to_v.\2",
r"^blocks\.(\d+)\.self_attn\.o\.(.*)$": r"blocks.\1.to_out.\2",
r"^blocks\.(\d+)\.self_attn\.norm_q\.(.*)$": r"blocks.\1.norm_q.\2",
r"^blocks\.(\d+)\.self_attn\.norm_k\.(.*)$": r"blocks.\1.norm_k.\2",
r"^blocks\.(\d+)\.norm3\.(.*)$":
r"blocks.\1.self_attn_residual_norm.norm.\2",
r"^blocks\.(\d+)\.cross_attn\.q\.(.*)$": r"blocks.\1.attn2.to_q.\2",
r"^blocks\.(\d+)\.cross_attn\.k\.(.*)$": r"blocks.\1.attn2.to_k.\2",
r"^blocks\.(\d+)\.cross_attn\.v\.(.*)$": r"blocks.\1.attn2.to_v.\2",
r"^blocks\.(\d+)\.cross_attn\.o\.(.*)$":
r"blocks.\1.attn2.to_out.\2",
r"^blocks\.(\d+)\.cross_attn\.norm_q\.(.*)$":
r"blocks.\1.attn2.norm_q.\2",
r"^blocks\.(\d+)\.cross_attn\.norm_k\.(.*)$":
r"blocks.\1.attn2.norm_k.\2",
r"^blocks\.(\d+)\.ffn\.0\.(.*)$": r"blocks.\1.ffn.fc_in.\2",
r"^blocks\.(\d+)\.ffn\.2\.(.*)$": r"blocks.\1.ffn.fc_out.\2",
r"^blocks\.(\d+)\.cam_injector_layer1\.(.*)$":
r"blocks.\1.cam_conditioner.cam_injector.fc_in.\2",
r"^blocks\.(\d+)\.cam_injector_layer2\.(.*)$":
r"blocks.\1.cam_conditioner.cam_injector.fc_out.\2",
r"^blocks\.(\d+)\.cam_scale_layer\.(.*)$":
r"blocks.\1.cam_conditioner.cam_scale_layer.\2",
r"^blocks\.(\d+)\.cam_shift_layer\.(.*)$":
r"blocks.\1.cam_conditioner.cam_shift_layer.\2",
r"^head\.modulation$": r"scale_shift_table",
r"^head\.head\.(.*)$": r"proj_out.\1",
})
# Reverse mapping for saving checkpoints: custom -> hf
reverse_param_names_mapping: dict = field(default_factory=lambda: {})
# Some LoRA adapters use the original official layer names instead of hf layer names,
# so apply this before the param_names_mapping
lora_param_names_mapping: dict = field(default_factory=lambda: {})
patch_size: tuple[int, int, int] = (1, 2, 2)
text_len: int = 512
num_attention_heads: int = 40
attention_head_dim: int = 128
in_channels: int = 16
out_channels: int = 16
text_dim: int = 4096
freq_dim: int = 256
ffn_dim: int = 13824
num_layers: int = 40
cross_attn_norm: bool = True
qk_norm: str = "rms_norm_across_heads"
eps: float = 1e-6
image_dim: int | None = None
added_kv_proj_dim: int | None = None
rope_max_seq_len: int = 1024
pos_embed_seq_len: int | None = None
exclude_lora_layers: list[str] = field(default_factory=lambda: ["embedder"])
# Wan MoE
boundary_ratio: float | None = None
# Causal Wan
local_attn_size: int = -1 # Window size for temporal local attention (-1 indicates global attention)
sink_size: int = 0 # Size of the attention sink, we keep the first `sink_size` frames unchanged when rolling the KV cache
num_frames_per_block: int = 3
sliding_window_num_frames: int = 21
def __post_init__(self):
super().__post_init__()
self.out_channels = self.out_channels or self.in_channels
self.hidden_size = self.num_attention_heads * self.attention_head_dim
self.num_channels_latents = self.out_channels
@dataclass
class LingBotWorldVideoConfig(DiTConfig):
arch_config: DiTArchConfig = field(default_factory=LingBotWorldArchConfig)
prefix: str = "Wan"
+2 -4
View File
@@ -7,12 +7,10 @@ from dataclasses import dataclass, field
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
import re
def is_ltx2_blocks(name: str, _module) -> bool:
res = re.search(r"(?:^|\.)transformer_blocks\.\d+$", name) is not None
return res
"""FSDP shard condition for LTX-2 transformer blocks."""
return "transformer_blocks" in name
@dataclass
-31
View File
@@ -1,31 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
@dataclass
class SD3Transformer2DArchConfig(DiTArchConfig):
# Diffusers SD3Transformer2DModel config fields.
sample_size: int = 128
patch_size: int = 2
num_layers: int = 24
attention_head_dim: int = 64
joint_attention_dim: int = 4096
caption_projection_dim: int = 1536
pooled_projection_dim: int = 2048
pos_embed_max_size: int = 384
dual_attention_layers: list[int] = field(
default_factory=lambda: [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12])
qk_norm: str = "rms_norm"
in_channels: int = 16
out_channels: int = 16
num_attention_heads: int = 24
@dataclass
class SD3DiTConfig(DiTConfig):
arch_config: DiTArchConfig = field(
default_factory=SD3Transformer2DArchConfig)
prefix: str = "sd3"
@@ -0,0 +1,64 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
@dataclass
class StepVideoArchConfig(DiTArchConfig):
_fsdp_shard_conditions: list = field(
default_factory=lambda:
[lambda n, m: "transformer_blocks" in n and n.split(".")[-1].isdigit()])
param_names_mapping: dict = field(
default_factory=lambda: {
# transformer block
r"^transformer_blocks\.(\d+)\.norm1\.(weight|bias)$":
r"transformer_blocks.\1.norm1.norm.\2",
r"^transformer_blocks\.(\d+)\.norm2\.(weight|bias)$":
r"transformer_blocks.\1.norm2.norm.\2",
r"^transformer_blocks\.(\d+)\.ff\.net\.0\.proj\.weight$":
r"transformer_blocks.\1.ff.fc_in.weight",
r"^transformer_blocks\.(\d+)\.ff\.net\.2\.weight$":
r"transformer_blocks.\1.ff.fc_out.weight",
# adanorm block
r"^adaln_single\.emb\.timestep_embedder\.linear_1\.(weight|bias)$":
r"adaln_single.emb.mlp.fc_in.\1",
r"^adaln_single\.emb\.timestep_embedder\.linear_2\.(weight|bias)$":
r"adaln_single.emb.mlp.fc_out.\1",
# caption projection
r"^caption_projection\.linear_1\.(weight|bias)$":
r"caption_projection.fc_in.\1",
r"^caption_projection\.linear_2\.(weight|bias)$":
r"caption_projection.fc_out.\1",
})
num_attention_heads: int = 48
attention_head_dim: int = 128
in_channels: int = 64
out_channels: int | None = 64
num_layers: int = 48
dropout: float = 0.0
patch_size: int = 1
norm_type: str = "ada_norm_single"
norm_elementwise_affine: bool = False
norm_eps: float = 1e-6
caption_channels: int | list[int] | tuple[int, ...] | None = field(
default_factory=lambda: [6144, 1024])
attention_type: str | None = "torch"
use_additional_conditions: bool | None = False
exclude_lora_layers: list[str] = field(default_factory=lambda: [])
def __post_init__(self):
self.hidden_size = self.num_attention_heads * self.attention_head_dim
self.out_channels = self.in_channels if self.out_channels is None else self.out_channels
self.num_channels_latents = self.out_channels
@dataclass
class StepVideoConfig(DiTConfig):
arch_config: DiTArchConfig = field(default_factory=StepVideoArchConfig)
prefix: str = "StepVideo"
@@ -7,19 +7,6 @@ from fastvideo.configs.models.encoders.base import (
)
def _is_feature_extractor_linear(n: str, m) -> bool:
return n.endswith("feature_extractor_linear")
def _is_embeddings(n: str, m) -> bool:
return n.endswith("embeddings_connector") or n.endswith(
"audio_embeddings_connector")
def _is_gemma_model(n: str, m) -> bool:
return "_gemma_model" in n
@dataclass
class LTX2GemmaArchConfig(TextEncoderArchConfig):
architectures: list[str] = field(
@@ -48,10 +35,6 @@ class LTX2GemmaArchConfig(TextEncoderArchConfig):
connector_double_precision_rope: bool = False
connector_num_learnable_registers: int | None = 128
_fsdp_shard_conditions: list = field(
default_factory=lambda:
[_is_feature_extractor_linear, _is_embeddings, _is_gemma_model])
def __post_init__(self) -> None:
super().__post_init__()
self.tokenizer_kwargs["padding"] = "max_length"
+2 -2
View File
@@ -1,15 +1,15 @@
from fastvideo.configs.models.vaes.cosmosvae import CosmosVAEConfig
from fastvideo.configs.models.vaes.cosmos2_5vae import Cosmos25VAEConfig
from fastvideo.configs.models.vaes.gamecraftvae import GameCraftVAEConfig
from fastvideo.configs.models.vaes.hunyuanvae import HunyuanVAEConfig
from fastvideo.configs.models.vaes.hunyuan15vae import Hunyuan15VAEConfig
from fastvideo.configs.models.vaes.ltx2vae import LTX2VAEConfig
from fastvideo.configs.models.vaes.stepvideovae import StepVideoVAEConfig
from fastvideo.configs.models.vaes.wanvae import WanVAEConfig
__all__ = [
"GameCraftVAEConfig",
"HunyuanVAEConfig",
"WanVAEConfig",
"StepVideoVAEConfig",
"CosmosVAEConfig",
"Cosmos25VAEConfig",
"Hunyuan15VAEConfig",
@@ -1,39 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
import torch
from fastvideo.configs.models.vaes.base import VAEArchConfig, VAEConfig
@dataclass
class AutoencoderKLArchConfig(VAEArchConfig):
_name_or_path: str = ""
act_fn: str = "silu"
block_out_channels: tuple[int, ...] | list[int] = field(
default_factory=list)
down_block_types: tuple[str, ...] | list[str] = field(default_factory=list)
up_block_types: tuple[str, ...] | list[str] = field(default_factory=list)
force_upcast: bool = True
in_channels: int = 3
latent_channels: int = 4
latents_mean: tuple[float, ...] | list[float] | None = None
latents_std: tuple[float, ...] | list[float] | None = None
layers_per_block: int = 1
mid_block_add_attention: bool = True
norm_num_groups: int = 32
out_channels: int = 3
sample_size: int = 32
scaling_factor: float | torch.Tensor = 0.18215
shift_factor: float | None = None
use_post_quant_conv: bool = True
use_quant_conv: bool = True
temporal_compression_ratio: int = 1
spatial_compression_ratio: int = 8
@dataclass
class AutoencoderKLVAEConfig(VAEConfig):
arch_config: VAEArchConfig = field(default_factory=AutoencoderKLArchConfig)
@@ -1,50 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""
GameCraft VAE config - matches official config.json from Hunyuan-GameCraft-1.0.
"""
from dataclasses import dataclass, field
from fastvideo.configs.models.vaes.base import VAEArchConfig, VAEConfig
@dataclass
class GameCraftVAEArchConfig(VAEArchConfig):
"""Architecture config matching official AutoencoderKLCausal3D config.json."""
in_channels: int = 3
out_channels: int = 3
latent_channels: int = 16
down_block_types: tuple[str, ...] = (
"DownEncoderBlockCausal3D",
"DownEncoderBlockCausal3D",
"DownEncoderBlockCausal3D",
"DownEncoderBlockCausal3D",
)
up_block_types: tuple[str, ...] = (
"UpDecoderBlockCausal3D",
"UpDecoderBlockCausal3D",
"UpDecoderBlockCausal3D",
"UpDecoderBlockCausal3D",
)
block_out_channels: tuple[int, ...] = (128, 256, 512, 512)
layers_per_block: int = 2
act_fn: str = "silu"
norm_num_groups: int = 32
scaling_factor: float = 0.476986
spatial_compression_ratio: int = 8
temporal_compression_ratio: int = 4
time_compression_ratio: int = 4 # alias for DecoderCausal3D
mid_block_add_attention: bool = True
mid_block_causal_attn: bool = True
sample_size: int = 256 # from config.json
sample_tsize: int = 64 # from config.json
def __post_init__(self):
self.spatial_compression_ratio = 2**(len(self.block_out_channels) - 1)
@dataclass
class GameCraftVAEConfig(VAEConfig):
"""Full config for GameCraft VAE."""
arch_config: VAEArchConfig = field(default_factory=GameCraftVAEArchConfig)
-1
View File
@@ -16,7 +16,6 @@ class LTX2VAEArchConfig(VAEArchConfig):
in_channels: int = 3
out_channels: int = 3
latent_channels: int = 128
z_dim: int = 128 # follow num_channels_latents
encoder_blocks: list = field(default_factory=list)
decoder_blocks: list = field(default_factory=list)
patch_size: int = 4
@@ -0,0 +1,29 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.configs.models.vaes.base import VAEArchConfig, VAEConfig
@dataclass
class StepVideoVAEArchConfig(VAEArchConfig):
in_channels: int = 3
out_channels: int = 3
z_channels: int = 64
num_res_blocks: int = 2
version: int = 2
frame_len: int = 17
world_size: int = 1
spatial_compression_ratio: int = 16
temporal_compression_ratio: int = 8
scaling_factor: float = 1.0
@dataclass
class StepVideoVAEConfig(VAEConfig):
arch_config: VAEArchConfig = field(default_factory=StepVideoVAEArchConfig)
use_tiling: bool = False
use_temporal_tiling: bool = False
use_parallel_tiling: bool = False
use_temporal_scaling_frames: bool = False
+5 -5
View File
@@ -4,19 +4,19 @@ from fastvideo.configs.pipelines.cosmos import CosmosConfig
from fastvideo.configs.pipelines.cosmos2_5 import Cosmos25Config
from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
from fastvideo.configs.pipelines.hunyuan15 import Hunyuan15T2V480PConfig, Hunyuan15T2V720PConfig
from fastvideo.configs.pipelines.hunyuangamecraft import HunyuanGameCraftPipelineConfig
from fastvideo.configs.pipelines.hyworld import HYWorldConfig
from fastvideo.configs.pipelines.ltx2 import LTX2T2VConfig
from fastvideo.registry import get_pipeline_config_cls_from_name
from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
from fastvideo.configs.pipelines.wan import (SelfForcingWanT2V480PConfig,
WanI2V480PConfig, WanI2V720PConfig,
WanT2V480PConfig, WanT2V720PConfig)
__all__ = [
"HunyuanConfig", "FastHunyuanConfig", "HunyuanGameCraftPipelineConfig",
"PipelineConfig", "Hunyuan15T2V480PConfig", "Hunyuan15T2V720PConfig",
"SlidingTileAttnConfig", "WanT2V480PConfig", "WanI2V480PConfig",
"WanT2V720PConfig", "WanI2V720PConfig", "SelfForcingWanT2V480PConfig",
"HunyuanConfig", "FastHunyuanConfig", "PipelineConfig",
"Hunyuan15T2V480PConfig", "Hunyuan15T2V720PConfig", "SlidingTileAttnConfig",
"WanT2V480PConfig", "WanI2V480PConfig", "WanT2V720PConfig",
"WanI2V720PConfig", "StepVideoT2VConfig", "SelfForcingWanT2V480PConfig",
"CosmosConfig", "Cosmos25Config", "LTX2T2VConfig", "HYWorldConfig",
"get_pipeline_config_cls_from_name"
]
+28
View File
@@ -76,6 +76,11 @@ class PipelineConfig:
...] = field(default_factory=lambda:
(postprocess_text, ))
# StepVideo specific parameters
pos_magic: str | None = None
neg_magic: str | None = None
timesteps_scale: bool | None = None
# STA (Sliding Tile Attention) parameters
mask_strategy_file_path: str | None = None
STA_mode: STA_Mode = STA_Mode.STA_INFERENCE
@@ -182,6 +187,29 @@ class PipelineConfig:
choices=["fp32", "fp16", "bf16"],
help="Precision for image encoder",
)
parser.add_argument(
f"--{prefix_with_dot}pos_magic",
type=str,
dest=f"{prefix_with_dot.replace('-', '_')}pos_magic",
default=PipelineConfig.pos_magic,
help="Positive magic prompt for sampling, used in stepvideo",
)
parser.add_argument(
f"--{prefix_with_dot}neg_magic",
type=str,
dest=f"{prefix_with_dot.replace('-', '_')}neg_magic",
default=PipelineConfig.neg_magic,
help="Negative magic prompt for sampling, used in stepvideo",
)
parser.add_argument(
f"--{prefix_with_dot}timesteps_scale",
type=bool,
dest=f"{prefix_with_dot.replace('-', '_')}timesteps_scale",
default=PipelineConfig.timesteps_scale,
help=
"Bool for applying scheduler scale in set_timesteps, used in stepvideo",
)
# DMD parameters
parser.add_argument(
f"--{prefix_with_dot}dmd-denoising-steps",
@@ -1,130 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""
Pipeline configuration for HunyuanGameCraft.
HunyuanGameCraft extends HunyuanVideo with:
1. CameraNet for camera/action conditioning (Plücker coordinates)
2. Mask-based conditioning for autoregressive generation
3. 33 input channels (16 latent + 16 gt_latent + 1 mask)
Text encoders are the same as HunyuanVideo:
- LLaVA-LLaMA-3-8B for primary text encoding (4096 dim)
- CLIP ViT-L/14 for secondary pooled embeddings (768 dim)
"""
from collections.abc import Callable
from dataclasses import dataclass, field
from typing import TypedDict
import torch
from fastvideo.configs.models import DiTConfig, EncoderConfig, VAEConfig
from fastvideo.configs.models.dits import HunyuanGameCraftConfig
from fastvideo.configs.models.encoders import (
BaseEncoderOutput,
CLIPTextConfig,
LlamaConfig,
)
from fastvideo.configs.models.vaes import GameCraftVAEConfig
from fastvideo.configs.pipelines.base import PipelineConfig
# GameCraft uses the same prompt template as HunyuanVideo
PROMPT_TEMPLATE_ENCODE_VIDEO = (
"<|start_header_id|>system<|end_header_id|>\n\nDescribe the video by detailing the following aspects: "
"1. The main content and theme of the video."
"2. The color, shape, size, texture, quantity, text, and spatial relationships of the objects."
"3. Actions, events, behaviors temporal relationships, physical movement changes of the objects."
"4. background environment, light, style and atmosphere."
"5. camera angles, movements, and transitions used in the video:<|eot_id|>"
"<|start_header_id|>user<|end_header_id|>\n\n{}<|eot_id|>")
class PromptTemplate(TypedDict):
template: str
crop_start: int
prompt_template_video: PromptTemplate = {
"template": PROMPT_TEMPLATE_ENCODE_VIDEO,
"crop_start": 95,
}
def llama_preprocess_text(prompt: str) -> str:
"""Apply prompt template for LLaMA encoder."""
return prompt_template_video["template"].format(prompt)
def llama_postprocess_text(outputs: BaseEncoderOutput) -> torch.Tensor:
"""Extract hidden states from LLaMA output, skipping instruction tokens."""
hidden_state_skip_layer = 2
hidden_states = outputs.hidden_states
if hidden_states is not None and len(
hidden_states) > hidden_state_skip_layer:
last_hidden_state: torch.Tensor = hidden_states[-(
hidden_state_skip_layer + 1)]
elif outputs.last_hidden_state is not None:
# Fallback for encoder outputs without hidden_states.
last_hidden_state = outputs.last_hidden_state
else:
raise ValueError(
"LLaMA encoder output must contain hidden_states or last_hidden_state."
)
crop_start = prompt_template_video.get("crop_start", -1)
last_hidden_state = last_hidden_state[:, crop_start:]
return last_hidden_state
def clip_preprocess_text(prompt: str) -> str:
"""No preprocessing for CLIP encoder."""
return prompt
def clip_postprocess_text(outputs: BaseEncoderOutput) -> torch.Tensor:
"""Extract pooled output from CLIP encoder."""
pooler_output: torch.Tensor = outputs.pooler_output
return pooler_output
@dataclass
class HunyuanGameCraftPipelineConfig(PipelineConfig):
"""Configuration for HunyuanGameCraft pipeline.
Inherits text encoding from HunyuanVideo but uses:
- GameCraft DiT with CameraNet
- Same VAE (HunyuanVAE)
- Same text encoders (LLaMA + CLIP)
"""
# DiT config - uses GameCraft config (33 input channels)
dit_config: DiTConfig = field(default_factory=HunyuanGameCraftConfig)
# VAE config - GameCraft VAE (mid_block_causal_attn=True, etc.)
vae_config: VAEConfig = field(default_factory=GameCraftVAEConfig)
# Denoising parameters
# Official GameCraft does NOT use embedded guidance (passes guidance=None)
# It uses standard CFG with guidance_scale=6.0 instead
embedded_cfg_scale = None
flow_shift: int = 5 # Official GameCraft uses flow_shift=5.0
# Text encoding stage - same as HunyuanVideo
# Uses LLaMA-3-8B (via LLaVA) + CLIP
text_encoder_configs: tuple[EncoderConfig, ...] = field(
default_factory=lambda: (LlamaConfig(), CLIPTextConfig()))
preprocess_text_funcs: tuple[Callable[[str], str], ...] = field(
default_factory=lambda: (llama_preprocess_text, clip_preprocess_text))
postprocess_text_funcs: tuple[
Callable[[BaseEncoderOutput], torch.Tensor],
...] = field(default_factory=lambda:
(llama_postprocess_text, clip_postprocess_text))
# Precision for each component
dit_precision: str = "bf16"
vae_precision: str = "fp16"
text_encoder_precisions: tuple[str, ...] = field(
default_factory=lambda: ("fp16", "fp16"))
def __post_init__(self):
# VAE only needs decoder for inference
self.vae_config.load_encoder = False
self.vae_config.load_decoder = True
@@ -1,13 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.configs.models import DiTConfig
from fastvideo.configs.pipelines.wan import Wan2_2_I2V_A14B_Config
from fastvideo.configs.models.dits.lingbotworld import LingBotWorldVideoConfig
@dataclass
class LingBotWorldI2V480PConfig(Wan2_2_I2V_A14B_Config):
dit_config: DiTConfig = field(default_factory=LingBotWorldVideoConfig)
flow_shift: float | None = 10.0
boundary_ratio: float | None = 0.947
-80
View File
@@ -1,80 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
from collections.abc import Callable
from dataclasses import dataclass, field
import torch
from fastvideo.configs.models import EncoderConfig
from fastvideo.configs.models.encoders import (
BaseEncoderOutput,
CLIPTextConfig,
T5Config,
)
from fastvideo.configs.models.dits.sd3 import SD3DiTConfig
from fastvideo.configs.models.vaes.autoencoder_kl import AutoencoderKLVAEConfig
from fastvideo.configs.pipelines.base import PipelineConfig, preprocess_text
def _sd35_clip_text_postprocess(outputs: BaseEncoderOutput) -> torch.Tensor:
hs = outputs.hidden_states
if hs is None:
raise RuntimeError(
"SD3.5 CLIP prompt embeddings require hidden_states. "
"Set output_hidden_states=True for CLIP encoders.")
return hs[-2]
def _sd35_t5_text_postprocess(outputs: BaseEncoderOutput) -> torch.Tensor:
assert outputs.last_hidden_state is not None
return outputs.last_hidden_state
@dataclass
class SD35Config(PipelineConfig):
scheduler_arch: str = "FlowMatchEulerDiscreteScheduler"
transformer_arch: str = "SD3Transformer2DModel"
vae_arch: str = "AutoencoderKL"
text_encoder_archs: tuple[str, ...] = (
"CLIPTextModelWithProjection",
"CLIPTextModelWithProjection",
"T5EncoderModel",
)
tokenizer_archs: tuple[str, ...] = (
"CLIPTokenizer",
"CLIPTokenizer",
"T5TokenizerFast",
)
dit_config: SD3DiTConfig = field(default_factory=SD3DiTConfig)
vae_config: AutoencoderKLVAEConfig = field(
default_factory=AutoencoderKLVAEConfig)
embedded_cfg_scale: float = 0.0
flow_shift: float | None = None
text_encoder_configs: tuple[EncoderConfig, ...] = field(
default_factory=lambda:
(CLIPTextConfig(), CLIPTextConfig(), T5Config()))
preprocess_text_funcs: tuple[Callable[[str], str], ...] = field(
default_factory=lambda:
(preprocess_text, preprocess_text, preprocess_text))
postprocess_text_funcs: tuple[
Callable[[BaseEncoderOutput], torch.Tensor],
...] = field(default_factory=lambda:
(_sd35_clip_text_postprocess, _sd35_clip_text_postprocess,
_sd35_t5_text_postprocess))
dit_precision: str = "bf16"
vae_precision: str = "fp32"
text_encoder_precisions: tuple[str, ...] = field(
default_factory=lambda: ("fp32", "fp32", "bf16"))
def __post_init__(self) -> None:
te_cfgs = list(self.text_encoder_configs)
for idx in (0, 1):
if idx < len(te_cfgs):
te_cfgs[idx].arch_config.output_hidden_states = True
+30
View File
@@ -0,0 +1,30 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.configs.models import DiTConfig, VAEConfig
from fastvideo.configs.models.dits import StepVideoConfig
from fastvideo.configs.models.vaes import StepVideoVAEConfig
from fastvideo.configs.pipelines.base import PipelineConfig
@dataclass
class StepVideoT2VConfig(PipelineConfig):
"""Base configuration for StepVideo pipeline architecture."""
# WanConfig-specific parameters with defaults
# DiT
dit_config: DiTConfig = field(default_factory=StepVideoConfig)
# VAE
vae_config: VAEConfig = field(default_factory=StepVideoVAEConfig)
vae_tiling: bool = False
vae_sp: bool = False
# Denoising stage
flow_shift: int = 13
timesteps_scale: bool = False
pos_magic: str = "超高清、HDR 视频、环境光、杜比全景声、画面稳定、流畅动作、逼真的细节、专业级构图、超现实主义、自然、生动、超细节、清晰。"
neg_magic: str = "画面暗、低分辨率、不良手、文本、缺少手指、多余的手指、裁剪、低质量、颗粒状、签名、水印、用户名、模糊。"
# Precision for each component
precision: str = "bf16"
vae_precision: str = "bf16"
+1 -11
View File
@@ -1,13 +1,3 @@
from fastvideo.configs.sample.base import SamplingParam
from fastvideo.configs.sample.hunyuangamecraft import (
HunyuanGameCraftSamplingParam,
HunyuanGameCraft65FrameSamplingParam,
HunyuanGameCraft129FrameSamplingParam,
)
__all__ = [
"SamplingParam",
"HunyuanGameCraftSamplingParam",
"HunyuanGameCraft65FrameSamplingParam",
"HunyuanGameCraft129FrameSamplingParam",
]
__all__ = ["SamplingParam"]
-3
View File
@@ -31,9 +31,6 @@ class SamplingParam:
# Camera control inputs (HYWorld)
pose: str | None = None # Camera trajectory: pose string (e.g., 'w-31') or JSON file path
# Camera control inputs (LingBotWorld)
c2ws_plucker_emb: Any | None = None # Plucker embedding: [B, C, F_lat, H_lat, W_lat]
# Refine inputs (LongCat 480p->720p upscaling)
# Path-based refine (load stage1 video from disk, e.g. MP4)
refine_from: str | None = None # Path to stage1 video (480p output from distill)
@@ -1,104 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""
Sampling parameters for HunyuanGameCraft video generation.
GameCraft generates game-like videos with camera/action control.
Default parameters are based on the official implementation.
"""
from dataclasses import dataclass, field
from typing import Any
from fastvideo.configs.sample.base import SamplingParam
from fastvideo.configs.sample.teacache import TeaCacheParams
@dataclass
class HunyuanGameCraftSamplingParam(SamplingParam):
"""Sampling parameters for HunyuanGameCraft video generation.
Supports camera/action conditioning via:
- camera_trajectory: Plücker coordinates for camera motion
- action_list: List of actions (e.g., ["forward", "left", "right"])
- action_speed_list: Speed multipliers for each action
Default resolution is 704x1280 (same as HunyuanVideo).
Default frame count is 33 video frames -> 9 latent frames.
"""
# Number of denoising steps
num_inference_steps: int = 50
# Video dimensions
# 33 video frames -> 9 latent frames (4x temporal compression)
num_frames: int = 33
height: int = 704
width: int = 1280
fps: int = 24
# Guidance scale - official GameCraft uses CFG with guidance_scale=6.0
guidance_scale: float = 6.0
# Negative prompt for CFG (empty string = unconditional)
negative_prompt: str = ""
# Camera/Action conditioning
# Camera states as Plücker coordinates [B, T_video, 6, H, W]
camera_states: Any | None = None
# Camera trajectory file/identifier (alternative to camera_states)
camera_trajectory: str | None = None
# Action list for camera motion (e.g., ["forward", "left"])
action_list: list[str] | None = None
# Speed multipliers for each action
action_speed_list: list[float] | None = None
# History frame conditioning (for autoregressive generation)
# Ground truth latents for conditioning [B, 16, T, H, W]
gt_latents: Any | None = None
# Mask for conditioning (1=use gt, 0=generate) [B, 1, T, H, W]
conditioning_mask: Any | None = None
# Number of conditioning frames (for autoregressive) - maps to num_cond_frames
num_cond_frames: int = 0
# TeaCache parameters (if enabled)
teacache_params: TeaCacheParams = field(
default_factory=lambda: TeaCacheParams(
teacache_thresh=0.15,
coefficients=[
7.33226126e+02, -4.01131952e+02, 6.75869174e+01,
-3.14987800e+00, 9.61237896e-02
],
))
def __post_init__(self) -> None:
super().__post_init__()
# Validate action lists
if (self.action_list is not None and self.action_speed_list is not None
and len(self.action_list) != len(self.action_speed_list)):
raise ValueError(
f"action_list length ({len(self.action_list)}) must match "
f"action_speed_list length ({len(self.action_speed_list)})")
@dataclass
class HunyuanGameCraft65FrameSamplingParam(HunyuanGameCraftSamplingParam):
"""Sampling parameters for 65-frame GameCraft generation.
65 video frames -> 17 latent frames (with first frame as key frame).
This is useful for longer video generation.
"""
num_frames: int = 65
@dataclass
class HunyuanGameCraft129FrameSamplingParam(HunyuanGameCraftSamplingParam):
"""Sampling parameters for 129-frame GameCraft generation.
129 video frames -> 33 latent frames.
This is the maximum supported by the official implementation.
"""
num_frames: int = 129
-21
View File
@@ -1,21 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass
from fastvideo.configs.sample.wan import Wan2_2_I2V_A14B_SamplingParam
@dataclass
class LingBotWorld_SamplingParam(Wan2_2_I2V_A14B_SamplingParam):
guidance_scale: float = 5.0 # high_noise
guidance_scale_2: float = 5.0 # low_noise
num_inference_steps: int = 70
boundary_ratio: float | None = 0.947
negative_prompt: str | None = (
"画面突变,色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,"
"最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,"
"畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走,"
"镜头晃动,画面闪烁,模糊,噪点,水印,签名,文字,变形,扭曲,液化,不合逻辑的结构,卡顿,"
"PPT幻灯片感,过暗,欠曝,低对比度,霓虹灯光感,过度锐化,3D渲染感,人物,行人,游客,身体,"
"皮肤,肢体,面部特征,汽车,电线")
fps: int = 16
# NOTE(will): default boundary timestep is tracked by PipelineConfig, but
# can be overridden during sampling
+2 -47
View File
@@ -5,51 +5,10 @@ from fastvideo.configs.sample.base import SamplingParam
@dataclass
class LTX2BaseSamplingParam(SamplingParam):
"""Default sampling parameters for LTX-2 base one-stage T2V.
Values follow the official LTX-2 one-stage defaults.
Multi-modal CFG params are read by ``LTX2DenoisingStage``.
class LTX2SamplingParam(SamplingParam):
"""Default sampling parameters for LTX-2 distilled T2V.
"""
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.")
# Official LTX-2 multi-modal CFG defaults.
ltx2_cfg_scale_video: float = 3.0
ltx2_cfg_scale_audio: float = 7.0
ltx2_modality_scale_video: float = 3.0
ltx2_modality_scale_audio: float = 3.0
ltx2_rescale_scale: float = 0.7
@dataclass
class LTX2DistilledSamplingParam(SamplingParam):
"""Default sampling parameters for LTX-2 distilled one-stage T2V."""
seed: int = 10
num_frames: int = 121
height: int = 1024
@@ -59,7 +18,3 @@ class LTX2DistilledSamplingParam(SamplingParam):
guidance_scale: float = 1.0
# No default negative_prompt for distilled models
negative_prompt: str = ""
# Backward compatibility alias.
LTX2SamplingParam = LTX2DistilledSamplingParam
-25
View File
@@ -1,25 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
from dataclasses import dataclass
from fastvideo.configs.sample.base import SamplingParam
@dataclass
class SD35SamplingParam(SamplingParam):
prompt: str | None = "a photo of a cat"
negative_prompt: str = ""
num_videos_per_prompt: int = 1
seed: int = 0
num_frames: int = 1
height: int = 512
width: int = 512
fps: int = 1
num_inference_steps: int = 28
guidance_scale: float = 6.0
+20
View File
@@ -0,0 +1,20 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass
from fastvideo.configs.sample.base import SamplingParam
@dataclass
class StepVideoT2VSamplingParam(SamplingParam):
# Video parameters
height: int = 720
width: int = 1280
num_frames: int = 81
# Denoising stage
guidance_scale: float = 9.0
num_inference_steps: int = 50
# neg magic and pos magic
# pos_magic: str = "超高清、HDR 视频、环境光、杜比全景声、画面稳定、流畅动作、逼真的细节、专业级构图、超现实主义、自然、生动、超细节、清晰。"
# neg_magic: str = "画面暗、低分辨率、不良手、文本、缺少手指、多余的手指、裁剪、低质量、颗粒状、签名、水印、用户名、模糊。"
+2 -8
View File
@@ -4,8 +4,6 @@ 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)
@@ -48,10 +46,6 @@ def gettextdataset(args) -> TextDataset:
__all__ = [
"build_parquet_map_style_dataloader",
"build_ltx2_precomputed_dataloader",
"LTX2PrecomputedDataset",
"ValidationDataset",
"VideoCaptionMergedDataset",
"TextDataset",
"build_parquet_map_style_dataloader", "ValidationDataset",
"VideoCaptionMergedDataset", "TextDataset"
]
@@ -1,210 +0,0 @@
# 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
+8 -52
View File
@@ -95,9 +95,8 @@ class DP_SP_BatchSampler(Sampler[list[int]]):
def get_parquet_files_and_length(path: str):
dataset_root = os.path.realpath(os.path.expanduser(path))
# Check if cached info exists
cache_dir = os.path.join(dataset_root, "map_style_cache")
cache_dir = os.path.join(path, "map_style_cache")
cache_file = os.path.join(cache_dir, "file_info.pkl")
# Only rank 0 checks for cache and scans files if needed
@@ -112,39 +111,8 @@ def get_parquet_files_and_length(path: str):
try:
with open(cache_file, "rb") as f:
file_names_sorted, lengths_sorted = pickle.load(f)
file_names_sorted = tuple(
os.path.realpath(
os.path.join(os.getcwd(), p)
if not os.path.isabs(p) else p)
for p in file_names_sorted)
files_outside_dataset_root = [
file_path for file_path in file_names_sorted
if os.path.commonpath([dataset_root, file_path
]) != dataset_root
]
missing_files = [
file_path for file_path in file_names_sorted
if not os.path.exists(file_path)
]
if files_outside_dataset_root:
logger.warning(
"Cached parquet file list points outside dataset root "
"(%s). Cache will be rebuilt. First out-of-root file: %s",
dataset_root,
files_outside_dataset_root[0],
)
cache_loaded = False
elif missing_files:
logger.warning(
"Cached parquet file list contains %d missing files. "
"Cache will be rebuilt. First missing file: %s",
len(missing_files),
missing_files[0],
)
cache_loaded = False
else:
cache_loaded = True
logger.info("Successfully loaded cached file info")
cache_loaded = True
logger.info("Successfully loaded cached file info")
except Exception as e:
logger.error("Error loading cached file info: %s", str(e))
logger.info("Falling back to scanning files")
@@ -155,17 +123,11 @@ def get_parquet_files_and_length(path: str):
logger.info("Scanning parquet files to get lengths")
lengths = []
file_names = []
for root, _, files in os.walk(dataset_root):
for root, _, files in os.walk(path):
for file in sorted(files):
if file.endswith('.parquet'):
file_path = os.path.realpath(os.path.join(root, file))
file_path = os.path.join(root, file)
file_names.append(file_path)
if len(file_names) == 0:
raise FileNotFoundError(
"No parquet files found under dataset path: "
f"{path}. "
"Please verify this path points to preprocessed parquet "
"data.")
for file_path in tqdm.tqdm(
file_names, desc="Reading parquet files to get lengths"):
num_rows = pq.ParquetFile(file_path).metadata.num_rows
@@ -176,6 +138,9 @@ def get_parquet_files_and_length(path: str):
strict=True),
key=lambda x: x[0]),
strict=True)
assert len(
file_names_sorted) != 0, "No parquet files found in the dataset"
# Save the cache
os.makedirs(cache_dir, exist_ok=True)
with open(cache_file, "wb") as f:
@@ -190,15 +155,6 @@ def get_parquet_files_and_length(path: str):
logger.info("Loading cached file info from %s after barrier", cache_file)
with open(cache_file, "rb") as f:
file_names_sorted, lengths_sorted = pickle.load(f)
if len(file_names_sorted) == 0:
raise RuntimeError(
"Cached parquet metadata is empty after synchronization at "
f"{cache_file}. "
"Please verify the dataset path and regenerate cache.")
if len(file_names_sorted) != len(lengths_sorted):
raise RuntimeError(
"Cached parquet metadata is corrupted at "
f"{cache_file}: file count and length count do not match.")
return file_names_sorted, lengths_sorted
+3 -4
View File
@@ -990,10 +990,9 @@ def maybe_init_distributed_environment_and_model_parallel(
# set device if we're on a CUDA/NPU platform
from fastvideo.platforms import current_platform
if current_platform.is_cuda_alike() or current_platform.is_npu():
device_type = current_platform.device_type
device = torch.device(f"{device_type}:{local_rank}")
current_platform.get_torch_device().set_device(device)
device_type = current_platform.device_type
device = torch.device(f"{device_type}:{local_rank}")
current_platform.get_torch_device().set_device(device)
def model_parallel_is_initialized() -> bool:
+43 -59
View File
@@ -6,6 +6,7 @@ This module provides a consolidated interface for generating videos using
diffusion models.
"""
import math
import os
import re
import threading
@@ -222,82 +223,64 @@ class VideoGenerator:
sampling_param=sampling_param,
**kwargs)
def _is_image_workload(self) -> bool:
"""Return True when the workload produces a single image (t2i, i2i …)."""
args = getattr(self, "fastvideo_args", None)
if args is None:
return False
return args.workload_type.value.endswith("2i")
def _prepare_output_path(
self,
output_path: str,
prompt: str,
) -> str:
"""Build a unique, sanitized output file path.
"""Build a unique, sanitized .mp4 output file path.
The file extension is chosen automatically based on the workload type:
``.png`` for image workloads (``t2i``, ``i2i``, …) and ``.mp4`` for
video workloads.
- If ``output_path`` already carries the correct extension, treat it
as a file path.
- Otherwise, treat ``output_path`` as a directory and derive the
filename from the prompt.
- If `output_path` ends with .mp4 (case-insensitive), treat it as a file path.
- Otherwise, treat `output_path` as a directory and derive the filename
from the prompt.
- Invalid filename characters are removed; if the name changes, a
warning is logged.
- If the target path already exists, a numeric suffix is appended.
"""
target_ext = ".png" if self._is_image_workload() else ".mp4"
def _sanitize_filename_component(name: str) -> str:
# Remove characters invalid on common filesystems, strip spaces/dots
sanitized = re.sub(r'[\\/:*?"<>|]', '', name)
sanitized = sanitized.strip().strip('.')
sanitized = re.sub(r'\s+', ' ', sanitized)
return sanitized or "output"
return sanitized or "video"
base_path, extension = os.path.splitext(output_path)
extension_lower = extension.lower()
if extension_lower == target_ext:
if extension_lower == ".mp4":
output_dir = os.path.dirname(output_path)
base_name = os.path.basename(
base_path) # filename without extension
sanitized_base = _sanitize_filename_component(base_name)
if sanitized_base != base_name:
logger.warning(
"The output name '%s' contained invalid characters. "
"It has been renamed to '%s%s'",
"The video name '%s' contained invalid characters. It has been renamed to '%s.mp4'",
os.path.basename(output_path),
sanitized_base,
target_ext,
)
out_name = f"{sanitized_base}{target_ext}"
video_name = f"{sanitized_base}.mp4"
else:
# Treat as directory; inform if an unexpected extension was
# provided.
# Treat as directory; inform if an unexpected extension was provided.
if extension:
logger.info(
"Output path '%s' has extension '%s' which does not "
"match the target '%s'; treating it as a directory",
"Output path '%s' has non-mp4 extension '%s'; treating it as a directory and using a .mp4 filename derived from the prompt",
output_path,
extension,
target_ext,
)
output_dir = output_path
prompt_component = _sanitize_filename_component(prompt[:100])
out_name = f"{prompt_component}{target_ext}"
video_name = f"{prompt_component}.mp4"
if output_dir:
os.makedirs(output_dir, exist_ok=True)
new_output_path = os.path.join(output_dir, out_name)
new_output_path = os.path.join(output_dir, video_name)
counter = 1
while os.path.exists(new_output_path):
name_part, ext_part = os.path.splitext(out_name)
new_name = f"{name_part}_{counter}{ext_part}"
new_output_path = os.path.join(output_dir, new_name)
name_part, ext_part = os.path.splitext(video_name)
new_video_name = f"{name_part}_{counter}{ext_part}"
new_output_path = os.path.join(output_dir, new_video_name)
counter += 1
return new_output_path
@@ -335,19 +318,30 @@ class VideoGenerator:
temporal_scale_factor = pipeline_config.vae_config.arch_config.temporal_compression_ratio
num_frames = sampling_param.num_frames
num_gpus = fastvideo_args.num_gpus
use_temporal_scaling_frames = pipeline_config.vae_config.use_temporal_scaling_frames
# Adjust number of frames based on number of GPUs
if not use_temporal_scaling_frames:
raise ValueError(
"Only temporal-scaling-frame VAE configs are supported.")
orig_latent_num_frames = (num_frames - 1) // temporal_scale_factor + 1
if use_temporal_scaling_frames:
orig_latent_num_frames = (num_frames -
1) // temporal_scale_factor + 1
else: # stepvideo only
orig_latent_num_frames = sampling_param.num_frames // 17 * 3
if orig_latent_num_frames % fastvideo_args.num_gpus != 0:
# Convert back to number of frames, ensuring num_frames-1 is a
# multiple of temporal_scale_factor.
new_num_frames = (orig_latent_num_frames -
1) * temporal_scale_factor + 1
if use_temporal_scaling_frames:
# Convert back to number of frames, ensuring num_frames-1 is a multiple of temporal_scale_factor
new_num_frames = (orig_latent_num_frames -
1) * temporal_scale_factor + 1
else: # stepvideo only
# Find the least common multiple of 3 and num_gpus
divisor = math.lcm(3, num_gpus)
# Round up to the nearest multiple of this LCM
orig_latent_num_frames = (
(orig_latent_num_frames + divisor - 1) // divisor) * divisor
# Convert back to actual frames using the StepVideo formula
new_num_frames = orig_latent_num_frames // 3 * 17
logger.info(
"Adjusting number of frames from %s to %s based on number of GPUs (%s)",
@@ -432,25 +426,15 @@ class VideoGenerator:
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
frames.append((x * 255).numpy().astype(np.uint8))
# Save output if requested
# Save video if requested
if batch.save_video:
if self._is_image_workload():
# Image workloads (t2i, i2i, …): save the first frame as PNG.
imageio.imwrite(output_path, frames[0])
logger.info("Saved image to %s", output_path)
else:
imageio.mimsave(output_path,
frames,
fps=batch.fps,
format="mp4")
logger.info("Saved video to %s", output_path)
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(output_path, audio,
audio_sample_rate)):
logger.warning(
"Audio mux failed; saved video without audio.")
imageio.mimsave(output_path, frames, fps=batch.fps, format="mp4")
logger.info("Saved video to %s", output_path)
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(output_path, audio, audio_sample_rate)):
logger.warning("Audio mux failed; saved video without audio.")
if batch.return_frames:
return frames
+5 -8
View File
@@ -257,6 +257,11 @@ class FastVideoArgs:
help=
"The path of the model weights. This can be a local folder or a Hugging Face repo ID.",
)
parser.add_argument(
"--model-dir",
type=str,
help="Directory containing StepVideo model",
)
# Running mode
parser.add_argument(
@@ -898,7 +903,6 @@ 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
@@ -1253,13 +1257,6 @@ 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(
+22 -61
View File
@@ -109,45 +109,26 @@ def _apply_rotary_emb(
"""
Args:
x: [num_tokens, num_heads, head_size]
cos: [num_tokens, head_size] or [num_tokens, head_size // 2]
sin: [num_tokens, head_size] or [num_tokens, head_size // 2]
cos: [num_tokens, head_size // 2]
sin: [num_tokens, head_size // 2]
is_neox_style: Whether to use the Neox-style or GPT-J-style rotary
positional embeddings.
The function auto-detects whether cos/sin are full or half head_size:
- If cos/sin have head_size: use rotate_half style (for HunyuanVideo/GameCraft)
- If cos/sin have head_size // 2: use Neox/GPT-J style
"""
head_size = x.shape[-1]
rope_dim = cos.shape[-1]
# Check if cos/sin are full head_dim (rotate_half style) or half (traditional style)
if rope_dim == head_size:
# Full head_dim - use rotate_half style (HunyuanVideo, GameCraft)
# x * cos + rotate_half(x) * sin
cos = cos.unsqueeze(-2) # [num_tokens, 1, head_size]
sin = sin.unsqueeze(-2) # [num_tokens, 1, head_size]
# rotate_half: split into pairs, negate and swap
x_real, x_imag = x.float().reshape(*x.shape[:-1], -1,
2).unbind(-1) # [B, H, D//2] each
x_rotated = torch.stack([-x_imag, x_real],
dim=-1).flatten(-2) # [B, H, D]
return (x.float() * cos + x_rotated * sin).type_as(x)
# cos = cos.unsqueeze(-2).to(x.dtype)
# sin = sin.unsqueeze(-2).to(x.dtype)
cos = cos.unsqueeze(-2)
sin = sin.unsqueeze(-2)
if is_neox_style:
x1, x2 = torch.chunk(x, 2, dim=-1)
else:
# Half head_dim - use traditional Neox/GPT-J style
cos = cos.unsqueeze(-2)
sin = sin.unsqueeze(-2)
if is_neox_style:
x1, x2 = torch.chunk(x, 2, dim=-1)
else:
x1 = x[..., ::2]
x2 = x[..., 1::2]
o1 = (x1.float() * cos - x2.float() * sin).type_as(x)
o2 = (x2.float() * cos + x1.float() * sin).type_as(x)
if is_neox_style:
return torch.cat((o1, o2), dim=-1)
else:
return torch.stack((o1, o2), dim=-1).flatten(-2)
x1 = x[..., ::2]
x2 = x[..., 1::2]
o1 = (x1.float() * cos - x2.float() * sin).type_as(x)
o2 = (x2.float() * cos + x1.float() * sin).type_as(x)
if is_neox_style:
return torch.cat((o1, o2), dim=-1)
else:
return torch.stack((o1, o2), dim=-1).flatten(-2)
@CustomOp.register("rotary_embedding")
@@ -297,7 +278,6 @@ def get_1d_rotary_pos_embed(
theta_rescale_factor: float = 1.0,
interpolation_factor: float = 1.0,
dtype: torch.dtype = torch.float32,
use_real: bool = True,
) -> tuple[torch.Tensor, torch.Tensor]:
"""
Precompute the frequency tensor for complex exponential (cis) with given dimensions.
@@ -312,12 +292,9 @@ def get_1d_rotary_pos_embed(
theta (float, optional): Scaling factor for frequency computation. Defaults to 10000.0.
theta_rescale_factor (float, optional): Rescale factor for theta. Defaults to 1.0.
interpolation_factor (float, optional): Factor to scale positions. Defaults to 1.0.
use_real (bool, optional): If True, output full head_dim with repeated cos/sin for
rotate_half style RoPE. If False, output half head_dim for complex style. Defaults to True.
Returns:
freqs_cos, freqs_sin: Precomputed frequency tensor with real and imaginary parts separately.
Shape is [S, D] if use_real=True, [S, D/2] if use_real=False.
freqs_cos, freqs_sin: Precomputed frequency tensor with real and imaginary parts separately. [S, D]
"""
if isinstance(pos, int):
pos = torch.arange(pos).float()
@@ -332,16 +309,6 @@ def get_1d_rotary_pos_embed(
freqs = torch.outer(pos * interpolation_factor, freqs) # [S, D/2]
freqs_cos = freqs.cos() # [S, D/2]
freqs_sin = freqs.sin() # [S, D/2]
if use_real:
# For rotate_half style RoPE (used by HunyuanVideo, GameCraft),
# we need to expand cos/sin to full head_dim using repeat_interleave.
# The rotate_half operation works on consecutive PAIRS: (x0,x1), (x2,x3)...
# so cos/sin must be interleaved: [c0,c0,c1,c1,...] to match the pairing.
# Using torch.cat would produce [c0,c1,...,c0,c1,...] which is WRONG.
freqs_cos = freqs_cos.repeat_interleave(2, dim=-1) # [S, D]
freqs_sin = freqs_sin.repeat_interleave(2, dim=-1) # [S, D]
return freqs_cos, freqs_sin
@@ -357,7 +324,6 @@ def get_nd_rotary_pos_embed(
sp_world_size: int = 1,
dtype: torch.dtype = torch.float32,
start_frame: int = 0,
use_real: bool = True,
) -> tuple[torch.Tensor, torch.Tensor]:
"""
This is a n-d version of precompute_freqs_cis, which is a RoPE for tokens with n-d structure.
@@ -375,10 +341,9 @@ def get_nd_rotary_pos_embed(
shard_dim (int): Which dimension to shard for sequence parallelism. Defaults to 0.
sp_rank (int): Rank in the sequence parallel group. Defaults to 0.
sp_world_size (int): World size of the sequence parallel group. Defaults to 1.
use_real (bool): If True, output full head_dim for rotate_half style. Defaults to True.
Returns:
Tuple[torch.Tensor, torch.Tensor]: (cos, sin) tensors of shape [HW, D] if use_real, [HW, D/2] otherwise
Tuple[torch.Tensor, torch.Tensor]: (cos, sin) tensors of shape [HW, D/2]
"""
# Get the full grid
full_grid = get_meshgrid_nd(
@@ -447,12 +412,11 @@ def get_nd_rotary_pos_embed(
theta_rescale_factor=theta_rescale_factor[i],
interpolation_factor=interpolation_factor[i],
dtype=dtype,
use_real=use_real,
) # 2 x [WHD, rope_dim_list[i]] or 2 x [WHD, rope_dim_list[i]*2] if use_real
) # 2 x [WHD, rope_dim_list[i]]
embs.append(emb)
cos = torch.cat([emb[0] for emb in embs], dim=1) # (WHD, D) or (WHD, D/2)
sin = torch.cat([emb[1] for emb in embs], dim=1) # (WHD, D) or (WHD, D/2)
cos = torch.cat([emb[0] for emb in embs], dim=1) # (WHD, D/2)
sin = torch.cat([emb[1] for emb in embs], dim=1) # (WHD, D/2)
return cos, sin
@@ -468,7 +432,6 @@ def get_rotary_pos_embed(
do_sp_sharding: bool = False,
dtype: torch.dtype = torch.float32,
start_frame: int = 0,
use_real: bool = True,
) -> tuple[torch.Tensor, torch.Tensor]:
"""
Generate rotary positional embeddings for the given sizes.
@@ -483,10 +446,9 @@ def get_rotary_pos_embed(
interpolation_factor: Factor to scale positions. Defaults to 1.0
shard_dim: Which dimension to shard for sequence parallelism. Defaults to 0.
do_sp_sharding: Whether to shard the positional embeddings for sequence parallelism. Defaults to False.
use_real: If True, output full head_dim for rotate_half style RoPE. Defaults to True.
Returns:
Tuple of (cos, sin) tensors for rotary embeddings. Shape [S, D] if use_real, [S, D/2] otherwise.
Tuple of (cos, sin) tensors for rotary embeddings
"""
target_ndim = 3
@@ -519,7 +481,6 @@ def get_rotary_pos_embed(
sp_world_size=sp_world_size,
dtype=dtype,
start_frame=start_frame,
use_real=use_real,
)
return freqs_cos, freqs_sin

Some files were not shown because too many files have changed in this diff Show More