Compare commits

...
Author SHA1 Message Date
Peiyuan Zhang e4de77692f add readme 2026-02-23 01:08:03 +00:00
Peiyuan Zhang 077232aa5f fix_sta 2026-02-23 00:54:48 +00:00
Zhang Peiyuan 03d9ce2edb [Misc] Remove StepVideo (#1118) 2026-02-21 17:15:42 -08:00
10fc92dba5 Upstream LTX2 Training (#1116)
Co-authored-by: RandNMR73 <notomatthew31@gmail.com>
Co-authored-by: JerryZhou54 <zhouw.jerry2017@outlook.com>
Co-authored-by: Davids048 <jundasu@ucsd.edu>
Co-authored-by: Peiyuan Zhang <a1286225768@slurm-h200-204-239.slurm-compute.tenant-slurm.svc.cluster.local>
2026-02-21 16:06:55 -08:00
6736dc06a5 Improve Docs (#1112)
Co-authored-by: Peiyuan Zhang <a1286225768@slurm-h200-204-227.slurm-compute.tenant-slurm.svc.cluster.local>
Co-authored-by: Peiyuan Zhang <a1286225768@slurm-login-0.slurm-login.tenant-slurm.svc.cluster.local>
2026-02-19 14:23:32 -08:00
William Lin 8c002c62af [misc] add hy-world link to readme (#1113) 2026-02-18 12:01:10 -08:00
Darren 7061313d04 [bugfix] get_torch_device and other device calls were being made on non-cuda platforms (#1107) 2026-02-18 11:43:46 -08:00
Zhang Peiyuanandgemini-code-assist[bot] 76d3ba69e0 [Misc] clean up VSA finetuning examples. (#1111)
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
2026-02-18 11:37:20 -08:00
8e39ce38c9 [Feat] Native dit implementation for SD3.5 (#1093)
Co-authored-by: Ishan Vaish <vaish.ishan@gmail.com>
Co-authored-by: Cursor <cursoragent@cursor.com>
2026-02-18 10:19:40 +08:00
Darren d4bd8bf2c0 Update README.md (#1110) 2026-02-17 14:54:34 -08:00
Darren e83d7bc50c [bugfix] fix import PreTrainedModel in stepllm.py (#1108) 2026-02-16 21:45:08 -08:00
Jinzhe Pan ff3d5aff75 [Fix] hunyuan postprecessing issue (#1104) 2026-02-15 12:26:35 -08:00
XOR-op 959dbcc8a2 [perf] causal MatrixGame optimization (#1078) 2026-02-15 09:51:00 +08:00
104 changed files with 2678 additions and 579625 deletions
+7 -1
View File
@@ -45,6 +45,12 @@ 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
@@ -63,4 +69,4 @@ jobs:
steps:
- name: Deploy to GitHub Pages
id: deployment
uses: actions/deploy-pages@v4
uses: actions/deploy-pages@v4
+2 -2
View File
@@ -33,11 +33,11 @@ env
**.txt
*.log
weights/
official_weights/
converted_weights/
# SSIM test outputs
fastvideo/tests/ssim/generated_videos/
**/.cache/**
# Distribution / packaging
build/
+72 -123
View File
@@ -1,152 +1,101 @@
<div align="center">
<img src=assets/logos/logo.svg width="30%"/>
</div>
# Sliding Tile Attention (STA) Branch
<p align="center">
| <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start/"><b> Quick Start</b></a> | <a href="https://github.com/hao-ai-lab/FastVideo/discussions/982" target="_blank"><b>Weekly Dev Meeting</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-3f4lao1uq-u~Ipx6Lt4J27AlD2y~IdLQ" target="_blank"> <b>Slack</b> </a> | 🟣💬 <a href="https://github.com/hao-ai-lab/FastVideo/discussions/1097" target="_blank"> <b> WeChat </b> </a> |
</p>
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.
**FastVideo is a unified post-training and inference framework for accelerated video generation.**
## What is STA
## NEWS
Sliding Tile Attention is an optimized attention backend for
window-based video generation.
- `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/).
- 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`
### More News
## Setup
- `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:
Install FastVideo from source:
```bash
# Create and activate a new conda environment
conda create -n fastvideo python=3.12
conda activate fastvideo
# Install FastVideo
pip install fastvideo
uv pip install -e .
```
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:
Build the STA kernel package:
```bash
python example.py
cd fastvideo-kernel
./build.sh
cd ..
```
For a more detailed guide, please see our [inference quick start](https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start/).
## Run STA Inference
## More Guides
STA backend:
- [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/)
```bash
export FASTVIDEO_ATTENTION_BACKEND=SLIDING_TILE_ATTN
```
## Awesome work using FastVideo or our research projects
Ready-to-run examples:
- [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.
- HunyuanVideo: `scripts/inference/v1_inference_hunyuan_STA.sh`
- Wan2.1-T2V-14B: `scripts/inference/v1_inference_wan_STA.sh`
## 🤝 Contributing
Run:
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).
```bash
bash scripts/inference/v1_inference_hunyuan_STA.sh
# or
bash scripts/inference/v1_inference_wan_STA.sh
```
## Acknowledgement
Both scripts already set STA-related env vars:
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.
- `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
```
## 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
@@ -0,0 +1 @@
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
@@ -0,0 +1 @@
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
@@ -0,0 +1 @@
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
@@ -0,0 +1 @@
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
@@ -0,0 +1 @@
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
@@ -0,0 +1 @@
fox in the forest close-up quickly turned its head to the left
+1
View File
@@ -0,0 +1 @@
Man walking his dog in the woods on a hot sunny day
+1
View File
@@ -0,0 +1 @@
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.
+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)
+18 -5
View File
@@ -2,6 +2,7 @@
# 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
@@ -126,8 +127,20 @@ class Example:
Raises:
IndexError: If no Markdown files are found in the directory.
""" # noqa: E501
return self.path if self.path.is_file() else list(
self.path.glob("*.md")).pop()
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]
def determine_other_files(self) -> list[Path]:
"""
@@ -521,11 +534,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 = category_index.path.relative_to(
main_index_dir.parent)
rel_path = os.path.relpath(category_index.path,
start=main_index_dir)
examples_index.documents.insert(
0,
str(rel_path).replace(".md", ""))
str(rel_path).replace("\\", "/").replace(".md", ""))
# Write the category index file
with open(category_index.path, "w+") as f:
+13
View File
@@ -18,12 +18,25 @@ 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:
+2 -1
View File
@@ -79,5 +79,6 @@ if __name__ == '__main__':
- [Installation Guide](installation.md) - Detailed installation instructions
- [Configuration](../inference/configuration.md) - Learn about configuration options
- [Examples](../inference/examples/) - Explore more examples
- [Examples](../inference/examples/examples_inference_index.md) - Explore more
examples
- [Optimizations](../inference/optimizations.md) - Performance optimization tips
+14 -4
View File
@@ -101,14 +101,17 @@ Replace standard attention with FastVideo's optimized attention:
```python
# Local attention patterns
from fastvideo.attention import LocalAttention
from fastvideo.attention.backends.abstract import _Backend
from fastvideo.platforms.interface import AttentionBackendEnum
self.attn = LocalAttention(
num_heads=num_heads,
head_size=head_dim,
dropout_rate=0.0,
softmax_scale=None,
causal=False,
supported_attention_backends=(_Backend.FLASH_ATTN, _Backend.TORCH_SDPA)
supported_attention_backends=(
AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA,
)
)
# Distributed attention for long sequences
@@ -119,14 +122,21 @@ self.attn = DistributedAttention(
dropout_rate=0.0,
softmax_scale=None,
causal=False,
supported_attention_backends=(_Backend.SLIDING_TILE_ATTN, _Backend.FLASH_ATTN, _Backend.TORCH_SDPA)
supported_attention_backends=(
AttentionBackendEnum.SLIDING_TILE_ATTN,
AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA,
)
)
```
#### Define supported backend selection
```python
_supported_attention_backends = (_Backend.FLASH_ATTN, _Backend.TORCH_SDPA)
_supported_attention_backends = (
AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA,
)
```
### Registering Models
+65 -94
View File
@@ -1,104 +1,85 @@
# FastVideo CLI Inference
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).
The FastVideo CLI exposes the same core inference controls as the Python API.
## Basic Usage
The basic command to generate a video is:
Use either:
1. `--model-path` + `--prompt`
2. `--model-path` + `--prompt-txt` (batch prompts, one line per prompt)
3. `--config` (JSON/YAML)
```bash
fastvideo generate --model-path {MODEL_PATH} --prompt {PROMPT}
fastvideo generate --model-path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
--prompt "A cat playing with a ball of yarn"
```
### Required Parameters
```bash
fastvideo generate --model-path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
--prompt-txt prompts.txt
```
- `--model-path {MODEL_PATH}`: Path to the model or model ID
- `--prompt {PROMPT}`: Text description for the video you want to generate
You cannot provide both `--prompt` and `--prompt-txt` in the same run.
## Common Arguments
To see all the options, you can use the `--help` flag:
## View All Arguments
```bash
fastvideo generate --help
```
### Hardware Configuration
Arguments come from:
- `--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)
- FastVideo runtime args (`FastVideoArgs`)
- Sampling args (`SamplingParam`)
- Pipeline config args (`PipelineConfig`)
#### Video Configuration
## Common Arguments
- `--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
### Parallelism
#### Generation Parameters
- `--num-gpus`
- `--sp-size`
- `--tp-size`
- `--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
### Sampling
#### Output Options
- `--num-frames`
- `--height` / `--width`
- `--num-inference-steps`
- `--guidance-scale`
- `--seed`
- `--negative-prompt`
- `--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
## Using Configuration Files
- `--output-path`
- `--save-video` / `--no-save-video`
- `--return-frames`
Instead of specifying all parameters on the command line, you can use a configuration file:
### 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
```bash
fastvideo generate --config {CONFIG_FILE_PATH}
fastvideo generate --config config.yaml
```
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.
Config files can be JSON or YAML. CLI flags override config-file values.
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):
Example `config.yaml`:
```yaml
model_path: "FastVideo/FastHunyuan-diffusers"
prompt: "A beautiful woman in a red dress walking down a street"
prompt: "A capybara lounging in a hammock"
output_path: "outputs/"
num_gpus: 2
sp_size: 2
@@ -108,44 +89,34 @@ height: 720
width: 1280
num_inference_steps: 6
seed: 1024
fps: 24
precision: "bf16"
dit_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
Generating a simple video:
Simple generation:
```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/
```
Using a negative prompt to avoid certain elements:
Config + CLI override:
```bash
fastvideo generate --model-path FastVideo/FastHunyuan-diffusers --prompt "A beautiful forest landscape" --negative-prompt "people, buildings, roads"
fastvideo generate --config config.yaml --prompt "A panda skiing at sunset"
```
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.
+31 -3
View File
@@ -1,4 +1,3 @@
# Configuration
## Multi-GPU Setup
@@ -18,7 +17,8 @@ generator = VideoGenerator.from_pretrained(
- `PipelineConfig`: Initialization time parameters
- `SamplingParam`: Generation time parameters
You can customize various parameters when generating videos using the `PipelineConfig` and `SamplingParam` class:
You can customize generation behavior using `PipelineConfig` and
`SamplingParam`:
```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,6 +72,34 @@ 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,7 +61,8 @@ python example.py
The generated video will be saved in the current directory under `my_videos/`
More inference example scripts can be found in `scripts/inference/`
More inference scripts and recipes can be found in `examples/inference/` and
`scripts/inference/`.
## Available Models
@@ -103,11 +104,10 @@ Common issues and their solutions:
If you encounter CUDA out of memory errors:
- Reduce `num_frames` or video resolution
- Enable memory optimization with `enable_model_cpu_offload`
- Enable FastVideo offloading options such as `dit_layerwise_offload=True`
(single GPU) or `use_fsdp_inference=True` (multi-GPU)
- 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
+13 -11
View File
@@ -25,6 +25,9 @@ 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
@@ -70,22 +73,15 @@ python setup.py install
**`SLIDING_TILE_ATTN`**
```bash
pip install st_attn==0.0.4
```
Please see [this page](../attention/sta/index.md) for more installation instructions.
Sliding Tile Attention is provided by `fastvideo-kernel`.
See [STA docs](../attention/sta/index.md) for installation details.
### Video Sparse Attention
**`VIDEO_SPARSE_ATTN`**
```bash
git submodule update --init --recursive
python setup_vsa.py install
```
Please see [this page](../attention/vsa/index.md) for more installation instructions.
Video Sparse Attention is provided by `fastvideo-kernel`.
See [VSA docs](../attention/vsa/index.md) for installation details.
### Sage Attention
@@ -115,6 +111,12 @@ 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.
+15 -6
View File
@@ -1,6 +1,10 @@
# Compatibility Matrix
The table below shows every supported model and optimizations supported for them.
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 symbols used have the following meanings:
@@ -10,7 +14,9 @@ The symbols used have the following meanings:
## Models x Optimization
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.
The `HuggingFace Model ID` can be passed directly to
`from_pretrained()`. FastVideo then uses model-specific default settings for
pipeline initialization and sampling.
<style>
/* Target tables in this section */
@@ -53,9 +59,9 @@ The `HuggingFace Model ID` can be directly pass to `from_pretrained()` methods a
| 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 | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
@@ -63,10 +69,13 @@ The `HuggingFace Model ID` can be directly pass to `from_pretrained()` methods a
**Note**: Wan2.2 TI2V 5B has some quality issues when performing I2V generation. We are working on fixing this issue.
## Special requirements
## Canonical Supported IDs
### StepVideo T2V
- The self-attention in text-encoder (step_llm) only supports CUDA capabilities sm_80 sm_86 and sm_90
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
### Sliding Tile Attention
- Currently only Hopper GPUs (H100s) are supported.
+75
View File
@@ -0,0 +1,75 @@
# 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`.
+3 -9
View File
@@ -20,10 +20,10 @@ def main() -> None:
# Uses FastVideo default sampling settings for LTX2 base.
generator = VideoGenerator.from_pretrained(
"Davids048/LTX2-Base-Diffusers",
num_gpus=8,
num_gpus=1,
)
output_path = "outputs_video/ltx2_basic/output_ltx2_base_t2v_1088_1920_1.4.mp4"
output_path = "outputs_video/ltx2_basic/output_ltx2_base_t2v_1088_1920_1.1.mp4"
generator.generate_video(
prompt=PROMPT,
output_path=output_path,
@@ -31,15 +31,9 @@ def main() -> None:
num_frames=121,
height=1088,
width=1920,
# LTX2 uses these parameters for multi-modal CFG instead of guidance_scale
# ltx2_cfg_scale_video=3.0,
# ltx2_cfg_scale_audio=7.0,
# ltx2_modality_scale_video=3.0,
# ltx2_modality_scale_audio=3.0,
# ltx2_rescale_scale=0.7,
)
generator.shutdown()
if __name__ == "__main__":
main()
main()
@@ -13,12 +13,13 @@ PROMPT = (
"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=1,
num_gpus=4,
)
output_path = "outputs_video/ltx2_basic/output_ltx2_distilled_t2v.mp4"
+2 -2
View File
@@ -32,13 +32,13 @@ def _safe_filename(text: str, max_len: int = 100) -> str:
def _remove_existing_outputs(out_dir: str, filename_base: str) -> None:
"""
Ensure deterministic naming by deleting any existing mp4s that would
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$")
pattern = re.compile(rf"^{re.escape(filename_base)}(_\d+)?\.(mp4|png)$")
for fn in os.listdir(out_dir):
if pattern.match(fn):
try:
@@ -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=1
num_gpu=8
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]
# Training arguments
training_args=(
# Core training arguments
core_training_args=(
--tracker_project_name wan_i2v_VSA
--output_dir "checkpoints/wan_i2v_finetune_VSA"
--max_train_steps 4000
@@ -53,10 +53,14 @@ 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
@@ -127,8 +131,9 @@ srun torchrun \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${core_training_args[@]}" \
"${validation_generation_shape_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${vsa_args[@]}"
"${vsa_args[@]}"
@@ -6,9 +6,7 @@ These are e2e example scripts for finetuning Wan2.1 T2V with VSA to accelerate i
## Make sure you have installed VSA
```bash
pip install vsa
```
Go to fastvideo-kernel/README.md for instructions.
### 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]
# Training arguments
training_args=(
# Core training arguments
core_training_args=(
--tracker_project_name wan_t2v_VSA
--output_dir "checkpoints/wan_t2v_finetune_VSA"
--max_train_steps 4000
@@ -53,10 +53,14 @@ 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
@@ -127,8 +131,9 @@ srun torchrun \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${core_training_args[@]}" \
"${validation_generation_shape_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${vsa_args[@]}"
"${vsa_args[@]}"
@@ -0,0 +1,114 @@
# 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]
# Training arguments
training_args=(
# Core training arguments
core_training_args=(
--tracker_project_name wan_t2v_VSA
--output_dir "checkpoints/wan_t2v_finetune_VSA"
--max_train_steps 4000
@@ -53,10 +53,14 @@ 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
@@ -127,8 +131,9 @@ srun torchrun \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${core_training_args[@]}" \
"${validation_generation_shape_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${vsa_args[@]}"
"${vsa_args[@]}"
@@ -1,7 +1,8 @@
#!/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 "FastVideo/Wan-Syn_77x448x832_600k" --repo_type "dataset"
python scripts/huggingface/download_hf.py --repo_id "FastVideo/Wan-Syn_77x448x832_600k" --local_dir "data/Wan-Syn_77x448x832_600k" --repo_type "dataset"
# 720P dataset
python scripts/huggingface/download_hf.py --repo_id "FastVideo/Wan-Syn_77x768x1280_250k" --local_dir "FastVideo/Wan-Syn_77x768x1280_250k" --repo_type "dataset"
python scripts/huggingface/download_hf.py --repo_id "FastVideo/Wan-Syn_77x768x1280_250k" --local_dir "data/Wan-Syn_77x768x1280_250k" --repo_type "dataset"
@@ -0,0 +1,36 @@
{
"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
}
]
}
@@ -0,0 +1,3 @@
*.err
*.out
*slurm_logs/*
@@ -7,15 +7,17 @@ These are e2e example scripts for finetuning LTX-2 on the crush-smol dataset.
### Download crush-smol dataset:
`bash examples/training/finetune/ltx2/overfit/download_dataset.sh`
`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/overfit/preprocess_ltx2_data_t2v_new.sh`
`bash examples/training/finetune/ltx2/preprocess_ltx2_data_t2v_new.sh`
### Edit the following file and run finetuning:
`bash examples/training/finetune/ltx2/overfit/finetune_t2v.sh`
`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`).
@@ -4,16 +4,14 @@ export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export TOKENIZERS_PARALLELISM=false
MODEL_PATH="Davids048/LTX2-Base-Diffusers"
# Also can use simple 1 video for overfitting experiments.
# DATA_DIR="/home/hal-jundas/codes/FastVideo/data/crush-smol"
DATA_DIR="<PATH_TO_PROCESSED_DATASET>"
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
OVERFIT_HEIGHT=480
OVERFIT_WIDTH=832
OVERFIT_FRAMES=73
HEIGHT=1088
WIDTH=1920
FRAMES=121
training_args=(
--tracker_project_name "ltx2_t2v_finetune"
@@ -22,18 +20,17 @@ training_args=(
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 10
--num_height $OVERFIT_HEIGHT
--num_width $OVERFIT_WIDTH
--num_frames $OVERFIT_FRAMES
--ltx2-first-frame-conditioning-p 0.1
--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 $NUM_GPUS
--sp_size 2
--tp_size 1
--hsdp_replicate_dim 1
--hsdp_shard_dim $NUM_GPUS
@@ -52,9 +49,9 @@ dataset_args=(
validation_args=(
--log_validation
--validation_dataset_file $VALIDATION_DATASET_FILE
--validation_steps 50
--validation_sampling_steps "50"
--validation_guidance_scale "3.0"
--validation_steps 20
--validation_sampling_steps "8"
--validation_guidance_scale "1.0"
)
optimizer_args=(
@@ -76,14 +73,15 @@ miscellaneous_args=(
--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.
export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
torchrun \
--nnodes 1 \
--master_port 29501 \
--nproc_per_node $NUM_GPUS \
fastvideo/training/ltx2_training_pipeline.py \
"${parallel_args[@]}" \
@@ -4,7 +4,7 @@ export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export TOKENIZERS_PARALLELISM=false
MODEL_PATH="/path/to/LTX-2"
MODEL_PATH="FastVideo/LTX2-Distilled-Diffusers"
DATA_DIR="data/crush-smol"
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
NUM_GPUS=1
@@ -1,13 +0,0 @@
{
"data": [
{
"caption": "The camera opens in a calm, sunlit frog yoga studio. Warm morning light washes over the wooden floor as incense smoke drifts lazily in the air. The senior frog instructor sits cross-legged at the center, eyes closed, voice deep and calm. “We are one with the pond.” All the frogs answer softly: “Ommm...” “We are one with the mud.” “Ommm...” He smiles faintly. “We are one with the flies.” A quiet pause. The camera slowly pans to the side — one frog twitches, eyes darting. Suddenly — *thwip!* — its tongue snaps out, catching a fly mid-air and pulling it into its mouth. The master exhales slowly, still serene.",
"image_path": null,
"video_path": null,
"num_inference_steps": 50,
"height": 1088,
"width": 1920,
"num_frames": 121
}
]
}
@@ -1,31 +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.",
"image_path": null,
"video_path": null,
"num_inference_steps": 50,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.",
"image_path": null,
"video_path": null,
"num_inference_steps": 50,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "A large metal cylinder is seen pressing down on a pile of colorful candies, flattening them as if they were under a hydraulic press. The candies are crushed and broken into small pieces, creating a mess on the table.",
"image_path": null,
"video_path": null,
"num_inference_steps": 50,
"height": 480,
"width": 832,
"num_frames": 77
}
]
}
@@ -1,21 +1,13 @@
#!/bin/bash
GPU_NUM=1
MODEL_PATH="Davids048/LTX2-Base-Diffusers"
# DATASET_PATH="data/overfit"
GPU_NUM=4
MODEL_PATH="FastVideo/LTX2-Distilled-Diffusers"
DATASET_PATH="data/crush-smol"
OUTPUT_DIR="$DATASET_PATH"
WITH_AUDIO=true
# Convert one-file overfit metadata into merged format if needed.
if [ ! -f "$DATASET_PATH/videos2caption.json" ] && [ -f "$DATASET_PATH/overfit.json" ]; then
python scripts/dataset_preparation/convert_to_merged_dataset.py \
--items-json "$DATASET_PATH/overfit.json" \
--output-dir "$DATASET_PATH"
fi
torchrun --nproc_per_node=$GPU_NUM \
torchrun --nproc_per_node=$GPU_NUM \
--master_port=29513 \
-m fastvideo.pipelines.preprocess.v1_preprocessing_new \
--model_path $MODEL_PATH \
@@ -28,8 +20,8 @@ torchrun --nproc_per_node=$GPU_NUM \
--preprocess.with_audio $WITH_AUDIO \
--preprocess.preprocess_video_batch_size 1 \
--preprocess.dataloader_num_workers 0 \
--preprocess.max_height 480 \
--preprocess.max_width 832 \
--preprocess.num_frames 73 \
--preprocess.train_fps 16 \
--preprocess.max_height 1088 \
--preprocess.max_width 1920 \
--preprocess.num_frames 121 \
--preprocess.train_fps 24 \
--preprocess.video_length_tolerance_range 5
@@ -0,0 +1,16 @@
{
"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."
}
]
}
@@ -350,7 +350,7 @@ def select_best_mask_strategy(
def save_mask_search_results(
mask_search_final_result: list[dict[str, list[float]]],
mask_search_final_result: list[Any],
prompt: str,
mask_strategies: list[str],
output_dir: str = 'output/mask_search_result/') -> str | None:
@@ -358,8 +358,9 @@ def save_mask_search_results(
print("No mask search results to save")
return None
# Create result dictionary with defaultdict for nested lists
mask_search_dict: dict[str, dict[str, list[list[float]]]] = {
# Create result dictionary with nested lists:
# [timesteps][layers][heads].
mask_search_dict: dict[str, dict[str, list[list[list[float]]]]] = {
"L2_loss": defaultdict(list),
"L1_loss": defaultdict(list)
}
@@ -371,23 +372,79 @@ 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)
# Process L2 loss
step_results: list[list[float]] = []
for step_data in mask_search_final_result:
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
l2_step_results: list[list[list[float]]] = []
l1_step_results: list[list[list[float]]] = []
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
l2_layer_losses = _extract_timestep_layer_losses(
step_data, "L2_loss", i)
if l2_layer_losses:
l2_step_results.append(l2_layer_losses)
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
# Create the output directory if it doesn't exist
os.makedirs(output_dir, exist_ok=True)
+49 -1
View File
@@ -92,11 +92,59 @@ 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
from fastvideo.attention.utils.flash_attn_no_pad import (
flash_attn_no_pad, flash_attn_varlen_qk_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,3 +97,60 @@ 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
+2 -4
View File
@@ -5,13 +5,11 @@ 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", "StepVideoConfig", "CosmosVideoConfig",
"Cosmos25VideoConfig", "LongCatVideoConfig", "LTX2VideoConfig",
"HYWorldConfig"
"WanVideoConfig", "CosmosVideoConfig", "Cosmos25VideoConfig",
"LongCatVideoConfig", "LTX2VideoConfig", "HYWorldConfig"
]
+1 -1
View File
@@ -21,7 +21,7 @@ class SD3Transformer2DArchConfig(DiTArchConfig):
qk_norm: str = "rms_norm"
in_channels: int = 16
out_channels: int = 16
num_attention_heads = 24
num_attention_heads: int = 24
@dataclass
@@ -1,64 +0,0 @@
# 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,6 +7,19 @@ 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(
@@ -35,6 +48,10 @@ 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"
@@ -4,14 +4,12 @@ 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
View File
@@ -16,6 +16,7 @@ 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
@@ -1,29 +0,0 @@
# 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
+3 -4
View File
@@ -8,7 +8,6 @@ from fastvideo.configs.pipelines.hunyuangamecraft import HunyuanGameCraftPipelin
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)
@@ -17,7 +16,7 @@ __all__ = [
"HunyuanConfig", "FastHunyuanConfig", "HunyuanGameCraftPipelineConfig",
"PipelineConfig", "Hunyuan15T2V480PConfig", "Hunyuan15T2V720PConfig",
"SlidingTileAttnConfig", "WanT2V480PConfig", "WanI2V480PConfig",
"WanT2V720PConfig", "WanI2V720PConfig", "StepVideoT2VConfig",
"SelfForcingWanT2V480PConfig", "CosmosConfig", "Cosmos25Config",
"LTX2T2VConfig", "HYWorldConfig", "get_pipeline_config_cls_from_name"
"WanT2V720PConfig", "WanI2V720PConfig", "SelfForcingWanT2V480PConfig",
"CosmosConfig", "Cosmos25Config", "LTX2T2VConfig", "HYWorldConfig",
"get_pipeline_config_cls_from_name"
]
-28
View File
@@ -76,11 +76,6 @@ 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
@@ -187,29 +182,6 @@ 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",
@@ -57,10 +57,18 @@ def llama_preprocess_text(prompt: str) -> str:
def llama_postprocess_text(outputs: BaseEncoderOutput) -> torch.Tensor:
"""Extract hidden states from LLaMA output, skipping instruction tokens."""
hidden_state_skip_layer = 2
assert outputs.hidden_states is not None
hidden_states: tuple[torch.Tensor, ...] = outputs.hidden_states
last_hidden_state: torch.Tensor = hidden_states[-(hidden_state_skip_layer +
1)]
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
+18 -3
View File
@@ -18,7 +18,16 @@ from fastvideo.configs.models.vaes.autoencoder_kl import AutoencoderKLVAEConfig
from fastvideo.configs.pipelines.base import PipelineConfig, preprocess_text
def _sd35_text_postprocess(outputs: BaseEncoderOutput) -> torch.Tensor:
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
@@ -56,10 +65,16 @@ class SD35Config(PipelineConfig):
postprocess_text_funcs: tuple[
Callable[[BaseEncoderOutput], torch.Tensor],
...] = field(default_factory=lambda:
(_sd35_text_postprocess, _sd35_text_postprocess,
_sd35_text_postprocess))
(_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
@@ -1,30 +0,0 @@
# 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"
-20
View File
@@ -1,20 +0,0 @@
# 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 = "画面暗、低分辨率、不良手、文本、缺少手指、多余的手指、裁剪、低质量、颗粒状、签名、水印、用户名、模糊。"
+52 -8
View File
@@ -95,8 +95,9 @@ 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(path, "map_style_cache")
cache_dir = os.path.join(dataset_root, "map_style_cache")
cache_file = os.path.join(cache_dir, "file_info.pkl")
# Only rank 0 checks for cache and scans files if needed
@@ -111,8 +112,39 @@ def get_parquet_files_and_length(path: str):
try:
with open(cache_file, "rb") as f:
file_names_sorted, lengths_sorted = pickle.load(f)
cache_loaded = True
logger.info("Successfully loaded cached file info")
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")
except Exception as e:
logger.error("Error loading cached file info: %s", str(e))
logger.info("Falling back to scanning files")
@@ -123,11 +155,17 @@ 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(path):
for root, _, files in os.walk(dataset_root):
for file in sorted(files):
if file.endswith('.parquet'):
file_path = os.path.join(root, file)
file_path = os.path.realpath(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
@@ -138,9 +176,6 @@ 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:
@@ -155,6 +190,15 @@ 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
+4 -3
View File
@@ -990,9 +990,10 @@ def maybe_init_distributed_environment_and_model_parallel(
# set device if we're on a CUDA/NPU platform
from fastvideo.platforms import current_platform
device_type = current_platform.device_type
device = torch.device(f"{device_type}:{local_rank}")
current_platform.get_torch_device().set_device(device)
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)
def model_parallel_is_initialized() -> bool:
+8 -20
View File
@@ -6,7 +6,6 @@ This module provides a consolidated interface for generating videos using
diffusion models.
"""
import math
import os
import re
import threading
@@ -336,30 +335,19 @@ 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 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 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 orig_latent_num_frames % fastvideo_args.num_gpus != 0:
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
# 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
logger.info(
"Adjusting number of frames from %s to %s based on number of GPUs (%s)",
-5
View File
@@ -257,11 +257,6 @@ 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(
+9 -9
View File
@@ -126,26 +126,26 @@ class CausalWanSelfAttention(nn.Module):
# If we are using local attention and the current KV cache size is larger than the local attention size, we need to truncate the KV cache
kv_cache_size = kv_cache["k"].shape[1]
num_new_tokens = roped_query.shape[1]
if self.local_attn_size != -1 and (current_end > kv_cache["global_end_index"].item()) and (
num_new_tokens + kv_cache["local_end_index"].item() > kv_cache_size):
if self.local_attn_size != -1 and (current_end > kv_cache["global_end_index"]) and (
num_new_tokens + kv_cache["local_end_index"] > kv_cache_size):
# Calculate the number of new tokens added in this step
# Shift existing cache content left to discard oldest tokens
# Clone the source slice to avoid overlapping memory error
num_evicted_tokens = num_new_tokens + kv_cache["local_end_index"].item() - kv_cache_size
num_rolled_tokens = kv_cache["local_end_index"].item() - num_evicted_tokens - sink_tokens
num_evicted_tokens = num_new_tokens + kv_cache["local_end_index"] - kv_cache_size
num_rolled_tokens = kv_cache["local_end_index"] - num_evicted_tokens - sink_tokens
kv_cache["k"][:, sink_tokens:sink_tokens + num_rolled_tokens] = \
kv_cache["k"][:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
kv_cache["v"][:, sink_tokens:sink_tokens + num_rolled_tokens] = \
kv_cache["v"][:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
# Insert the new keys/values at the end
local_end_index = kv_cache["local_end_index"].item() + current_end - \
kv_cache["global_end_index"].item() - num_evicted_tokens
local_end_index = kv_cache["local_end_index"] + current_end - \
kv_cache["global_end_index"] - num_evicted_tokens
local_start_index = local_end_index - num_new_tokens
kv_cache["k"][:, local_start_index:local_end_index] = roped_key
kv_cache["v"][:, local_start_index:local_end_index] = v
else:
# Assign new keys/values directly up to current_end
local_end_index = kv_cache["local_end_index"].item() + current_end - kv_cache["global_end_index"].item()
local_end_index = kv_cache["local_end_index"] + current_end - kv_cache["global_end_index"]
local_start_index = local_end_index - num_new_tokens
kv_cache["k"] = kv_cache["k"].detach()
kv_cache["v"] = kv_cache["v"].detach()
@@ -157,8 +157,8 @@ class CausalWanSelfAttention(nn.Module):
kv_cache["k"][:, max(0, local_end_index - self.max_attention_size):local_end_index],
kv_cache["v"][:, max(0, local_end_index - self.max_attention_size):local_end_index]
)
kv_cache["global_end_index"].fill_(current_end)
kv_cache["local_end_index"].fill_(local_end_index)
kv_cache["global_end_index"] = current_end
kv_cache["local_end_index"] = local_end_index
return x
+1 -1
View File
@@ -602,7 +602,7 @@ class LongCatCrossAttention(nn.Module):
return out
# === Standard cross-attention ===
# Project Q, K, V (standard cross-attention like WanVideo/StepVideo/Cosmos)
# Project Q, K, V (standard cross-attention like WanVideo/Cosmos)
q, _ = self.to_q(x)
k, _ = self.to_k(context)
v, _ = self.to_v(context)
+120 -13
View File
@@ -855,7 +855,13 @@ class TransformerArgsPreprocessor:
)
def prepare(self, modality: Modality) -> TransformerArgs:
x = self.patchify_proj(modality.latent)
# `rearrange`-produced latent tokens can be non-contiguous with large
# leading strides. Materialize a contiguous view before Linear to avoid
# invalid BF16 GEMM descriptors on some CUDA/cuBLAS stacks.
latent = modality.latent
if not latent.is_contiguous():
latent = latent.contiguous()
x = self.patchify_proj(latent)
timestep, embedded_timestep = self._prepare_timestep(modality.timesteps, x.shape[0], modality.latent.dtype)
context, attention_mask = self._prepare_context(modality.context, x, modality.context_mask)
attention_mask = self._prepare_attention_mask(attention_mask, modality.latent.dtype)
@@ -1017,6 +1023,7 @@ class LTXDistributedAttention(DistributedAttention):
replicated_q: torch.Tensor | None = None,
replicated_k: torch.Tensor | None = None,
replicated_v: torch.Tensor | None = None,
gate_compress: torch.Tensor | None = None,
attention_mask: torch.Tensor | None = None,
ltx_freqs_cis: tuple[torch.Tensor, torch.Tensor] | None = None,
) -> tuple[torch.Tensor, torch.Tensor | None]:
@@ -1043,6 +1050,56 @@ class LTXDistributedAttention(DistributedAttention):
forward_context: ForwardContext = get_forward_context()
ctx_attn_metadata = forward_context.attn_metadata
use_vsa = self.backend == AttentionBackendEnum.VIDEO_SPARSE_ATTN
if use_vsa:
assert (
replicated_q is None and replicated_k is None
and replicated_v is None
), "Replicated QKV is not supported for VSA now"
if gate_compress is None:
raise ValueError(
"gate_compress must be provided when using VIDEO_SPARSE_ATTN"
)
qkvg = torch.cat([q, k, v, gate_compress],
dim=0) # [4*batch, seq_len, heads, head_dim]
qkvg = sequence_model_parallel_all_to_all_4D(qkvg,
scatter_dim=2,
gather_dim=1)
valid_seq_len = None
if attention_mask is not None:
valid_seq_len = (attention_mask[0] == 1).sum().item()
qkvg = qkvg[:, :valid_seq_len, :, :]
if ltx_freqs_cis is not None:
cos, sin = ltx_freqs_cis
heads_per_rank = num_heads // world_size
head_start = local_rank * heads_per_rank
head_end = head_start + heads_per_rank
cos_local = cos[:, head_start:head_end, :valid_seq_len, :]
sin_local = sin[:, head_start:head_end, :valid_seq_len, :]
qk_part = qkvg[:batch_size * 2].transpose(1, 2)
qk_part = apply_ltx_rotary_emb_4d(
qk_part, (cos_local, sin_local), self.rope_type)
qkvg[:batch_size * 2] = qk_part.transpose(1, 2)
qkvg = self.attn_impl.preprocess_qkv(qkvg, ctx_attn_metadata)
q, k, v, gate_compress = qkvg.chunk(4, dim=0)
output = self.attn_impl.forward( # type: ignore[call-arg]
q, k, v, gate_compress, ctx_attn_metadata)
output = self.attn_impl.postprocess_output(output, ctx_attn_metadata)
if attention_mask is not None:
pad_len = (attention_mask[0] == 0).sum().item()
output = torch.nn.functional.pad(output, (0, 0, 0, 0, 0, pad_len))
output = sequence_model_parallel_all_to_all_4D(output,
scatter_dim=1,
gather_dim=2)
return output, None
# Stack QKV
qkv = torch.cat([q, k, v], dim=0) # [3*batch, seq_len, num_heads, head_dim]
@@ -1143,6 +1200,7 @@ class LTXLocalAttention(LocalAttention):
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
gate_compress: torch.Tensor | None = None,
ltx_freqs_cis: tuple[torch.Tensor, torch.Tensor] | None = None,
ltx_k_freqs_cis: tuple[torch.Tensor, torch.Tensor] | None = None,
) -> torch.Tensor:
@@ -1176,7 +1234,15 @@ class LTXLocalAttention(LocalAttention):
q = q.transpose(1, 2)
k = k.transpose(1, 2)
output = self.attn_impl.forward(q, k, v, ctx_attn_metadata)
if self.backend == AttentionBackendEnum.VIDEO_SPARSE_ATTN:
if gate_compress is None:
raise ValueError(
"gate_compress must be provided when using VIDEO_SPARSE_ATTN"
)
output = self.attn_impl.forward( # type: ignore[call-arg]
q, k, v, gate_compress, ctx_attn_metadata)
else:
output = self.attn_impl.forward(q, k, v, ctx_attn_metadata)
return output
@@ -1216,6 +1282,9 @@ class LTXSelfAttention(nn.Module):
causal=False,
supported_attention_backends=supported_attention_backends,
)
self.to_gate_compress: nn.Linear | None = None
if self.attn.backend == AttentionBackendEnum.VIDEO_SPARSE_ATTN:
self.to_gate_compress = nn.Linear(context_dim, inner_dim, bias=True)
self.attn_masked = LTXLocalAttention(
num_heads=heads,
head_size=dim_head,
@@ -1238,6 +1307,8 @@ class LTXSelfAttention(nn.Module):
context = x if context is None else context
k = self.to_k(context)
v = self.to_v(context)
gate_compress = (self.to_gate_compress(context)
if self.to_gate_compress is not None else None)
q = self.q_norm(q)
k = self.k_norm(k)
@@ -1248,6 +1319,9 @@ class LTXSelfAttention(nn.Module):
q = q.view(b, q_len, self.heads, self.dim_head)
k = k.view(b, k_len, self.heads, self.dim_head)
v = v.view(b, k_len, self.heads, self.dim_head)
if gate_compress is not None:
gate_compress = gate_compress.view(b, k_len, self.heads,
self.dim_head)
if mask is not None:
if mask.ndim == 2:
@@ -1270,9 +1344,18 @@ class LTXSelfAttention(nn.Module):
attn_metadata=attn_metadata,
forward_batch=forward_batch,
):
out = self.attn_masked(q, k, v, ltx_freqs_cis=pe, ltx_k_freqs_cis=k_pe)
out = self.attn_masked(q,
k,
v,
ltx_freqs_cis=pe,
ltx_k_freqs_cis=k_pe)
else:
out = self.attn(q, k, v, ltx_freqs_cis=pe, ltx_k_freqs_cis=k_pe)
out = self.attn(q,
k,
v,
gate_compress=gate_compress,
ltx_freqs_cis=pe,
ltx_k_freqs_cis=k_pe)
out = out.reshape(b, q_len, -1)
return self.to_out(out)
@@ -1314,6 +1397,9 @@ class LTXDistributedSelfAttention(nn.Module):
supported_attention_backends=supported_attention_backends,
prefix=f"{prefix}.attn",
)
self.to_gate_compress: nn.Linear | None = None
if self.attn.backend == AttentionBackendEnum.VIDEO_SPARSE_ATTN:
self.to_gate_compress = nn.Linear(context_dim, inner_dim, bias=True)
def forward(
self,
@@ -1337,6 +1423,8 @@ class LTXDistributedSelfAttention(nn.Module):
context = x if context is None else context
k = self.to_k(context)
v = self.to_v(context)
gate_compress = (self.to_gate_compress(context)
if self.to_gate_compress is not None else None)
q = self.q_norm(q)
k = self.k_norm(k)
@@ -1348,11 +1436,15 @@ class LTXDistributedSelfAttention(nn.Module):
q = q.view(b, q_len, self.heads, self.dim_head)
k = k.view(b, k_len, self.heads, self.dim_head)
v = v.view(b, k_len, self.heads, self.dim_head)
if gate_compress is not None:
gate_compress = gate_compress.view(b, k_len, self.heads,
self.dim_head)
# Pass full RoPE to distributed attention - it will apply after all-to-all
# and slice to this rank's heads
out, _ = self.attn(
q, k, v,
gate_compress=gate_compress,
attention_mask=attention_mask,
ltx_freqs_cis=pe,
)
@@ -1383,6 +1475,15 @@ class BasicAVTransformerBlock(torch.nn.Module):
# Text cross-attention uses LocalAttention (text embeddings are replicated)
SelfAttnCls = LTXDistributedSelfAttention if use_distributed_attention else LTXSelfAttention
CrossAttnCls = LTXSelfAttention # Text cross-attention is always local
video_self_attn_backends = (
AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA,
AttentionBackendEnum.VIDEO_SPARSE_ATTN,
)
dense_attn_backends = (
AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA,
)
if video is not None:
# Video self-attention - use distributed when SP > 1
@@ -1393,7 +1494,7 @@ class BasicAVTransformerBlock(torch.nn.Module):
dim_head=video.d_head,
norm_eps=norm_eps,
rope_type=rope_type,
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA),
supported_attention_backends=video_self_attn_backends,
prefix=f"{prefix}.blocks.{idx}.attn1" if use_distributed_attention else "",
) if use_distributed_attention else LTXSelfAttention(
query_dim=video.dim,
@@ -1402,7 +1503,7 @@ class BasicAVTransformerBlock(torch.nn.Module):
dim_head=video.d_head,
norm_eps=norm_eps,
rope_type=rope_type,
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA),
supported_attention_backends=video_self_attn_backends,
)
# Text cross-attention - always local (text is replicated)
self.attn2 = CrossAttnCls(
@@ -1412,7 +1513,7 @@ class BasicAVTransformerBlock(torch.nn.Module):
dim_head=video.d_head,
norm_eps=norm_eps,
rope_type=rope_type,
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA),
supported_attention_backends=dense_attn_backends,
)
self.ff = FeedForward(video.dim, dim_out=video.dim)
self.scale_shift_table = torch.nn.Parameter(torch.empty(6, video.dim))
@@ -1426,7 +1527,7 @@ class BasicAVTransformerBlock(torch.nn.Module):
dim_head=audio.d_head,
norm_eps=norm_eps,
rope_type=rope_type,
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA),
supported_attention_backends=dense_attn_backends,
prefix=f"{prefix}.blocks.{idx}.audio_attn1" if use_distributed_attention else "",
) if use_distributed_attention else LTXSelfAttention(
query_dim=audio.dim,
@@ -1435,7 +1536,7 @@ class BasicAVTransformerBlock(torch.nn.Module):
dim_head=audio.d_head,
norm_eps=norm_eps,
rope_type=rope_type,
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA),
supported_attention_backends=dense_attn_backends,
)
# Text cross-attention - always local (text is replicated)
self.audio_attn2 = CrossAttnCls(
@@ -1445,7 +1546,7 @@ class BasicAVTransformerBlock(torch.nn.Module):
dim_head=audio.d_head,
norm_eps=norm_eps,
rope_type=rope_type,
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA),
supported_attention_backends=dense_attn_backends,
)
self.audio_ff = FeedForward(audio.dim, dim_out=audio.dim)
self.audio_scale_shift_table = torch.nn.Parameter(torch.empty(6, audio.dim))
@@ -1460,7 +1561,7 @@ class BasicAVTransformerBlock(torch.nn.Module):
dim_head=audio.d_head,
norm_eps=norm_eps,
rope_type=rope_type,
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA),
supported_attention_backends=dense_attn_backends,
)
# Video-to-audio cross-attention
# Uses local attention - context is gathered from all SP ranks in forward()
@@ -1471,7 +1572,7 @@ class BasicAVTransformerBlock(torch.nn.Module):
dim_head=audio.d_head,
norm_eps=norm_eps,
rope_type=rope_type,
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA),
supported_attention_backends=dense_attn_backends,
)
self.scale_shift_table_a2v_ca_audio = torch.nn.Parameter(torch.empty(5, audio.dim))
self.scale_shift_table_a2v_ca_video = torch.nn.Parameter(torch.empty(5, video.dim))
@@ -2120,7 +2221,9 @@ class LTX2Transformer3DModel(CachableDiT):
# Get SP world size for distributed attention
sp_world_size = get_sp_world_size()
use_distributed_attention = sp_world_size > 1
use_vsa_backend = os.getenv("FASTVIDEO_ATTENTION_BACKEND",
"") == "VIDEO_SPARSE_ATTN"
use_distributed_attention = sp_world_size > 1 or use_vsa_backend
# Validate that attention heads are divisible by SP world size
if sp_world_size > 1:
@@ -2135,6 +2238,10 @@ class LTX2Transformer3DModel(CachableDiT):
logger.info(
f"LTX2 sequence parallelism enabled with SP world size {sp_world_size}"
)
elif use_vsa_backend:
logger.info(
"LTX2 VSA enabled with SP world size 1; using distributed "
"attention path for VSA compatibility")
model_type = LTXModelType.AudioVideo
self.model = LTXModel(
@@ -95,7 +95,7 @@ def _update_kv_cache_and_attend(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
kv_cache: dict[str, torch.Tensor],
kv_cache: dict[str, torch.Tensor | int],
attn_layer: LocalAttention,
start_frame: int,
num_frame_per_block: int,
@@ -139,15 +139,18 @@ def _update_kv_cache_and_attend(
kv_cache_size = kv_cache["k"].shape[1]
num_new_tokens = k.shape[1] if use_k_for_num_tokens else q.shape[1]
original_global_end_index = kv_cache["global_end_index"]
original_local_end_index = kv_cache["local_end_index"]
# Check if we need to evict tokens
if (current_end > kv_cache["global_end_index"].item()) and (
num_new_tokens + kv_cache["local_end_index"].item() > kv_cache_size
if (current_end > original_global_end_index) and (
num_new_tokens + original_local_end_index > kv_cache_size
):
num_evicted_tokens = (
num_new_tokens + kv_cache["local_end_index"].item() - kv_cache_size
num_new_tokens + original_local_end_index - kv_cache_size
)
num_rolled_tokens = (
kv_cache["local_end_index"].item()
original_local_end_index
- num_evicted_tokens
- sink_tokens
)
@@ -171,18 +174,18 @@ def _update_kv_cache_and_attend(
)
# Calculate indices with eviction adjustment
local_end_index = (
kv_cache["local_end_index"].item()
original_local_end_index
+ current_end
- kv_cache["global_end_index"].item()
- original_global_end_index
- num_evicted_tokens
)
local_start_index = local_end_index - num_new_tokens
else:
# Calculate indices without eviction
local_end_index = (
kv_cache["local_end_index"].item()
original_local_end_index
+ current_end
- kv_cache["global_end_index"].item()
- original_global_end_index
)
local_start_index = local_end_index - num_new_tokens
@@ -206,8 +209,8 @@ def _update_kv_cache_and_attend(
attn = attn_layer(q, cached_k, cached_v)
# Update indices
kv_cache["global_end_index"].fill_(current_end)
kv_cache["local_end_index"].fill_(local_end_index)
kv_cache["global_end_index"] = current_end
kv_cache["local_end_index"] = local_end_index
return attn
@@ -49,7 +49,9 @@ from .model import MatrixGameCrossAttention
logger = init_logger(__name__)
def causal_rope_apply(x, grid_sizes, freqs, start_frame=0):
def causal_rope_apply(
x: torch.Tensor, grid_sizes: tuple[int, int, int], freqs: torch.Tensor, start_frame: int = 0
):
n, c = x.size(2), x.size(3) // 2
# split freqs
@@ -57,7 +59,7 @@ def causal_rope_apply(x, grid_sizes, freqs, start_frame=0):
# loop over samples
output = []
f, h, w = grid_sizes.tolist()
f, h, w = grid_sizes
for i in range(len(x)):
seq_len = f * h * w
@@ -176,16 +178,16 @@ class CausalMatrixGameSelfAttention(nn.Module):
v: torch.Tensor,
freqs_cis: tuple[torch.Tensor, torch.Tensor],
block_mask: BlockMask,
grid_sizes: tuple[int, int, int],
kv_cache: dict | None = None,
current_start: int = 0,
cache_start: int | None = None,
grid_sizes: torch.Tensor | None = None,
):
if cache_start is None:
cache_start = current_start
# Calculate start_frame for causal mode
if kv_cache is not None and grid_sizes is not None:
if kv_cache is not None:
frame_seqlen = int(grid_sizes[1] * grid_sizes[2])
start_frame = current_start // frame_seqlen
else:
@@ -279,17 +281,14 @@ class CausalMatrixGameSelfAttention(nn.Module):
kv_cache_size = kv_cache["k"].shape[1]
num_new_tokens = roped_query.shape[1]
if (current_end > kv_cache["global_end_index"].item()) and (
num_new_tokens + kv_cache["local_end_index"].item()
> kv_cache_size
if (current_end > kv_cache["global_end_index"]) and (
num_new_tokens + kv_cache["local_end_index"] > kv_cache_size
):
num_evicted_tokens = (
num_new_tokens
+ kv_cache["local_end_index"].item()
- kv_cache_size
num_new_tokens + kv_cache["local_end_index"] - kv_cache_size
)
num_rolled_tokens = (
kv_cache["local_end_index"].item()
kv_cache["local_end_index"]
- num_evicted_tokens
- sink_tokens
)
@@ -310,9 +309,9 @@ class CausalMatrixGameSelfAttention(nn.Module):
+ num_rolled_tokens,
].clone()
local_end_index = (
kv_cache["local_end_index"].item()
kv_cache["local_end_index"]
+ current_end
- kv_cache["global_end_index"].item()
- kv_cache["global_end_index"]
- num_evicted_tokens
)
local_start_index = local_end_index - num_new_tokens
@@ -320,9 +319,9 @@ class CausalMatrixGameSelfAttention(nn.Module):
kv_cache["v"][:, local_start_index:local_end_index] = v
else:
local_end_index = (
kv_cache["local_end_index"].item()
kv_cache["local_end_index"]
+ current_end
- kv_cache["global_end_index"].item()
- kv_cache["global_end_index"]
)
local_start_index = local_end_index - num_new_tokens
@@ -339,8 +338,8 @@ class CausalMatrixGameSelfAttention(nn.Module):
v_for_attn.transpose(1, 2),
).transpose(1, 2)
kv_cache["global_end_index"].fill_(current_end)
kv_cache["local_end_index"].fill_(local_end_index)
kv_cache["global_end_index"] = current_end
kv_cache["local_end_index"] = local_end_index
return x
@@ -460,7 +459,7 @@ class CausalMatrixGameTransformerBlock(nn.Module):
temb: torch.Tensor,
freqs_cis: tuple[torch.Tensor, torch.Tensor],
block_mask: BlockMask,
grid_sizes: torch.Tensor,
grid_sizes: tuple[int, int, int],
mouse_cond: torch.Tensor | None = None,
keyboard_cond: torch.Tensor | None = None,
block_mask_mouse: BlockMask | None = None,
@@ -516,10 +515,10 @@ class CausalMatrixGameTransformerBlock(nn.Module):
value,
freqs_cis,
block_mask,
grid_sizes,
kv_cache,
current_start,
cache_start,
grid_sizes,
)
attn_output = attn_output.flatten(2)
attn_output, _ = self.to_out(attn_output)
@@ -546,9 +545,9 @@ class CausalMatrixGameTransformerBlock(nn.Module):
hidden_states = self.action_model(
hidden_states,
int(grid_sizes[0]),
int(grid_sizes[1]),
int(grid_sizes[2]),
grid_sizes[0],
grid_sizes[1],
grid_sizes[2],
mouse_cond,
keyboard_cond,
block_mask_mouse,
@@ -981,10 +980,7 @@ class CausalMatrixGameWanModel(BaseDiT):
freqs_cis = (freqs_cos, freqs_sin) if freqs_cos is not None else None
hidden_states = self.patch_embedding(hidden_states)
grid_sizes = torch.tensor(
[post_patch_num_frames, post_patch_height, post_patch_width],
device=hidden_states.device,
)
grid_sizes = (post_patch_num_frames, post_patch_height, post_patch_width)
hidden_states = hidden_states.flatten(2).transpose(1, 2)
if timestep.dim() == 2:
@@ -1200,9 +1196,9 @@ class CausalMatrixGameWanModel(BaseDiT):
hidden_states.shape
)
p_t, p_h, p_w = self.patch_size
post_patch_num_frames = num_frames // p_t
post_patch_height = height // p_h
post_patch_width = width // p_w
post_patch_num_frames: int = num_frames // p_t
post_patch_height: int = height // p_h
post_patch_width: int = width // p_w
d = self.hidden_size // self.num_attention_heads
rope_dim_list = [d - 4 * (d // 6), 2 * (d // 6), 2 * (d // 6)]
@@ -1263,10 +1259,7 @@ class CausalMatrixGameWanModel(BaseDiT):
)
hidden_states = self.patch_embedding(hidden_states)
grid_sizes = torch.tensor(
[post_patch_num_frames, post_patch_height, post_patch_width],
device=hidden_states.device,
)
grid_sizes = (post_patch_num_frames, post_patch_height, post_patch_width)
hidden_states = hidden_states.flatten(2).transpose(1, 2)
if timestep.dim() == 2:
@@ -1393,10 +1386,10 @@ class CausalMatrixGameWanModel(BaseDiT):
else:
return self._forward_train(*args, **kwargs)
def unpatchify(self, x, grid_sizes):
def unpatchify(self, x: torch.Tensor, grid_sizes: tuple[int, int, int]) -> torch.Tensor:
c = self.proj_out.out_features // math.prod(self.patch_size)
p_t, p_h, p_w = self.patch_size
f, h, w = grid_sizes.tolist()
f, h, w = grid_sizes
x = x[:, : f * h * w].view(-1, f, h, w, p_t, p_h, p_w, c)
x = x.permute(0, 7, 1, 4, 2, 5, 3, 6)
File diff suppressed because it is too large Load Diff
-690
View File
@@ -1,690 +0,0 @@
# Copyright 2025 StepFun Inc. All Rights Reserved.
#
# Permission is hereby granted, free of charge, to any person obtaining a copy
# of this software and associated documentation files (the "Software"), to deal
# in the Software without restriction, including without limitation the rights
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
# copies of the Software, and to permit persons to whom the Software is
# furnished to do so, subject to the following conditions:
#
# The above copyright notice and this permission notice shall be included in all
# copies or substantial portions of the Software.
# ==============================================================================
from typing import Any
import torch
from einops import rearrange, repeat
from torch import nn
from fastvideo.attention import DistributedAttention, LocalAttention
from fastvideo.configs.models.dits import StepVideoConfig
from fastvideo.distributed.parallel_state import get_sp_world_size
from fastvideo.layers.layernorm import LayerNormScaleShift
from fastvideo.layers.linear import ReplicatedLinear
from fastvideo.layers.mlp import MLP
from fastvideo.layers.rotary_embedding import (_apply_rotary_emb,
get_rotary_pos_embed)
from fastvideo.layers.visual_embedding import TimestepEmbedder
from fastvideo.models.dits.base import BaseDiT
from fastvideo.platforms import AttentionBackendEnum
class PatchEmbed2D(nn.Module):
"""2D Image to Patch Embedding
Image to Patch Embedding using Conv2d
A convolution based approach to patchifying a 2D image w/ embedding projection.
Based on the impl in https://github.com/google-research/vision_transformer
Hacked together by / Copyright 2020 Ross Wightman
Remove the _assert function in forward function to be compatible with multi-resolution images.
"""
def __init__(self,
patch_size=16,
in_chans=3,
embed_dim=768,
norm_layer=None,
flatten=True,
bias=True,
dtype=None,
prefix: str = ""):
super().__init__()
# Convert patch_size to 2-tuple
if isinstance(patch_size, list | tuple):
if len(patch_size) == 1:
patch_size = (patch_size[0], patch_size[0])
else:
patch_size = (patch_size, patch_size)
self.patch_size = patch_size
self.flatten = flatten
self.proj = nn.Conv2d(in_chans,
embed_dim,
kernel_size=patch_size,
stride=patch_size,
bias=bias,
dtype=dtype)
self.norm = norm_layer(embed_dim) if norm_layer else nn.Identity()
def forward(self, x):
x = self.proj(x)
if self.flatten:
x = x.flatten(2).transpose(1, 2) # BCHW -> BNC
x = self.norm(x)
return x
class StepVideoRMSNorm(nn.Module):
def __init__(
self,
dim: int,
elementwise_affine=True,
eps: float = 1e-6,
device=None,
dtype=None,
):
"""
Initialize the RMSNorm normalization layer.
Args:
dim (int): The dimension of the input tensor.
eps (float, optional): A small value added to the denominator for numerical stability. Default is 1e-6.
Attributes:
eps (float): A small value added to the denominator for numerical stability.
weight (nn.Parameter): Learnable scaling parameter.
"""
factory_kwargs = {"device": device, "dtype": dtype}
super().__init__()
self.eps = eps
if elementwise_affine:
self.weight = nn.Parameter(torch.ones(dim, **factory_kwargs))
def _norm(self, x) -> torch.Tensor:
"""
Apply the RMSNorm normalization to the input tensor.
Args:
x (torch.Tensor): The input tensor.
Returns:
torch.Tensor: The normalized tensor.
"""
return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)
def forward(self, x):
"""
Forward pass through the RMSNorm layer.
Args:
x (torch.Tensor): The input tensor.
Returns:
torch.Tensor: The output tensor after applying RMSNorm.
"""
output = self._norm(x.float()).type_as(x)
if hasattr(self, "weight"):
output = output * self.weight
return output
class SelfAttention(nn.Module):
def __init__(
self,
hidden_dim,
head_dim,
rope_split: tuple[int, int, int] = (64, 32, 32),
bias: bool = False,
with_rope: bool = True,
with_qk_norm: bool = True,
attn_type: str = "torch",
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA)):
super().__init__()
self.head_dim = head_dim
self.hidden_dim = hidden_dim
self.rope_split = list(rope_split)
self.n_heads = hidden_dim // head_dim
self.wqkv = ReplicatedLinear(hidden_dim, hidden_dim * 3, bias=bias)
self.wo = ReplicatedLinear(hidden_dim, hidden_dim, bias=bias)
self.with_rope = with_rope
self.with_qk_norm = with_qk_norm
if self.with_qk_norm:
self.q_norm = StepVideoRMSNorm(head_dim, elementwise_affine=True)
self.k_norm = StepVideoRMSNorm(head_dim, elementwise_affine=True)
# self.core_attention = self.attn_processor(attn_type=attn_type)
self.parallel = attn_type == 'parallel'
self.attn = DistributedAttention(
num_heads=self.n_heads,
head_size=head_dim,
causal=False,
supported_attention_backends=supported_attention_backends)
def _apply_rope(self, x: torch.Tensor, cos: torch.Tensor,
sin: torch.Tensor):
"""
x: [B, S, H, D]
cos: [S, D/2] where D = head_dim = sum(self.rope_split)
sin: [S, D/2]
returns x with rotary applied exactly as v0 did
"""
B, S, H, D = x.shape
# 1) split cos/sin per chunk
half_splits = [c // 2
for c in self.rope_split] # [32,16,16] for [64,32,32]
cos_splits = cos.split(half_splits, dim=1)
sin_splits = sin.split(half_splits, dim=1)
outs = []
idx = 0
for (chunk_size, cos_i, sin_i) in zip(self.rope_split,
cos_splits,
sin_splits,
strict=True):
# slice the corresponding channels
x_chunk = x[..., idx:idx + chunk_size] # [B,S,H,chunk_size]
idx += chunk_size
# flatten to [S, B*H, chunk_size]
x_flat = rearrange(x_chunk, 'b s h d -> s (b h) d')
# apply rotary on *that* chunk
out_flat = _apply_rotary_emb(x_flat,
cos_i,
sin_i,
is_neox_style=True)
# restore [B,S,H,chunk_size]
out = rearrange(out_flat, 's (b h) d -> b s h d', b=B, h=H)
outs.append(out)
# concatenate back to [B,S,H,D]
return torch.cat(outs, dim=-1)
def forward(self,
x,
cu_seqlens=None,
max_seqlen=None,
rope_positions=None,
cos_sin=None,
attn_mask=None,
mask_strategy=None):
B, S, _ = x.shape
xqkv, _ = self.wqkv(x)
xqkv = xqkv.view(*x.shape[:-1], self.n_heads, 3 * self.head_dim)
q, k, v = torch.split(xqkv, [self.head_dim] * 3, dim=-1) # [B,S,H,D]
if self.with_qk_norm:
q = self.q_norm(q)
k = self.k_norm(k)
if self.with_rope:
if rope_positions is not None:
F, Ht, W = rope_positions
assert F * Ht * W == S, "rope_positions mismatches sequence length"
cos, sin = cos_sin
cos = cos.to(x.device, dtype=x.dtype)
sin = sin.to(x.device, dtype=x.dtype)
q = self._apply_rope(q, cos, sin)
k = self._apply_rope(k, cos, sin)
output, _ = self.attn(q, k, v) # [B,heads,S,D]
output = rearrange(output, 'b s h d -> b s (h d)')
output, _ = self.wo(output)
return output
class CrossAttention(nn.Module):
def __init__(
self,
hidden_dim,
head_dim,
bias=False,
with_qk_norm=True,
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA)
) -> None:
super().__init__()
self.head_dim = head_dim
self.n_heads = hidden_dim // head_dim
self.wq = ReplicatedLinear(hidden_dim, hidden_dim, bias=bias)
self.wkv = ReplicatedLinear(hidden_dim, hidden_dim * 2, bias=bias)
self.wo = ReplicatedLinear(hidden_dim, hidden_dim, bias=bias)
self.with_qk_norm = with_qk_norm
if self.with_qk_norm:
self.q_norm = StepVideoRMSNorm(head_dim, elementwise_affine=True)
self.k_norm = StepVideoRMSNorm(head_dim, elementwise_affine=True)
self.attn = LocalAttention(
num_heads=self.n_heads,
head_size=head_dim,
causal=False,
supported_attention_backends=supported_attention_backends)
def forward(self,
x: torch.Tensor,
encoder_hidden_states: torch.Tensor,
attn_mask=None) -> torch.Tensor:
xq, _ = self.wq(x)
xq = xq.view(*xq.shape[:-1], self.n_heads, self.head_dim)
xkv, _ = self.wkv(encoder_hidden_states)
xkv = xkv.view(*xkv.shape[:-1], self.n_heads, 2 * self.head_dim)
xk, xv = torch.split(xkv, [self.head_dim] * 2,
dim=-1) ## seq_len, n, dim
if self.with_qk_norm:
xq = self.q_norm(xq)
xk = self.k_norm(xk)
output = self.attn(xq, xk, xv)
output = rearrange(output, 'b s h d -> b s (h d)')
output, _ = self.wo(output)
return output
class AdaLayerNormSingle(nn.Module):
r"""
Norm layer adaptive layer norm single (adaLN-single).
As proposed in PixArt-Alpha (see: https://arxiv.org/abs/2310.00426; Section 2.3).
Parameters:
embedding_dim (`int`): The size of each embedding vector.
use_additional_conditions (`bool`): To use additional conditions for normalization or not.
"""
def __init__(self, embedding_dim: int, time_step_rescale=1000):
super().__init__()
self.emb = TimestepEmbedder(embedding_dim)
self.silu = nn.SiLU()
self.linear = ReplicatedLinear(embedding_dim,
6 * embedding_dim,
bias=True)
self.time_step_rescale = time_step_rescale ## timestep usually in [0, 1], we rescale it to [0,1000] for stability
def forward(
self,
timestep: torch.Tensor,
added_cond_kwargs: dict[str, torch.Tensor] | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
embedded_timestep = self.emb(timestep * self.time_step_rescale)
out, _ = self.linear(self.silu(embedded_timestep))
return out, embedded_timestep
class StepVideoTransformerBlock(nn.Module):
r"""
A basic Transformer block.
Parameters:
dim (`int`): The number of channels in the input and output.
num_attention_heads (`int`): The number of heads to use for multi-head attention.
attention_head_dim (`int`): The number of channels in each head.
dropout (`float`, *optional*, defaults to 0.0): The dropout probability to use.
cross_attention_dim (`int`, *optional*): The size of the encoder_hidden_states vector for cross attention.
activation_fn (`str`, *optional*, defaults to `"geglu"`): Activation function to be used in feed-forward.
num_embeds_ada_norm (:
obj: `int`, *optional*): The number of diffusion steps used during training. See `Transformer2DModel`.
attention_bias (:
obj: `bool`, *optional*, defaults to `False`): Configure if the attentions should contain a bias parameter.
only_cross_attention (`bool`, *optional*):
Whether to use only cross-attention layers. In this case two cross attention layers are used.
double_self_attention (`bool`, *optional*):
Whether to use two self-attention layers. In this case no cross attention layers are used.
upcast_attention (`bool`, *optional*):
Whether to upcast the attention computation to float32. This is useful for mixed precision training.
norm_elementwise_affine (`bool`, *optional*, defaults to `True`):
Whether to use learnable elementwise affine parameters for normalization.
norm_type (`str`, *optional*, defaults to `"layer_norm"`):
The normalization layer to use. Can be `"layer_norm"`, `"ada_norm"` or `"ada_norm_zero"`.
final_dropout (`bool` *optional*, defaults to False):
Whether to apply a final dropout after the last feed-forward layer.
positional_embeddings (`str`, *optional*, defaults to `None`):
The type of positional embeddings to apply to.
num_positional_embeddings (`int`, *optional*, defaults to `None`):
The maximum number of positional embeddings to apply.
"""
def __init__(self,
dim: int,
attention_head_dim: int,
norm_eps: float = 1e-5,
ff_inner_dim: int | None = None,
ff_bias: bool = False,
attention_type: str = 'torch'):
super().__init__()
self.dim = dim
self.norm1 = LayerNormScaleShift(dim,
norm_type="layer",
elementwise_affine=True,
eps=norm_eps)
self.attn1 = SelfAttention(
dim,
attention_head_dim,
bias=False,
with_rope=True,
with_qk_norm=True,
)
self.norm2 = LayerNormScaleShift(dim,
norm_type="layer",
elementwise_affine=True,
eps=norm_eps)
self.attn2 = CrossAttention(dim,
attention_head_dim,
bias=False,
with_qk_norm=True)
self.ff = MLP(input_dim=dim,
mlp_hidden_dim=dim *
4 if ff_inner_dim is None else ff_inner_dim,
act_type="gelu_pytorch_tanh",
bias=ff_bias)
self.scale_shift_table = nn.Parameter(torch.randn(6, dim) / dim**0.5)
@torch.no_grad()
def forward(self,
q: torch.Tensor,
kv: torch.Tensor,
t_expand: torch.LongTensor,
attn_mask=None,
rope_positions: list | None = None,
cos_sin=None,
mask_strategy=None) -> torch.Tensor:
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = (
torch.clone(chunk)
for chunk in (self.scale_shift_table[None] +
t_expand.reshape(-1, 6, self.dim)).chunk(6, dim=1))
scale_shift_q = self.norm1(q,
scale=scale_msa.squeeze(1),
shift=shift_msa.squeeze(1))
attn_q = self.attn1(scale_shift_q,
rope_positions=rope_positions,
cos_sin=cos_sin,
mask_strategy=mask_strategy)
q = attn_q * gate_msa + q
attn_q = self.attn2(q, kv, attn_mask)
q = attn_q + q
scale_shift_q = self.norm2(q,
scale=scale_mlp.squeeze(1),
shift=shift_mlp.squeeze(1))
ff_output = self.ff(scale_shift_q)
q = ff_output * gate_mlp + q
return q
class StepVideoModel(BaseDiT):
# (Optional) Keep the same attribute for compatibility with splitting, etc.
_fsdp_shard_conditions = [
lambda n, m: "transformer_blocks" in n and n.split(".")[-1].isdigit(),
# lambda n, m: "pos_embed" in n # If needed for the patch embedding.
]
param_names_mapping = StepVideoConfig().param_names_mapping
reverse_param_names_mapping = StepVideoConfig().reverse_param_names_mapping
lora_param_names_mapping = StepVideoConfig().lora_param_names_mapping
_supported_attention_backends = StepVideoConfig(
)._supported_attention_backends
def __init__(self, config: StepVideoConfig, hf_config: dict[str,
Any]) -> None:
super().__init__(config=config, hf_config=hf_config)
self.num_attention_heads = config.num_attention_heads
self.attention_head_dim = config.attention_head_dim
self.in_channels = config.in_channels
self.out_channels = config.out_channels
self.num_layers = config.num_layers
self.dropout = config.dropout
self.patch_size = config.patch_size
self.norm_type = config.norm_type
self.norm_elementwise_affine = config.norm_elementwise_affine
self.norm_eps = config.norm_eps
self.use_additional_conditions = config.use_additional_conditions
self.caption_channels = config.caption_channels
self.attention_type = config.attention_type
self.num_channels_latents = config.num_channels_latents
# Compute inner dimension.
self.hidden_size = config.hidden_size
# Image/video patch embedding.
self.pos_embed = PatchEmbed2D(
patch_size=self.patch_size,
in_chans=self.in_channels,
embed_dim=self.hidden_size,
)
self._rope_cache: dict[tuple, tuple[torch.Tensor, torch.Tensor]] = {}
# Transformer blocks.
self.transformer_blocks = nn.ModuleList([
StepVideoTransformerBlock(
dim=self.hidden_size,
attention_head_dim=self.attention_head_dim,
attention_type=self.attention_type)
for _ in range(self.num_layers)
])
# Output blocks.
self.norm_out = LayerNormScaleShift(
self.hidden_size,
norm_type="layer",
eps=self.norm_eps,
elementwise_affine=self.norm_elementwise_affine)
self.scale_shift_table = nn.Parameter(
torch.randn(2, self.hidden_size) / (self.hidden_size**0.5))
self.proj_out = ReplicatedLinear(
self.hidden_size,
self.patch_size * self.patch_size * self.out_channels)
# Time modulation via adaptive layer norm.
self.adaln_single = AdaLayerNormSingle(self.hidden_size)
# Set up caption conditioning.
if isinstance(self.caption_channels, int):
caption_channel = self.caption_channels
else:
caption_channel, clip_channel = self.caption_channels
self.clip_projection = ReplicatedLinear(clip_channel,
self.hidden_size)
self.caption_norm = nn.LayerNorm(
caption_channel,
eps=self.norm_eps,
elementwise_affine=self.norm_elementwise_affine)
self.caption_projection = MLP(input_dim=caption_channel,
mlp_hidden_dim=self.hidden_size,
act_type="gelu_pytorch_tanh")
# Flag to indicate if using parallel attention.
self.parallel = (self.attention_type == "parallel")
self.__post_init__()
def patchfy(self, hidden_states) -> torch.Tensor:
hidden_states = rearrange(hidden_states, 'b f c h w -> (b f) c h w')
hidden_states = self.pos_embed(hidden_states)
return hidden_states
def prepare_attn_mask(self, encoder_attention_mask, encoder_hidden_states,
q_seqlen) -> tuple[torch.Tensor, torch.Tensor]:
kv_seqlens = encoder_attention_mask.sum(dim=1).int()
mask = torch.zeros([len(kv_seqlens), q_seqlen,
max(kv_seqlens)],
dtype=torch.bool,
device=encoder_attention_mask.device)
encoder_hidden_states = encoder_hidden_states[:, :max(kv_seqlens)]
for i, kv_len in enumerate(kv_seqlens):
mask[i, :, :kv_len] = 1
return encoder_hidden_states, mask
def block_forward(self,
hidden_states,
encoder_hidden_states=None,
t_expand=None,
rope_positions=None,
cos_sin=None,
attn_mask=None,
parallel=True,
mask_strategy=None) -> torch.Tensor:
for i, block in enumerate(self.transformer_blocks):
hidden_states = block(hidden_states,
encoder_hidden_states,
t_expand=t_expand,
attn_mask=attn_mask,
rope_positions=rope_positions,
cos_sin=cos_sin,
mask_strategy=mask_strategy[i])
return hidden_states
def _get_rope(self, rope_positions: tuple[int, int, int],
dtype: torch.dtype, device: torch.device):
F, Ht, W = rope_positions
key = (F, Ht, W, dtype)
if key not in self._rope_cache:
cos, sin = get_rotary_pos_embed(
rope_sizes=(F * get_sp_world_size(), Ht, W),
hidden_size=self.hidden_size,
heads_num=self.hidden_size // self.attention_head_dim,
rope_dim_list=(64, 32, 32), # same split you used
rope_theta=1.0e4,
dtype=torch.float32 # build once in fp32
)
# move & cast once
self._rope_cache[key] = (cos.to(device, dtype=dtype),
sin.to(device, dtype=dtype))
return self._rope_cache[key]
@torch.inference_mode()
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor | None = None,
t_expand: torch.LongTensor | None = None,
encoder_hidden_states_2: torch.Tensor | None = None,
added_cond_kwargs: dict[str, torch.Tensor] | None = None,
encoder_attention_mask: torch.Tensor | None = None,
fps: torch.Tensor | None = None,
return_dict: bool = True,
mask_strategy=None,
guidance=None,
):
assert hidden_states.ndim == 5
"hidden_states's shape should be (bsz, f, ch, h ,w)"
frame = hidden_states.shape[2]
hidden_states = rearrange(hidden_states,
'b c f h w -> b f c h w',
f=frame)
if mask_strategy is None:
mask_strategy = [None, None]
bsz, frame, _, height, width = hidden_states.shape
height, width = height // self.patch_size, width // self.patch_size
hidden_states = self.patchfy(hidden_states)
len_frame = hidden_states.shape[1]
t_expand, embedded_timestep = self.adaln_single(t_expand)
encoder_hidden_states = self.caption_projection(
self.caption_norm(encoder_hidden_states))
if encoder_hidden_states_2 is not None and hasattr(
self, 'clip_projection'):
clip_embedding, _ = self.clip_projection(encoder_hidden_states_2)
encoder_hidden_states = torch.cat(
[clip_embedding, encoder_hidden_states], dim=1)
hidden_states = rearrange(hidden_states,
'(b f) l d-> b (f l) d',
b=bsz,
f=frame,
l=len_frame).contiguous()
encoder_hidden_states, attn_mask = self.prepare_attn_mask(
encoder_attention_mask,
encoder_hidden_states,
q_seqlen=frame * len_frame)
cos_sin = self._get_rope((frame, height, width), hidden_states.dtype,
hidden_states.device)
hidden_states = self.block_forward(
hidden_states,
encoder_hidden_states,
t_expand=t_expand,
rope_positions=[frame, height, width],
cos_sin=cos_sin,
attn_mask=attn_mask,
parallel=self.parallel,
mask_strategy=mask_strategy)
hidden_states = rearrange(hidden_states,
'b (f l) d -> (b f) l d',
b=bsz,
f=frame,
l=len_frame)
embedded_timestep = repeat(embedded_timestep, 'b d -> (b f) d',
f=frame).contiguous()
shift, scale = (self.scale_shift_table[None] +
embedded_timestep[:, None]).chunk(2, dim=1)
hidden_states = self.norm_out(hidden_states,
shift=shift.squeeze(1),
scale=scale.squeeze(1))
# Modulation
hidden_states, _ = self.proj_out(hidden_states)
# unpatchify
hidden_states = hidden_states.reshape(shape=(-1, height, width,
self.patch_size,
self.patch_size,
self.out_channels))
hidden_states = rearrange(hidden_states, 'n h w p q c -> n c h p w q')
output = hidden_states.reshape(shape=(-1, self.out_channels,
height * self.patch_size,
width * self.patch_size))
output = rearrange(output, '(b f) c h w -> b c f h w', f=frame)
return output
# Entry point for model registry
EntryClass = StepVideoModel
+7 -3
View File
@@ -21,6 +21,7 @@ from fastvideo.models.dits.ltx2 import (
)
from fastvideo.models.loader.weight_utils import default_weight_loader
from fastvideo.platforms import AttentionBackendEnum
from fastvideo.distributed import get_local_torch_device
def _debug_log_line(message: str) -> None:
@@ -516,15 +517,18 @@ class LTX2GemmaTextEncoderModel(TextEncoder):
attention_mask = torch.ones_like(input_ids)
model = self.gemma_model
input_ids = input_ids.to(device=model.device)
attention_mask = attention_mask.to(device=model.device)
orig_device = model.device
model.to(device=get_local_torch_device())
# input_ids = input_ids.to(device=model.device)
# attention_mask = attention_mask.to(device=model.device)
outputs = model(
input_ids=input_ids,
attention_mask=attention_mask,
output_hidden_states=True,
return_dict=True,
)
model.to(device=orig_device)
encoded_inputs = self._run_feature_extractor(
outputs.hidden_states,
attention_mask,
-588
View File
@@ -1,588 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# type: ignore
# Copyright 2025 StepFun Inc. All Rights Reserved.
#
# Permission is hereby granted, free of charge, to any person obtaining a copy
# of this software and associated documentation files (the "Software"), to deal
# in the Software without restriction, including without limitation the rights
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
# copies of the Software, and to permit persons to whom the Software is
# furnished to do so, subject to the following conditions:
#
# The above copyright notice and this permission notice shall be included in all
# copies or substantial portions of the Software.
# ==============================================================================
import os
from functools import wraps
import torch
import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange
from transformers import PretrainedConfig
from transformers.modeling_utils import PreTrainedModel
from fastvideo.models.dits.stepvideo import StepVideoRMSNorm
class EmptyInitOnDevice(torch.overrides.TorchFunctionMode):
def __init__(self, device=None):
self.device = device
def __torch_function__(self, func, types, args=(), kwargs=None):
kwargs = kwargs or {}
if getattr(func, '__module__', None) == 'torch.nn.init':
if 'tensor' in kwargs:
return kwargs['tensor']
else:
return args[0]
if self.device is not None and func in torch.utils._device._device_constructors(
) and kwargs.get('device') is None:
kwargs['device'] = self.device
return func(*args, **kwargs)
def with_empty_init(func):
@wraps(func)
def wrapper(*args, **kwargs):
with EmptyInitOnDevice('cpu'):
return func(*args, **kwargs)
return wrapper
class LLaMaEmbedding(nn.Module):
"""Language model embeddings.
Arguments:
hidden_size: hidden size
vocab_size: vocabulary size
max_sequence_length: maximum size of sequence. This
is used for positional embedding
embedding_dropout_prob: dropout probability for embeddings
init_method: weight initialization method
num_tokentypes: size of the token-type embeddings. 0 value
will ignore this embedding
"""
def __init__(
self,
cfg,
):
super().__init__()
self.hidden_size = cfg.hidden_size
self.params_dtype = cfg.params_dtype
self.fp32_residual_connection = cfg.fp32_residual_connection
self.embedding_weights_in_fp32 = cfg.embedding_weights_in_fp32
self.word_embeddings = torch.nn.Embedding(
cfg.padded_vocab_size,
self.hidden_size,
)
self.embedding_dropout = torch.nn.Dropout(cfg.hidden_dropout)
def forward(self, input_ids):
# Embeddings.
if self.embedding_weights_in_fp32:
self.word_embeddings = self.word_embeddings.to(torch.float32)
embeddings = self.word_embeddings(input_ids)
if self.embedding_weights_in_fp32:
embeddings = embeddings.to(self.params_dtype)
self.word_embeddings = self.word_embeddings.to(self.params_dtype)
# Data format change to avoid explicit transposes : [b s h] --> [s b h].
embeddings = embeddings.transpose(0, 1).contiguous()
# If the input flag for fp32 residual connection is set, convert for float.
if self.fp32_residual_connection:
embeddings = embeddings.float()
# Dropout.
embeddings = self.embedding_dropout(embeddings)
return embeddings
class StepChatTokenizer:
"""Step Chat Tokenizer"""
def __init__(
self,
model_file,
name="StepChatTokenizer",
bot_token="<|BOT|>", # Begin of Turn
eot_token="<|EOT|>", # End of Turn
call_start_token="<|CALL_START|>", # Call Start
call_end_token="<|CALL_END|>", # Call End
think_start_token="<|THINK_START|>", # Think Start
think_end_token="<|THINK_END|>", # Think End
mask_start_token="<|MASK_1e69f|>", # Mask start
mask_end_token="<|UNMASK_1e69f|>", # Mask end
):
import sentencepiece
self._tokenizer = sentencepiece.SentencePieceProcessor(
model_file=model_file)
self._vocab = {}
self._inv_vocab = {}
self._special_tokens = {}
self._inv_special_tokens = {}
self._t5_tokens = []
for idx in range(self._tokenizer.get_piece_size()):
text = self._tokenizer.id_to_piece(idx)
self._inv_vocab[idx] = text
self._vocab[text] = idx
if self._tokenizer.is_control(idx) or self._tokenizer.is_unknown(
idx):
self._special_tokens[text] = idx
self._inv_special_tokens[idx] = text
self._unk_id = self._tokenizer.unk_id()
self._bos_id = self._tokenizer.bos_id()
self._eos_id = self._tokenizer.eos_id()
for token in [
bot_token, eot_token, call_start_token, call_end_token,
think_start_token, think_end_token
]:
assert token in self._vocab, f"Token '{token}' not found in tokenizer"
assert token in self._special_tokens, f"Token '{token}' is not a special token"
for token in [mask_start_token, mask_end_token]:
assert token in self._vocab, f"Token '{token}' not found in tokenizer"
self._bot_id = self._tokenizer.piece_to_id(bot_token)
self._eot_id = self._tokenizer.piece_to_id(eot_token)
self._call_start_id = self._tokenizer.piece_to_id(call_start_token)
self._call_end_id = self._tokenizer.piece_to_id(call_end_token)
self._think_start_id = self._tokenizer.piece_to_id(think_start_token)
self._think_end_id = self._tokenizer.piece_to_id(think_end_token)
self._mask_start_id = self._tokenizer.piece_to_id(mask_start_token)
self._mask_end_id = self._tokenizer.piece_to_id(mask_end_token)
self._underline_id = self._tokenizer.piece_to_id("\u2581")
@property
def vocab(self):
return self._vocab
@property
def inv_vocab(self):
return self._inv_vocab
@property
def vocab_size(self):
return self._tokenizer.vocab_size()
def tokenize(self, text: str) -> list[int]:
return self._tokenizer.encode_as_ids(text)
def detokenize(self, token_ids: list[int]) -> str:
return self._tokenizer.decode_ids(token_ids)
class Tokens:
def __init__(self, input_ids, cu_input_ids, attention_mask, cu_seqlens,
max_seq_len) -> None:
self.input_ids = input_ids
self.attention_mask = attention_mask
self.cu_input_ids = cu_input_ids
self.cu_seqlens = cu_seqlens
self.max_seq_len = max_seq_len
def to(self, device):
self.input_ids = self.input_ids.to(device)
self.attention_mask = self.attention_mask.to(device)
self.cu_input_ids = self.cu_input_ids.to(device)
self.cu_seqlens = self.cu_seqlens.to(device)
return self
class Wrapped_StepChatTokenizer(StepChatTokenizer):
def __call__(self,
text,
max_length=320,
padding="max_length",
truncation=True,
return_tensors="pt"):
# [bos, ..., eos, pad, pad, ..., pad]
self.BOS = 1
self.EOS = 2
self.PAD = 2
out_tokens = []
attn_mask = []
if len(text) == 0:
part_tokens = [self.BOS] + [self.EOS]
valid_size = len(part_tokens)
if len(part_tokens) < max_length:
part_tokens += [self.PAD] * (max_length - valid_size)
out_tokens.append(part_tokens)
attn_mask.append([1] * valid_size + [0] * (max_length - valid_size))
else:
for part in text:
part_tokens = self.tokenize(part)
part_tokens = part_tokens[:(max_length -
2)] # leave 2 space for bos and eos
part_tokens = [self.BOS] + part_tokens + [self.EOS]
valid_size = len(part_tokens)
if len(part_tokens) < max_length:
part_tokens += [self.PAD] * (max_length - valid_size)
out_tokens.append(part_tokens)
attn_mask.append([1] * valid_size + [0] *
(max_length - valid_size))
out_tokens = torch.tensor(out_tokens, dtype=torch.long)
attn_mask = torch.tensor(attn_mask, dtype=torch.long)
# padding y based on tp size
padded_len = 0
padded_flag = False
if padded_len > 0:
padded_flag = True
if padded_flag:
pad_tokens = torch.tensor([[self.PAD] * max_length],
device=out_tokens.device)
pad_attn_mask = torch.tensor([[1] * padded_len + [0] *
(max_length - padded_len)],
device=attn_mask.device)
out_tokens = torch.cat([out_tokens, pad_tokens], dim=0)
attn_mask = torch.cat([attn_mask, pad_attn_mask], dim=0)
# cu_seqlens
cu_out_tokens = out_tokens.masked_select(attn_mask != 0).unsqueeze(0)
seqlen = attn_mask.sum(dim=1).tolist()
cu_seqlens = torch.cumsum(torch.tensor([0] + seqlen),
0).to(device=out_tokens.device,
dtype=torch.int32)
max_seq_len = max(seqlen)
return Tokens(out_tokens, cu_out_tokens, attn_mask, cu_seqlens,
max_seq_len)
def flash_attn_func(q,
k,
v,
dropout_p=0.0,
softmax_scale=None,
causal=True,
return_attn_probs=False,
tp_group_rank=0,
tp_group_size=1):
softmax_scale = q.size(-1)**(
-0.5) if softmax_scale is None else softmax_scale
return torch.ops.Optimus.fwd(q, k, v, None, dropout_p, softmax_scale,
causal, return_attn_probs, None, tp_group_rank,
tp_group_size)[0]
class FlashSelfAttention(torch.nn.Module):
def __init__(
self,
attention_dropout=0.0,
):
super().__init__()
self.dropout_p = attention_dropout
def forward(self, q, k, v, cu_seqlens=None, max_seq_len=None):
if cu_seqlens is None:
output = flash_attn_func(q, k, v, dropout_p=self.dropout_p)
else:
raise ValueError('cu_seqlens is not supported!')
return output
def safediv(n, d):
q, r = divmod(n, d)
assert r == 0
return q
class MultiQueryAttention(nn.Module):
def __init__(self, cfg, layer_id=None):
super().__init__()
self.head_dim = cfg.hidden_size // cfg.num_attention_heads
self.max_seq_len = cfg.seq_length
self.use_flash_attention = cfg.use_flash_attn
assert self.use_flash_attention, 'FlashAttention is required!'
self.n_groups = cfg.num_attention_groups
self.tp_size = 1
self.n_local_heads = cfg.num_attention_heads
self.n_local_groups = self.n_groups
self.wqkv = nn.Linear(
cfg.hidden_size,
cfg.hidden_size + self.head_dim * 2 * self.n_groups,
bias=False,
)
self.wo = nn.Linear(
cfg.hidden_size,
cfg.hidden_size,
bias=False,
)
# assert self.use_flash_attention, 'non-Flash attention not supported yet.'
self.core_attention = FlashSelfAttention(
attention_dropout=cfg.attention_dropout)
# self.core_attention = LocalAttention(
# num_heads = self.n_local_heads,
# head_size = self.head_dim,
# # num_kv_heads = self.n_local_groups,
# casual = True,
# supported_attention_backends = [_Backend.FLASH_ATTN, _Backend.TORCH_SDPA], # RIVER TODO
# )
self.layer_id = layer_id
def forward(
self,
x: torch.Tensor,
mask: torch.Tensor | None,
cu_seqlens: torch.Tensor | None,
max_seq_len: torch.Tensor | None,
):
seqlen, bsz, dim = x.shape
xqkv = self.wqkv(x)
xq, xkv = torch.split(
xqkv,
(dim // self.tp_size,
self.head_dim * 2 * self.n_groups // self.tp_size),
dim=-1,
)
# gather on 1st dimension
xq = xq.view(seqlen, bsz, self.n_local_heads, self.head_dim)
xkv = xkv.view(seqlen, bsz, self.n_local_groups, 2 * self.head_dim)
xk, xv = xkv.chunk(2, -1)
# rotary embedding + flash attn
xq = rearrange(xq, "s b h d -> b s h d")
xk = rearrange(xk, "s b h d -> b s h d")
xv = rearrange(xv, "s b h d -> b s h d")
# q_per_kv = self.n_local_heads // self.n_local_groups
# if q_per_kv > 1:
# b, s, h, d = xk.size()
# if h == 1:
# xk = xk.expand(b, s, q_per_kv, d)
# xv = xv.expand(b, s, q_per_kv, d)
# else:
# ''' To cover the cases where h > 1, we have
# the following implementation, which is equivalent to:
# xk = xk.repeat_interleave(q_per_kv, dim=-2)
# xv = xv.repeat_interleave(q_per_kv, dim=-2)
# but can avoid calling aten::item() that involves cpu.
# '''
# idx = torch.arange(q_per_kv * h, device=xk.device).reshape(q_per_kv, -1).permute(1, 0).flatten()
# xk = torch.index_select(xk.repeat(1, 1, q_per_kv, 1), 2, idx).contiguous()
# xv = torch.index_select(xv.repeat(1, 1, q_per_kv, 1), 2, idx).contiguous()
if self.use_flash_attention:
output = self.core_attention(xq, xk, xv)
# reduce-scatter only support first dimension now
output = rearrange(output, "b s h d -> s b (h d)").contiguous()
else:
xq, xk, xv = [
rearrange(x, "b s ... -> s b ...").contiguous()
for x in (xq, xk, xv)
]
output = self.core_attention(xq, xk, xv) #, mask)
output = self.wo(output)
return output
class FeedForward(nn.Module):
def __init__(
self,
cfg,
dim: int,
hidden_dim: int,
layer_id: int,
multiple_of: int = 256,
):
super().__init__()
hidden_dim = multiple_of * (
(hidden_dim + multiple_of - 1) // multiple_of)
def swiglu(x):
x = torch.chunk(x, 2, dim=-1)
return F.silu(x[0]) * x[1]
self.swiglu = swiglu
self.w1 = nn.Linear(
dim,
2 * hidden_dim,
bias=False,
)
self.w2 = nn.Linear(
hidden_dim,
dim,
bias=False,
)
def forward(self, x):
x = self.swiglu(self.w1(x))
output = self.w2(x)
return output
class TransformerBlock(nn.Module):
def __init__(self, cfg, layer_id: int):
super().__init__()
self.n_heads = cfg.num_attention_heads
self.dim = cfg.hidden_size
self.head_dim = cfg.hidden_size // cfg.num_attention_heads
self.attention = MultiQueryAttention(
cfg,
layer_id=layer_id,
)
self.feed_forward = FeedForward(
cfg,
dim=cfg.hidden_size,
hidden_dim=cfg.ffn_hidden_size,
layer_id=layer_id,
)
self.layer_id = layer_id
self.attention_norm = StepVideoRMSNorm(
cfg.hidden_size,
eps=cfg.layernorm_epsilon,
)
self.ffn_norm = StepVideoRMSNorm(
cfg.hidden_size,
eps=cfg.layernorm_epsilon,
)
def forward(
self,
x: torch.Tensor,
mask: torch.Tensor | None,
cu_seqlens: torch.Tensor | None,
max_seq_len: torch.Tensor | None,
):
residual = self.attention.forward(self.attention_norm(x), mask,
cu_seqlens, max_seq_len)
h = x + residual
ffn_res = self.feed_forward.forward(self.ffn_norm(h))
out = h + ffn_res
return out
class Transformer(nn.Module):
def __init__(
self,
config,
max_seq_size=8192,
):
super().__init__()
self.num_layers = config.num_layers
self.layers = self._build_layers(config)
def _build_layers(self, config):
layers = torch.nn.ModuleList()
for layer_id in range(self.num_layers):
layers.append(TransformerBlock(
config,
layer_id=layer_id + 1,
))
return layers
def forward(
self,
hidden_states,
attention_mask,
cu_seqlens=None,
max_seq_len=None,
):
if max_seq_len is not None and not isinstance(max_seq_len,
torch.Tensor):
max_seq_len = torch.tensor(max_seq_len,
dtype=torch.int32,
device="cpu")
for lid, layer in enumerate(self.layers):
hidden_states = layer(
hidden_states,
attention_mask,
cu_seqlens,
max_seq_len,
)
return hidden_states
class Step1Model(PreTrainedModel):
config_class = PretrainedConfig
@with_empty_init
def __init__(
self,
config,
):
super().__init__(config)
self.tok_embeddings = LLaMaEmbedding(config)
self.transformer = Transformer(config)
def forward(
self,
input_ids=None,
attention_mask=None,
):
hidden_states = self.tok_embeddings(input_ids)
hidden_states = self.transformer(
hidden_states,
attention_mask,
)
return hidden_states
class STEP1TextEncoder(torch.nn.Module):
def __init__(self, model_dir, max_length=320):
super().__init__()
self.max_length = max_length
self.text_tokenizer = Wrapped_StepChatTokenizer(
os.path.join(model_dir, 'step1_chat_tokenizer.model'))
text_encoder = Step1Model.from_pretrained(model_dir)
self.text_encoder = text_encoder.eval().to(torch.bfloat16)
@torch.no_grad
def forward(self, prompts, with_mask=True, max_length=None):
self.device = next(self.text_encoder.parameters()).device
with torch.no_grad(), torch.amp.autocast('cuda', dtype=torch.bfloat16):
if type(prompts) is str:
prompts = [prompts]
txt_tokens = self.text_tokenizer(prompts,
max_length=max_length
or self.max_length,
padding="max_length",
truncation=True,
return_tensors="pt")
y = self.text_encoder(txt_tokens.input_ids.to(self.device),
attention_mask=txt_tokens.attention_mask.to(
self.device) if with_mask else None)
y_mask = txt_tokens.attention_mask
return y.transpose(0, 1), y_mask
# Entry point for model registry
EntryClass = STEP1TextEncoder
+18 -6
View File
@@ -518,10 +518,25 @@ class TokenizerLoader(ComponentLoader):
def load(self, model_path: str, fastvideo_args: FastVideoArgs):
"""Load the tokenizer based on the model path, and inference args."""
logger.info("Loading tokenizer from %s", model_path)
resolved_model_path = model_path
# LTX2 checkpoints may not ship a top-level tokenizer/ directory.
# In that case, tokenizer assets live under text_encoder/gemma/.
if not os.path.isdir(resolved_model_path):
ltx2_gemma_path = os.path.normpath(
os.path.join(resolved_model_path, "..", "text_encoder",
"gemma"))
if os.path.isdir(ltx2_gemma_path):
logger.info(
"Tokenizer directory %s missing; falling back to %s",
resolved_model_path,
ltx2_gemma_path,
)
resolved_model_path = ltx2_gemma_path
# Cosmos2.5 stores an AutoProcessor config in `tokenizer/config.json` (not a tokenizer
# config). Use its `_name_or_path` (e.g. Qwen/Qwen2.5-VL-7B-Instruct) as the source.
tokenizer_cfg_path = os.path.join(model_path, "config.json")
tokenizer_cfg_path = os.path.join(resolved_model_path, "config.json")
if os.path.exists(tokenizer_cfg_path):
try:
with open(tokenizer_cfg_path, "r") as f:
@@ -547,10 +562,11 @@ class TokenizerLoader(ComponentLoader):
pass
tokenizer = AutoTokenizer.from_pretrained(
model_path, # "<path to model>/tokenizer"
resolved_model_path, # "<path to model>/tokenizer"
# in v0, this was same string as encoder_name "ClipTextModel"
# TODO(will): pass these tokenizer kwargs from inference args? Maybe
# other method of config?
local_files_only=os.path.isdir(resolved_model_path),
)
padding_side = None
if hasattr(fastvideo_args.pipeline_config, "text_encoder_configs"):
@@ -892,10 +908,6 @@ class SchedulerLoader(ComponentLoader):
scheduler = scheduler_cls(**config)
if fastvideo_args.pipeline_config.flow_shift is not None:
scheduler.set_shift(fastvideo_args.pipeline_config.flow_shift)
if fastvideo_args.pipeline_config.timesteps_scale is not None:
scheduler.set_timesteps_scale(
fastvideo_args.pipeline_config.timesteps_scale
)
return scheduler
+1 -4
View File
@@ -33,7 +33,6 @@ _TEXT_TO_VIDEO_DIT_MODELS = {
("dits", "hyworld", "HYWorldTransformer3DModel"),
"WanTransformer3DModel": ("dits", "wanvideo", "WanTransformer3DModel"),
"CausalWanTransformer3DModel": ("dits", "causal_wanvideo", "CausalWanTransformer3DModel"),
"StepVideoModel": ("dits", "stepvideo", "StepVideoModel"),
"CosmosTransformer3DModel": ("dits", "cosmos", "CosmosTransformer3DModel"),
"Cosmos25Transformer3DModel": ("dits", "cosmos2_5", "Cosmos25Transformer3DModel"),
"LongCatVideoTransformer3DModel": ("dits", "longcat_video_dit", "LongCatVideoTransformer3DModel"), # Wrapper (Phase 1)
@@ -58,7 +57,6 @@ _TEXT_ENCODER_MODELS = {
"LlamaModel": ("encoders", "llama", "LlamaModel"),
"UMT5EncoderModel": ("encoders", "t5", "UMT5EncoderModel"),
"T5EncoderModel": ("encoders", "t5_hf", "T5EncoderModel"),
"STEP1TextEncoder": ("encoders", "stepllm", "STEP1TextEncoder"),
"BertModel": ("encoders", "clip", "CLIPTextModel"),
"Qwen2_5_VLTextModel": ("encoders", "qwen2_5", "Qwen2_5_VLTextModel"),
"Reason1TextEncoder": ("encoders", "reason1", "Reason1TextEncoder"),
@@ -81,7 +79,6 @@ _VAE_MODELS = {
"AutoencoderKLHYWorld": ("vaes", "hyworldvae", "AutoencoderKLHYWorld"),
"AutoencoderKLHunyuanVideo15": ("vaes", "hunyuan15vae", "AutoencoderKLHunyuanVideo15"),
"AutoencoderKLWan": ("vaes", "wanvae", "AutoencoderKLWan"),
"AutoencoderKLStepvideo": ("vaes", "stepvideovae", "AutoencoderKLStepvideo"),
"AutoencoderKL": ("vaes", "autoencoder_kl", "AutoencoderKL"),
"CausalVideoAutoencoder": ("vaes", "ltx2vae", "LTX2CausalVideoAutoencoder"),
}
@@ -454,4 +451,4 @@ ModelRegistry = _ModelRegistry({
)
for model_arch, (component_name, mod_relname,
cls_name) in _FAST_VIDEO_MODELS.items()
})
})
File diff suppressed because it is too large Load Diff
@@ -8,7 +8,7 @@ from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.pipelines.stages.input_validation import InputValidationStage
from fastvideo.pipelines.stages.text_encoding import TextEncodingStage
from fastvideo.pipelines.stages.timestep_preparation import (
TimestepPreparationStage, )
SD35TimestepPreparationStage, )
from fastvideo.pipelines.stages.sd35_conditioning import (
SD35ConditioningStage,
SD35DecodingStage,
@@ -71,7 +71,7 @@ class SD35Pipeline(ComposedPipelineBase):
self.add_stage(
stage_name="timestep_preparation_stage",
stage=TimestepPreparationStage(
stage=SD35TimestepPreparationStage(
scheduler=self.get_module("scheduler")),
)
@@ -1,151 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# type: ignore
# SPDX-License-Identifier: Apache-2.0
"""
Hunyuan video diffusion pipeline implementation.
This module contains an implementation of the Hunyuan video diffusion pipeline
using the modular pipeline architecture.
"""
import os
from typing import Any
import torch
from huggingface_hub import hf_hub_download
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.models.encoders.bert import HunyuanClip # type: ignore
from fastvideo.models.encoders.stepllm import STEP1TextEncoder
from fastvideo.models.loader.component_loader import PipelineComponentLoader
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.pipelines.lora_pipeline import LoRAPipeline
from fastvideo.pipelines.stages import (DecodingStage, DenoisingStage,
InputValidationStage,
LatentPreparationStage,
StepvideoPromptEncodingStage,
TimestepPreparationStage)
logger = init_logger(__name__)
class StepVideoPipeline(LoRAPipeline, ComposedPipelineBase):
_required_config_modules = ["transformer", "scheduler", "vae"]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
self.add_stage(stage_name="input_validation_stage",
stage=InputValidationStage())
self.add_stage(stage_name="prompt_encoding_stage",
stage=StepvideoPromptEncodingStage(
stepllm=self.get_module("text_encoder"),
clip=self.get_module("text_encoder_2"),
))
self.add_stage(stage_name="timestep_preparation_stage",
stage=TimestepPreparationStage(
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer"),
))
self.add_stage(stage_name="denoising_stage",
stage=DenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
def build_llm(self, model_dir, device) -> torch.nn.Module:
text_encoder = STEP1TextEncoder(
model_dir, max_length=320).to(device).to(torch.bfloat16).eval()
return text_encoder
def build_clip(self, model_dir, device) -> HunyuanClip:
clip = HunyuanClip(model_dir, max_length=77).to(device).eval()
return clip
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
"""
Initialize the pipeline.
"""
target_device = get_local_torch_device()
llm_dir = os.path.join(self.model_path, "step_llm")
clip_dir = os.path.join(self.model_path, "hunyuan_clip")
text_enc = self.build_llm(llm_dir, target_device)
clip_enc = self.build_clip(clip_dir, target_device)
self.add_module("text_encoder", text_enc)
self.add_module("text_encoder_2", clip_enc)
lib_path = (
os.path.join(
fastvideo_args.model_path,
'lib/liboptimus_ths-torch2.5-cu124.cpython-310-x86_64-linux-gnu.so'
) if os.path.isdir(fastvideo_args.model_path) # local checkout
else hf_hub_download(
repo_id=fastvideo_args.model_path,
filename=
'lib/liboptimus_ths-torch2.5-cu124.cpython-310-x86_64-linux-gnu.so'
))
torch.ops.load_library(lib_path)
def load_modules(self, fastvideo_args: FastVideoArgs) -> dict[str, Any]:
"""
Load the modules from the config.
"""
model_index = self._load_config(self.model_path)
logger.info("Loading pipeline modules from config: %s", model_index)
# remove keys that are not pipeline modules
model_index.pop("_class_name")
model_index.pop("_diffusers_version")
# some sanity checks
assert len(
model_index
) > 1, "model_index.json must contain at least one pipeline module"
required_modules = ["transformer", "scheduler", "vae"]
for module_name in required_modules:
if module_name not in model_index:
raise ValueError(
f"model_index.json must contain a {module_name} module")
logger.info("Diffusers config passed sanity checks")
# all the component models used by the pipeline
modules = {}
for module_name, (transformers_or_diffusers,
architecture) in model_index.items():
component_model_path = os.path.join(self.model_path, module_name)
module = PipelineComponentLoader.load_module(
module_name=module_name,
component_model_path=component_model_path,
transformers_or_diffusers=transformers_or_diffusers,
fastvideo_args=fastvideo_args,
)
logger.info("Loaded module %s from %s", module_name,
component_model_path)
if module_name in modules:
logger.warning("Overwriting module %s", module_name)
modules[module_name] = module
required_modules = self.required_config_modules
# Check if all required modules were loaded
for module_name in required_modules:
if module_name not in modules or modules[module_name] is None:
raise ValueError(
f"Required module {module_name} was not loaded properly")
return modules
EntryClass = StepVideoPipeline
+15 -2
View File
@@ -374,8 +374,21 @@ class ComposedPipelineBase(ABC):
logger.info("Loading required modules: %s", required_modules)
modules = {}
for module_name, (transformers_or_diffusers,
architecture) in model_index.items():
for module_name, module_spec in model_index.items():
if not isinstance(module_spec, list | tuple):
logger.info(
"Skipping non-module config entry %s=%s",
module_name,
module_spec,
)
continue
if len(module_spec) < 1:
logger.warning(
"Skipping module %s due to invalid empty spec in model_index.json",
module_name,
)
continue
transformers_or_diffusers = module_spec[0]
if transformers_or_diffusers is None:
logger.warning(
"Module %s in model_index.json has null value, removing from required_config_modules",
-3
View File
@@ -36,8 +36,6 @@ from fastvideo.pipelines.stages.matrixgame_denoising import (
MatrixGameCausalDenoisingStage)
from fastvideo.pipelines.stages.hyworld_denoising import HYWorldDenoisingStage
from fastvideo.pipelines.stages.gamecraft_denoising import GameCraftDenoisingStage
from fastvideo.pipelines.stages.stepvideo_encoding import (
StepvideoPromptEncodingStage)
from fastvideo.pipelines.stages.text_encoding import (Cosmos25TextEncodingStage,
TextEncodingStage)
from fastvideo.pipelines.stages.timestep_preparation import (
@@ -88,7 +86,6 @@ __all__ = [
"GameCraftImageVAEEncodingStage",
"TextEncodingStage",
"Cosmos25TextEncodingStage",
"StepvideoPromptEncodingStage",
# LongCat stages
"LongCatVideoVAEEncodingStage",
"LongCatKVCacheInitStage",
+3
View File
@@ -148,16 +148,19 @@ class PipelineStage(ABC):
# Execute the actual stage logic
if envs.FASTVIDEO_STAGE_LOGGING:
logger.info("[%s] Starting execution", stage_name)
torch.cuda.synchronize()
start_time = time.perf_counter()
try:
result = self.forward(batch, fastvideo_args)
torch.cuda.synchronize()
execution_time = time.perf_counter() - start_time
logger.info("[%s] Execution completed in %s ms", stage_name,
execution_time * 1000)
batch.logging_info.add_stage_execution_time(
stage_name, execution_time)
except Exception as e:
torch.cuda.synchronize()
execution_time = time.perf_counter() - start_time
logger.error("[%s] Error during execution after %s ms: %s",
stage_name, execution_time * 1000, e)
@@ -436,9 +436,9 @@ class CausalDMDDenosingStage(DenoisingStage):
dtype=dtype,
device=device),
"global_end_index":
torch.tensor([0], dtype=torch.long, device=device),
0,
"local_end_index":
torch.tensor([0], dtype=torch.long, device=device),
0,
})
return kv_cache1
@@ -494,4 +494,4 @@ class CausalDMDDenosingStage(DenoisingStage):
result.add_check(
"negative_prompt_embeds", batch.negative_prompt_embeds, lambda x:
not batch.do_classifier_free_guidance or V.list_not_empty(x))
return result
return result
+2 -8
View File
@@ -722,20 +722,14 @@ class DenoisingStage(PipelineStage):
from fastvideo.attention.backends.STA_configuration import save_mask_search_results
if batch.mask_search_final_result_pos is not None and batch.prompt is not None:
save_mask_search_results(
[
dict(layer_data)
for layer_data in batch.mask_search_final_result_pos
],
batch.mask_search_final_result_pos,
prompt=str(batch.prompt),
mask_strategies=sparse_mask_candidates_searching,
output_dir=f'output/mask_search_result_pos_{size[0]}x{size[1]}/'
)
if batch.mask_search_final_result_neg is not None and batch.prompt is not None:
save_mask_search_results(
[
dict(layer_data)
for layer_data in batch.mask_search_final_result_neg
],
batch.mask_search_final_result_neg,
prompt=str(batch.prompt),
mask_strategies=sparse_mask_candidates_searching,
output_dir=f'output/mask_search_result_neg_{size[0]}x{size[1]}/'
@@ -179,11 +179,11 @@ class LatentPreparationStage(PipelineStage):
video_length = batch.num_frames
use_temporal_scaling_frames = fastvideo_args.pipeline_config.vae_config.use_temporal_scaling_frames
if use_temporal_scaling_frames:
temporal_scale_factor = fastvideo_args.pipeline_config.vae_config.arch_config.temporal_compression_ratio
latent_num_frames = (video_length - 1) // temporal_scale_factor + 1
else: # stepvideo only
latent_num_frames = video_length // 17 * 3
if not use_temporal_scaling_frames:
raise ValueError(
"Only temporal-scaling-frame VAE configs are supported.")
temporal_scale_factor = fastvideo_args.pipeline_config.vae_config.arch_config.temporal_compression_ratio
latent_num_frames = (video_length - 1) // temporal_scale_factor + 1
return int(latent_num_frames)
def verify_input(self, batch: ForwardBatch,
@@ -687,11 +687,11 @@ class Cosmos25LatentPreparationStage(CosmosLatentPreparationStage):
video_length = batch.num_frames
use_temporal_scaling_frames = fastvideo_args.pipeline_config.vae_config.use_temporal_scaling_frames
if use_temporal_scaling_frames:
temporal_scale_factor = fastvideo_args.pipeline_config.vae_config.arch_config.temporal_compression_ratio
latent_num_frames = (video_length - 1) // temporal_scale_factor + 1
else: # stepvideo only
latent_num_frames = video_length // 17 * 3
if not use_temporal_scaling_frames:
raise ValueError(
"Only temporal-scaling-frame VAE configs are supported.")
temporal_scale_factor = fastvideo_args.pipeline_config.vae_config.arch_config.temporal_compression_ratio
latent_num_frames = (video_length - 1) // temporal_scale_factor + 1
return int(latent_num_frames)
def verify_input(self, batch: ForwardBatch,
+24 -1
View File
@@ -11,6 +11,9 @@ import os
import torch
from tqdm.auto import tqdm
import fastvideo.envs as envs
from fastvideo.attention.backends.video_sparse_attn import (
VideoSparseAttentionMetadataBuilder)
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
@@ -23,6 +26,7 @@ from fastvideo.models.dits.ltx2 import (
DEFAULT_LTX2_AUDIO_DOWNSAMPLE, DEFAULT_LTX2_AUDIO_HOP_LENGTH,
DEFAULT_LTX2_AUDIO_MEL_BINS, DEFAULT_LTX2_AUDIO_SAMPLE_RATE,
VideoLatentShape)
from fastvideo.utils import is_vsa_available
BASE_SHIFT_ANCHOR = 1024
MAX_SHIFT_ANCHOR = 4096
@@ -35,6 +39,11 @@ DISTILLED_SIGMA_VALUES = [
logger = init_logger(__name__)
try:
vsa_available = is_vsa_available()
except ImportError:
vsa_available = False
def _ltx2_sigmas(
steps: int,
@@ -234,6 +243,10 @@ class LTX2DenoisingStage(PipelineStage):
tuple(sigmas.shape),
tuple(latents.shape),
)
use_vsa = (vsa_available
and envs.FASTVIDEO_ATTENTION_BACKEND == "VIDEO_SPARSE_ATTN")
vsa_metadata_builder = (VideoSparseAttentionMetadataBuilder()
if use_vsa else None)
for step_index in tqdm(range(len(sigmas) - 1)):
sigma = sigmas[step_index]
@@ -241,6 +254,16 @@ class LTX2DenoisingStage(PipelineStage):
timestep = timestep_template * sigma
audio_timestep = (audio_timestep_template * sigma
if audio_timestep_template is not None else None)
attn_metadata = None
if vsa_metadata_builder is not None:
attn_metadata = vsa_metadata_builder.build(
current_timestep=step_index,
raw_latent_shape=latents.shape[2:5],
patch_size=fastvideo_args.pipeline_config.dit_config.
patch_size,
VSA_sparsity=fastvideo_args.VSA_sparsity,
device=latents.device,
)
with torch.autocast(
device_type="cuda",
@@ -248,7 +271,7 @@ class LTX2DenoisingStage(PipelineStage):
enabled=autocast_enabled,
), set_forward_context(
current_timestep=sigma.item(),
attn_metadata=None,
attn_metadata=attn_metadata,
forward_batch=batch,
):
# Pass 1: Full conditioning (text + cross-modal)
@@ -335,9 +335,9 @@ class MatrixGameCausalDenoisingStage(DenoisingStage):
dtype=dtype,
device=device),
"global_end_index":
torch.tensor([0], dtype=torch.long, device=device),
0,
"local_end_index":
torch.tensor([0], dtype=torch.long, device=device),
0,
})
return kv_cache
@@ -373,9 +373,9 @@ class MatrixGameCausalDenoisingStage(DenoisingStage):
dtype=dtype,
device=device),
"global_end_index":
torch.tensor([0], dtype=torch.long, device=device),
0,
"local_end_index":
torch.tensor([0], dtype=torch.long, device=device),
0,
})
kv_cache_mouse.append({
"k":
@@ -393,9 +393,9 @@ class MatrixGameCausalDenoisingStage(DenoisingStage):
dtype=dtype,
device=device),
"global_end_index":
torch.tensor([0], dtype=torch.long, device=device),
0,
"local_end_index":
torch.tensor([0], dtype=torch.long, device=device),
0,
})
return kv_cache_mouse, kv_cache_keyboard
@@ -52,6 +52,15 @@ class SD35LatentPreparationStage(PipelineStage):
in_channels = fastvideo_args.pipeline_config.dit_config.arch_config.in_channels
spatial_ratio = fastvideo_args.pipeline_config.vae_config.arch_config.spatial_compression_ratio
patch_size = fastvideo_args.pipeline_config.dit_config.arch_config.patch_size
if not isinstance(patch_size, int):
raise TypeError(
f"SD3.5 expects integer patch_size, got {type(patch_size)}")
required_divisor = spatial_ratio * patch_size
if batch.height % required_divisor != 0 or batch.width % required_divisor != 0:
raise ValueError(
f"height/width must be divisible by {required_divisor} for SD3.5 "
f"(got height={batch.height}, width={batch.width})")
h_lat = batch.height // spatial_ratio
w_lat = batch.width // spatial_ratio
shape = (batch_size, in_channels, 1, h_lat, w_lat)
@@ -1,80 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
import torch
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.base import PipelineStage
from fastvideo.pipelines.stages.validators import StageValidators as V
from fastvideo.pipelines.stages.validators import VerificationResult
logger = init_logger(__name__)
# The dedicated stepvideo prompt encoding stage.
class StepvideoPromptEncodingStage(PipelineStage):
"""
Stage for encoding prompts using the remote caption API.
This stage applies the magic string transformations and calls
the remote caption service asynchronously to get:
- primary prompt embeddings,
- an attention mask,
- and a clip embedding.
"""
def __init__(self, stepllm, clip) -> None:
super().__init__()
# self.caption_client = caption_client # This should have a call_caption(prompts: List[str]) method.
self.stepllm = stepllm
self.clip = clip
@torch.no_grad()
def forward(self, batch: ForwardBatch, fastvideo_args) -> ForwardBatch:
prompts = [batch.prompt + fastvideo_args.pipeline_config.pos_magic]
bs = len(prompts)
prompts += [fastvideo_args.pipeline_config.neg_magic] * bs
with set_forward_context(current_timestep=0, attn_metadata=None):
y, y_mask = self.stepllm(prompts)
clip_emb, _ = self.clip(prompts)
len_clip = clip_emb.shape[1]
y_mask = torch.nn.functional.pad(y_mask, (len_clip, 0), value=1)
pos_clip, neg_clip = clip_emb[:bs], clip_emb[bs:]
# split positive vs negative text
batch.prompt_embeds = y[:bs] # [bs, seq_len, dim]
batch.negative_prompt_embeds = y[bs:2 * bs] # [bs, seq_len, dim]
batch.prompt_attention_mask = y_mask[:bs] # [bs, seq_len]
batch.negative_attention_mask = y_mask[bs:2 * bs] # [bs, seq_len]
batch.clip_embedding_pos = pos_clip
batch.clip_embedding_neg = neg_clip
return batch
def verify_input(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
"""Verify stepvideo encoding stage inputs."""
result = VerificationResult()
result.add_check("prompt", batch.prompt, V.string_not_empty)
return result
def verify_output(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
"""Verify stepvideo encoding stage outputs."""
result = VerificationResult()
result.add_check("prompt_embeds", batch.prompt_embeds,
[V.is_tensor, V.with_dims(3)])
result.add_check("negative_prompt_embeds", batch.negative_prompt_embeds,
[V.is_tensor, V.with_dims(3)])
result.add_check("prompt_attention_mask", batch.prompt_attention_mask,
[V.is_tensor, V.with_dims(2)])
result.add_check("negative_attention_mask",
batch.negative_attention_mask,
[V.is_tensor, V.with_dims(2)])
result.add_check("clip_embedding_pos", batch.clip_embedding_pos,
[V.is_tensor, V.with_dims(2)])
result.add_check("clip_embedding_neg", batch.clip_embedding_neg,
[V.is_tensor, V.with_dims(2)])
return result
@@ -152,3 +152,78 @@ class Cosmos25TimestepPreparationStage(TimestepPreparationStage):
**extra_kwargs)
batch.timesteps = scheduler.timesteps
return batch
class SD35TimestepPreparationStage(TimestepPreparationStage):
"""SD3/SD3.5 timestep preparation with optional dynamic shifting (mu).
When the scheduler supports `use_dynamic_shifting`, this stage computes a
resolution-dependent `mu` value and passes it to `set_timesteps()`.
"""
@staticmethod
def _calculate_mu(
image_seq_len: int,
base_seq_len: int = 256,
max_seq_len: int = 4096,
base_shift: float = 0.5,
max_shift: float = 1.15,
) -> float:
m = (max_shift - base_shift) / (max_seq_len - base_seq_len)
b = base_shift - m * base_seq_len
return float(image_seq_len) * m + b
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
sig = inspect.signature(self.scheduler.set_timesteps)
if "mu" in sig.parameters:
cfg = getattr(self.scheduler, "config", None)
use_dynamic = bool(getattr(cfg, "use_dynamic_shifting",
False)) if cfg is not None else False
if use_dynamic:
arch = fastvideo_args.pipeline_config.dit_config.arch_config
vae_arch = fastvideo_args.pipeline_config.vae_config.arch_config
patch_size = getattr(arch, "patch_size", None)
spatial_ratio = getattr(vae_arch, "spatial_compression_ratio",
None)
if not isinstance(patch_size, int) or not isinstance(
spatial_ratio, int):
raise TypeError(
"SD3.5 dynamic shifting requires integer patch_size "
"and spatial_compression_ratio.")
if batch.height is None or batch.width is None:
raise ValueError(
"height/width must be set before timesteps.")
h_lat = batch.height // spatial_ratio
w_lat = batch.width // spatial_ratio
image_seq_len = (h_lat // patch_size) * (w_lat // patch_size)
base_seq_len = int(getattr(cfg, "base_image_seq_len", 256))
max_seq_len = int(getattr(cfg, "max_image_seq_len", 4096))
base_shift = float(getattr(cfg, "base_shift", 0.5))
max_shift = float(getattr(cfg, "max_shift", 1.15))
batch.n_tokens = None
device = get_local_torch_device()
self.scheduler.set_timesteps(
batch.num_inference_steps,
device=device,
mu=self._calculate_mu(
image_seq_len=image_seq_len,
base_seq_len=base_seq_len,
max_seq_len=max_seq_len,
base_shift=base_shift,
max_shift=max_shift,
),
)
batch.timesteps = self.scheduler.timesteps
return batch
return super().forward(batch, fastvideo_args)
+1 -13
View File
@@ -26,7 +26,6 @@ from fastvideo.configs.pipelines.hyworld import HYWorldConfig
from fastvideo.configs.pipelines.lingbotworld import LingBotWorldI2V480PConfig
from fastvideo.configs.pipelines.longcat import LongCatT2V480PConfig
from fastvideo.configs.pipelines.ltx2 import LTX2T2VConfig
from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
from fastvideo.configs.pipelines.turbodiffusion import (
TurboDiffusionI2V_A14B_Config,
TurboDiffusionT2V_14B_Config,
@@ -64,7 +63,6 @@ from fastvideo.configs.sample.hunyuangamecraft import HunyuanGameCraftSamplingPa
from fastvideo.configs.sample.lingbotworld import LingBotWorld_SamplingParam
from fastvideo.configs.sample.ltx2 import (LTX2BaseSamplingParam,
LTX2DistilledSamplingParam)
from fastvideo.configs.sample.stepvideo import StepVideoT2VSamplingParam
from fastvideo.configs.sample.turbodiffusion import (
TurboDiffusionI2V_A14B_SamplingParam,
TurboDiffusionT2V_14B_SamplingParam,
@@ -389,16 +387,6 @@ def _register_configs() -> None:
],
)
# StepVideo
register_configs(
sampling_param_cls=StepVideoT2VSamplingParam,
pipeline_config_cls=StepVideoT2VConfig,
hf_model_paths=[
"FastVideo/stepvideo-t2v-diffusers",
],
model_detectors=[lambda path: "stepvideo" in path.lower()],
)
# MatrixGame
register_configs(
sampling_param_cls=MatrixGame2_SamplingParam,
@@ -679,4 +667,4 @@ __all__ = [
"get_pipeline_config_cls_from_name",
"get_sampling_param_cls_for_name",
"get_pipeline_config_classes",
]
]
+12 -1
View File
@@ -42,6 +42,10 @@ TEST_PROMPTS = [
"a photo of a cat",
]
pytestmark = pytest.mark.filterwarnings(
"ignore:.*torch.jit.script_method.*:DeprecationWarning",
)
@pytest.mark.skipif(not torch.cuda.is_available(), reason="SD3.5 SSIM test requires CUDA")
@pytest.mark.parametrize("ATTENTION_BACKEND", ["TORCH_SDPA"])
@@ -59,7 +63,9 @@ def test_sd35_similarity(prompt: str, ATTENTION_BACKEND: str) -> None:
)
old_backend = os.environ.get("FASTVIDEO_ATTENTION_BACKEND")
old_transformers_verbosity = os.environ.get("TRANSFORMERS_VERBOSITY")
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = ATTENTION_BACKEND
os.environ["TRANSFORMERS_VERBOSITY"] = "error"
try:
script_dir = os.path.dirname(os.path.abspath(__file__))
output_dir = os.path.join(
@@ -80,7 +86,8 @@ def test_sd35_similarity(prompt: str, ATTENTION_BACKEND: str) -> None:
init_kwargs = {
"num_gpus": 1,
"workload_type": "t2i",
# Save a single-frame MP4 to reuse the existing video-SSIM harness.
"workload_type": "t2v",
"sp_size": 1,
"tp_size": 1,
"dit_cpu_offload": False,
@@ -191,3 +198,7 @@ def test_sd35_similarity(prompt: str, ATTENTION_BACKEND: str) -> None:
os.environ.pop("FASTVIDEO_ATTENTION_BACKEND", None)
else:
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = old_backend
if old_transformers_verbosity is None:
os.environ.pop("TRANSFORMERS_VERBOSITY", None)
else:
os.environ["TRANSFORMERS_VERBOSITY"] = old_transformers_verbosity
@@ -51,6 +51,7 @@ def apply_activation_checkpointing(
def _apply_activation_checkpointing_blocks(module: torch.nn.Module,
n_layer: int | None = None
) -> torch.nn.Module:
applied = False
for transformer_block_name in TRANSFORMER_BLOCK_NAMES:
blocks: torch.nn.Module = getattr(module, transformer_block_name, None)
if blocks is None:
@@ -59,6 +60,9 @@ def _apply_activation_checkpointing_blocks(module: torch.nn.Module,
if n_layer is None or index % n_layer == 0:
block = checkpoint_wrapper(block, preserve_rng_state=False)
blocks.register_module(layer_id, block)
applied = True
if not applied:
raise ValueError("Activation checkpointing is not applied successfully")
return module
+105 -36
View File
@@ -75,6 +75,30 @@ class DistillationPipeline(TrainingPipeline):
video_latent_shape_sp: tuple[int, ...]
train_fake_score_transformer_2: bool = False
@staticmethod
def _clone_batch_value(value: Any) -> Any:
"""Clone values in a TrainingBatch without tensor deepcopy."""
if isinstance(value, torch.Tensor):
# Avoid torch tensor deepcopy limitations on non-leaf tensors.
return value.detach().clone()
if isinstance(value, dict):
return {
k: DistillationPipeline._clone_batch_value(v)
for k, v in value.items()
}
if isinstance(value, list):
return [DistillationPipeline._clone_batch_value(v) for v in value]
if isinstance(value, tuple):
return tuple(
DistillationPipeline._clone_batch_value(v) for v in value)
return copy.deepcopy(value)
def _clone_training_batch(self, batch: TrainingBatch) -> TrainingBatch:
cloned_batch = TrainingBatch()
for key, value in batch.__dict__.items():
setattr(cloned_batch, key, self._clone_batch_value(value))
return cloned_batch
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
raise RuntimeError(
"create_pipeline_stages should not be called for training pipeline")
@@ -971,13 +995,13 @@ class DistillationPipeline(TrainingPipeline):
batch.attn_metadata.VSA_sparsity = 0.0 # type: ignore
batches.append(batch)
self.optimizer.zero_grad()
self.optimizer.zero_grad(set_to_none=True)
total_dmd_loss = 0.0
dmd_latent_vis_dict = {}
fake_score_latent_vis_dict = {}
if (self.current_trainstep % self.generator_update_interval == 0):
for batch in batches:
batch_gen = copy.deepcopy(batch)
batch_gen = self._clone_training_batch(batch)
with set_forward_context(
current_timestep=batch_gen.timesteps,
@@ -1022,12 +1046,12 @@ class DistillationPipeline(TrainingPipeline):
else:
training_batch.generator_loss = 0.0
self.fake_score_optimizer.zero_grad()
self.fake_score_optimizer.zero_grad(set_to_none=True)
if self.fake_score_transformer_2 is not None:
self.fake_score_optimizer_2.zero_grad()
self.fake_score_optimizer_2.zero_grad(set_to_none=True)
total_fake_score_loss = 0.0
for batch in batches:
batch_fake = copy.deepcopy(batch)
batch_fake = self._clone_training_batch(batch)
batch_fake, fake_score_loss = self.faker_score_forward(batch_fake)
with set_forward_context(current_timestep=batch_fake.timesteps,
attn_metadata=batch_fake.attn_metadata):
@@ -1236,9 +1260,12 @@ class DistillationPipeline(TrainingPipeline):
# Helper function to run validation with optional EMA contexts
def run_validation_with_ema(
steps: int) -> tuple[list[np.ndarray], list[str]]:
steps: int
) -> tuple[list[np.ndarray], list[str], list[Any], list[Any]]:
videos: list[np.ndarray] = []
captions: list[str] = []
audios: list[Any] = []
audio_sample_rates: list[Any] = []
for validation_batch in validation_dataloader:
batch = self._prepare_validation_batch(
sampling_param, training_args, validation_batch, steps)
@@ -1282,25 +1309,32 @@ class DistillationPipeline(TrainingPipeline):
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
frames.append((x * 255).numpy().astype(np.uint8))
videos.append(frames)
audios.append(output_batch.extra.get("audio"))
audio_sample_rates.append(
output_batch.extra.get("audio_sample_rate"))
return videos, captions
return videos, captions, audios, audio_sample_rates
# Apply EMA contexts if available (nested context managers)
if ema_context is not None and ema_2_context is not None:
with ema_context, ema_2_context:
step_videos, step_captions = run_validation_with_ema(
num_inference_steps)
(step_videos, step_captions, step_audios,
step_audio_sample_rates
) = run_validation_with_ema(num_inference_steps)
elif ema_context is not None:
with ema_context:
step_videos, step_captions = run_validation_with_ema(
num_inference_steps)
(step_videos, step_captions, step_audios,
step_audio_sample_rates
) = run_validation_with_ema(num_inference_steps)
elif ema_2_context is not None:
with ema_2_context:
step_videos, step_captions = run_validation_with_ema(
num_inference_steps)
(step_videos, step_captions, step_audios,
step_audio_sample_rates
) = run_validation_with_ema(num_inference_steps)
else:
step_videos, step_captions = run_validation_with_ema(
num_inference_steps)
(step_videos, step_captions, step_audios,
step_audio_sample_rates
) = run_validation_with_ema(num_inference_steps)
# Log validation results for this step
world_group = get_world_group()
@@ -1313,24 +1347,42 @@ class DistillationPipeline(TrainingPipeline):
# Global rank 0 collects results from all sp_group leaders
all_videos = step_videos # Start with own results
all_captions = step_captions
all_audios = step_audios
all_audio_sample_rates = step_audio_sample_rates
# Receive from other sp_group leaders
for sp_group_idx in range(1, num_sp_groups):
src_rank = sp_group_idx * self.sp_world_size # Global rank of other sp_group leaders
recv_videos = world_group.recv_object(src=src_rank)
recv_captions = world_group.recv_object(src=src_rank)
recv_audios = world_group.recv_object(src=src_rank)
recv_audio_sample_rates = world_group.recv_object(
src=src_rank)
all_videos.extend(recv_videos)
all_captions.extend(recv_captions)
all_audios.extend(recv_audios)
all_audio_sample_rates.extend(recv_audio_sample_rates)
video_filenames = []
for i, (video, caption) in enumerate(
zip(all_videos, all_captions, strict=True)):
for i, (video, caption, audio,
audio_sample_rate) in enumerate(
zip(all_videos,
all_captions,
all_audios,
all_audio_sample_rates,
strict=True)):
os.makedirs(training_args.output_dir, exist_ok=True)
filename = os.path.join(
training_args.output_dir,
f"validation_step_{global_step}_inference_steps_{num_inference_steps}_video_{i}.mp4"
)
imageio.mimsave(filename, video, fps=sampling_param.fps)
if (audio is not None and audio_sample_rate is not None
and not self._mux_audio(filename, audio,
audio_sample_rate)):
logger.warning(
"Audio mux failed for validation video %s; saved video without audio.",
filename)
video_filenames.append(filename)
artifacts = []
@@ -1351,16 +1403,38 @@ class DistillationPipeline(TrainingPipeline):
# Other sp_group leaders send their results to global rank 0
world_group.send_object(step_videos, dst=0)
world_group.send_object(step_captions, dst=0)
world_group.send_object(step_audios, dst=0)
world_group.send_object(step_audio_sample_rates, dst=0)
# Re-enable gradients for training - set both transformers back to train mode
transformer.train()
if hasattr(self, 'transformer_2') and self.transformer_2 is not None:
self.transformer_2.train()
gc.collect()
torch.cuda.empty_cache()
def visualize_intermediate_latents(self, training_batch: TrainingBatch,
training_args: TrainingArgs, step: int):
"""Add visualization data to tracker logging and save frames to disk."""
def _prepare_vae_latents(latents: torch.Tensor) -> torch.Tensor:
# Most training paths store latents as [B, C, T, H, W]. Older
# visualization code permuted to [B, T, C, H, W] for WAN-like VAEs.
# LTX2 VAE expects [B, C, T, H, W], so keep native layout there.
if hasattr(self.vae, "TIME_SCALE") and hasattr(
self.vae, "SPATIAL_SCALE"):
return latents
return latents.permute(0, 2, 1, 3, 4)
def _apply_vae_scale(latents: torch.Tensor) -> torch.Tensor:
scaling_factor = getattr(self.vae, "scaling_factor", None)
if scaling_factor is None:
return latents
if isinstance(scaling_factor, torch.Tensor):
return latents / scaling_factor.to(latents.device,
latents.dtype)
return latents / scaling_factor
tracker_loss_dict: dict[str, Any] = {}
dmd_latents_vis_dict = training_batch.dmd_latent_vis_dict
fake_score_latents_vis_dict = training_batch.fake_score_latent_vis_dict
@@ -1369,13 +1443,8 @@ class DistillationPipeline(TrainingPipeline):
for latent_key in fake_score_log_keys:
latents = fake_score_latents_vis_dict[latent_key]
latents = latents.permute(0, 2, 1, 3, 4)
if isinstance(self.vae.scaling_factor, torch.Tensor):
latents = latents / self.vae.scaling_factor.to(
latents.device, latents.dtype)
else:
latents = latents / self.vae.scaling_factor
latents = _prepare_vae_latents(latents)
latents = _apply_vae_scale(latents)
# Apply shifting if needed
if (hasattr(self.vae, "shift_factor")
@@ -1402,13 +1471,9 @@ class DistillationPipeline(TrainingPipeline):
if 'generator_pred_video' in dmd_latents_vis_dict:
for latent_key in dmd_log_keys:
latents = dmd_latents_vis_dict[latent_key]
latents = latents.permute(0, 2, 1, 3, 4)
latents = _prepare_vae_latents(latents)
# decoded_latent = decode_stage(ForwardBatch(data_type="video", latents=latents), training_args)
if isinstance(self.vae.scaling_factor, torch.Tensor):
latents = latents / self.vae.scaling_factor.to(
latents.device, latents.dtype)
else:
latents = latents / self.vae.scaling_factor
latents = _apply_vae_scale(latents)
# Apply shifting if needed
if (hasattr(self.vae, "shift_factor")
@@ -1451,13 +1516,14 @@ class DistillationPipeline(TrainingPipeline):
set_random_seed(seed + self.global_rank)
# Set random seeds for deterministic training
self.noise_random_generator = torch.Generator(device="cpu").manual_seed(
self.seed)
self.noise_gen_cuda = torch.Generator(device="cuda").manual_seed(
self.seed)
self.noise_random_generator = torch.Generator(
device="cpu").manual_seed(self.seed + self.global_rank)
self.noise_gen_cuda = torch.Generator(
device="cuda").manual_seed(self.seed + self.global_rank)
self.validation_random_generator = torch.Generator(
device="cpu").manual_seed(self.seed)
logger.info("Initialized random seeds with seed: %s", seed)
device="cpu").manual_seed(self.seed + self.global_rank)
logger.info("Initialized random seeds with seed: %s",
seed + self.global_rank)
# Initialize current_trainstep for EMA ready checks
#TODO: check if needed
@@ -1488,6 +1554,9 @@ class DistillationPipeline(TrainingPipeline):
use_vsa = vsa_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VIDEO_SPARSE_ATTN"
for step in range(self.init_steps + 1,
self.training_args.max_train_steps + 1):
if step % 5 == 0:
gc.collect()
torch.cuda.empty_cache()
start_time = time.perf_counter()
if use_vsa:
vsa_sparsity = self.training_args.VSA_sparsity
+142 -58
View File
@@ -6,7 +6,9 @@ from pathlib import Path
import torch
import torch.distributed as dist
from fastvideo.dataset import build_ltx2_precomputed_dataloader
from fastvideo.dataset import (build_ltx2_precomputed_dataloader,
build_parquet_map_style_dataloader)
from fastvideo.dataset.dataloader.schema import pyarrow_schema_text_only
from fastvideo.distributed import get_local_torch_device, get_sp_group, get_world_group
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.forward_context import set_forward_context
@@ -64,14 +66,14 @@ class LTX2TrainingPipeline(TrainingPipeline):
self.seed = training_args.seed
assert self.seed is not None, "seed must be set"
set_random_seed(self.seed)
set_random_seed(self.seed + self.global_rank)
self.transformer.train()
if training_args.enable_gradient_checkpointing_type is not None:
from fastvideo.training.activation_checkpoint import (
apply_activation_checkpointing)
self.transformer = apply_activation_checkpointing(
self.transformer,
self.transformer.model = apply_activation_checkpointing(
self.transformer.model,
checkpointing_type=training_args.
enable_gradient_checkpointing_type)
@@ -103,17 +105,37 @@ class LTX2TrainingPipeline(TrainingPipeline):
last_epoch=self.init_steps - 1,
)
data_sources = self._get_ltx2_data_sources(training_args.data_path)
self.with_audio = "audio_latents" in data_sources
self.train_dataset, self.train_dataloader = (
build_ltx2_precomputed_dataloader(
training_args.data_path,
training_args.train_batch_size,
num_data_workers=training_args.dataloader_num_workers,
data_sources=data_sources,
drop_last=True,
seed=self.seed,
))
if self._has_precomputed_pt_data(training_args.data_path):
data_sources = self._get_ltx2_data_sources(training_args.data_path)
self.with_audio = "audio_latents" in data_sources
self.train_dataset, self.train_dataloader = (
build_ltx2_precomputed_dataloader(
training_args.data_path,
training_args.train_batch_size,
num_data_workers=training_args.dataloader_num_workers,
data_sources=data_sources,
drop_last=True,
seed=self.seed,
))
self._use_parquet_data = False
logger.info("Using precomputed .pt data")
else:
text_padding_length = (training_args.pipeline_config.
text_encoder_configs[0].arch_config.text_len)
self.with_audio = False
self.train_dataset, self.train_dataloader = (
build_parquet_map_style_dataloader(
training_args.data_path,
training_args.train_batch_size,
num_data_workers=training_args.dataloader_num_workers,
parquet_schema=pyarrow_schema_text_only,
cfg_rate=training_args.training_cfg_rate,
drop_last=True,
text_padding_length=text_padding_length,
seed=self.seed,
))
self._use_parquet_data = True
logger.info("No precomputed .pt data found, using parquet data")
self.num_update_steps_per_epoch = max(
1,
@@ -168,6 +190,17 @@ class LTX2TrainingPipeline(TrainingPipeline):
)
self.validation_pipeline = validation_pipeline
def _has_precomputed_pt_data(self, data_path: str) -> bool:
"""Check whether precomputed .pt data exists for LTX2."""
data_root = Path(data_path).expanduser().resolve()
if (data_root / ".precomputed").exists():
data_root = data_root / ".precomputed"
latents_dir = data_root / "latents"
conditions_dir = data_root / "conditions"
return (latents_dir.exists() and conditions_dir.exists()
and any(latents_dir.rglob("*.pt"))
and any(conditions_dir.rglob("*.pt")))
def _get_ltx2_data_sources(self, data_path: str) -> dict[str, str]:
data_root = Path(data_path).expanduser().resolve()
if (data_root / ".precomputed").exists():
@@ -187,52 +220,104 @@ class LTX2TrainingPipeline(TrainingPipeline):
self.train_loader_iter = iter(self.train_dataloader)
batch = next(self.train_loader_iter)
latents = batch["latents"]["latents"].to(get_local_torch_device(),
dtype=torch.bfloat16)
latents = latents[:, :, :self.training_args.num_latent_t]
conditions = batch["conditions"]
if ("video_prompt_embeds" in conditions
and "audio_prompt_embeds" in conditions
and "prompt_attention_mask" in conditions):
video_embeds = conditions["video_prompt_embeds"].to(
get_local_torch_device())
audio_embeds = conditions["audio_prompt_embeds"].to(
get_local_torch_device())
attention_mask = conditions["prompt_attention_mask"].to(
get_local_torch_device(), dtype=torch.int64)
if self._use_parquet_data:
training_batch = self._get_next_batch_parquet(
batch, training_batch)
else:
prompt_embeds = conditions["prompt_embeds"].to(
get_local_torch_device())
prompt_attention_mask = conditions["prompt_attention_mask"].to(
get_local_torch_device(), dtype=torch.int64)
training_batch = self._get_next_batch_pt(batch, training_batch)
video_embeds, audio_embeds, attention_mask = (
self.text_encoder.run_connectors(prompt_embeds,
prompt_attention_mask))
return training_batch
training_batch.latents = latents.to(get_local_torch_device(),
dtype=torch.bfloat16)
training_batch.encoder_hidden_states = video_embeds.to(
get_local_torch_device(), dtype=torch.bfloat16)
training_batch.encoder_attention_mask = attention_mask.to(
def _get_next_batch_parquet(
self,
batch: dict,
training_batch: TrainingBatch,
) -> TrainingBatch:
"""Extract training data from a parquet-format batch."""
device = get_local_torch_device()
prompt_embeds = batch["text_embedding"].to(device)
prompt_attention_mask = batch["text_attention_mask"].to(
device, dtype=torch.int64)
video_embeds, audio_embeds, attention_mask = (
self.text_encoder.run_connectors(prompt_embeds,
prompt_attention_mask))
# Parquet text-only data has no latents; create zeros.
batch_size = video_embeds.shape[0]
vae_cfg = (self.training_args.pipeline_config.vae_config.arch_config)
num_channels = vae_cfg.z_dim
scr = vae_cfg.spatial_compression_ratio
latent_h = self.training_args.num_height // scr
latent_w = self.training_args.num_width // scr
latents = torch.zeros(batch_size,
num_channels,
self.training_args.num_latent_t,
latent_h,
latent_w,
device=device,
dtype=torch.bfloat16)
training_batch.latents = latents
training_batch.encoder_hidden_states = video_embeds.to(
device, dtype=torch.bfloat16)
training_batch.encoder_attention_mask = attention_mask.to(
device, dtype=torch.bfloat16)
training_batch.infos = []
training_batch.raw_latent_shape = latents.shape
return training_batch
def _get_next_batch_pt(
self,
batch: dict,
training_batch: TrainingBatch,
) -> TrainingBatch:
"""Extract training data from a precomputed .pt batch."""
latents = batch["latents"]["latents"].to(get_local_torch_device(),
dtype=torch.bfloat16)
latents = latents[:, :, :self.training_args.num_latent_t]
conditions = batch["conditions"]
if ("video_prompt_embeds" in conditions
and "audio_prompt_embeds" in conditions
and "prompt_attention_mask" in conditions):
video_embeds = conditions["video_prompt_embeds"].to(
get_local_torch_device())
audio_embeds = conditions["audio_prompt_embeds"].to(
get_local_torch_device())
attention_mask = conditions["prompt_attention_mask"].to(
get_local_torch_device(), dtype=torch.int64)
else:
prompt_embeds = conditions["prompt_embeds"].to(
get_local_torch_device())
prompt_attention_mask = conditions["prompt_attention_mask"].to(
get_local_torch_device(), dtype=torch.int64)
video_embeds, audio_embeds, attention_mask = (
self.text_encoder.run_connectors(prompt_embeds,
prompt_attention_mask))
training_batch.latents = latents.to(get_local_torch_device(),
dtype=torch.bfloat16)
training_batch.encoder_hidden_states = video_embeds.to(
get_local_torch_device(), dtype=torch.bfloat16)
training_batch.encoder_attention_mask = attention_mask.to(
get_local_torch_device(), dtype=torch.bfloat16)
if self.with_audio and "audio_latents" in batch:
audio_latents = batch["audio_latents"]["latents"].to(
get_local_torch_device(), dtype=torch.bfloat16)
training_batch.audio_latents = audio_latents
training_batch.audio_encoder_hidden_states = (audio_embeds.to(
get_local_torch_device(), dtype=torch.bfloat16))
training_batch.audio_encoder_attention_mask = (attention_mask.to(
get_local_torch_device()))
if self.with_audio and "audio_latents" in batch:
audio_latents = batch["audio_latents"]["latents"].to(
get_local_torch_device(), dtype=torch.bfloat16)
training_batch.audio_latents = audio_latents
training_batch.audio_encoder_hidden_states = audio_embeds.to(
get_local_torch_device(), dtype=torch.bfloat16)
training_batch.audio_encoder_attention_mask = attention_mask.to(
get_local_torch_device())
idxs = batch.get("idx")
if idxs is not None and torch.is_tensor(idxs):
training_batch.infos = [{"idx": int(i)} for i in idxs.tolist()]
else:
training_batch.infos = []
training_batch.raw_latent_shape = latents.shape
idxs = batch.get("idx")
if idxs is not None and torch.is_tensor(idxs):
training_batch.infos = [{"idx": int(i)} for i in idxs.tolist()]
else:
training_batch.infos = []
training_batch.raw_latent_shape = latents.shape
return training_batch
def _normalize_dit_input(self,
@@ -339,8 +424,7 @@ class LTX2TrainingPipeline(TrainingPipeline):
def _build_attention_metadata(
self, training_batch: TrainingBatch) -> TrainingBatch:
training_batch.attn_metadata = None
return training_batch
return super()._build_attention_metadata(training_batch)
def _build_input_kwargs(self,
training_batch: TrainingBatch) -> TrainingBatch:
@@ -582,9 +582,9 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
dtype=dtype,
device=device),
"global_end_index":
torch.tensor([0], dtype=torch.long, device=device),
0,
"local_end_index":
torch.tensor([0], dtype=torch.long, device=device)
0
})
# Initialize cross-attention cache
@@ -616,8 +616,8 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
"""Reset KV cache and cross-attention cache to clean state."""
if kv_cache is not None:
for cache_dict in kv_cache:
cache_dict["global_end_index"].fill_(0)
cache_dict["local_end_index"].fill_(0)
cache_dict["global_end_index"] = 0
cache_dict["local_end_index"] = 0
cache_dict["k"].zero_()
cache_dict["v"].zero_()
+30 -7
View File
@@ -856,6 +856,9 @@ class TrainingPipeline(LoRAPipeline, ABC):
step_videos: list[np.ndarray] = []
step_captions: list[str] = []
step_audio: list[np.ndarray | None] = []
step_sample_rates: list[int | None] = []
for validation_batch in validation_dataloader:
batch = self._prepare_validation_batch(sampling_param,
training_args,
@@ -876,6 +879,16 @@ class TrainingPipeline(LoRAPipeline, ABC):
batch, training_args)
samples = output_batch.output.cpu()
# Capture audio if available
audio = output_batch.extra.get("audio")
sample_rate = output_batch.extra.get("audio_sample_rate")
if audio is not None and torch.is_tensor(audio):
audio = audio.detach().cpu().float().numpy()
step_audio.append(audio)
step_sample_rates.append(sample_rate)
if self.rank_in_sp_group != 0:
continue
@@ -894,18 +907,29 @@ class TrainingPipeline(LoRAPipeline, ABC):
# Global rank 0 collects results from all sp_group leaders
all_videos = step_videos # Start with own results
all_captions = step_captions
all_audios = step_audio
all_sample_rates = step_sample_rates
# Receive from other sp_group leaders
for sp_group_idx in range(1, num_sp_groups):
src_rank = sp_group_idx * self.sp_world_size # Global rank of other sp_group leaders
recv_videos = world_group.recv_object(src=src_rank)
recv_captions = world_group.recv_object(src=src_rank)
recv_audios = world_group.recv_object(src=src_rank)
recv_sample_rates = world_group.recv_object(src=src_rank)
all_videos.extend(recv_videos)
all_captions.extend(recv_captions)
all_audios.extend(recv_audios)
all_sample_rates.extend(recv_sample_rates)
video_filenames = []
for i, (video, caption) in enumerate(
zip(all_videos, all_captions, strict=True)):
for i, (video, caption, audio, sample_rate) in enumerate(
zip(all_videos,
all_captions,
all_audios,
all_sample_rates,
strict=True)):
os.makedirs(training_args.output_dir, exist_ok=True)
filename = os.path.join(
training_args.output_dir,
@@ -913,14 +937,11 @@ class TrainingPipeline(LoRAPipeline, ABC):
)
imageio.mimsave(filename, video, fps=sampling_param.fps)
# Mux audio if available
audio = output_batch.extra.get("audio")
audio_sample_rate = output_batch.extra.get(
"audio_sample_rate")
if (audio is not None and audio_sample_rate is not None
if (audio is not None and sample_rate is not None
and not self._mux_audio(
filename,
audio,
audio_sample_rate,
sample_rate,
)):
logger.warning(
"Audio mux failed for validation video %s; saved video without audio.",
@@ -945,6 +966,8 @@ class TrainingPipeline(LoRAPipeline, ABC):
# Other sp_group leaders send their results to global rank 0
world_group.send_object(step_videos, dst=0)
world_group.send_object(step_captions, dst=0)
world_group.send_object(step_audio, dst=0)
world_group.send_object(step_sample_rates, dst=0)
# Re-enable gradients for training
training_args.inference_mode = False
+1
View File
@@ -125,6 +125,7 @@ nav:
- Inference:
- Quick Start: inference/inference_quick_start.md
- Configuration: inference/configuration.md
- Offloading: inference/offloading.md
- Optimizations: inference/optimizations.md
- ComfyUI: inference/comfyui.md
- Support Matrix: inference/support_matrix.md
+1 -5
View File
@@ -76,7 +76,7 @@ dependencies = [
"av",
# Preprocessing Dependencies
"torchcodec==0.5.0",
"torchcodec",
"ray>=2.49.1",
"ftfy==6.3.1",
]
@@ -179,11 +179,7 @@ ignore = [
]
[tool.ruff.lint.per-file-ignores]
"fastvideo/models/stepvideo/diffusion/video_pipeline.py" = ["F821"]
"fastvideo/sample/call_remote_server_stepvideo.py" = ["E722"]
"csrc/sliding_tile_attention/test/bench.py" = ["F841"]
"fastvideo/models/stepvideo/__init__.py" = ["F403"]
"fastvideo/models/stepvideo/utils/__init__.py" = ["F403"]
# Ignore all files that end in `_test.py`.
"fastvideo/models/hunyuan/diffusion/pipelines/pipeline_hunyuan_video.py" = ["E741"]
+1 -5
View File
@@ -75,7 +75,7 @@ dependencies = [
"av",
# Preprocessing Dependencies
"torchcodec==0.5.0",
"torchcodec",
"ray>=2.49.1",
]
@@ -157,11 +157,7 @@ ignore = [
]
[tool.ruff.lint.per-file-ignores]
"fastvideo/models/stepvideo/diffusion/video_pipeline.py" = ["F821"]
"fastvideo/sample/call_remote_server_stepvideo.py" = ["E722"]
"csrc/sliding_tile_attention/test/bench.py" = ["F841"]
"fastvideo/models/stepvideo/__init__.py" = ["F403"]
"fastvideo/models/stepvideo/utils/__init__.py" = ["F403"]
# Ignore all files that end in `_test.py`.
"fastvideo/models/hunyuan/diffusion/pipelines/pipeline_hunyuan_video.py" = ["E741"]

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