Compare commits
13
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
e4de77692f | ||
|
|
077232aa5f | ||
|
|
03d9ce2edb | ||
|
|
10fc92dba5 | ||
|
|
6736dc06a5 | ||
|
|
8c002c62af | ||
|
|
7061313d04 | ||
|
|
76d3ba69e0 | ||
|
|
8e39ce38c9 | ||
|
|
d4bd8bf2c0 | ||
|
|
e83d7bc50c | ||
|
|
ff3d5aff75 | ||
|
|
959dbcc8a2 |
@@ -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
@@ -33,11 +33,11 @@ env
|
||||
**.txt
|
||||
*.log
|
||||
weights/
|
||||
official_weights/
|
||||
converted_weights/
|
||||
|
||||
# SSIM test outputs
|
||||
fastvideo/tests/ssim/generated_videos/
|
||||
**/.cache/**
|
||||
|
||||
|
||||
# Distribution / packaging
|
||||
build/
|
||||
|
||||
@@ -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
@@ -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.
|
||||
@@ -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.
|
||||
@@ -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.
|
||||
@@ -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.
|
||||
@@ -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
|
||||
@@ -0,0 +1 @@
|
||||
fox in the forest close-up quickly turned its head to the left
|
||||
@@ -0,0 +1 @@
|
||||
Man walking his dog in the woods on a hot sunny day
|
||||
@@ -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.
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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.
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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`.
|
||||
@@ -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"
|
||||
|
||||
@@ -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[@]}"
|
||||
|
||||
+114
@@ -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/*
|
||||
+5
-3
@@ -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`).
|
||||
+15
-17
@@ -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[@]}" \
|
||||
+1
-1
@@ -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
|
||||
}
|
||||
]
|
||||
}
|
||||
+7
-15
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"
|
||||
]
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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
|
||||
@@ -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"
|
||||
]
|
||||
|
||||
@@ -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,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
|
||||
|
||||
@@ -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"
|
||||
@@ -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 = "画面暗、低分辨率、不良手、文本、缺少手指、多余的手指、裁剪、低质量、颗粒状、签名、水印、用户名、模糊。"
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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)",
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
+1056
-19
File diff suppressed because it is too large
Load Diff
@@ -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
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -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",
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
@@ -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",
|
||||
]
|
||||
]
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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_()
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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"]
|
||||
|
||||
|
||||
@@ -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
Reference in New Issue
Block a user