Compare commits

..
25 Commits
Author SHA1 Message Date
Will Lin 141a1140f6 refactor sampling pipeline 2026-01-20 15:40:53 -08:00
Shijie Wang 6294015389 Debug transformer output misalignment 2026-01-20 14:38:47 -08:00
Shijie Wang e31b6c9e90 Fix OCR Rewards 2026-01-20 14:38:22 -08:00
Tamoghno Kandar bf0ff21eeb Fix OCR Rewards 2026-01-20 14:38:22 -08:00
Shijie Wang 6f937102ad Enable validation videos 2026-01-20 14:38:22 -08:00
Tamoghno Kandar 0164e93019 Add Validation Loop 2026-01-20 14:38:21 -08:00
Shijie Wang 67e457aa92 resolved cuda OOM error 2026-01-20 14:38:21 -08:00
Shijie Wang d795f0c443 remove additional sampling pipeline 2026-01-20 14:38:21 -08:00
loaydatrain 02452dd6e7 fixed dtype mismatch 2026-01-20 14:38:21 -08:00
Shijie Wang 91ef24bc14 update run script 2026-01-20 14:38:20 -08:00
Shijie Wang e76e9fda15 minor fix 2026-01-20 14:38:20 -08:00
Shijie Wang 3b17f5a621 fix trajectory collection & reward computation 2026-01-20 14:38:20 -08:00
Shijie Wang d758878705 minor fix 2026-01-20 14:38:19 -08:00
Shijie Wang 689e629420 Add entry point script 2026-01-20 14:38:19 -08:00
Shijie Wang 873dc9695f Complete train_one_step and grpo policy loss 2026-01-20 14:38:19 -08:00
Shijie Wang bfc0f46d61 Implement trajectories collection, reward and advantage computing 2026-01-20 14:38:18 -08:00
Shijie Wang 39907dbe4d Port per-prompt stat tracker 2026-01-20 14:38:18 -08:00
Shijie Wang abdd0c9b6a Implement SDE step & SDE pipeline with log prob 2026-01-20 14:38:18 -08:00
Shijie (Jacob) Wang f32a12200d Refactor and trim down unnecessary RL args 2026-01-20 14:38:18 -08:00
Shijie Wang f1d2c9e6b7 Add RL dataset & dataloader 2026-01-20 14:38:18 -08:00
Jiali Chen 450579cb42 init algorithm backbone and refactor rl_pipeline 2026-01-20 14:38:17 -08:00
Jiali Chen 26d7d6cc08 minor bug fix 2026-01-20 14:38:17 -08:00
Jiali Chen 44f0124eaa refactor and add ocr reward model 2026-01-20 14:38:17 -08:00
Jiali Chen d3ace51394 Phase 1 minor fixes 2026-01-20 14:38:17 -08:00
Jiali Chen 58954c660b implement Phase 1 backbone code 2026-01-20 14:38:16 -08:00
228 changed files with 7156 additions and 27904 deletions
+52 -30
View File
@@ -1,29 +1,35 @@
<div align="center">
<img src=assets/logos/logo.svg width="30%"/>
</div>
| **[Documentation](https://hao-ai-lab.github.io/FastVideo)** | **[Quick Start](https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start/)** | **[Weekly Dev Meeting](https://github.com/hao-ai-lab/FastVideo/discussions/982)** | 🟣💬 **[Slack](https://join.slack.com/t/fastvideo/shared_invite/zt-3f4lao1uq-u~Ipx6Lt4J27AlD2y~IdLQ)** |
<p align="center">
| <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start/"><b> Quick Start</b></a> | <a href="https://github.com/hao-ai-lab/FastVideo/discussions/982" target="_blank"><b>Weekly Dev Meeting</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-3f4lao1uq-u~Ipx6Lt4J27AlD2y~IdLQ" target="_blank"> <b>Slack</b> </a> | 🟣💬 <a href="https://ibb.co/sv3MMKyv" target="_blank"> <b> WeChat </b> </a> |
</p>
**FastVideo is a unified post-training and inference framework for accelerated video generation.**
## NEWS
- ```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/).
- `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/).
<details>
<summary>More</summary>
### More News
- ```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/).
- `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/).
</details>
## 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
- [Sparse distillation](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/) to achineve >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.
@@ -38,7 +44,6 @@ FastVideo has the following features:
- See this [page](https://hao-ai-lab.github.io/FastVideo/inference/hardware_support/) for full list of supported hardware and OS.
## Getting Started
We recommend using an environment manager such as `Conda` to create a clean environment:
```bash
@@ -53,20 +58,18 @@ pip install fastvideo
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) |
| 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.1-T2V-14B-Preview](https://huggingface.co/FastVideo/FastWan2.1-T2V-14B-Diffusers) | Coming soon! | [FastVideo Synthetic Wan2.1 720P](https://huggingface.co/datasets/FastVideo/Wan-Syn_77x768x1280_250k) |
| [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
@@ -105,36 +108,55 @@ python example.py
For a more detailed guide, please see our [inference quick start](https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start/).
## More Guides
### Other docs:
- [Design Overview](https://hao-ai-lab.github.io/FastVideo/design/overview/)
- [Contribution Guide](https://hao-ai-lab.github.io/FastVideo/getting_started/installation/)
## Distillation and Finetuning
- [Distillation Guide](https://hao-ai-lab.github.io/FastVideo/distillation/dmd/)
- [Contribution Guide](https://hao-ai-lab.github.io/FastVideo/contributing/overview/)
<!-- - [Finetuning Guide](https://hao-ai-lab.github.io/FastVideo/training/finetune.html) -->
## Awesome work using FastVideo or our research projects
- [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.
- [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. [![Star](https://img.shields.io/github/stars/sgl-project/sglang.svg?style=social&label=Star)](https://github.com/sgl-project/sglang)
- [DanceGRPO](https://github.com/XueZeyue/DanceGRPO): A unified framework to adapt Group Relative Policy Optimization (GRPO) to visual generation paradigms. Code based on FastVideo. [![Star](https://img.shields.io/github/stars/XueZeyue/DanceGRPO.svg?style=social&label=Star)](https://github.com/XueZeyue/DanceGRPO)
- [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. [![Star](https://img.shields.io/github/stars/Tencent-Hunyuan/SRPO.svg?style=social&label=Star)](https://github.com/Tencent-Hunyuan/SRPO)
- [DCM](https://github.com/Vchitect/DCM): Dual-expert consistency model for efficient and high-quality video generation. Code based on FastVideo. [![Star](https://img.shields.io/github/stars/Vchitect/DCM.svg?style=social&label=Star)](https://github.com/Vchitect/DCM)
- [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. [![Star](https://img.shields.io/github/stars/Tencent-Hunyuan/HunyuanVideo-1.5.svg?style=social&label=Star)](https://github.com/Tencent-Hunyuan/HunyuanVideo-1.5)
- [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. [![Star](https://img.shields.io/github/stars/kandinskylab/kandinsky-5.svg?style=social&label=Star)](https://github.com/kandinskylab/kandinsky-5)
- [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. [![Star](https://img.shields.io/github/stars/meituan-longcat/LongCat-Video.svg?style=social&label=Star)](https://github.com/meituan-longcat/LongCat-Video)
## 🤝 Contributing
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).
## Acknowledgement
We learned and reused code from the following projects:
- [Wan-Video](https://github.com/Wan-Video)
- [ThunderKittens](https://github.com/HazyResearch/ThunderKittens)
- [Triton](https://github.com/triton-lang/triton)
- [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 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.
We thank [MBZUAI](https://ifm.mbzuai.ac.ae/), [Anyscale](https://www.anyscale.com/), and [GMI Cloud](https://www.gmicloud.ai/) for their support throughout this project.
## Citation
If you find FastVideo useful, please consider citing our research work:
If you find FastVideo useful, please considering citing our work:
```bibtex
@software{fastvideo2024,
title = {FastVideo: A Unified Framework for Accelerated Video Generation},
author = {The FastVideo Team},
url = {https://github.com/hao-ai-lab/FastVideo},
month = apr,
year = {2024},
}
@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},
+1 -1
View File
@@ -43,7 +43,7 @@ RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
uv pip install --no-cache-dir --upgrade pip && \
uv pip install --no-cache-dir .[dev] && \
uv pip install --no-cache-dir https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.7.16/flash_attn-2.8.3+cu128torch2.10-cp310-cp310-linux_x86_64.whl
uv pip install --no-cache-dir https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.5.4/flash_attn-2.8.3%2Bcu128torch2.9-cp310-cp310-linux_x86_64.whl
COPY . .
+1 -1
View File
@@ -43,7 +43,7 @@ RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
uv pip install --no-cache-dir --upgrade pip && \
uv pip install --no-cache-dir .[dev] && \
uv pip install --no-cache-dir https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.7.16/flash_attn-2.8.3+cu128torch2.10-cp311-cp311-linux_x86_64.whl
uv pip install --no-cache-dir https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.5.4/flash_attn-2.8.3%2Bcu128torch2.9-cp311-cp311-linux_x86_64.whl
COPY . .
+1 -1
View File
@@ -43,7 +43,7 @@ RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
uv pip install --no-cache-dir --upgrade pip && \
uv pip install --no-cache-dir .[dev] && \
uv pip install --no-cache-dir https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.7.16/flash_attn-2.8.3+cu128torch2.10-cp312-cp312-linux_x86_64.whl
uv pip install --no-cache-dir https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.5.4/flash_attn-2.8.3%2Bcu128torch2.9-cp312-cp312-linux_x86_64.whl
COPY . .
-9
View File
@@ -121,15 +121,6 @@ This page contains the complete API reference for the FastVideo library.
show_root_toc_entry: true
heading_level: 4
## fastvideo.registry
::: fastvideo.registry
options:
show_source: true
show_root_heading: true
show_root_toc_entry: true
heading_level: 3
## fastvideo.pipelines
::: fastvideo.pipelines
Binary file not shown.

Before

Width:  |  Height:  |  Size: 211 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 117 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 461 KiB

After

Width:  |  Height:  |  Size: 18 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 210 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 437 KiB

After

Width:  |  Height:  |  Size: 27 KiB

-2
View File
@@ -6,8 +6,6 @@ FastVideo provides highly optimized custom attention kernels to accelerate video
* **[Video Sparse Attention (VSA)](vsa/index.md)**: Sparse attention mechanism selecting top-k blocks.
* **[Sliding Tile Attention (STA)](sta/index.md)**: Optimized attention for window-based video generation.
* **Backend development guide**: See the developer guide at
[Attention Backend Development](../contributing/attention_backend.md).
## General Build Instructions
-176
View File
@@ -1,176 +0,0 @@
# Attention Backend Development
This guide is for contributors adding a new attention backend (kernel or
implementation) to FastVideo. If you just want to use existing kernels or build
fastvideo-kernel, see [Attention overview](../attention/index.md).
## When you need this guide
Use this guide when you are:
- Adding a new attention kernel or algorithm.
- Wiring an existing kernel into FastVideo's attention selection.
- Extending attention support to a new platform.
## 0) Choose a backend name and scope
Pick a backend name in `UPPER_SNAKE_CASE` and decide where it should run.
Example: `MY_NEW_ATTN`.
You will use this name in:
- `AttentionBackendEnum` (global list of backends).
- `get_name()` in your backend class (match the enum name).
- Platform selectors (CUDA/ROCm/MPS/NPU) to return your backend class.
## 1) Add enum + platform selection
1) Add your backend to `fastvideo/platforms/interface.py`:
```python
class AttentionBackendEnum(enum.Enum):
...
MY_NEW_ATTN = enum.auto()
```
1) Register it in platform selection (example: CUDA). Update
`fastvideo/platforms/cuda.py` inside `get_attn_backend_cls`:
```python
elif selected_backend == AttentionBackendEnum.MY_NEW_ATTN:
try:
from fastvideo.attention.backends.my_new_attn import MyNewAttnBackend
return "fastvideo.attention.backends.my_new_attn.MyNewAttnBackend"
except ImportError as e:
logger.error("Failed to import MY_NEW_ATTN backend: %s", str(e))
raise
```
If you want support on other platforms, add a similar branch in
`fastvideo/platforms/rocm.py`, `fastvideo/platforms/mps.py`, or `fastvideo/platforms/npu.py`.
## 2) Implement the backend
Create `fastvideo/attention/backends/my_new_attn.py` and implement the required
classes.
Minimal skeleton (no custom metadata):
```python
import torch
from dataclasses import dataclass
from fastvideo.attention.backends.abstract import (
AttentionBackend,
AttentionImpl,
AttentionMetadata,
AttentionMetadataBuilder,
)
class MyNewAttnBackend(AttentionBackend):
@staticmethod
def get_name() -> str:
return "MY_NEW_ATTN"
@staticmethod
def get_impl_cls() -> type["MyNewAttnImpl"]:
return MyNewAttnImpl
@staticmethod
def get_metadata_cls() -> type["AttentionMetadata"]:
return AttentionMetadata
@staticmethod
def get_builder_cls() -> type["AttentionMetadataBuilder"]:
return MyNewAttnMetadataBuilder
@dataclass
class MyNewAttnMetadata(AttentionMetadata):
current_timestep: int
class MyNewAttnMetadataBuilder(AttentionMetadataBuilder):
def __init__(self) -> None:
pass
def prepare(self) -> None:
pass
def build(self, current_timestep: int, **kwargs):
return MyNewAttnMetadata(current_timestep=current_timestep)
class MyNewAttnImpl(AttentionImpl):
def __init__(
self,
num_heads: int,
head_size: int,
softmax_scale: float,
causal: bool = False,
num_kv_heads: int | None = None,
prefix: str = "",
**extra_impl_args,
) -> None:
self.softmax_scale = softmax_scale
self.causal = causal
def forward(
self,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
attn_metadata: MyNewAttnMetadata,
) -> torch.Tensor:
# Implement attention
return torch.nn.functional.scaled_dot_product_attention(
query.transpose(1, 2),
key.transpose(1, 2),
value.transpose(1, 2),
is_causal=self.causal,
scale=self.softmax_scale,
).transpose(1, 2)
```
Optional:
- Implement `preprocess_qkv` / `postprocess_output` if your kernel needs tiling
or reshaping.
- Use `fastvideo.forward_context.get_forward_context()` if you need dynamic
per-step data (e.g., window sizes).
- Set `accept_output_buffer = True` if your backend writes into a provided
output buffer.
## 3) Wire into attention layers
Backends are used by `LocalAttention` and `DistributedAttention`. These layers
accept a `supported_attention_backends` tuple. If your backend should be
eligible, update the call sites that construct these layers (search for
`supported_attention_backends=`).
## 4) Add compiled kernels (optional)
If you have a custom CUDA kernel:
1) Add sources in `fastvideo-kernel/csrc/attention/`.
2) Register bindings in `fastvideo-kernel/csrc/common_extension.cpp`.
3) Add to `fastvideo-kernel/CMakeLists.txt` (and any feature flags).
4) Expose in `fastvideo-kernel/python/fastvideo_kernel/ops.py`.
5) Export in `fastvideo-kernel/python/fastvideo_kernel/__init__.py`.
Keep a Python/Triton fallback so the backend runs even when the kernel is not
available.
## 5) Testing and debugging
- Add a small parity test or microbenchmark comparing to SDPA.
- Force your backend with the env var:
`FASTVIDEO_ATTENTION_BACKEND=MY_NEW_ATTN`.
- Check logs from `fastvideo/attention/selector.py` to confirm selection.
## Checklist
- [ ] Added enum entry in `fastvideo/platforms/interface.py`.
- [ ] Implemented backend in `fastvideo/attention/backends/`.
- [ ] Registered selection in platform(s).
- [ ] Updated layer call sites to include the backend where appropriate.
- [ ] Added tests and documentation.
-480
View File
@@ -1,480 +0,0 @@
# FastVideo + Coding Agents
Coding agents are now strong at navigating large codebases and iterating fast
with parity tests and examples. This guide shows how to use them to add new
model pipelines and ship PRs in a production-grade video diffusion framework.
FastVideo is a great project to contribute to, with production-grade
infrastructure, active collaborations (including NVIDIA), and a pipeline design
and inference architecture that has been forked by [SGLang’s
multimodal generation stack](https://github.com/sgl-project/sglang/tree/main/python/sglang/multimodal_gen).
Goal: run the new pipeline with a minimal script like
`examples/inference/basic/basic.py`. In production, FastVideo can download
models automatically via `HF_HOME`; for development, use local directories so
agents can run scripts and tests deterministically. We standardize local paths
as:
- `official_weights/<model_name>/` for official checkpoints
- `converted_weights/<model_name>/` if conversion is required
## Tips when prompting the agent
When prompting the agent, include:
- This guide and the [FastVideo design overview](../design/overview.md).
- Exact file paths to edit.
- A closest reference example file in FastVideo.
- Expected behavior and acceptance criteria.
- Repro steps (command, inputs, logs).
- Constraints (performance, memory, compatibility).
- Local paths (e.g., `official_weights/<model_name>/` or
`converted_weights/<model_name>/`) for parity tests.
## FastVideo structure at a glance
Before diving in, scan these references:
- [Contributing overview](overview.md) for environment/setup context.
- [FastVideo design overview](../design/overview.md) for pipeline architecture, configs, and HF layout.
FastVideo maps a Diffusers-style repo into a pipeline like:
- `fastvideo/models/*`: model implementations (DiT, VAE, encoders, upsamplers).
- `fastvideo/configs/models/*`: arch configs and `param_names_mapping` for
weight name translation.
- `fastvideo/configs/pipelines/*`: pipeline wiring (component classes + names).
- `fastvideo/configs/sample/*`: default runtime sampling parameters.
- `fastvideo/pipelines/basic/*`: end-to-end pipeline logic built from stages.
- `model_index.json`: the HF repo entrypoint that maps component names to
classes and weight files.
- Component loading happens in `VideoGenerator.from_pretrained`, which reads
`model_index.json`, resolves configs, and loads weights.
Minimal usage example (based on `examples/inference/basic/basic.py`):
```python
from fastvideo import VideoGenerator
from fastvideo.configs.sample import SamplingParam
model_id = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # or official_weights/<model_name>/
generator = VideoGenerator.from_pretrained(model_id, num_gpus=1)
sampling = SamplingParam.from_pretrained(model_id)
sampling.num_frames = 45
video = generator.generate_video(
"A vibrant city street at sunset.",
sampling_param=sampling,
output_path="video_samples",
save_video=True,
)
```
## Some questions to ask yourself before starting
Answering these upfront clarifies the work and speeds up implementation.
### Is the model already supported by SGLang's multimodal generation stack?
If yes, you can port many components from SGLang. It is a FastVideo fork, so
interfaces line up, but you still need to swap layers/modules to match
FastVideo's architecture and attention stack.
If not, implement the model directly in FastVideo.
### Is there an official implementation of the model you are adding?
If yes, use it as the numerical reference. For example, LTX‑2 has an official
implementation here: https://github.com/Lightricks/LTX-2. Prefer official code
even if Diffusers also has one.
### Is there a HuggingFace repo for the model you are adding? Is it in Diffusers format?
If yes, load it directly in FastVideo after setting tensor mapping rules in the
config. Otherwise, convert the weights to Diffusers format. See [Weights and
Diffusers format](../design/overview.md#weights-and-diffusers-format) for details.
### Do I have official weights + local paths ready?
Standardize local paths as:
- `official_weights/<model_name>/` for official checkpoints
- `converted_weights/<model_name>/` if conversion is required (can be created later)
### What pipeline components are required for the model you are adding?
Usually you need a transformer (DiT), VAE, text encoder, and tokenizer. Some
models add extra components.
### What tasks does the model support?
Usually a video diffusion model supports text‑to‑video (T2V),
image‑to‑video (I2V), and video‑to‑video (V2V). Some add extra tasks (two‑stage
generation, keyframe interpolation), which require extra components.
It's usually easiest to start with a T2V pipeline and add the other tasks later.
You can refer to the [Pipeline system](../design/overview.md#pipeline-system)
section for more details.
### Am I able to generate videos with the official implementation?
These videos and prompts are your reference. Once the FastVideo pipeline works,
compare outputs to the official implementation. Due to seeding and other
factors, outputs may not match exactly, but they should be comparable.
## Workflow: adding a full pipeline
This is an example workflow for adding a full model pipeline (model + configs +
examples + tests). This guide is in active development; feedback is welcome.
!!! note
If you get stuck, refer to existing models/pipelines in FastVideo or ask in Slack.
### 0) Fetch official model's code and weights
Purpose:
- Keep official checkpoints and source code local so conversion, parity tests,
and reference runs are reproducible.
- Clone the official repo so you can use it as a numerical reference.
Action:
- Download official weights into `official_weights/<model_name>/`
(Diffusers format or not).
- Clone the official repo under the project root (e.g., `FastVideo/LTX-2/`).
- If a Diffusers-format HF repo already exists, you can skip manual weight
handling and download it directly with
`scripts/huggingface/download_hf.py`.
!!! note
This step is best done manually because large downloads can time out.
Example:
```bash
python scripts/huggingface/download_hf.py \
--repo_id Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
--local_dir official_weights/Wan2.1-T2V-1.3B-Diffusers \
--repo_type model
```
### 1) Implement the model + config mapping
Purpose:
- Model weights are a dictionary of named tensors (`state_dict`). If the names
don’t line up with FastVideo’s module names, weights won’t load correctly (or
will silently load into the wrong layer).
- Official checkpoints often use different prefixes or module layouts than
FastVideo, so we translate names via the mapping (during load or conversion).
- Mapping aligns three things:
1. the official implementation’s module names,
2. the checkpoint `state_dict` keys,
3. FastVideo’s model classes and layer naming conventions.
- If names don’t align, weights won’t load; implement the FastVideo model and
define mapping rules first.
Action:
- Implement the FastVideo model + config mapping.
- Add/extend the model in `fastvideo/models/...` and config in
`fastvideo/configs/models/...` (including `param_names_mapping`).
- Reuse existing FastVideo layers/modules where possible.
- Use FastVideo’s attention layers:
- `DistributedAttention` only for full‑sequence self‑attention in the DiT.
- `LocalAttention` for cross‑attention and other attention layers.
- See the “Configuration System” and “Weights and Diffusers format” sections
in `docs/design/overview.md` for how these pieces connect.
- If you are using an agent, ask it to implement the model, config mapping,
and a parity test together so you can validate numerics immediately.
!!! note
After the first component is aligned and parity‑tested, open a **DRAFT PR**
on FastVideo so the rest of the pipeline work can build on top of it.
!!! note
If a Diffusers-format HF repo already exists and loads correctly, you can
skip conversion entirely (no conversion script needed) and just download it
with `scripts/huggingface/download_hf.py`. Otherwise, you may need a
conversion script + a `converted_weights/<model>/` staging directory.
Example (key renaming via arch config mapping, Wan2.1‑style):
```python
# Official model (simplified) in the upstream repo.
class OfficialWanTransformer(torch.nn.Module):
def __init__(self):
super().__init__()
self.patch_embedding = torch.nn.Conv3d(16, 1536, kernel_size=2, padding=0)
def forward(self, x):
return self.patch_embedding(x)
# FastVideo model (simplified) in fastvideo/models/dits/wanvideo.py
class PatchEmbed(torch.nn.Module):
def __init__(self):
super().__init__()
self.proj = torch.nn.Conv3d(16, 1536, kernel_size=2, padding=0)
def forward(self, x):
return self.proj(x)
class WanTransformer3DModel(torch.nn.Module):
def __init__(self):
super().__init__()
self.patch_embedding = PatchEmbed()
def forward(self, x):
return self.patch_embedding(x)
# Mapping defined in a config (simplified; see the real mapping in
# fastvideo/configs/models/dits/wanvideo.py)
param_names_mapping = {
r"^patch_embedding\.(.*)$": r"patch_embedding.proj.\1",
r"^blocks\.(\d+)\.attn1\.to_q\.(.*)$": r"blocks.\1.to_q.\2",
}
def apply_regex_map(state_dict, mapping):
# Pseudocode: apply regex substitutions in order
...
# Official checkpoint keys (example)
official = {
"patch_embedding.weight": ...,
"blocks.0.attn1.to_q.weight": ...,
}
# Apply mapping so keys match FastVideo modules
converted = apply_regex_map(official, param_names_mapping)
```
Optional helper (print a few checkpoint keys quickly):
```bash
python - <<'PY'
import safetensors.torch as st
keys = list(st.load_file("official_weights/<model>/transformer/diffusion_pytorch_model.safetensors").keys())
print(keys[:20])
PY
```
Example agent prompt (task request):
```
Please add the Wan2.1 T2V 1.3B Diffusers pipeline to FastVideo:
- Add a FastVideo native Wan2.1 DiT implementation + config mapping.
- Make sure to use the existing FastVideo layers and attention modules where possible.
- Add a parity test that loads the official model alongside the FastVideo model and compares outputs numerically with fixed seeds and inputs.
Paths:
- Official repo: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
- Local download: official_weights/Wan2.1-T2V-1.3B-Diffusers
Mapping steps:
- Load the official DiT weights from
official_weights/Wan2.1-T2V-1.3B-Diffusers/transformer/diffusion_pytorch_model.safetensors.
- Instantiate the FastVideo DiT (`WanTransformer3DModel`) and compare
its `state_dict().keys()` to the official keys.
- Update `param_names_mapping` in
fastvideo/configs/models/dits/wanvideo.py to resolve missing/unexpected keys.
- Use `load_state_dict(strict=False)` during iteration to surface mismatches.
```
External examples of the same pattern:
- SGLang uses prefix-based routing in its weight loader to map checkpoint keys
into internal submodules (e.g., stripping a top-level prefix before delegating).
- vLLM includes model-specific renamers for certain checkpoints that adjust
key prefixes so weights match its internal naming.
### 2) Test numerical alignment with the official implementation
Purpose:
- Verify that the FastVideo component is numerically aligned with the official
implementation.
Action:
- Add or reuse a numerical parity test that loads the official model and the
FastVideo model and compares outputs.
- See examples in `tests/local_tests/` (e.g., `tests/local_tests/upsamplers/`)
and the commands in `tests/local_tests/README.md`.
- If there are discrepancies, add opt‑in logging to both models and compare
activation summaries (layer output sums, per‑stage logs).
- First align the loaded weights (validate `param_names_mapping`).
- Then align forward outputs using fixed seeds and inputs.
- Start with `atol=1e-4, rtol=1e-4` in `assert_close`.
- Keep dtype consistent (bf16 if available; otherwise fp32).
- If attention parity is unstable, align backends (e.g.,
`FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA`).
### 3) Repeat the process for each component
If the model requires additional components, repeat Steps 1–2 for each one.
For example, implement the VAE in `fastvideo/models/vaes/` and its config in
`fastvideo/configs/models/vaes/`, then add parity coverage for it.
### 4) Add a pipeline config + sample defaults
Purpose:
- `fastvideo/configs/pipelines/` describes pipeline wiring and model module
names.
- `fastvideo/configs/sample/` defines default runtime parameters.
Action:
- Add a new pipeline config + sampling params.
- Register them in `fastvideo/registry.py` using explicit
`register_configs(...)` blocks (this file is the single source of truth now).
### 5) Wire pipeline stages
Purpose:
- `fastvideo/pipelines/basic/<pipeline>/` contains the actual pipeline logic.
- `fastvideo/pipelines/stages/` holds reusable, testable stages.
Action:
- Build the pipeline using stages; keep new stages isolated and documented.
- Prefer opt‑in flags for expensive or optional steps.
### 6) Add pipeline‑level tests
Purpose:
- Ensure the end‑to‑end pipeline works and stays aligned as pieces evolve.
Action:
- Add a pipeline parity test under `tests/local_tests/pipelines/`.
- See the [Testing Guide](testing.md) for test conventions.
### 7) Add user‑facing examples
Purpose:
- `examples/inference/basic/` is the entry point for simple, runnable scripts.
Action:
- Provide a minimal “hello world” example plus advanced variations.
- Use fixed seeds and stable prompts.
- Run the example locally to confirm end‑to‑end behavior.
### 8) Add SSIM tests for CI checks
Purpose:
- Ensure visual similarity stays within expected bounds for regression testing.
- SSIM tests act as a higher‑level guardrail beyond unit/parity tests.
Action:
- Add SSIM tests under `fastvideo/tests/ssim/` and include reference videos
(see the structure in the Testing Guide).
- Use stable prompts/seeds and document any GPU‑specific requirements.
- Follow the [Testing Guide](testing.md) for reference video placement and
execution details.
### 9) Document it
Purpose:
- `docs/` is where users find the new pipeline usage and limitations.
Action:
- Add a short doc page or update an existing one.
- Mention any caveats (memory, speed, constraints).
## Common pitfalls when porting models
- **Attention backend mismatch**: parity can fail if the official model uses a
different attention backend (e.g., SDPA vs custom). Align backends before
debugging deeper issues.
- **Patchifier shape mistakes**: wrong patchification or reshape lengths can
silently corrupt outputs. Validate patch shapes early.
- **Mask handling**: attention masks must match the official behavior (padding,
causal masks, and broadcast shapes).
- **Scheduler / sigma schedule mismatch**: even small differences in schedules
or timestep shapes can cause noticeable drift.
## Diffusers vs manual conversion
If a model already ships in Diffusers format (with a proper `model_index.json`),
prefer downloading it directly and loading it via FastVideo. In that case:
- You usually **do not need** a conversion script.
- You still need a correct `param_names_mapping` if the internal module names
differ from FastVideo’s implementation.
If the model does **not** have a Diffusers-format repo:
- You will need a conversion script to rewrite `state_dict` keys into FastVideo
naming and stage the result (e.g., under `converted_weights/<model>/`).
- You may still use the official repo for reference parity and debugging.
In both cases, parity testing is required to validate correctness.
If you want to publish a Diffusers‑style repo after conversion, use
`scripts/checkpoint_conversion/create_hf_repo.py` to assemble a HuggingFace‑ready
directory before uploading.
## FAQ
**Q: Why do we implement the FastVideo model before conversion?**
A: You can’t define the key‑mapping rules until the FastVideo module names are
known. The implementation determines the target `state_dict` schema.
**Q: Do we always need a conversion script?**
A: No. If a Diffusers‑format repo exists and loads correctly, download it and
skip conversion.
**Q: How do I figure out `param_names_mapping` quickly?**
A: Load the official weights, instantiate the FastVideo model, and diff
`state_dict().keys()` on both sides. Add regex rules until missing/unexpected
keys are resolved. Agents can help you with this.
**Q: What if parity fails even after mapping?**
A: Align attention backends, sigma schedules, and timestep shapes first. Then
add opt‑in activation logging to locate the first divergent layer.
## Case study: LTX‑2 port (from PLAN.md)
The LTX‑2 port in `PLAN.md` shows the real sequence of steps and backtracking
that happened during integration. Use it as a reference for how parity work
actually unfolds:
- Ported components first (transformer, VAE, audio, text encoder).
- Added parity tests per component; used SDPA for reference parity.
- Added debug logging to compare per‑block activations and isolate divergence.
- Fixed cross‑attention reshape and patch grid bounds issues after logging.
- Aligned sigma schedule and masking behavior to match the official pipeline.
Recommendation:
- Keep raw step‑by‑step logs in your own local `PLAN.md` for large ports.
## Worked example: Wan2.1 T2V 1.3B pipeline
The Wan2.1 T2V 1.3B Diffusers pipeline is a good “standard” example for
FastVideo integration.
1. Verify model config + mapping.
- DiT mapping: `fastvideo/configs/models/dits/wanvideo.py`
- VAE: `fastvideo/models/vaes/wanvae.py`
- Text encoder: `fastvideo/models/encoders/t5.py`
2. Parity test the core components.
- Example tests: `fastvideo/tests/transformers/test_wanvideo.py`,
`fastvideo/tests/vaes/test_wan_vae.py`,
`fastvideo/tests/encoders/test_t5_encoder.py`
3. Pipeline wiring.
- Pipeline: `fastvideo/pipelines/basic/wan/wan_pipeline.py`
- Pipeline config: `fastvideo/configs/pipelines/wan.py`
- Sampling defaults: `fastvideo/configs/sample/wan.py`
4. Minimal example.
- Script: `examples/inference/basic/basic.py`
+31 -43
View File
@@ -1,55 +1,16 @@
# 📦 Developing FastVideo on RunPod
You can easily use the FastVideo Pod Template on [RunPod](https://www.runpod.io) for development or experimentation.
You can easily use the FastVideo Docker image as a custom container on [RunPod](https://www.runpod.io) for development or experimentation.
## Creating a new pod
- Make sure you are using the correct RunPod account.
![RunPod Account Selection](../../assets/images/runpod_account.png)
Choose a GPU that supports CUDA 12.8
Pick 1 or 2 L40S GPU(s)
- Use "Additional Filters" to select CUDA 12.8.
![RunPod CUDA selection](../../assets/images/runpod_cuda.png)
- Click "Deploy" and Pick a single A40 or RTX 4090 GPU.
![RunPod GPU Selection](../../assets/images/runpod_deploy.png)
- Select the "FastVideo" or "fastvideo-dev" Pod Template.
![RunPod Pod Template Selection](../../assets/images/runpod_create.png)
- Set the Pod name to "`<name>-<FastVideo>-<date>`".
- Finally, once the pod is deployed (will take a few minutes as the image is being pulled), you can SSH into it using "SSH exposed over TCP". You'll need to use the matching private ssh key you provided.
![RunPod SSH](../../assets/images/runpod_ssh.png)
## Working with the pod
After SSH'ing into your pod, you'll find the correct `uv` environment already activated and you should be in /FastVideo/ directory. Make sure to use /FastVideo/ for all your work.
To pull in the latest changes from the GitHub repo:
```bash
cd /FastVideo
git pull
```
Run your development workflows as usual:
```bash
# Run linters
pre-commit run --all-files
# Run tests
pytest tests/
```
Make sure to push your changes back to the GitHub repo as nothing will be saved to the pod when it is terminated.
After you are done with your work, you can terminate the pod by clicking the "Terminate" and "Delete" buttons. Remember if the pod is not completely deleted, Runpod will keep charging you for it.
## Extra Information:
If you need to customize the pod template this section has some useful information. For the most part you can leave the defaults of the FastVideo Pod Template.
When creating your pod template, use this image:
```
@@ -63,3 +24,30 @@ bash -c "apt update;DEBIAN_FRONTEND=noninteractive apt-get install openssh-serve
```
![RunPod template configuration](../../assets/images/runpod_template.png)
After deploying, the pod will take a few minutes to pull the image and start the SSH service.
![RunPod ssh](../../assets/images/runpod_ssh.png)
## Working with the pod
After SSH'ing into your pod, you'll find the `fastvideo-dev` Conda environment already activated.
To pull in the latest changes from the GitHub repo:
```bash
cd /FastVideo
git pull
```
`If you have a persistent volume and want to keep your code changes, you can move /FastVideo to /workspace/FastVideo, or simply clone the repository there.`
Run your development workflows as usual:
```bash
# Run linters
pre-commit run --all-files
# Run tests
pytest tests/
```
+30 -38
View File
@@ -1,84 +1,76 @@
# 🛠️ Contributing to FastVideo
Thank you for your interest in contributing to FastVideo. We want the process
to be smooth and beginner‑friendly, whether you are adding a new pipeline,
improving performance, or fixing a bug.
Thank you for your interest in contributing to FastVideo. We want to make the process as smooth for you as possible and this is a guide to help get you started!
## Quick prerequisites
Our community is open to everyone and welcomes any contributions no matter how large or small.
- **OS**: Linux is the primary development target (WSL can work).
- **GPU**: NVIDIA GPU recommended for inference and training workflows.
- **CUDA**: Use a recent CUDA 12.x toolchain (see the installation guide for
the current recommendation).
# Developer Environment:
Do make sure you have CUDA 12.4 installed and supported. FastVideo currently only supports Linux and CUDA GPUs, but we hope to support other platforms in the future.
For a full install checklist, see `docs/getting_started/installation/gpu.md`.
## Local development (Conda + editable install)
We recommend using a fresh Python 3.10 Conda environment to develop FastVideo:
Install Miniconda:
```bash
```
wget https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh
bash Miniconda3-latest-Linux-x86_64.sh
source ~/.bashrc
```
Create and activate a Conda environment:
Create and activate a Conda environment for FastVideo:
```bash
```
conda create -n fastvideo python=3.12 -y
conda activate fastvideo
```
Install `uv` (optional, but recommended):
```bash
From instructions on [uv](https://astral.sh/uv/):
```
curl -LsSf https://astral.sh/uv/install.sh | sh
# or
# or
wget -qO- https://astral.sh/uv/install.sh | sh
```
Clone the repo:
Clone the FastVideo repository and go to the FastVideo directory:
```bash
```
git clone https://github.com/hao-ai-lab/FastVideo.git && cd FastVideo
```
Install FastVideo in editable mode and set up hooks:
Now you can install FastVideo and setup git hooks for running linting. By using `pre-commit`, the linters will run and have to pass before you'll be able to make a commit.
```bash
uv pip install -e .[dev]
# Optional: FlashAttention (builds native kernels)
uv pip install flash-attn --no-build-isolation
# Can also install flash-attn (optional)
uv pip install flash-attn --no-build-isolation
# Linting, formatting, static typing
# Linting, formatting and static type checking
pre-commit install --hook-type pre-commit --hook-type commit-msg
# You can manually run pre-commit with
pre-commit run --all-files
# Unit tests
pytest tests/
```
If you are on a Hopper GPU, installing FlashAttention 3 can improve
performance (see `docs/inference/optimizations.md`).
If you are on a Hopper GPU, you should also install [FA3](https://github.com/Dao-AILab/flash-attention) for much better performance:
## Docker development (optional)
```
git clone https://github.com/Dao-AILab/flash-attention.git && cd flash-attention/hopper
If you prefer a containerized environment, use the dev image documented in
`docs/contributing/developer_env/docker.md`.
# make sure you have ninja installed
uv pip install ninja
python setup.py install
```
## Testing
See the [Testing Guide](testing.md) for how to add and run tests in FastVideo.
## Attention backend development
If you are adding a new attention kernel or backend, follow
[Attention Backend Development](attention_backend.md).
## Contributing with coding agents
For a step‑by‑step workflow on adding pipelines or components with coding
agents, see `docs/contributing/coding_agents.md`.
Please refer to the [Testing Guide](testing.md) for more information on how to add and run tests in FastVideo.
+387 -142
View File
@@ -1,179 +1,424 @@
# FastVideo Architecture Overview
# 🔍 FastVideo Overview
This document summarizes how FastVideo is structured and how a Diffusers-style
model repo maps into a runnable pipeline. It is intended for contributors who
need the high-level layout and key entrypoints, not every internal detail.
This document outlines FastVideo's architecture for developers interested in framework internals or contributions. It serves as an onboarding guide for new contributors by providing an overview of the most important directories and files within the `fastvideo/` codebase.
## FastVideo structure at a glance
## Table of Contents - Directory Structure and Files
FastVideo maps a Diffusers-style repo into a pipeline like this:
- [`fastvideo/pipelines/`](#pipeline-system) - Core diffusion pipeline components
- [`fastvideo/models/`](#model-components) - Model implementations
- [`dits/`](#transformer-models) - Transformer-based diffusion models
- [`vaes/`](#vae-variational-auto-encoder) - Variational autoencoders
- [`encoders/`](#text-and-image-encoders) - Text and image encoders
- [`schedulers/`](#schedulers) - Diffusion schedulers
- [`fastvideo/attention/`](#optimized-attention) - Optimized attention implementations
- [`fastvideo/distributed/`](#distributed-processing) - Distributed computing utilities
- [`fastvideo/layers/`](#tensor-parallelism) - Custom neural network layers
- [`fastvideo/platforms/`](#platforms) - Hardware platform abstractions
- [`fastvideo/worker/`](#executor-and-worker-system) - Multi-GPU process management
- [`fastvideo/fastvideo_args.py`](#fastvideoargs) - Argument handling
- [`fastvideo/forward_context.py`](#forward-context-management) - Forward pass context management
- `fastvideo/utils.py` - Utility functions
- [`fastvideo/logger.py`](#logger) - Logging infrastructure
- `fastvideo/models/*`: model implementations (DiT, VAE, encoders, upsamplers).
- `fastvideo/configs/models/*`: arch configs and `param_names_mapping` for
weight name translation.
- `fastvideo/configs/pipelines/*`: pipeline wiring (component classes + names).
- `fastvideo/configs/sample/*`: default runtime sampling parameters.
- `fastvideo/pipelines/basic/*`: end-to-end pipelines.
- `fastvideo/pipelines/stages/*`: reusable pipeline stages.
- `fastvideo/models/loader/*`: component loaders for Diffusers-style repos.
- `model_index.json`: HF repo entrypoint mapping component names to classes.
## Core Architecture
Flow:
`model_index.json` -> component loaders -> model modules -> pipeline stages ->
sampling params.
FastVideo separates model components from execution logic with these principles:
Minimal usage (from `examples/inference/basic/basic.py`):
- **Component Isolation**: Models (encoders, VAEs, transformers) are isolated from execution (pipelines, stages, distributed processing)
- **Modular Design**: Components can be independently replaced
- **Distributed Execution**: Supports various parallelism strategies (Tensor, Sequence)
- **Custom Attention Backends**: Components can support and use different Attention implementations
- **Pipeline Abstraction**: Consistent interface across diffusion models
## FastVideoArgs
The `FastVideoArgs` class in `fastvideo/fastvideo_args.py` serves as the central configuration system for FastVideo. It contains all parameters needed to control model loading, inference configuration, performance optimization settings, and more.
Key features include:
- **Command-line Interface**: Automatic conversion between CLI arguments and dataclass fields
- **Configuration Groups**: Organized by functional areas (model loading, video params, optimization settings)
- **Context Management**: Global access to current settings via `get_current_fastvideo_args()`
- **Parameter Validation**: Ensures valid combinations of settings
Common configuration areas:
- **Model paths and loading options**: `model_path`, `trust_remote_code`, `revision`
- **Distributed execution settings**: `num_gpus`, `tp_size`, `sp_size`
- **Video generation parameters**: `height`, `width`, `num_frames`, `num_inference_steps`
- **Precision settings**: Control computation precision for different components
Example usage:
```python
from fastvideo import VideoGenerator
from fastvideo.configs.sample import SamplingParam
# Load arguments from command line
fastvideo_args = prepare_fastvideo_args(sys.argv[1:])
model_id = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # or official_weights/<model_name>/
generator = VideoGenerator.from_pretrained(model_id, num_gpus=1)
# Access parameters
model = load_model(fastvideo_args.model_path)
sampling = SamplingParam.from_pretrained(model_id)
sampling.num_frames = 45
video = generator.generate_video(
"A vibrant city street at sunset.",
sampling_param=sampling,
output_path="video_samples",
save_video=True,
# Set as global context
with set_current_fastvideo_args(fastvideo_args):
# Code that requires access to these arguments
result = generate_video()
```
## Pipeline System
### `ComposedPipelineBase`
This foundational class provides:
- **Model Loading**: Automatically loads components from HuggingFace-Diffusers-compatible model directories
- **Stage Management**: Creates and orchestrates processing stages
- **Data Flow Coordination**: Ensures proper state flow between stages
```python
class MyCustomPipeline(ComposedPipelineBase):
_required_config_modules = [
"text_encoder", "tokenizer", "vae", "transformer", "scheduler"
]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
# Pipeline-specific initialization
pass
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
self.add_stage("input_validation_stage", InputValidationStage())
self.add_stage("text_encoding_stage", CLIPTextEncodingStage(
text_encoder=self.get_module("text_encoder"),
tokenizer=self.get_module("tokenizer")
))
# Additional stages...
```
### Pipeline Stages
Each stage handles a specific diffusion process component:
- **Input Validation**: Parameter verification
- **Text Encoding**: CLIP, LLaMA, or T5-based encoding
- **Image Encoding**: Image input processing
- **Timestep & Latent Preparation**: Setup for diffusion
- **Denoising**: Core diffusion loop
- **Decoding**: Latent-to-pixel conversion
Each stage implements a standard interface:
```python
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
# Process batch and update state
return batch
```
![Pipeline execution and data flow](../assets/images/pipeline.png)
### ForwardBatch
Defined in `fastvideo/pipelines/pipeline_batch_info.py`, `ForwardBatch` encapsulates the data payload passed between pipeline stages. It typically holds:
- **Input Data**: Prompts, images, generation parameters
- **Intermediate State**: Embeddings, latents, timesteps, accumulated during stage execution
- **Output Storage**: Generated results and metadata
- **Configuration**: Sampling parameters, precision settings
This structure facilitates clear state transitions between stages.
## Model Components
The `fastvideo/models/` directory contains implementations of the core neural network models used in video diffusion:
### Transformer Models
Transformer networks perform the actual denoising during diffusion:
- **Location**: `fastvideo/models/dits/`
- **Examples**:
- `WanTransformer3DModel`
- `HunyuanVideoTransformer3DModel`
Features include:
- Text/image conditioning
- Standardized interface for model-specific optimizations
```python
def forward(
self,
latents, # [B, T, C, H, W]
encoder_hidden_states, # Text embeddings
timestep, # Current diffusion timestep
encoder_hidden_states_image=None, # Optional image embeddings
**kwargs
):
# Perform denoising computation
return noise_pred # Predicted noise residual
```
### VAE (Variational Auto-Encoder)
VAEs handle conversion between pixel space and latent space:
- **Location**: `fastvideo/models/vaes/`
- **Examples**:
- `AutoencoderKLWan`
- `AutoencoderKLHunyuanVideo`
These models compress image/video data to a more efficient latent representation (typically 4x-8x smaller in each dimension).
FastVideo's VAE implementations include:
- Efficient video batch processing
- Memory optimization
- Optional tiling for large frames
- Distributed weight support
### Text and Image Encoders
Encoders process conditioning inputs into embeddings:
- **Location**: `fastvideo/models/encoders/`
- **Text Encoders**:
- `CLIPTextModel`
- `LlamaModel`
- `UMT5EncoderModel`
- **Image Encoders**:
- `CLIPVisionModel`
FastVideo implements optimizations such as:
- Vocab parallelism for distributed processing
- Caching for common prompts
- Precision-tuned computation
### Schedulers
Schedulers manage the diffusion sampling process:
- **Location**: `fastvideo/models/schedulers/`
- **Examples**:
- `UniPCMultistepScheduler`
- `FlowMatchEulerDiscreteScheduler`
These components control:
- Diffusion timestep sequences
- Noise prediction to latent update conversions
- Quality/speed trade-offs
```python
def step(
self,
model_output: torch.Tensor,
timestep: torch.LongTensor,
sample: torch.Tensor,
**kwargs
) -> torch.Tensor:
# Process model output and update latents
# Return updated latents
return prev_sample
```
This diagram shows how models are discovered, validated, and loaded across entrypoints, executors, pipelines, and model loaders.
![Model loading flow](../assets/images/load_models.png)
## Optimized Attention
The `fastvideo/attention/` directory contains optimized attention implementations crucial for efficient video diffusion:
### Attention Backends
Multiple implementations with automatic selection:
- **FLASH_ATTN**: Optimized for supporting hardware
- **TORCH_SDPA**: Built-in PyTorch scaled dot-product attention
- **SLIDING_TILE_ATTN**: For very long sequences
```python
# Configure available attention backends for this layer
self.attn = LocalAttention(
num_heads=num_heads,
head_size=head_dim,
causal=False,
supported_attention_backends=(_Backend.FLASH_ATTN, _Backend.TORCH_SDPA)
)
# Override via environment variable
# export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
```
![Attention backend selector design](../assets/images/attention_backend.png)
### Attention Patterns
Supports various patterns with memory optimization techniques:
- **Cross/Self/Temporal/Global-Local Attention**
- Chunking, progressive computation, optimized masking
## Distributed Processing
The `fastvideo/distributed/` directory contains implementations for distributed model execution:
### Tensor Parallelism
Tensor parallelism splits model weights across devices:
- **Implementation**: Through `RowParallelLinear` and `ColumnParallelLinear` layers
- **Use cases**: Will be used by encoder models as their sequence lengths are shorter and enables efficient sharding.
```python
# Tensor-parallel layers in a transformer block
from fastvideo.layers.linear import ColumnParallelLinear, RowParallelLinear
# Split along output dimension
self.qkv_proj = ColumnParallelLinear(
input_size=hidden_size,
output_size=3 * hidden_size,
bias=True,
gather_output=False
)
# Split along input dimension
self.out_proj = RowParallelLinear(
input_size=hidden_size,
output_size=hidden_size,
bias=True,
input_is_parallel=True
)
```
## Configuration system
### Sequence Parallelism
FastVideo uses typed configs to keep model definitions, pipeline wiring, and
runtime parameters consistent:
Sequence parallelism splits sequences across devices:
- `fastvideo/configs/models/`: architecture definitions, layer shapes, and
`param_names_mapping` rules for key renaming.
- `fastvideo/configs/pipelines/`: pipeline wiring and required components.
- `fastvideo/configs/sample/`: default sampling parameters (steps, frames,
guidance scale, resolution, fps).
- `fastvideo/registry.py`: unified registry for pipeline config + sampling
defaults and model metadata resolution, defined via explicit
`register_configs(...)` blocks (no separate dict registries).
- **Implementation**: Through `DistributedAttention` and sequence splitting
- **Use cases**: Long video sequences or high-resolution processing. Used by DiT models.
`FastVideoArgs` (in `fastvideo/fastvideo_args.py`) provides runtime settings and
is passed into pipeline construction and stages.
```python
# Distributed attention for long sequences
from fastvideo.attention import DistributedAttention
## Weights and Diffusers format
FastVideo follows the HuggingFace Diffusers repo layout. This keeps loaders
compatible with HF repos and makes it easy to add new components.
Typical Diffusers repo:
```
<model-repo>/
model_index.json
scheduler/
scheduler_config.json
transformer/ # or unet/ for image models
config.json
diffusion_pytorch_model.safetensors
vae/
config.json
diffusion_pytorch_model.safetensors
text_encoder/
config.json
model.safetensors
tokenizer/
tokenizer_config.json
tokenizer.json
self.attn = DistributedAttention(
num_heads=num_heads,
head_size=head_dim,
causal=False,
supported_attention_backends=(_Backend.SLIDING_TILE_ATTN, _Backend.FLASH_ATTN)
)
```
Key points:
### Communication Primitives
- `model_index.json` is the root map that tells FastVideo which components to
load and which classes implement them.
- Each component lives in its own folder with a `config.json` and weights.
- Weights are usually in `diffusion_pytorch_model.safetensors`.
Efficient distributed operations via AllGather, AllReduce, and synchronization mechanisms.
Note on tensor names:
Efficient communication primitives minimize distributed overhead:
Official checkpoints often use different `state_dict` names than FastVideo's
module layout. We translate tensor names via the DiT arch config mapping
(`param_names_mapping` under `fastvideo/configs/models/dits/`). This is similar
in spirit to name-translation layers used in systems like vLLM and SGLang.
- **Sequence-Parallel AllGather**: Collects sequence chunks
- **Tensor-Parallel AllReduce**: Combines partial results
- **Distributed Synchronization**: Coordinates execution
Example HF repo (Wan 2.1 T2V 1.3B Diffusers):
## Forward Context Management
```
https://huggingface.co/Wan-AI/Wan2.1-T2V-1.3B-Diffusers/tree/main
### ForwardContext
Defined in `fastvideo/forward_context.py`, `ForwardContext` manages execution-specific state *within* a forward pass, particularly for low-level optimizations. It is accessed via `get_forward_context()`.
- **Attention Metadata**: Configuration for optimized attention kernels (`attn_metadata`)
- **Profiling Data**: Potential hooks for performance metrics collection
This context-based approach enables:
- Dynamic optimization based on execution state (e.g., attention backend selection)
- Step-specific customizations within model components
Usage example:
```python
with set_forward_context(current_timestep, attn_metadata, fastvideo_args):
# During this forward pass, components can access context
# through get_forward_context()
output = model(inputs)
```
Example `model_index.json` from that repo:
## Executor and Worker System
```json
{
"_class_name": "WanPipeline",
"_diffusers_version": "0.33.0.dev0",
"scheduler": [
"diffusers",
"UniPCMultistepScheduler"
],
"text_encoder": [
"transformers",
"UMT5EncoderModel"
],
"tokenizer": [
"transformers",
"T5TokenizerFast"
],
"transformer": [
"diffusers",
"WanTransformer3DModel"
],
"vae": [
"diffusers",
"AutoencoderKLWan"
]
}
The `fastvideo/worker/` directory contains the distributed execution framework:
### Executor Abstraction
FastVideo implements a flexible execution model for distributed processing:
- **Executor Base Class**: An abstract base class defining the interface for all executors
- **MultiProcExecutor**: Primary implementation that spawns and manages worker processes
- **GPU Workers**: Handle actual model execution on individual GPUs
The MultiProcExecutor implementation:
1. Spawns worker processes for each GPU
2. Establishes communication channels via pipes
3. Coordinates distributed operations across workers
4. Handles graceful startup and shutdown of the process group
Each GPU worker:
1. Initializes the distributed environment
2. Builds the pipeline for the specified model
3. Executes requested operations on its assigned GPU
4. Manages local resources and communicates results back to the executor
This design allows FastVideo to efficiently utilize multiple GPUs while providing a simple, unified interface for model execution.
## Platforms
The `fastvideo/platforms/` directory provides hardware platform abstractions that enable FastVideo to run efficiently on different hardware configurations:
### Platform Abstraction
FastVideo's platform abstraction layer enables:
- **Hardware Detection**: Automatic detection of available hardware
- **Backend Selection**: Appropriate selection of compute kernels
- **Memory Management**: Efficient utilization of hardware-specific memory features
The primary components include:
- **Platform Interface**: Defines the common API for all platform implementations
- **CUDA Platform**: Optimized implementation for NVIDIA GPUs
- **Backend Enum**: Used throughout the codebase for feature selection
Usage example:
```python
from fastvideo.platforms import current_platform, _Backend
# Check hardware capabilities
if current_platform.supports_backend(_Backend.FLASH_ATTN):
# Use FlashAttention implementation
else:
# Fall back to standard implementation
```
How this maps to FastVideo:
The platform system is designed to be extensible for future hardware targets.
- `WanPipeline` -> `fastvideo/pipelines/basic/wan/wan_pipeline.py`
- `WanTransformer3DModel` -> `fastvideo/models/dits/wanvideo.py`
- `AutoencoderKLWan` -> `fastvideo/models/vaes/wanvae.py`
- `UMT5EncoderModel` -> `fastvideo/models/encoders/t5.py`
- `T5TokenizerFast` -> loaded via HF in `fastvideo/models/loader/`
- `UniPCMultistepScheduler` -> loaded via Diffusers scheduler utilities
- Pipeline defaults -> `fastvideo/configs/pipelines/wan.py`
- Sampling defaults -> `fastvideo/configs/sample/wan.py`
## Logger
## Pipeline system
See [PR](https://github.com/hao-ai-lab/FastVideo/pull/356)
- `fastvideo/pipelines/basic/*` contains end-to-end pipelines for each model
family.
- `fastvideo/pipelines/stages/*` contains reusable, testable stages.
- Pipelines subclass `ComposedPipelineBase` and declare required components via
`_required_config_modules`.
- `ForwardBatch` (in `fastvideo/pipelines/pipeline_batch_info.py`) carries
prompts, latents, timesteps, and intermediate state across stages.
*TODO*: (help wanted) Add an environment variable that disables process-aware logging.
## Model components
## Contributing to FastVideo
- DiT models: `fastvideo/models/dits/`
- VAEs: `fastvideo/models/vaes/`
- Text/image encoders: `fastvideo/models/encoders/`
- Schedulers: `fastvideo/models/schedulers/`
- Upsamplers: `fastvideo/models/upsamplers/`
- Optional audio models: `fastvideo/models/audio/`
If you're a new contributor, here are some common areas to explore:
## Attention and distributed execution
1. **Adding a new model**: Implement new model types in the appropriate subdirectory of `fastvideo/models/`
2. **Optimizing performance**: Look at attention implementations or memory management
3. **Adding a new pipeline**: Create a new pipeline subclass in `fastvideo/pipelines/`
4. **Hardware support**: Extend the `platforms` module for new hardware targets
- Attention backends live in `fastvideo/attention/` and can be selected via
`FASTVIDEO_ATTENTION_BACKEND`.
- `LocalAttention` is used for cross-attention and most attention layers.
- `DistributedAttention` is used for full-sequence self-attention in the DiT.
- Tensor-parallel layers live in `fastvideo/layers/`.
- Sequence/tensor parallel utilities live in `fastvideo/distributed/`.
When adding code, follow these practices:
## Related docs
- [Contributing overview](../contributing/overview.md)
- [Coding agents workflow](../contributing/coding_agents.md)
- [Testing guide](../contributing/testing.md)
- Use type hints for better code readability
- Add appropriate docstrings
- Maintain the separation between model components and execution logic
- Follow existing patterns for distributed processing
+8 -32
View File
@@ -131,41 +131,17 @@ self.attn = DistributedAttention(
### Registering Models
Register implemented modules for auto‑discovery by adding `EntryClass` in each
model module (the registry scans for it):
Register implemented modules in the model registry:
```python
# In fastvideo/models/dits/your_module.py
class YourTransformerModel(...):
...
# In fastvideo/models/registry.py
_TEXT_TO_VIDEO_DIT_MODELS = {
"YourTransformerModel": ("dits", "yourmodule", "YourTransformerClass"),
}
# Entry point for model registry
EntryClass = YourTransformerModel
```
```python
# In fastvideo/models/vaes/your_vae.py
class YourVAEModel(...):
...
# Entry point for model registry
EntryClass = YourVAEModel
```
Register pipeline config + sampling defaults in the unified registry:
```python
# In fastvideo/registry.py
register_configs(
sampling_param_cls=YourSamplingParam,
pipeline_config_cls=YourPipelineConfig,
hf_model_paths=[
"org/your-model-id",
],
model_detectors=[
lambda path: "your-model" in path.lower(),
],
)
_VAE_MODELS = {
"YourVAEModel": ("vaes", "yourvae", "YourVAEClass"),
}
```
## Step 2: Directory Structure
-140
View File
@@ -1,140 +0,0 @@
# Offloading
This page describes how to use offloading techniques for inference to reduce GPU memory usage while maintaining acceptable performance.
## Default Behavior
```python
dit_cpu_offload: bool = True
use_fsdp_inference: bool = False
dit_layerwise_offload: bool = True
text_encoder_cpu_offload: bool = True
image_encoder_cpu_offload: bool = True
vae_cpu_offload: bool = True
pin_cpu_memory: bool = True
```
## Behavior Explanation
!!! note
For CLI usage, replace underscores (`_`) with hyphens (`-`).
### `use_fsdp_inference`
Enables [FSDP](https://docs.pytorch.org/tutorials/intermediate/FSDP_tutorial.html) for inference. The model weights are sharded across multiple GPUs to reduce memory usage per GPU, and weights are broadcast to all GPUs layer by layer during inference.
#### Performance Impact
FSDP inference introduces negligible performance overhead due to weight prefetching. Performance overhead may be visible when GPU interconnect is slow (e.g., multiple consumer-level GPUs connected by slow PCIe without GPU P2P support).
#### Usage Recommendation
We recommend enabling this option when multiple GPUs are available.
### `dit_cpu_offload`
Enables CPU offloading for FSDP inference. When enabled, the model weights are offloaded to CPU memory, and the weight of each layer is moved to GPU memory only when that layer is being computed.
#### Performance Impact
The PyTorch FSDP implementation does not overlap computation and data transfer perfectly for inference, so enabling this option will harm performance.
#### Usage Recommendation
This option only takes effect when FSDP is enabled. For single GPU usage, we recommend using `dit_layerwise_offload` instead.
### `dit_layerwise_offload`
This option is similar to `dit_cpu_offload`, but with two key differences:
1. It overlaps computation and PCIe data transfer
2. It only works for single GPU inference
#### Performance Impact
This option introduces negligible performance overhead.
#### Usage Recommendation
We recommend enabling this option for single GPU usage. This option is not compatible with FSDP.
### `text_encoder_cpu_offload`
When enabled, the text encoder model weights are offloaded to CPU memory, and text encoding is computed on CPU.
#### Performance Impact
This option significantly slows down text encoding computation, but text encoding is usually not the bottleneck.
#### Usage Recommendation
We recommend enabling this option only when OOM happens.
### `image_encoder_cpu_offload` and `vae_cpu_offload`
When enabled, the weights are stored in CPU memory and moved to GPU memory when the corresponding module is being computed. After computation, the weights are moved back to CPU memory.
#### Performance Impact
These options introduce performance overhead due to PCIe data transfer.
#### Usage Recommendation
We recommend enabling these options when OOM happens.
## General Recommendations
### Single GPU Inference
We recommend enabling `dit_layerwise_offload`. If OOM happens, also enable `image_encoder_cpu_offload` and `vae_cpu_offload`. If OOM still happens, consider enabling `text_encoder_cpu_offload`.
### Multi-GPU Inference
We recommend enabling `use_fsdp_inference` and disabling both `dit_layerwise_offload` and `dit_cpu_offload`. If OOM happens, consider enabling `text_encoder_cpu_offload`, `image_encoder_cpu_offload`, and `vae_cpu_offload`. If OOM still happens, consider enabling `dit_cpu_offload`.
## Examples
### Single GPU with Layerwise Offloading
```python
from fastvideo import VideoGenerator
generator = VideoGenerator.from_pretrained(
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
num_gpus=1,
# Recommended for single GPU
dit_layerwise_offload=True,
# Enable if OOM happens
vae_cpu_offload=True,
image_encoder_cpu_offload=True,
text_encoder_cpu_offload=True,
# Speeds up CPU-GPU transfer
pin_cpu_memory=True,
)
prompt = "A curious raccoon peers through a vibrant field of yellow sunflowers."
video = generator.generate_video(prompt, output_path="output/", save_video=True)
```
### Multi-GPU with FSDP
```python
from fastvideo import VideoGenerator
generator = VideoGenerator.from_pretrained(
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
num_gpus=2,
# Recommended for multi-GPU
use_fsdp_inference=True,
dit_layerwise_offload=False,
dit_cpu_offload=False,
# Enable if OOM happens
vae_cpu_offload=True,
image_encoder_cpu_offload=True,
text_encoder_cpu_offload=True,
pin_cpu_memory=True,
)
prompt = "A majestic lion strides across the golden savanna."
video = generator.generate_video(prompt, output_path="output/", save_video=True)
```
+6 -2
View File
@@ -103,7 +103,7 @@ python setup.py install # or pip install -e .
**`SAGE_ATTN_THREE`**
[SageAttention 3](https://github.com/thu-ml/SageAttention/tree/main/sageattention3_blackwell) is an advanced attention mechanism that leverages FP4 quantization and Blackwell GPU Tensor Cores for significant performance improvements.
[SageAttention 3](https://huggingface.co/jt-zhang/SageAttention3) is an advanced attention mechanism that leverages FP4 quantization and Blackwell GPU Tensor Cores for significant performance improvements.
#### Hardware Requirements
@@ -113,7 +113,11 @@ python setup.py install # or pip install -e .
Note that Sage Attention 3 requires `python>=3.13`, `torch>=2.8.0`, `CUDA >=12.8`. If you are using `uv` and using `torch==2.8.0` make sure that `sentencepiece==0.2.1` in the pyproject.toml file.
To use Sage Attention 3 in FastVideo, follow the `README.md` in the linked repository to install the package from source.
To use Sage Attention 3 in FastVideo, first get access to the SageAttention3 code, then move `sageattn/` and `setup.py` to the directory `fastvideo/attention/backends`, then install from using:
```bash
python setup.py install
```
## Teacache
@@ -1,51 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from fastvideo import VideoGenerator
from fastvideo.configs.sample import SamplingParam
def main():
# Point this to your local diffusers model dir (or replace with a HF model ID).
model_path = "KyleShao/Cosmos-Predict2.5-2B-Diffusers"
generator = VideoGenerator.from_pretrained(
model_path,
num_gpus=1,
use_fsdp_inference=False, # set True if GPU is out of memory
dit_cpu_offload=False,
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
pin_cpu_memory=True,
)
sampling_param = SamplingParam.from_pretrained(model_path)
# image2world example from official repo
image_path = "images/bus_terminal.jpg"
prompt = (
"A nighttime city bus terminal gradually shifts from stillness to subtle movement. "
"At first, multiple double-decker buses are parked under the glow of overhead lights, "
"with a central bus labeled '87D' facing forward and stationary. "
"As the video progresses, the bus in the middle moves ahead slowly, its headlights brightening the surrounding area "
"and casting reflections onto adjacent vehicles. "
"The motion creates space in the lineup, signaling activity within the otherwise quiet station. "
"It then comes to a smooth stop, resuming its position in line. "
"Overhead signage in Chinese characters remains illuminated, enhancing the vibrant, urban night scene."
)
generator.generate_video(
prompt,
sampling_param=sampling_param,
image_path=str(image_path),
num_cond_frames=1,
output_path="outputs_video/cosmos2_5_i2w.mp4",
save_video=True,
)
generator.shutdown()
if __name__ == "__main__":
main()
@@ -1,5 +1,4 @@
from fastvideo import VideoGenerator
from fastvideo.configs.sample import SamplingParam
def main():
@@ -16,27 +15,19 @@ def main():
pin_cpu_memory=True,
)
# Load default sampling parameters (negative_prompt, resolution, steps, etc.)
sampling_param = SamplingParam.from_pretrained(model_path)
prompt = (
"A high-definition video captures the precision of robotic welding in an industrial setting. "
"The first frame showcases a robotic arm, equipped with a welding torch, positioned over a large metal structure. "
"The welding process is in full swing, with bright sparks and intense light illuminating the scene, "
"creating a vivid display of blue and white hues. "
"A significant amount of smoke billows around the welding area, partially obscuring the view but emphasizing the heat and activity. "
"The background reveals parts of the workshop environment, including a ventilation system and various pieces of machinery, "
"indicating a busy and functional industrial workspace. "
"As the video progresses, the robotic arm maintains its steady position, continuing the welding process and moving to its left. "
"The welding torch consistently emits sparks and light, and the smoke continues to rise, diffusing slightly as it moves upward. "
"The metal surface beneath the torch shows ongoing signs of heating and melting. "
"The scene retains its industrial ambiance, with the welding sparks and smoke dominating the visual field, "
"underscoring the ongoing nature of the welding operation."
"A high-definition video captures the precision of robotic welding in an industrial setting. The first frame showcases a robotic arm, equipped with a welding torch, positioned over a large metal structure. The welding process is in full swing, with bright sparks and intense light illuminating the scene, creating a vivid display of blue and white hues. A significant amount of smoke billows around the welding area, partially obscuring the view but emphasizing the heat and activity. The background reveals parts of the workshop environment, including a ventilation system and various pieces of machinery, indicating a busy and functional industrial workspace. As the video progresses, the robotic arm maintains its steady position, continuing the welding process and moving to its left. The welding torch consistently emits sparks and light, and the smoke continues to rise, diffusing slightly as it moves upward. The metal surface beneath the torch shows ongoing signs of heating and melting. The scene retains its industrial ambiance, with the welding sparks and smoke dominating the visual field, underscoring the ongoing nature of the welding operation."
)
generator.generate_video(
video = generator.generate_video(
prompt,
sampling_param=sampling_param,
negative_prompt="",
height=704,
width=1280,
num_frames=77,
num_inference_steps=35,
guidance_scale=7.0,
fps=24,
output_path="outputs_video/cosmos2_5_t2w.mp4",
save_video=True,
)
@@ -1,54 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from fastvideo import VideoGenerator
from fastvideo.configs.sample import SamplingParam
def main():
# Point this to your local diffusers model dir (or replace with a HF model ID).
model_path = "KyleShao/Cosmos-Predict2.5-2B-Diffusers"
generator = VideoGenerator.from_pretrained(
model_path,
num_gpus=1,
use_fsdp_inference=False, # set True if GPU is out of memory
dit_cpu_offload=False,
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
pin_cpu_memory=True,
)
sampling_param = SamplingParam.from_pretrained(model_path)
# video2world example from official repo
video_path = "videos/robot_pouring.mp4"
prompt = (
"A robotic arm, primarily white with black joints and cables, is shown in a clean, modern indoor setting with a white tabletop. "
"The arm, equipped with a gripper holding a small, light green pitcher, is positioned above a clear glass containing a reddish-brown liquid and a spoon. "
"The robotic arm is in the process of pouring a transparent liquid into the glass. "
"To the left of the pitcher, there is an opened jar with a similar reddish-brown substance visible through its transparent body. "
"In the background, a vase with white flowers and a brown couch are partially visible, adding to the contemporary ambiance. "
"The lighting is bright, casting soft shadows on the table. "
"The robotic arm's movements are smooth and controlled, demonstrating precision in its task. "
"As the video progresses, the robotic arm completes the pour, leaving the glass half-filled with the reddish-brown liquid. "
"The jar remains untouched throughout the sequence, and the spoon inside the glass remains stationary. "
"The other robotic arm on the right side also stays stationary throughout the video. "
"The final frame captures the robotic arm with the pitcher finishing the pour, with the glass now filled to a higher level, while the pitcher is slightly tilted but still held securely by the gripper."
)
generator.generate_video(
prompt,
sampling_param=sampling_param,
video_path=str(video_path),
num_cond_frames=1,
output_path="outputs_video/cosmos2_5_v2w.mp4",
save_video=True,
)
generator.shutdown()
if __name__ == "__main__":
main()
+1 -1
View File
@@ -39,4 +39,4 @@ def main():
if __name__ == "__main__":
main()
main()
@@ -1,43 +0,0 @@
from fastvideo import VideoGenerator
import json
# from fastvideo.configs.sample import SamplingParam
OUTPUT_PATH = "video_samples_hy15_1080p"
def main():
# FastVideo will automatically use the optimal default arguments for the
# model.
# If a local path is provided, FastVideo will make a best effort
# attempt to identify the optimal arguments.
generator = VideoGenerator.from_pretrained(
"weizhou03/HunyuanVideo-1.5-Diffusers-1080p-2SR", # 480p -> 720p -> 1080p
# or "weizhou03/HunyuanVideo-1.5-Diffusers-1080p" # 720p -> 1080p
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
dit_cpu_offload=True,
vae_cpu_offload=True,
text_encoder_cpu_offload=True,
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
# image_encoder_cpu_offload=False,
)
prompt = (
"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."
)
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, negative_prompt="")
prompt2 = (
"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.")
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True, negative_prompt="")
if __name__ == "__main__":
main()
-66
View File
@@ -1,66 +0,0 @@
from fastvideo import VideoGenerator
from fastvideo.models.dits.hyworld.resolution_utils import get_resolution_from_image
# Default prompt from HY-WorldPlay run.sh
DEFAULT_PROMPT = 'A paved pathway leads towards a stone arch bridge spanning a calm body of water. Lush green trees and foliage line the path and the far bank of the water. A traditional-style pavilion with a tiered, reddish-brown roof sits on the far shore. The water reflects the surrounding greenery and the sky. The scene is bathed in soft, natural light, creating a tranquil and serene atmosphere. The pathway is composed of large, rectangular stones, and the bridge is constructed of light gray stone. The overall composition emphasizes the peaceful and harmonious nature of the landscape.'
DEFAULT_IMAGE = 'https://raw.githubusercontent.com/Tencent-Hunyuan/HY-WorldPlay/main/assets/img/test.png'
OUTPUT_PATH = "video_samples_hyworld"
def main():
import argparse
# pose: (a, w, s, d) - (15, 31)
# num_frames: (61, 125)
parser = argparse.ArgumentParser(description="HYWorld video generation with FastVideo")
parser.add_argument("--prompt", type=str, default=DEFAULT_PROMPT, help="Text prompt for video generation")
parser.add_argument("--image", type=str, default=DEFAULT_IMAGE, help="Path or URL to input image")
parser.add_argument("--pose", type=str, default='w-31', help="Pose string (e.g., 'a-31', 'w-31', 's-31', 'd-31')")
parser.add_argument("--output_path", type=str, default=OUTPUT_PATH, help="Output video path")
parser.add_argument("--num-frames", type=int, default=125, help="Number of frames")
parser.add_argument("--seed", type=int, default=1, help="Random seed")
parser.add_argument("--resolution", type=str, default="480p", help="Only support 480p for now")
args = parser.parse_args()
# Automatically determine resolution from input image
HEIGHT, WIDTH = get_resolution_from_image(args.image, args.resolution)
print(f"Image: {args.image}")
print(f"Pose: {args.pose}")
print(f"Resolution: {HEIGHT}x{WIDTH} (from {args.resolution} buckets)")
print(f"Num frames: {args.num_frames}")
print(f"Output path: {args.output_path}")
# Initialize generator
print("\nInitializing VideoGenerator for HYWorld...")
generator = VideoGenerator.from_pretrained(
"FastVideo/HY-WorldPlay-Bidirectional-Diffusers",
num_gpus=1,
use_fsdp_inference=True,
dit_cpu_offload=True,
vae_cpu_offload=True,
text_encoder_cpu_offload=True,
pin_cpu_memory=True,
image_encoder_cpu_offload=True,
)
# Generate video
# The pose string is automatically converted to camera matrices by the pipeline
print("\nGenerating video...")
generator.generate_video(
prompt=args.prompt,
image_path=args.image,
pose=args.pose, # Camera trajectory control
output_path=args.output_path,
save_video=True,
negative_prompt="",
num_frames=args.num_frames,
fps=24,
height=HEIGHT,
width=WIDTH,
seed=args.seed,
)
print(f"\nVideo saved to: {args.output_path}")
if __name__ == "__main__":
main()
-34
View File
@@ -1,34 +0,0 @@
from fastvideo import VideoGenerator
PROMPT = (
"A warm sunny backyard. The camera starts in a tight cinematic close-up "
"of a woman and a man in their 30s, facing each other with serious "
"expressions. The woman, emotional and dramatic, says softly, \"That's "
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
"annoyed: \"Stop being so dramatic, Jess.\" A beat. He glances aside, "
"then mutters defensively, \"He's just having fun.\" The camera slowly "
"pans right, revealing the grandfather in the garden wearing enormous "
"butterfly wings, waving his arms in the air like he's trying to take "
"off. He shouts, \"Wheeeew!\" as he flaps his wings with full commitment. "
"The woman covers her face, on the verge of tears. The tone is deadpan, "
"absurd, and quietly tragic."
)
def main() -> None:
generator = VideoGenerator.from_pretrained(
"FastVideo/LTX2-Distilled-Diffusers",
num_gpus=1,
)
output_path = "outputs_video/ltx2_basic/output_ltx2_distilled_t2v.mp4"
generator.generate_video(
prompt=PROMPT,
output_path=output_path,
save_video=True,
)
generator.shutdown()
if __name__ == "__main__":
main()
+2 -1
View File
@@ -1,5 +1,6 @@
from fastvideo import VideoGenerator
from fastvideo.models.dits.matrixgame.utils import create_action_presets
from fastvideo.configs.pipelines.wan import MatrixGameI2V480PConfig
from fastvideo.models.dits.matrix_game.utils import create_action_presets
import torch
@@ -1,5 +1,5 @@
from fastvideo.entrypoints.streaming_generator import StreamingVideoGenerator
from fastvideo.models.dits.matrixgame.utils import get_current_action_async, expand_action_to_frames
from fastvideo.models.dits.matrix_game.utils import get_current_action_async, expand_action_to_frames
import torch
import asyncio
@@ -10,7 +10,7 @@ from fastapi import FastAPI, Request, HTTPException
from fastapi.responses import HTMLResponse, FileResponse
from fastvideo.entrypoints.streaming_generator import StreamingVideoGenerator
from fastvideo.models.dits.matrixgame.utils import expand_action_to_frames
from fastvideo.models.dits.matrix_game.utils import expand_action_to_frames
VARIANT_CONFIG = {
@@ -4,7 +4,7 @@ export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export TOKENIZERS_PARALLELISM=false
MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATA_DIR="data/crush-smol_processed_t2v_1_3b_ode_init/"
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
NUM_GPUS=1
@@ -14,6 +14,7 @@ NUM_GPUS=1
training_args=(
--tracker_project_name "wan_ode_init"
--output_dir "wan_ode_init_crush_smol"
--override_transformer_cls_name "CausalWanTransformer3DModel"
--wandb_run_name "wan_ode_init_crush_smol"
--max_train_steps 6000
--train_batch_size 1
@@ -33,7 +34,7 @@ parallel_args=(
--sp_size 1
--tp_size 1
--hsdp_replicate_dim 1
--hsdp_shard_dim $NUM_GPUS
--hsdp_shard_dim 1
)
# Model arguments
@@ -50,17 +51,20 @@ dataset_args=(
# Validation arguments
validation_args=(
--log-visualization
--visualization-steps 100
--log_validation
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 50
--validation_sampling_steps "50"
--validation_guidance_scale "6.0"
)
# Optimizer arguments
optimizer_args=(
--learning_rate 1e-5
--learning_rate 6e-6
--mixed_precision "bf16"
--weight_only_checkpointing_steps 1000
--training_state_checkpointing_steps 1000
--weight_decay 0.01
--weight_decay 1e-4
--max_grad_norm 1.0
)
@@ -1,93 +0,0 @@
#!/bin/bash
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export TOKENIZERS_PARALLELISM=false
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
MODEL_PATH="FastVideo/Matrix-Game-2.0-Foundation-Diffusers"
DATA_DIR="footsies-dataset/preprocessed/combined_parquet_dataset"
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
NUM_GPUS=8
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name "matrixgame_finetune"
--output_dir "checkpoints/matrixgame_finetune"
--max_train_steps 1500
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 4
--num_latent_t 20
--num_height 352
--num_width 640
--num_frames 77
--enable_gradient_checkpointing_type "full"
)
# Parallel arguments
parallel_args=(
--num_gpus $NUM_GPUS
--sp_size 2
--tp_size 1
--hsdp_replicate_dim 4
--hsdp_shard_dim 2
)
# 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 1
)
# Validation arguments
validation_args=(
--log_validation
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 100
--validation_sampling_steps "40"
--validation_guidance_scale "6.0"
)
# Optimizer arguments
optimizer_args=(
--learning_rate 2e-5
--mixed_precision "bf16"
--weight_only_checkpointing_steps 100
--training_state_checkpointing_steps 100
--weight_decay 1e-4
--max_grad_norm 1.0
)
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
--training_cfg_rate 0.1
--multi_phased_distill_schedule "4000-1"
--not_apply_cfg_solver
--dit_precision "fp32"
--num_euler_timesteps 50
--ema_start_step 0
)
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
torchrun \
--nnodes 1 \
--nproc_per_node $NUM_GPUS \
fastvideo/training/matrixgame_training_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}"
@@ -1,27 +0,0 @@
#!/bin/bash
GPU_NUM=1 # 2,4,8
MODEL_PATH="Matrix-Game-2.0-Foundation-Diffusers"
DATA_MERGE_PATH="footsies-dataset/merge.txt"
OUTPUT_DIR="footsies-dataset/preprocessed/"
# export CUDA_VISIBLE_DEVICES=0
export MASTER_ADDR=localhost
export MASTER_PORT=29500
export RANK=0
export WORLD_SIZE=1
python fastvideo/pipelines/preprocess/v1_preprocess.py \
--model_path $MODEL_PATH \
--data_merge_path $DATA_MERGE_PATH \
--preprocess_video_batch_size 4 \
--seed 42 \
--max_height 352 \
--max_width 640 \
--num_frames 77 \
--dataloader_num_workers 0 \
--output_dir=$OUTPUT_DIR \
--samples_per_file 4 \
--train_fps 25 \
--flush_frequency 4 \
--preprocess_task matrixgame
File diff suppressed because it is too large Load Diff
+129
View File
@@ -0,0 +1,129 @@
#!/bin/bash
# Change to FastVideo root directory (3 levels up from this script)
SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
FASTVIDEO_ROOT="$(cd "$SCRIPT_DIR/../../.." && pwd)"
cd "$FASTVIDEO_ROOT"
# Add FastVideo root to PYTHONPATH so Python can find the fastvideo package
export PYTHONPATH="$FASTVIDEO_ROOT${PYTHONPATH:+:$PYTHONPATH}"
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
RL_DATASET_DIR="data/ocr/" # Path to RL prompt dataset directory (should contain train.txt and test.txt)
VALIDATION_DATASET_FILE="$SCRIPT_DIR/validation.json"
NUM_GPUS=1
# use GPU 3
export CUDA_VISIBLE_DEVICES=3
# Training arguments
training_args=(
--tracker_project_name "wan_t2v_grpo"
--output_dir "checkpoints/wan_t2v_grpo"
--max_train_steps 5000
--train_batch_size 4
# --train_sp_batch_size 4
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 5
--num_height 240
--num_width 416
--num_frames 33
--lora_rank 32
--lora_training True
)
# Parallel arguments
parallel_args=(
--num_gpus $NUM_GPUS
--sp_size $NUM_GPUS
--tp_size $NUM_GPUS
--hsdp_replicate_dim 1
--hsdp_shard_dim $NUM_GPUS
# --use-fsdp-inference False
)
# Model arguments
model_args=(
--model_path $MODEL_PATH
--pretrained_model_name_or_path $MODEL_PATH
)
# Dataset arguments (for RL prompt dataset)
dataset_args=(
--data_path $RL_DATASET_DIR # Used as fallback if rl_dataset_path not set
--rl_dataset_path $RL_DATASET_DIR # RL prompt dataset directory
--rl_dataset_type "text" # "text" or "geneval"
--rl_num_image_per_prompt 4 # k parameter (number of samples per prompt)
--dataloader_num_workers 1
)
# Validation arguments
validation_args=(
--log_validation True
--validation_dataset_file $VALIDATION_DATASET_FILE
--validation_steps 5
--validation_sampling_steps "50"
--validation_guidance_scale "6.0"
)
# Optimizer arguments
optimizer_args=(
--learning_rate 5e-5
--mixed_precision "bf16"
--weight_only_checkpointing_steps 10
--training_state_checkpointing_steps 10
--weight_decay 1e-4
--max_grad_norm 1.0
)
# RL-specific arguments
rl_args=(
--inference_mode False
--rl_mode True
--rl_algorithm "grpo"
--rl_kl_beta 0.004 # KL regularization coefficient
--rl_policy_clip_range 0.2 # Policy clipping range for GRPO
--rl_kl_reward 0.0 # KL reward coefficient (typically 0)
--rl_global_std False # Use per-prompt std (recommended for GRPO)
--rl_per_prompt_stat_tracking True # Enable per-prompt stat tracking
--rl_warmup_steps 0 # Number of warmup steps (SFT before RL)
--reward-models "{\"paddle_ocr\": 1.0}" # use video_ocr reward function
)
# CFG arguments
cfg_args=(
--guidance_scale 1.0 # use guidance_scale > 1.0 to enable CFG
)
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
--training_cfg_rate 0.0 # No CFG during training (CFG used in sampling)
--dit_precision "fp32"
# --dit_precision "bf16"
--num_euler_timesteps 50
--ema_start_step 0
# --resume_from_checkpoint "checkpoints/wan_t2v_grpo/checkpoint-XXX"
--enable-gradient-checkpointing-type "full"
)
torchrun \
--nnodes 1 \
--nproc_per_node $NUM_GPUS \
--master_port 29501 \
"$FASTVIDEO_ROOT/fastvideo/training/wan_rl_training_pipeline.py" \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${rl_args[@]}" \
"${miscellaneous_args[@]}"
-11
View File
@@ -40,17 +40,6 @@ out = video_sparse_attn(q, k, v, block_sizes, block_sizes, topk=5)
out = moba_attn_varlen(q, k, v, cu_seqlens_q, cu_seqlens_k, ...)
```
## Benchmark
### VSA (block-sparse) TFLOPs
After building/installing `fastvideo-kernel`, run:
```bash
cd fastvideo-kernel
python benchmarks/bench_vsa.py --batch_size 1 --num_heads 16 --head_dim 128 --q_seq_lens 49152 --topk 64
```
### TurboDiffusion Kernels
This package also includes kernels from [TurboDiffusion](https://github.com/thu-ml/TurboDiffusion), including INT8 GEMM, Quantization, RMSNorm and LayerNorm.
-166
View File
@@ -1,166 +0,0 @@
#!/usr/bin/env python3
"""
Benchmark VSA *wrapper* performance (forward + backward) and report TFLOPs.
This script benchmarks the autograd-enabled wrapper:
- fastvideo_kernel.block_sparse_attn.block_sparse_attn
So measured time includes wrapper overhead (map->index conversion, dispatch) plus kernel time.
"""
from __future__ import annotations
import argparse
import os
import random
from typing import Tuple, Callable
import numpy as np
import torch
try:
from triton.testing import do_bench
except Exception as e: # pragma: no cover
raise ImportError("This benchmark requires triton (for triton.testing.do_bench).") from e
BLOCK_M = 64
BLOCK_N = 64
def set_seed(seed: int = 42) -> None:
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
torch.cuda.manual_seed(seed)
torch.cuda.manual_seed_all(seed)
def parse_arguments() -> argparse.Namespace:
p = argparse.ArgumentParser(description="Benchmark FastVideo VSA block-sparse attention")
p.add_argument("--batch_size", type=int, default=1)
p.add_argument("--num_heads", type=int, default=12)
p.add_argument("--head_dim", type=int, default=128, choices=[64, 128])
p.add_argument("--topk", type=int, default=None, help="KV blocks per Q block (default: ~90%% sparsity)")
p.add_argument("--q_seq_lens", type=int, nargs="+", default=[49152], help="Q sequence lengths (must be /64)")
p.add_argument("--kv_seq_lens", type=int, nargs="+", default=None, help="KV sequence lengths (defaults to q_seq_len)")
p.add_argument("--warmup", type=int, default=5)
p.add_argument("--rep", type=int, default=20)
p.add_argument("--seed", type=int, default=42)
p.add_argument("--dtype", type=str, default="bf16", choices=["bf16", "fp16"])
p.add_argument("--force_triton", action="store_true", help="Force wrapper to use Triton path (if supported by shapes).")
return p.parse_args()
def create_qkv(batch: int, heads: int, q_len: int, kv_len: int, d: int, dtype: torch.dtype) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
q = torch.randn(batch, heads, q_len, d, dtype=dtype, device="cuda")
k = torch.randn(batch, heads, kv_len, d, dtype=dtype, device="cuda")
v = torch.randn(batch, heads, kv_len, d, dtype=dtype, device="cuda")
return q, k, v
def make_block_map(bs: int, h: int, num_q_blocks: int, num_kv_blocks: int, topk: int) -> torch.Tensor:
# block_map: [bs, h, num_q_blocks, num_kv_blocks] bool
scores = torch.rand(bs, h, num_q_blocks, num_kv_blocks, device="cuda")
topk = min(max(1, topk), num_kv_blocks)
idx = torch.topk(scores, topk, dim=-1).indices
block_map = torch.zeros(bs, h, num_q_blocks, num_kv_blocks, dtype=torch.bool, device="cuda")
block_map.scatter_(-1, idx, True)
return block_map
def flops_sparse_attention(bs: int, h: int, d: int, q_len: int, topk_blocks: int, block_n: int) -> float:
# Approx: QK^T + PV, each is ~2*bs*h*q_len*(topk_blocks*block_n)*d
return 4.0 * bs * h * d * q_len * (topk_blocks * block_n)
def bench_ms(fn: Callable[[], object], warmup: int, rep: int) -> float:
return do_bench(fn, warmup=warmup, rep=rep, quantiles=None)
def main() -> None:
args = parse_arguments()
set_seed(args.seed)
dtype = torch.bfloat16 if args.dtype == "bf16" else torch.float16
if args.force_triton:
os.environ["FASTVIDEO_KERNEL_VSA_FORCE_TRITON"] = "1"
from fastvideo_kernel.block_sparse_attn import block_sparse_attn
bs, h, d = args.batch_size, args.num_heads, args.head_dim
kv_seq_lens = args.kv_seq_lens
if kv_seq_lens is None:
kv_seq_lens = args.q_seq_lens
if len(kv_seq_lens) != len(args.q_seq_lens):
raise ValueError("kv_seq_lens must have the same number of entries as q_seq_lens (or be omitted).")
print("VSA Block-Sparse Attention Benchmark (WRAPPER)")
print(f"device: {torch.cuda.get_device_name(0)}")
print(f"batch={bs}, heads={h}, head_dim={d}, dtype={args.dtype}")
print(f"BLOCK_M={BLOCK_M}, BLOCK_N={BLOCK_N}")
print("NOTE: timings include wrapper overhead (map->index + dispatch).")
if args.force_triton:
print("dispatch: forced Triton (FASTVIDEO_KERNEL_VSA_FORCE_TRITON=1)")
else:
print("dispatch: SM90 if available, else Triton")
for q_len, kv_len in zip(args.q_seq_lens, kv_seq_lens):
if q_len % BLOCK_M != 0 or kv_len % BLOCK_N != 0:
print(f"[skip] q_len={q_len}, kv_len={kv_len} must be divisible by 64")
continue
num_q_blocks = q_len // BLOCK_M
num_kv_blocks = kv_len // BLOCK_N
topk = args.topk if args.topk is not None else max(1, num_kv_blocks // 10)
topk = min(topk, num_kv_blocks)
print("\n" + "=" * 80)
print(f"q_len={q_len}, kv_len={kv_len}, num_q_blocks={num_q_blocks}, num_kv_blocks={num_kv_blocks}, topk={topk}")
q, k, v = create_qkv(bs, h, q_len, kv_len, d, dtype)
block_map = make_block_map(bs, h, num_q_blocks, num_kv_blocks, topk)
# Variable block sizes: default full blocks (64 tokens per KV block)
variable_block_sizes = torch.full((num_kv_blocks,), BLOCK_N, dtype=torch.int32, device="cuda")
def _fwd():
return block_sparse_attn(q, k, v, block_map, variable_block_sizes)
fwd_ms = bench_ms(_fwd, warmup=args.warmup, rep=args.rep)
# Backward benchmark (wrapper autograd). We build the graph once, then repeatedly run backward
# on the retained graph so bwd timing excludes the forward compute.
q_ = q.detach().requires_grad_(True)
k_ = k.detach().requires_grad_(True)
v_ = v.detach().requires_grad_(True)
o_, _aux_ = block_sparse_attn(q_, k_, v_, block_map, variable_block_sizes)
og = torch.randn_like(o_)
loss = (o_ * og).sum()
for _ in range(max(1, args.warmup // 2)):
torch.autograd.grad(loss, (q_, k_, v_), retain_graph=True)
torch.cuda.synchronize()
bwd_ms = bench_ms(
lambda: torch.autograd.grad(loss, (q_, k_, v_), retain_graph=True),
warmup=0,
rep=max(5, args.rep // 2),
)
flops = flops_sparse_attention(bs, h, d, q_len, topk, BLOCK_N)
fwd_tflops = flops / fwd_ms * 1e-12 * 1e3
# Rough backward multiplier (attention backward typically ~2-3x forward)
bwd_tflops = (2.5 * flops) / bwd_ms * 1e-12 * 1e3
print(f"fwd(wrapper): {fwd_ms:.3f} ms | {fwd_tflops:.2f} TFLOPs (approx)")
print(f"bwd(wrapper): {bwd_ms:.3f} ms | {bwd_tflops:.2f} TFLOPs (approx)")
if __name__ == "__main__":
if not torch.cuda.is_available():
raise RuntimeError("CUDA is required for this benchmark.")
main()
+1 -1
View File
@@ -9,7 +9,7 @@ build-backend = "scikit_build_core.build"
[project]
name = "fastvideo-kernel"
version = "0.2.5"
version = "0.2.4"
description = "Unified CUDA kernels for FastVideo"
readme = "README.md"
requires-python = ">=3.10"
@@ -30,12 +30,13 @@ def _force_triton() -> bool:
return os.environ.get("FASTVIDEO_KERNEL_VSA_FORCE_TRITON", "0") == "1"
def _map_to_index(block_map: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
def _map_to_index_torch(block_map: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
"""
Preferred map->index conversion used by the wrapper.
This wrapper **requires** the Triton implementation.
If Triton (or the Triton map_to_index module) is not available, it raises.
Pure-torch (no triton) conversion:
block_map: [B, H, Q, KV] bool (or [H, Q, KV] which will be treated as B=1)
returns:
index: [B, H, Q, KV] int32 (packed KV indices, -1 padding)
num: [B, H, Q] int32 (#kv blocks per q block)
"""
if block_map.dim() == 3:
block_map = block_map.unsqueeze(0)
@@ -44,17 +45,20 @@ def _map_to_index(block_map: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
if block_map.dtype != torch.bool:
block_map = block_map.to(torch.bool)
if not block_map.is_cuda:
raise RuntimeError("block_map must be a CUDA tensor (Triton map_to_index required).")
B, H, Q, KV = block_map.shape
index = torch.full((B, H, Q, KV), -1, dtype=torch.int32, device=block_map.device)
num = torch.zeros((B, H, Q), dtype=torch.int32, device=block_map.device)
try:
from fastvideo_kernel.triton_kernels.index import map_to_index as triton_map_to_index # local import
except Exception as e:
raise ImportError(
"Triton map_to_index is required but not available. "
"Ensure Triton is installed and fastvideo_kernel.triton_kernels.index is importable."
) from e
return triton_map_to_index(block_map)
# Small sizes in practice (B=1, H<=16, Q/KV<=64), so a Python loop is fine.
for b in range(B):
for h in range(H):
for q in range(Q):
kv_idx = torch.nonzero(block_map[b, h, q], as_tuple=False).flatten().to(torch.int32)
n = int(kv_idx.numel())
if n:
index[b, h, q, :n] = kv_idx
num[b, h, q] = n
return index, num
@torch.library.custom_op(
@@ -73,7 +77,7 @@ def block_sparse_attn_triton(
k = k.contiguous()
v = v.contiguous()
block_map = block_map.to(torch.bool)
q2k_idx, q2k_num = _map_to_index(block_map)
q2k_idx, q2k_num = _map_to_index_torch(block_map)
from fastvideo_kernel.triton_kernels.block_sparse_attn_triton import ( # local import
triton_block_sparse_attn_forward,
@@ -83,7 +87,6 @@ def block_sparse_attn_triton(
return o, M
@torch.library.register_fake("fastvideo_kernel::block_sparse_attn_triton")
def _block_sparse_attn_triton_fake(
q: torch.Tensor,
@@ -114,8 +117,8 @@ def block_sparse_attn_backward_triton(
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
grad_output = grad_output.contiguous()
block_map = block_map.to(torch.bool)
q2k_idx, q2k_num = _map_to_index(block_map)
k2q_idx, k2q_num = _map_to_index(block_map.transpose(-1, -2).contiguous())
q2k_idx, q2k_num = _map_to_index_torch(block_map)
k2q_idx, k2q_num = _map_to_index_torch(block_map.transpose(-1, -2).contiguous())
from fastvideo_kernel.triton_kernels.block_sparse_attn_triton import ( # local import
triton_block_sparse_attn_backward,
@@ -179,7 +182,7 @@ def block_sparse_attn_sm90(
k_padded = k_padded.contiguous()
v_padded = v_padded.contiguous()
block_map = block_map.to(torch.bool)
q2k_idx, q2k_num = _map_to_index(block_map)
q2k_idx, q2k_num = _map_to_index_torch(block_map)
o_padded, lse_padded = block_sparse_fwd(
q_padded, k_padded, v_padded, q2k_idx, q2k_num, variable_block_sizes.int()
@@ -221,7 +224,7 @@ def block_sparse_attn_backward_sm90(
grad_output_padded = grad_output_padded.contiguous()
block_map = block_map.to(torch.bool)
k2q_idx, k2q_num = _map_to_index(block_map.transpose(-1, -2).contiguous())
k2q_idx, k2q_num = _map_to_index_torch(block_map.transpose(-1, -2).contiguous())
dq, dk, dv = block_sparse_bwd(
q_padded,
@@ -1 +1 @@
__version__ = "0.2.5"
__version__ = "0.2.4"
+2 -2
View File
@@ -1,7 +1,7 @@
# SPDX-License-Identifier: Apache-2.0
import torch
from sageattn3 import sageattn3_blackwell
from fastvideo.attention.backends.sageattn.api import sageattn_blackwell
from fastvideo.attention.backends.abstract import (AttentionBackend,
AttentionImpl,
@@ -67,6 +67,6 @@ class SageAttention3Impl(AttentionImpl):
query = query.transpose(1, 2)
key = key.transpose(1, 2)
value = value.transpose(1, 2)
output = sageattn3_blackwell(query, key, value, is_causal=self.causal)
output = sageattn_blackwell(query, key, value, is_causal=self.causal)
output = output.transpose(1, 2)
return output
+1 -14
View File
@@ -2,18 +2,5 @@ from fastvideo.configs.models.base import ModelConfig
from fastvideo.configs.models.dits.base import DiTConfig
from fastvideo.configs.models.encoders.base import EncoderConfig
from fastvideo.configs.models.vaes.base import VAEConfig
from fastvideo.configs.models.audio import (LTX2AudioDecoderConfig,
LTX2AudioEncoderConfig,
LTX2VocoderConfig)
from fastvideo.configs.models.upsamplers.base import UpsamplerConfig
__all__ = [
"ModelConfig",
"VAEConfig",
"DiTConfig",
"EncoderConfig",
"LTX2AudioEncoderConfig",
"LTX2AudioDecoderConfig",
"LTX2VocoderConfig",
"UpsamplerConfig",
]
__all__ = ["ModelConfig", "VAEConfig", "DiTConfig", "EncoderConfig"]
@@ -1,13 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from fastvideo.configs.models.audio.ltx2_audio_vae import (
LTX2AudioDecoderConfig,
LTX2AudioEncoderConfig,
LTX2VocoderConfig,
)
__all__ = [
"LTX2AudioEncoderConfig",
"LTX2AudioDecoderConfig",
"LTX2VocoderConfig",
]
@@ -1,31 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""
LTX-2 audio VAE and vocoder configuration.
"""
from dataclasses import dataclass, field
from fastvideo.configs.models.base import ArchConfig, ModelConfig
@dataclass
class LTX2AudioArchConfig(ArchConfig):
architectures: list[str] = field(default_factory=list)
@dataclass
class LTX2AudioEncoderConfig(ModelConfig):
arch_config: ArchConfig = field(default_factory=lambda: LTX2AudioArchConfig(
architectures=["LTX2AudioEncoder"]))
@dataclass
class LTX2AudioDecoderConfig(ModelConfig):
arch_config: ArchConfig = field(default_factory=lambda: LTX2AudioArchConfig(
architectures=["LTX2AudioDecoder"]))
@dataclass
class LTX2VocoderConfig(ModelConfig):
arch_config: ArchConfig = field(default_factory=lambda: LTX2AudioArchConfig(
architectures=["LTX2Vocoder"]))
+1 -3
View File
@@ -3,13 +3,11 @@ from fastvideo.configs.models.dits.cosmos2_5 import Cosmos25VideoConfig
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", "WanVideoConfig",
"StepVideoConfig", "CosmosVideoConfig", "Cosmos25VideoConfig",
"LongCatVideoConfig", "LTX2VideoConfig", "HYWorldConfig"
"LongCatVideoConfig"
]
@@ -55,8 +55,6 @@ class HunyuanVideo15ArchConfig(DiTArchConfig):
r"txt_in.refiner_blocks.\1.mlp.fc_out.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.norm_out\.linear\.(.*)$":
r"txt_in.refiner_blocks.\1.adaLN_modulation.linear.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.self_attn_qkv\.(.*)$":
r"txt_in.refiner_blocks.\1.self_attn_qkv.\2",
# 2. txt_in_2 mapping:
r"^context_embedder_2\.(.*)$":
-202
View File
@@ -1,202 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
def is_double_block(n: str, m) -> bool:
return "double" in n and str.isdigit(n.split(".")[-1])
def is_refiner_block(n: str, m) -> bool:
return "refiner" in n and str.isdigit(n.split(".")[-1])
def is_txt_in(n: str, m) -> bool:
return n.split(".")[-1] == "txt_in"
@dataclass
class HYWorldArchConfig(DiTArchConfig):
_fsdp_shard_conditions: list = field(
default_factory=lambda: [is_double_block, is_refiner_block])
_compile_conditions: list = field(
default_factory=lambda: [is_double_block, is_refiner_block, is_txt_in])
param_names_mapping: dict = field(
default_factory=lambda: {
# 1. txt_in submodules (text embedder, refiner blocks):
r"^txt_in\.t_embedder\.mlp\.0\.(.*)$":
r"txt_in.t_embedder.mlp.fc_in.\1",
r"^txt_in\.t_embedder\.mlp\.2\.(.*)$":
r"txt_in.t_embedder.mlp.fc_out.\1",
r"^txt_in\.c_embedder\.linear_1\.(.*)$":
r"txt_in.c_embedder.fc_in.\1",
r"^txt_in\.c_embedder\.linear_2\.(.*)$":
r"txt_in.c_embedder.fc_out.\1",
r"^txt_in\.individual_token_refiner\.blocks\.(\d+)\.norm1\.(.*)$":
r"txt_in.refiner_blocks.\1.norm1.\2",
r"^txt_in\.individual_token_refiner\.blocks\.(\d+)\.norm2\.(.*)$":
r"txt_in.refiner_blocks.\1.norm2.\2",
r"^txt_in\.individual_token_refiner\.blocks\.(\d+)\.self_attn_qkv\.(.*)$":
r"txt_in.refiner_blocks.\1.self_attn_qkv.\2",
r"^txt_in\.individual_token_refiner\.blocks\.(\d+)\.self_attn_proj\.(.*)$":
r"txt_in.refiner_blocks.\1.self_attn_proj.\2",
r"^txt_in\.individual_token_refiner\.blocks\.(\d+)\.mlp\.fc1\.(.*)$":
r"txt_in.refiner_blocks.\1.mlp.fc_in.\2",
r"^txt_in\.individual_token_refiner\.blocks\.(\d+)\.mlp\.fc2\.(.*)$":
r"txt_in.refiner_blocks.\1.mlp.fc_out.\2",
r"^txt_in\.individual_token_refiner\.blocks\.(\d+)\.adaLN_modulation\.1\.(.*)$":
r"txt_in.refiner_blocks.\1.adaLN_modulation.linear.\2",
# 2. time_in mappings:
r"^time_in\.mlp\.0\.(.*)$":
r"time_in.timestep_embedder.mlp.fc_in.\1",
r"^time_in\.mlp\.2\.(.*)$":
r"time_in.timestep_embedder.mlp.fc_out.\1",
# 3. action_in mappings:
r"^action_in\.mlp\.0\.(.*)$":
r"action_in.mlp.fc_in.\1",
r"^action_in\.mlp\.2\.(.*)$":
r"action_in.mlp.fc_out.\1",
# 4. byt5_in -> txt_in_2 mappings:
r"^byt5_in\.layernorm\.(.*)$":
r"txt_in_2.norm.\1",
r"^byt5_in\.fc1\.(.*)$":
r"txt_in_2.linear_1.\1",
r"^byt5_in\.fc2\.(.*)$":
r"txt_in_2.linear_2.\1",
r"^byt5_in\.fc3\.(.*)$":
r"txt_in_2.linear_3.\1",
# 5. cond_type_embedding -> cond_type_embed:
r"^cond_type_embedding\.(.*)$":
r"cond_type_embed.\1",
# 6. vision_in -> image_embedder mappings:
r"^vision_in\.proj\.0\.(.*)$":
r"image_embedder.norm_in.\1",
r"^vision_in\.proj\.1\.(.*)$":
r"image_embedder.linear_1.\1",
r"^vision_in\.proj\.3\.(.*)$":
r"image_embedder.linear_2.\1",
r"^vision_in\.proj\.4\.(.*)$":
r"image_embedder.norm_out.\1",
# 7. double_blocks mapping:
r"^double_blocks\.(\d+)\.img_attn_q\.(.*)$":
(r"double_blocks.\1.img_attn_qkv.\2", 0, 3),
r"^double_blocks\.(\d+)\.img_attn_k\.(.*)$":
(r"double_blocks.\1.img_attn_qkv.\2", 1, 3),
r"^double_blocks\.(\d+)\.img_attn_v\.(.*)$":
(r"double_blocks.\1.img_attn_qkv.\2", 2, 3),
r"^double_blocks\.(\d+)\.txt_attn_q\.(.*)$":
(r"double_blocks.\1.txt_attn_qkv.\2", 0, 3),
r"^double_blocks\.(\d+)\.txt_attn_k\.(.*)$":
(r"double_blocks.\1.txt_attn_qkv.\2", 1, 3),
r"^double_blocks\.(\d+)\.txt_attn_v\.(.*)$":
(r"double_blocks.\1.txt_attn_qkv.\2", 2, 3),
r"^double_blocks\.(\d+)\.img_mlp\.fc1\.(.*)$":
r"double_blocks.\1.img_mlp.fc_in.\2",
r"^double_blocks\.(\d+)\.img_mlp\.fc2\.(.*)$":
r"double_blocks.\1.img_mlp.fc_out.\2",
r"^double_blocks\.(\d+)\.txt_mlp\.fc1\.(.*)$":
r"double_blocks.\1.txt_mlp.fc_in.\2",
r"^double_blocks\.(\d+)\.txt_mlp\.fc2\.(.*)$":
r"double_blocks.\1.txt_mlp.fc_out.\2",
# 8. Final layer mapping:
r"^final_layer\.adaLN_modulation\.1\.(.*)$":
r"final_layer.adaLN_modulation.linear.\1",
})
# Reverse mapping for saving checkpoints: custom -> hf
reverse_param_names_mapping: dict = field(default_factory=lambda: {})
# Parameters from HY-WorldPlay config.json (loaded from checkpoint)
patch_size: list | tuple | int = field(default_factory=lambda: [1, 1, 1])
# Base latent channels - will be expanded in __post_init__ if concat_condition=True
in_channels: int = 32
concat_condition: bool = True
out_channels: int = 32
hidden_size: int = 2048
heads_num: int = 16
mlp_width_ratio: float = 4.0
mlp_act_type: str = "gelu_tanh"
mm_double_blocks_depth: int = 54
mm_single_blocks_depth: int = 0
rope_dim_list: list | tuple = field(default_factory=lambda: [16, 56, 56])
qkv_bias: bool = True
qk_norm: bool | str = True
qk_norm_type: str = "rms"
guidance_embed: bool = False
use_meanflow: bool = False
text_projection: str = "single_refiner"
use_attention_mask: bool = True
text_states_dim: int = 3584
text_states_dim_2: int | None = None
text_pool_type: str | None = None
rope_theta: float = 256.0
attn_mode: str = "flash"
attn_param: str | None = None
glyph_byT5_v2: bool = True
vision_projection: str = "linear"
vision_states_dim: int = 1152
is_reshape_temporal_channels: bool = False
use_cond_type_embedding: bool = True
ideal_resolution: str = "480p"
ideal_task: str = "i2v"
task_type: str = "i2v"
exclude_lora_layers: list[str] = field(
default_factory=lambda: ["img_in", "txt_in", "time_in", "vector_in"])
def __post_init__(self):
super().__post_init__()
# Convert HY-WorldPlay naming to FastVideo naming conventions
self.num_attention_heads: int = self.heads_num
self.attention_head_dim: int = self.hidden_size // self.heads_num
self.num_layers: int = self.mm_double_blocks_depth
self.num_single_layers: int = self.mm_single_blocks_depth
self.num_refiner_layers: int = 2 # Default for HYWorld
self.mlp_ratio: float = float(self.mlp_width_ratio)
self.text_embed_dim: int = self.text_states_dim
self.text_embed_2_dim: int = self.text_states_dim_2 if self.text_states_dim_2 else 1472
self.image_embed_dim: int = self.vision_states_dim
self.rope_axes_dim: tuple[int, ...] = tuple(self.rope_dim_list)
self.num_channels_latents: int = self.out_channels
self.target_size: int = 640
# Handle concat_condition: when True, actual in_channels = base * 2 + 1
# (base latent + condition latent + mask channel)
# config.json has base in_channels (32), but img_in needs full (65)
if self.concat_condition and self.in_channels == 32:
if self.is_reshape_temporal_channels:
self.in_channels = self.in_channels + self.in_channels // 2 + 1
else:
self.in_channels = self.in_channels * 2 + 1 # 32 * 2 + 1 = 65
# Handle patch_size (can be list/tuple or int)
if isinstance(self.patch_size, list | tuple):
self.patch_size_t: int = self.patch_size[0]
# assume square patch size for height and width
patch_size_hw: int = self.patch_size[1]
object.__setattr__(self, 'patch_size', patch_size_hw)
else:
self.patch_size_t = 1
# Convert qk_norm to string format
if isinstance(self.qk_norm, bool):
if self.qk_norm:
self.qk_norm = "rms_norm" if self.qk_norm_type == "rms" else self.qk_norm_type
else:
self.qk_norm = "none"
@dataclass
class HYWorldConfig(DiTConfig):
arch_config: DiTArchConfig = field(default_factory=HYWorldArchConfig)
prefix: str = "HYWorld"
-84
View File
@@ -1,84 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""
LTX-2 Transformer configuration for native FastVideo integration.
"""
from dataclasses import dataclass, field
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
def is_ltx2_blocks(name: str, _module) -> bool:
"""FSDP shard condition for LTX-2 transformer blocks."""
return "transformer_blocks" in name
@dataclass
class LTX2VideoArchConfig(DiTArchConfig):
"""Architecture configuration for LTX-2 video transformer."""
_fsdp_shard_conditions: list = field(
default_factory=lambda: [is_ltx2_blocks])
_compile_conditions: list = field(default_factory=lambda: [is_ltx2_blocks])
# Parameter name mapping for weight conversion (hf/comfy -> FastVideo)
param_names_mapping: dict = field(
default_factory=lambda: {
r"^model\.diffusion_model\.(.*)$": r"model.\1",
r"^diffusion_model\.(.*)$": r"model.\1",
r"^model\.(.*)$": r"model.\1",
r"^(.*)$": r"model.\1",
})
reverse_param_names_mapping: dict = field(default_factory=lambda: {})
lora_param_names_mapping: dict = field(default_factory=lambda: {})
# Core transformer settings (defaults from LTX-2 metadata)
num_attention_heads: int = 32
attention_head_dim: int = 128
num_layers: int = 48
cross_attention_dim: int = 4096
caption_channels: int = 3840
norm_eps: float = 1e-6
attention_type: str = "default"
rope_type: str = "split"
double_precision_rope: bool = True
positional_embedding_theta: float = 10000.0
positional_embedding_max_pos: list[int] = field(
default_factory=lambda: [20, 2048, 2048])
timestep_scale_multiplier: int = 1000
use_middle_indices_grid: bool = True
# Patchification (video-only path)
patch_size: tuple[int, int, int] = (1, 1, 1)
num_channels_latents: int = 128
in_channels: int | None = None
out_channels: int | None = None
# Audio defaults (reserved for joint AV ports)
audio_num_attention_heads: int = 32
audio_attention_head_dim: int = 64
audio_in_channels: int = 128
audio_out_channels: int = 128
audio_cross_attention_dim: int = 2048
audio_positional_embedding_max_pos: list[int] = field(
default_factory=lambda: [20])
av_ca_timestep_scale_multiplier: int = 1
def __post_init__(self):
super().__post_init__()
patch_volume = self.patch_size[0] * self.patch_size[
1] * self.patch_size[2]
if self.in_channels is None:
self.in_channels = self.num_channels_latents * patch_volume
if self.out_channels is None:
self.out_channels = self.in_channels
@dataclass
class LTX2VideoConfig(DiTConfig):
"""Main configuration for LTX-2 transformer."""
arch_config: DiTArchConfig = field(default_factory=LTX2VideoArchConfig)
prefix: str = "ltx2"
+2 -10
View File
@@ -1,6 +1,4 @@
from dataclasses import dataclass, field
import torch
from fastvideo.configs.models.dits.wanvideo import WanVideoArchConfig, WanVideoConfig
@@ -10,8 +8,8 @@ class MatrixGameWanVideoArchConfig(WanVideoArchConfig):
# because MatrixGame checkpoints already have patch_embedding.proj format
param_names_mapping: dict = field(
default_factory=lambda: {
r"^patch_embedding\.(?!proj\.)(.*)$":
r"patch_embedding.proj.\1",
# Removed: r"^patch_embedding\.(.*)$": r"patch_embedding.proj.\1"
# because checkpoint already has correct format
r"^condition_embedder\.text_embedder\.linear_1\.(.*)$":
r"condition_embedder.text_embedder.fc_in.\1",
r"^condition_embedder\.text_embedder\.linear_2\.(.*)$":
@@ -78,14 +76,8 @@ class MatrixGameWanVideoArchConfig(WanVideoArchConfig):
image_dim: int = 1280
def _is_transformer_block(param_name: str, module: torch.nn.Module) -> bool:
return bool("blocks" in param_name and param_name.split(".")[-1].isdigit())
@dataclass
class MatrixGameWanVideoConfig(WanVideoConfig):
arch_config: MatrixGameWanVideoArchConfig = field(
default_factory=MatrixGameWanVideoArchConfig)
prefix: str = "Wan"
_compile_conditions: list = field(
default_factory=lambda: [_is_transformer_block])
@@ -7,14 +7,11 @@ from fastvideo.configs.models.encoders.clip import (
from fastvideo.configs.models.encoders.llama import LlamaConfig
from fastvideo.configs.models.encoders.t5 import T5Config, T5LargeConfig
from fastvideo.configs.models.encoders.qwen2_5 import Qwen2_5_VLConfig
from fastvideo.configs.models.encoders.siglip import SiglipVisionConfig
from fastvideo.configs.models.encoders.reason1 import Reason1ArchConfig, Reason1Config
from fastvideo.configs.models.encoders.gemma import LTX2GemmaConfig
__all__ = [
"EncoderConfig", "TextEncoderConfig", "ImageEncoderConfig",
"BaseEncoderOutput", "CLIPTextConfig", "CLIPVisionConfig",
"WAN2_1ControlCLIPVisionConfig", "LlamaConfig", "T5Config", "T5LargeConfig",
"Qwen2_5_VLConfig", "Reason1ArchConfig", "Reason1Config", "LTX2GemmaConfig",
"SiglipVisionConfig"
"Qwen2_5_VLConfig", "Reason1ArchConfig", "Reason1Config"
]
+3 -3
View File
@@ -87,10 +87,10 @@ class CLIPVisionConfig(ImageEncoderConfig):
arch_config: ImageEncoderArchConfig = field(
default_factory=CLIPVisionArchConfig)
num_hidden_layers_override: int | None = 31
num_hidden_layers_override: int | None = None
require_post_norm: bool | None = None
enable_scale: bool = False
is_causal: bool = False
enable_scale: bool = True
is_causal: bool = True
prefix: str = "clip"
@@ -1,48 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.configs.models.encoders.base import (
TextEncoderArchConfig,
TextEncoderConfig,
)
@dataclass
class LTX2GemmaArchConfig(TextEncoderArchConfig):
architectures: list[str] = field(
default_factory=lambda: ["LTX2GemmaTextEncoderModel"])
hidden_size: int = 3840
num_hidden_layers: int = 48
num_attention_heads: int = 30
text_len: int = 1024
pad_token_id: int = 0
eos_token_id: int = 2
gemma_model_path: str = ""
gemma_dtype: str = "bfloat16"
padding_side: str = "left"
feature_extractor_in_features: int = 3840 * 49
feature_extractor_out_features: int = 3840
connector_num_attention_heads: int = 30
connector_attention_head_dim: int = 128
connector_num_layers: int = 2
connector_positional_embedding_theta: float = 10000.0
connector_positional_embedding_max_pos: list[int] = field(
default_factory=lambda: [4096])
connector_rope_type: str = "split"
connector_double_precision_rope: bool = False
connector_num_learnable_registers: int | None = 128
def __post_init__(self) -> None:
super().__post_init__()
self.tokenizer_kwargs["padding"] = "max_length"
@dataclass
class LTX2GemmaConfig(TextEncoderConfig):
arch_config: TextEncoderArchConfig = field(
default_factory=LTX2GemmaArchConfig)
prefix: str = "ltx2_gemma"
@@ -1,53 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""SigLIP vision encoder configuration for FastVideo."""
from dataclasses import dataclass, field
from fastvideo.configs.models.encoders.base import (ImageEncoderArchConfig,
ImageEncoderConfig)
@dataclass
class SiglipVisionArchConfig(ImageEncoderArchConfig):
"""Architecture configuration for SigLIP vision encoder.
Fields match the config.json from HuggingFace SigLIP checkpoints.
"""
# From config.json
architectures: list[str] = field(
default_factory=lambda: ["SiglipVisionModel"])
attention_dropout: float = 0.0
dtype: str | None = None
hidden_act: str = "gelu_pytorch_tanh"
hidden_size: int = 1152
image_size: int = 384
intermediate_size: int = 4304
layer_norm_eps: float = 1e-6
model_type: str = "siglip_vision_model"
num_attention_heads: int = 16
num_channels: int = 3
num_hidden_layers: int = 27
patch_size: int = 14
# FastVideo specific - QKV fusion mapping
stacked_params_mapping: list = field(default_factory=lambda: [
("qkv_proj", "q_proj", "q"),
("qkv_proj", "k_proj", "k"),
("qkv_proj", "v_proj", "v"),
])
@dataclass
class SiglipVisionConfig(ImageEncoderConfig):
"""Configuration for SigLIP vision encoder."""
arch_config: ImageEncoderArchConfig = field(
default_factory=SiglipVisionArchConfig)
# FastVideo specific
num_hidden_layers_override: int | None = None
require_post_norm: bool | None = None
enable_scale: bool = True
is_causal: bool = False
prefix: str = "siglip"
@@ -1,6 +0,0 @@
from fastvideo.configs.models.upsamplers.hunyuan15 import SRTo720pUpsamplerConfig, SRTo1080pUpsamplerConfig
from fastvideo.configs.models.upsamplers.base import UpsamplerConfig
__all__ = [
"SRTo720pUpsamplerConfig", "SRTo1080pUpsamplerConfig", "UpsamplerConfig"
]
@@ -1,7 +0,0 @@
from dataclasses import dataclass
from fastvideo.configs.models.base import ModelConfig
@dataclass
class UpsamplerConfig(ModelConfig):
pass
@@ -1,20 +0,0 @@
from dataclasses import dataclass
from fastvideo.configs.models.upsamplers.base import UpsamplerConfig
@dataclass
class SRTo720pUpsamplerConfig(UpsamplerConfig):
in_channels: int = 0
out_channels: int = 0
hidden_channels: int = 64
num_blocks: int = 6
global_residual: bool = False
@dataclass
class SRTo1080pUpsamplerConfig(UpsamplerConfig):
z_channels: int = 0
out_channels: int = 0
block_out_channels: tuple[int, ...] = (0, 0)
num_res_blocks: int = 2
is_residual: bool = False
@@ -2,7 +2,6 @@ from fastvideo.configs.models.vaes.cosmosvae import CosmosVAEConfig
from fastvideo.configs.models.vaes.cosmos2_5vae import Cosmos25VAEConfig
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
@@ -13,5 +12,4 @@ __all__ = [
"CosmosVAEConfig",
"Cosmos25VAEConfig",
"Hunyuan15VAEConfig",
"LTX2VAEConfig",
]
-45
View File
@@ -1,45 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""
LTX-2 VAE configuration.
"""
from dataclasses import dataclass, field
from fastvideo.configs.models.vaes.base import VAEArchConfig, VAEConfig
@dataclass
class LTX2VAEArchConfig(VAEArchConfig):
# Mirrors LTX-2 safetensors metadata config under "vae"
_class_name: str = "CausalVideoAutoencoder"
dims: int = 3
in_channels: int = 3
out_channels: int = 3
latent_channels: int = 128
encoder_blocks: list = field(default_factory=list)
decoder_blocks: list = field(default_factory=list)
patch_size: int = 4
norm_layer: str = "pixel_norm"
latent_log_var: str = "uniform"
encoder_spatial_padding_mode: str = "zeros"
decoder_spatial_padding_mode: str = "reflect"
causal_decoder: bool = False
timestep_conditioning: bool = True
use_quant_conv: bool = False
scaling_factor: float = 1.0
normalize_latent_channels: bool = False
# Match FastVideo naming for compression ratios (LTX-2 default)
temporal_compression_ratio: int = 8
spatial_compression_ratio: int = 32
@dataclass
class LTX2VAEConfig(VAEConfig):
arch_config: VAEArchConfig = field(default_factory=LTX2VAEArchConfig)
# LTX-2 tiling defaults (match ltx_core.video_vae.TilingConfig.default()).
ltx2_spatial_tile_size_in_pixels: int = 512
ltx2_spatial_tile_overlap_in_pixels: int = 64
ltx2_temporal_tile_size_in_frames: int = 64
ltx2_temporal_tile_overlap_in_frames: int = 24
+3 -5
View File
@@ -4,9 +4,8 @@ from fastvideo.configs.pipelines.cosmos import CosmosConfig
from fastvideo.configs.pipelines.cosmos2_5 import Cosmos25Config
from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
from fastvideo.configs.pipelines.hunyuan15 import Hunyuan15T2V480PConfig, Hunyuan15T2V720PConfig
from fastvideo.configs.pipelines.hyworld import HYWorldConfig
from fastvideo.configs.pipelines.ltx2 import LTX2T2VConfig
from fastvideo.registry import get_pipeline_config_cls_from_name
from fastvideo.configs.pipelines.registry import (
get_pipeline_config_cls_from_name)
from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
from fastvideo.configs.pipelines.wan import (SelfForcingWanT2V480PConfig,
WanI2V480PConfig, WanI2V720PConfig,
@@ -17,6 +16,5 @@ __all__ = [
"Hunyuan15T2V480PConfig", "Hunyuan15T2V720PConfig", "SlidingTileAttnConfig",
"WanT2V480PConfig", "WanI2V480PConfig", "WanT2V720PConfig",
"WanI2V720PConfig", "StepVideoT2VConfig", "SelfForcingWanT2V480PConfig",
"CosmosConfig", "Cosmos25Config", "LTX2T2VConfig", "HYWorldConfig",
"get_pipeline_config_cls_from_name"
"CosmosConfig", "Cosmos25Config", "get_pipeline_config_cls_from_name"
]
+5 -30
View File
@@ -8,7 +8,7 @@ from typing import Any, cast
import torch
from fastvideo.configs.models import (DiTConfig, EncoderConfig, ModelConfig,
VAEConfig, UpsamplerConfig)
VAEConfig)
from fastvideo.configs.models.encoders import BaseEncoderOutput
from fastvideo.configs.utils import update_config_from_args
from fastvideo.logger import init_logger
@@ -44,15 +44,12 @@ class PipelineConfig:
# Video generation parameters
embedded_cfg_scale: float = 6.0
flow_shift: float | None = None
flow_shift_sr: float | None = None
disable_autocast: bool = False
is_causal: bool = False
# Model configuration
dit_config: DiTConfig = field(default_factory=DiTConfig)
dit_precision: str = "bf16"
upsampler_config: UpsamplerConfig = field(default_factory=UpsamplerConfig)
upsampler_precision: str = "fp32"
# VAE configuration
vae_config: VAEConfig = field(default_factory=VAEConfig)
@@ -219,24 +216,6 @@ class PipelineConfig:
"Comma-separated list of denoising steps (e.g., '1000,757,522')",
)
# STA (Sliding Tile Attention) parameters
parser.add_argument(
f"--{prefix_with_dot}STA-mode",
type=str,
dest=f"{prefix_with_dot.replace('-', '_')}STA_mode",
default=PipelineConfig.STA_mode.value,
choices=[mode.value for mode in STA_Mode],
help=
"STA mode: STA_inference, STA_searching, STA_tuning, STA_tuning_cfg, None",
)
parser.add_argument(
f"--{prefix_with_dot}skip-time-steps",
type=int,
dest=f"{prefix_with_dot.replace('-', '_')}skip_time_steps",
default=PipelineConfig.skip_time_steps,
help="Number of time steps to warmup (full attention) for STA",
)
# Add VAE configuration arguments
from fastvideo.configs.models.vaes.base import VAEConfig
VAEConfig.add_cli_args(parser, prefix=f"{prefix_with_dot}vae-config")
@@ -266,7 +245,8 @@ class PipelineConfig:
"""
use the pipeline class setting from model_path to match the pipeline config
"""
from fastvideo.registry import get_pipeline_config_cls_from_name
from fastvideo.configs.pipelines.registry import (
get_pipeline_config_cls_from_name)
pipeline_config_cls = get_pipeline_config_cls_from_name(model_path)
return cast(PipelineConfig, pipeline_config_cls(model_path=model_path))
@@ -280,7 +260,8 @@ class PipelineConfig:
kwargs: dictionary of kwargs
config_cli_prefix: prefix of CLI arguments for this PipelineConfig instance
"""
from fastvideo.registry import get_pipeline_config_cls_from_name
from fastvideo.configs.pipelines.registry import (
get_pipeline_config_cls_from_name)
prefix_with_dot = f"{config_cli_prefix}." if (config_cli_prefix.strip()
!= "") else ""
@@ -317,12 +298,6 @@ class PipelineConfig:
# 4. Update PipelineConfig from CLI arguments if provided
kwargs[prefix_with_dot + 'model_path'] = model_path
pipeline_config.update_config_from_dict(kwargs, config_cli_prefix)
# Convert STA_mode string to enum if necessary
if isinstance(pipeline_config.STA_mode, str) and not isinstance(
pipeline_config.STA_mode, STA_Mode):
pipeline_config.STA_mode = STA_Mode(pipeline_config.STA_mode)
return pipeline_config
def check_pipeline_config(self) -> None:
+1 -28
View File
@@ -11,8 +11,7 @@ from fastvideo.configs.models.dits import HunyuanVideo15Config
from fastvideo.configs.models.encoders import (BaseEncoderOutput,
Qwen2_5_VLConfig, T5Config)
from fastvideo.configs.models.vaes import Hunyuan15VAEConfig
from fastvideo.configs.models.upsamplers import SRTo720pUpsamplerConfig, SRTo1080pUpsamplerConfig
from fastvideo.configs.pipelines.base import PipelineConfig, UpsamplerConfig
from fastvideo.configs.pipelines.base import PipelineConfig
PROMPT_TEMPLATE_TOKEN_LENGTH = 108
@@ -132,35 +131,9 @@ class Hunyuan15T2V480PConfig(PipelineConfig):
self.vae_config.load_decoder = True
@dataclass
class Hunyuan15I2V480PStepDistilledConfig(Hunyuan15T2V480PConfig):
flow_shift: int = 7
@dataclass
class Hunyuan15T2V720PConfig(Hunyuan15T2V480PConfig):
"""Base configuration for HunYuan pipeline architecture."""
# HunyuanConfig-specific parameters with defaults
flow_shift: int = 9
@dataclass
class Hunyuan15I2V720PConfig(Hunyuan15T2V720PConfig):
"""Base configuration for HunYuan pipeline architecture."""
# HunyuanConfig-specific parameters with defaults
flow_shift: int = 7
@dataclass
class Hunyuan15SR1080PConfig(Hunyuan15T2V720PConfig):
"""Base configuration for HunYuan pipeline architecture."""
# HunyuanConfig-specific parameters with defaults
flow_shift: int = 7
flow_shift_sr: int = 2
upsampler_config: tuple[UpsamplerConfig, ...] = field(
default_factory=lambda:
(SRTo720pUpsamplerConfig(), SRTo1080pUpsamplerConfig()))
upsampler_precision: str = "fp32"
-29
View File
@@ -1,29 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.configs.models import DiTConfig, EncoderConfig
from fastvideo.configs.models.dits import HYWorldConfig as HYWorldDiTConfig
from fastvideo.configs.models.encoders import SiglipVisionConfig
from fastvideo.configs.pipelines.hunyuan15 import Hunyuan15T2V480PConfig
@dataclass
class HYWorldConfig(Hunyuan15T2V480PConfig):
"""Base configuration for HYWorld pipeline architecture."""
# HYWorldConfig-specific parameters with defaults
dit_config: DiTConfig = field(default_factory=HYWorldDiTConfig)
# SigLIP image encoder for I2V
image_encoder_config: EncoderConfig = field(
default_factory=SiglipVisionConfig)
image_encoder_precision: str = "fp16"
# vae_precision: str = "fp32"
# Text encoding
text_encoder_precisions: tuple[str, ...] = field(
default_factory=lambda: ("fp16", "fp32"))
def __post_init__(self):
self.vae_config.load_encoder = True
self.vae_config.load_decoder = True
-50
View File
@@ -1,50 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from collections.abc import Callable
from dataclasses import dataclass, field
import torch
from fastvideo.configs.models import (DiTConfig, EncoderConfig, ModelConfig,
LTX2AudioDecoderConfig, LTX2VocoderConfig,
VAEConfig)
from fastvideo.configs.models.dits import LTX2VideoConfig
from fastvideo.configs.models.encoders import BaseEncoderOutput, LTX2GemmaConfig
from fastvideo.configs.models.vaes import LTX2VAEConfig
from fastvideo.configs.pipelines.base import PipelineConfig, preprocess_text
def ltx2_postprocess_text(outputs: BaseEncoderOutput) -> torch.Tensor:
return outputs.last_hidden_state
@dataclass
class LTX2T2VConfig(PipelineConfig):
"""Configuration for LTX-2 T2V pipeline."""
dit_config: DiTConfig = field(default_factory=LTX2VideoConfig)
vae_config: VAEConfig = field(default_factory=LTX2VAEConfig)
vae_tiling: bool = True
vae_sp: bool = False
text_encoder_configs: tuple[EncoderConfig, ...] = field(
default_factory=lambda: (LTX2GemmaConfig(), ))
preprocess_text_funcs: tuple[Callable[[str], str], ...] = field(
default_factory=lambda: (preprocess_text, ))
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor],
...] = field(default_factory=lambda:
(ltx2_postprocess_text, ))
dit_precision: str = "bf16"
vae_precision: str = "bf16"
text_encoder_precisions: tuple[str, ...] = field(
default_factory=lambda: ("bf16", ))
audio_decoder_config: ModelConfig = field(
default_factory=LTX2AudioDecoderConfig)
vocoder_config: ModelConfig = field(default_factory=LTX2VocoderConfig)
audio_decoder_precision: str = "bf16"
vocoder_precision: str = "bf16"
def __post_init__(self) -> None:
self.vae_config.load_encoder = False
self.vae_config.load_decoder = True
+201
View File
@@ -0,0 +1,201 @@
# SPDX-License-Identifier: Apache-2.0
"""Registry for pipeline weight-specific configurations."""
import os
from collections.abc import Callable
from fastvideo.configs.pipelines.base import PipelineConfig
from fastvideo.configs.pipelines.cosmos import CosmosConfig
from fastvideo.configs.pipelines.cosmos2_5 import Cosmos25Config
from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
from fastvideo.configs.pipelines.hunyuan15 import Hunyuan15T2V480PConfig, Hunyuan15T2V720PConfig
from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
from fastvideo.configs.pipelines.longcat import LongCatT2V480PConfig
from fastvideo.configs.pipelines.turbodiffusion import (
TurboDiffusionT2V_1_3B_Config, TurboDiffusionT2V_14B_Config,
TurboDiffusionI2V_A14B_Config)
# isort: off
from fastvideo.configs.pipelines.wan import (
FastWan2_1_T2V_480P_Config, FastWan2_2_TI2V_5B_Config,
Wan2_2_I2V_A14B_Config, Wan2_2_T2V_A14B_Config, Wan2_2_TI2V_5B_Config,
WanI2V480PConfig, WanI2V720PConfig, WanT2V480PConfig, WanT2V720PConfig,
SelfForcingWanT2V480PConfig, WANV2VConfig, SelfForcingWan2_2_T2V480PConfig,
MatrixGameI2V480PConfig)
# isort: on
from fastvideo.logger import init_logger
from fastvideo.utils import (maybe_download_model_index,
verify_model_config_and_directory)
logger = init_logger(__name__)
# Registry maps specific model weights to their config classes
PIPE_NAME_TO_CONFIG: dict[str, type[PipelineConfig]] = {
"FastVideo/FastHunyuan-diffusers": FastHunyuanConfig,
"hunyuanvideo-community/HunyuanVideo": HunyuanConfig,
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v":
Hunyuan15T2V480PConfig,
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-720p_t2v":
Hunyuan15T2V720PConfig,
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WanT2V480PConfig,
"weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers": WanI2V480PConfig,
"IRMChen/Wan2.1-Fun-1.3B-Control-Diffusers": WANV2VConfig,
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers": WanI2V480PConfig,
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers": WanI2V720PConfig,
"Wan-AI/Wan2.1-T2V-14B-Diffusers": WanT2V720PConfig,
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers": FastWan2_1_T2V_480P_Config,
"FastVideo/FastWan2.1-T2V-14B-480P-Diffusers": FastWan2_1_T2V_480P_Config,
"FastVideo/FastWan2.2-TI2V-5B-Diffusers": FastWan2_2_TI2V_5B_Config,
"FastVideo/stepvideo-t2v-diffusers": StepVideoT2VConfig,
"FastVideo/Wan2.1-VSA-T2V-14B-720P-Diffusers": WanT2V720PConfig,
"wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers": SelfForcingWanT2V480PConfig,
"rand0nmr/SFWan2.2-T2V-A14B-Diffusers": SelfForcingWan2_2_T2V480PConfig,
"FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers":
SelfForcingWan2_2_T2V480PConfig,
"Wan-AI/Wan2.2-TI2V-5B-Diffusers": Wan2_2_TI2V_5B_Config,
"Wan-AI/Wan2.2-T2V-A14B-Diffusers": Wan2_2_T2V_A14B_Config,
"Wan-AI/Wan2.2-I2V-A14B-Diffusers": Wan2_2_I2V_A14B_Config,
"nvidia/Cosmos-Predict2-2B-Video2World": CosmosConfig,
"KyleShao/Cosmos-Predict2.5-2B-Diffusers": Cosmos25Config,
"FastVideo/Matrix-Game-2.0-Base-Diffusers": MatrixGameI2V480PConfig,
"FastVideo/Matrix-Game-2.0-GTA-Diffusers": MatrixGameI2V480PConfig,
"FastVideo/Matrix-Game-2.0-TempleRun-Diffusers": MatrixGameI2V480PConfig,
# LongCat Video models
"FastVideo/LongCat-Video-T2V-Diffusers": LongCatT2V480PConfig,
"FastVideo/LongCat-Video-I2V-Diffusers": LongCatT2V480PConfig,
"FastVideo/LongCat-Video-VC-Diffusers": LongCatT2V480PConfig,
# TurboDiffusion models
"loayrashid/TurboWan2.1-T2V-1.3B-Diffusers": TurboDiffusionT2V_1_3B_Config,
"loayrashid/TurboWan2.1-T2V-14B-Diffusers": TurboDiffusionT2V_14B_Config,
"loayrashid/TurboWan2.2-I2V-A14B-Diffusers": TurboDiffusionI2V_A14B_Config,
# Add other specific weight variants
}
# For determining pipeline type from model ID
PIPELINE_DETECTOR: dict[str, Callable[[str], bool]] = {
"longcatimagetovideo":
lambda id: "longcatimagetovideo" in id.lower(),
"longcatvideocontinuation":
lambda id: "longcatvideocontinuation" in id.lower(),
"longcat":
lambda id: "longcat" in id.lower(),
"hunyuan":
lambda id: "hunyuan" in id.lower(),
"hunyuan15":
lambda id: "hunyuan15" in id.lower(),
"matrixgame":
lambda id: "matrix-game" in id.lower() or "matrixgame" in id.lower(),
"wanpipeline":
lambda id: "wanpipeline" in id.lower(),
"wanimagetovideo":
lambda id: "wanimagetovideo" in id.lower(),
"wandmdpipeline":
lambda id: "wandmdpipeline" in id.lower(),
"wancausaldmdpipeline":
lambda id: "wancausaldmdpipeline" in id.lower(),
"stepvideo":
lambda id: "stepvideo" in id.lower(),
"cosmos":
lambda id: "cosmos" in id.lower() and ("2.5" not in id.lower(
) and "2_5" not in id.lower() and "25" not in id.lower()),
"cosmos25":
lambda id: "cosmos25" in id.lower(),
"turbodiffusion":
lambda id: "turbodiffusion" in id.lower() or "turbowan" in id.lower(),
# Add other pipeline architecture detectors
}
# Fallback configs when exact match isn't found but architecture is detected
PIPELINE_FALLBACK_CONFIG: dict[str, type[PipelineConfig]] = {
"longcatimagetovideo": LongCatT2V480PConfig,
"longcatvideocontinuation": LongCatT2V480PConfig,
"longcat": LongCatT2V480PConfig,
"cosmos25": Cosmos25Config,
"hunyuan":
HunyuanConfig, # Base Hunyuan config as fallback for any Hunyuan variant
"matrixgame": MatrixGameI2V480PConfig,
"hunyuan15":
Hunyuan15T2V480PConfig, # Base Hunyuan15 config as fallback for any Hunyuan15 variant
"wanpipeline":
WanT2V480PConfig, # Base Wan config as fallback for any Wan variant
"wanimagetovideo": WanI2V480PConfig,
"wandmdpipeline": FastWan2_1_T2V_480P_Config,
"wancausaldmdpipeline": SelfForcingWanT2V480PConfig,
"stepvideo": StepVideoT2VConfig,
"turbodiffusion": TurboDiffusionT2V_1_3B_Config,
# Other fallbacks by architecture
}
def get_pipeline_config_cls_from_name(
pipeline_name_or_path: str) -> type[PipelineConfig]:
"""Get the appropriate configuration class for a given pipeline name or path.
This function implements a multi-step lookup process to find the most suitable
configuration class for a given pipeline. It follows this order:
1. Exact match in the PIPE_NAME_TO_CONFIG
2. Partial match in the PIPE_NAME_TO_CONFIG
3. Fallback to class name in the model_index.json
4. else raise an error
Args:
pipeline_name_or_path (str): The name or path of the pipeline. This can be:
- A registered model ID (e.g., "FastVideo/FastHunyuan-diffusers")
- A local path to a model directory
- A model ID that will be downloaded
Returns:
Type[PipelineConfig]: The configuration class that best matches the pipeline.
This will be one of:
- A specific weight configuration class if an exact match is found
- A fallback configuration class based on the pipeline architecture
- The base PipelineConfig class if no matches are found
Note:
- For local paths, the function will verify the model configuration
- For remote models, it will attempt to download the model index
- Warning messages are logged when falling back to less specific configurations
"""
pipeline_config_cls: type[PipelineConfig] | None = None
# First try exact match for specific weights
if pipeline_name_or_path in PIPE_NAME_TO_CONFIG:
pipeline_config_cls = PIPE_NAME_TO_CONFIG[pipeline_name_or_path]
return pipeline_config_cls
# Try partial matches (for local paths that might include the weight ID)
for registered_id, config_class in PIPE_NAME_TO_CONFIG.items():
if registered_id in pipeline_name_or_path:
pipeline_config_cls = config_class
break
# If no match, try to use the fallback config
if pipeline_config_cls is None:
if os.path.exists(pipeline_name_or_path):
config = verify_model_config_and_directory(pipeline_name_or_path)
else:
config = maybe_download_model_index(pipeline_name_or_path)
logger.warning(
"Trying to use the config from the model_index.json. FastVideo may not correctly identify the optimal config for this model in this situation."
)
pipeline_name = config["_class_name"]
# Try to determine pipeline architecture for fallback
for pipeline_type, detector in PIPELINE_DETECTOR.items():
if detector(pipeline_name.lower()):
pipeline_config_cls = PIPELINE_FALLBACK_CONFIG.get(
pipeline_type)
break
if pipeline_config_cls is not None:
logger.warning(
"No match found for pipeline %s, using fallback config %s.",
pipeline_name_or_path, pipeline_config_cls)
if pipeline_config_cls is None:
raise ValueError(
f"No match found for pipeline {pipeline_name_or_path}, please check the pipeline name or path."
)
return pipeline_config_cls
-6
View File
@@ -196,12 +196,6 @@ class SelfForcingWan2_2_T2V480PConfig(Wan2_2_T2V_A14B_Config):
# =============================================
# ============= Matrix Game ===================
# =============================================
@dataclass
class MatrixGameBaseI2V480PConfig(WanI2V480PConfig):
dit_config: DiTConfig = field(default_factory=MatrixGameWanVideoConfig)
flow_shift: float | None = 5.0
@dataclass
class MatrixGameI2V480PConfig(WanI2V480PConfig):
dit_config: DiTConfig = field(default_factory=MatrixGameWanVideoConfig)
+2 -7
View File
@@ -28,9 +28,6 @@ class SamplingParam:
keyboard_cond: Any | None = None # Shape: (B, T, K)
grid_sizes: Any | None = None # Shape: (3,) [F,H,W]
# Camera control inputs (HYWorld)
pose: str | None = None # Camera trajectory: pose string (e.g., 'w-31') or JSON file path
# Refine inputs (LongCat 480p->720p upscaling)
# Path-based refine (load stage1 video from disk, e.g. MP4)
refine_from: str | None = None # Path to stage1 video (480p output from distill)
@@ -58,13 +55,10 @@ class SamplingParam:
num_frames_round_down: bool = False # Whether to round down num_frames if it's not divisible by num_gpus
height: int = 720
width: int = 1280
height_sr: int = 1072
width_sr: int = 1920
fps: int = 24
# Denoising parameters
num_inference_steps: int = 50
num_inference_steps_sr: int = 50
guidance_scale: float = 1.0
guidance_rescale: float = 0.0
boundary_ratio: float | None = None
@@ -98,7 +92,8 @@ class SamplingParam:
@classmethod
def from_pretrained(cls, model_path: str) -> "SamplingParam":
from fastvideo.registry import get_sampling_param_cls_for_name
from fastvideo.configs.sample.registry import (
get_sampling_param_cls_for_name)
sampling_cls = get_sampling_param_cls_for_name(model_path)
if sampling_cls is not None:
sampling_param: SamplingParam = sampling_cls()
+8 -12
View File
@@ -5,19 +5,15 @@ from fastvideo.configs.sample.base import SamplingParam
@dataclass
class Cosmos25SamplingParamBase(SamplingParam):
height: int = 704
width: int = 1280
num_frames: int = 77
class Cosmos_Predict2_5_2B_Diffusers_SamplingParam(SamplingParam):
"""Defaults for Cosmos 2.5 (Predict2.5) text-to-video diffusers-format model."""
height: int = 480
width: int = 832
num_frames: int = 121
fps: int = 24
seed: int = 0
guidance_scale: float = 7.0
negative_prompt: str = (
"The video captures a series of frames showing ugly scenes, static with no motion, motion blur, "
"over-saturation, shaky footage, low resolution, grainy texture, pixelated images, poorly lit areas, "
"underexposed and overexposed scenes, poor color balance, washed out colors, choppy sequences, jerky movements, "
"low frame rate, artifacting, color banding, unnatural transitions, outdated special effects, fake elements, "
"unconvincing visuals, poorly edited content, jump cuts, visual noise, and flickering. "
"Overall, the video is of poor quality.")
# Official Cosmos2.5 sampling uses empty string as unconditional.
negative_prompt: str = ""
num_inference_steps: int = 35
-31
View File
@@ -22,39 +22,8 @@ class Hunyuan15_480P_SamplingParam(SamplingParam):
negative_prompt: str = ""
def __post_init__(self):
super().__post_init__()
self.sigmas = list(
np.linspace(1.0, 0.0, self.num_inference_steps + 1)[:-1])
@dataclass
class Hunyuan15_480P_StepDistilled_I2V_SamplingParam(
Hunyuan15_480P_SamplingParam):
num_inference_steps: int = 12
height: int = 720
width: int = 1280
guidance_scale: float = 1.0
@dataclass
class Hunyuan15_720P_SamplingParam(Hunyuan15_480P_SamplingParam):
height: int = 720
width: int = 1280
@dataclass
class Hunyuan15_720P_Distilled_I2V_SamplingParam(Hunyuan15_720P_SamplingParam):
guidance_scale: float = 1.0
@dataclass
class Hunyuan15_SR_1080P_SamplingParam(Hunyuan15_480P_SamplingParam):
height_sr: int = 1072
width_sr: int = 1920
num_inference_steps: int = 12
num_inference_steps_sr: int = 8
guidance_scale: float = 1.0
-26
View File
@@ -1,26 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.configs.sample.base import SamplingParam
import numpy as np
@dataclass
class HYWorld_SamplingParam(SamplingParam):
num_inference_steps: int = 50
num_frames: int = 125
height: int = 480
width: int = 832
fps: int = 24
# Camera trajectory: pose string (e.g., 'w-31' means generating [1 + 31] latents) or JSON file path
pose: str = 'w-31'
guidance_scale: float = 6.0
prompt_attention_mask: list = field(default_factory=list)
negative_attention_mask: list = field(default_factory=list)
sigmas: list[float] | None = field(
default_factory=lambda: list(np.linspace(1.0, 0.0, 50 + 1)[:-1]))
negative_prompt: str = ""
-20
View File
@@ -1,20 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass
from fastvideo.configs.sample.base import SamplingParam
@dataclass
class LTX2SamplingParam(SamplingParam):
"""Default sampling parameters for LTX-2 distilled T2V.
"""
seed: int = 10
num_frames: int = 121
height: int = 1024
width: int = 1536
fps: int = 24
num_inference_steps: int = 8
guidance_scale: float = 1.0
# No default negative_prompt for distilled models
negative_prompt: str = ""
+206
View File
@@ -0,0 +1,206 @@
# SPDX-License-Identifier: Apache-2.0
import os
from collections.abc import Callable
from typing import Any
from fastvideo.configs.sample.hunyuan import (FastHunyuanSamplingParam,
HunyuanSamplingParam)
from fastvideo.configs.sample.hunyuan15 import Hunyuan15_480P_SamplingParam, Hunyuan15_720P_SamplingParam
from fastvideo.configs.sample.stepvideo import StepVideoT2VSamplingParam
from fastvideo.configs.sample.cosmos import Cosmos_Predict2_2B_Video2World_SamplingParam
from fastvideo.configs.sample.cosmos2_5 import Cosmos_Predict2_5_2B_Diffusers_SamplingParam
# isort: off
from fastvideo.configs.sample.wan import (
FastWanT2V480P_SamplingParam,
Wan2_1_Fun_1_3B_InP_SamplingParam,
Wan2_2_I2V_A14B_SamplingParam,
Wan2_2_T2V_A14B_SamplingParam,
Wan2_2_TI2V_5B_SamplingParam,
WanI2V_14B_480P_SamplingParam,
WanI2V_14B_720P_SamplingParam,
WanT2V_1_3B_SamplingParam,
WanT2V_14B_SamplingParam,
Wan2_1_Fun_1_3B_Control_SamplingParam,
SelfForcingWan2_1_T2V_1_3B_480P_SamplingParam,
SelfForcingWan2_2_T2V_A14B_480P_SamplingParam,
MatrixGame2_SamplingParam,
)
from fastvideo.configs.sample.turbodiffusion import (
TurboDiffusionT2V_1_3B_SamplingParam,
TurboDiffusionT2V_14B_SamplingParam,
TurboDiffusionI2V_A14B_SamplingParam,
)
# isort: on
from fastvideo.logger import init_logger
from fastvideo.utils import (maybe_download_model_index,
verify_model_config_and_directory)
logger = init_logger(__name__)
# Registry maps specific model weights to their config classes
SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
"FastVideo/FastHunyuan-diffusers":
FastHunyuanSamplingParam,
"hunyuanvideo-community/HunyuanVideo":
HunyuanSamplingParam,
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v":
Hunyuan15_480P_SamplingParam,
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-720p_t2v":
Hunyuan15_720P_SamplingParam,
"FastVideo/stepvideo-t2v-diffusers":
StepVideoT2VSamplingParam,
# Wan2.1
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers":
WanT2V_1_3B_SamplingParam,
"Wan-AI/Wan2.1-T2V-14B-Diffusers":
WanT2V_14B_SamplingParam,
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers":
WanI2V_14B_480P_SamplingParam,
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers":
WanI2V_14B_720P_SamplingParam,
"weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers":
Wan2_1_Fun_1_3B_InP_SamplingParam,
"IRMChen/Wan2.1-Fun-1.3B-Control-Diffusers":
Wan2_1_Fun_1_3B_Control_SamplingParam,
# Wan2.2
"Wan-AI/Wan2.2-TI2V-5B-Diffusers":
Wan2_2_TI2V_5B_SamplingParam,
"FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers":
Wan2_2_TI2V_5B_SamplingParam,
"Wan-AI/Wan2.2-T2V-A14B-Diffusers":
Wan2_2_T2V_A14B_SamplingParam,
"Wan-AI/Wan2.2-I2V-A14B-Diffusers":
Wan2_2_I2V_A14B_SamplingParam,
# FastWan2.1
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers":
FastWanT2V480P_SamplingParam,
# FastWan2.2
"FastVideo/FastWan2.2-TI2V-5B-Diffusers":
Wan2_2_TI2V_5B_SamplingParam,
# Causal Self-Forcing Wan2.1
"wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers":
SelfForcingWan2_1_T2V_1_3B_480P_SamplingParam,
# Causal Self-Forcing Wan2.2
"rand0nmr/SFWan2.2-T2V-A14B-Diffusers":
SelfForcingWan2_2_T2V_A14B_480P_SamplingParam,
"FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers":
SelfForcingWan2_2_T2V_A14B_480P_SamplingParam,
# Cosmos2
"nvidia/Cosmos-Predict2-2B-Video2World":
Cosmos_Predict2_2B_Video2World_SamplingParam,
# Cosmos2.5
"KyleShao/Cosmos-Predict2.5-2B-Diffusers":
Cosmos_Predict2_5_2B_Diffusers_SamplingParam,
# MatrixGame2.0 models
"FastVideo/Matrix-Game-2.0-Base-Diffusers":
MatrixGame2_SamplingParam,
"FastVideo/Matrix-Game-2.0-GTA-Diffusers":
MatrixGame2_SamplingParam,
"FastVideo/Matrix-Game-2.0-TempleRun-Diffusers":
MatrixGame2_SamplingParam,
# TurboDiffusion models
"loayrashid/TurboWan2.1-T2V-1.3B-Diffusers":
TurboDiffusionT2V_1_3B_SamplingParam,
"loayrashid/TurboWan2.1-T2V-14B-Diffusers":
TurboDiffusionT2V_14B_SamplingParam,
"loayrashid/TurboWan2.2-I2V-A14B-Diffusers":
TurboDiffusionI2V_A14B_SamplingParam,
# Add other specific weight variants
}
# For determining pipeline type from model ID
SAMPLING_PARAM_DETECTOR: dict[str, Callable[[str], bool]] = {
"hunyuan":
lambda id: "hunyuan" in id.lower(),
"hunyuan15":
lambda id: "hunyuan15" in id.lower(),
"wanpipeline":
lambda id: "wanpipeline" in id.lower(),
"wanimagetovideo":
lambda id: "wanimagetovideo" in id.lower(),
"stepvideo":
lambda id: "stepvideo" in id.lower(),
"wandmdpipeline":
lambda id: "wandmdpipeline" in id.lower(),
"wancausaldmdpipeline":
lambda id: "wancausaldmdpipeline" in id.lower(),
"matrixgame":
lambda id: "matrixgame" in id.lower() or "matrix-game" in id.lower(),
"turbodiffusion":
lambda id: "turbodiffusion" in id.lower() or "turbowan" in id.lower(),
"cosmos25":
lambda id: "cosmos2_5" in id.lower(),
"cosmos":
lambda id: "cosmos" in id.lower() and "2_5" not in id.lower(),
# Add other pipeline architecture detectors
}
# Fallback configs when exact match isn't found but architecture is detected
SAMPLING_FALLBACK_PARAM: dict[str, Any] = {
"hunyuan":
HunyuanSamplingParam, # Base Hunyuan config as fallback for any Hunyuan variant
"hunyuan15":
Hunyuan15_480P_SamplingParam, # Base Hunyuan15 config as fallback for any Hunyuan15 variant
"wanpipeline":
WanT2V_1_3B_SamplingParam, # Base Wan config as fallback for any Wan variant
"wanimagetovideo": WanI2V_14B_480P_SamplingParam,
"wandmdpipeline": FastWanT2V480P_SamplingParam,
"wancausaldmdpipeline": SelfForcingWan2_1_T2V_1_3B_480P_SamplingParam,
"stepvideo": StepVideoT2VSamplingParam,
"matrixgame": MatrixGame2_SamplingParam,
"turbodiffusion":
TurboDiffusionT2V_1_3B_SamplingParam, # Default to T2V for fallback
"cosmos25": Cosmos_Predict2_5_2B_Diffusers_SamplingParam,
"cosmos": Cosmos_Predict2_2B_Video2World_SamplingParam,
# Other fallbacks by architecture
}
def get_sampling_param_cls_for_name(pipeline_name_or_path: str) -> Any | None:
"""Get the appropriate sampling param for specific pretrained weights."""
# First try exact match for specific weights
if pipeline_name_or_path in SAMPLING_PARAM_REGISTRY:
return SAMPLING_PARAM_REGISTRY[pipeline_name_or_path]
# Try partial matches (for local paths that might include the weight ID)
for registered_id, config_class in SAMPLING_PARAM_REGISTRY.items():
if registered_id in pipeline_name_or_path:
return config_class
matrixgame_patterns = ["Matrix-Game", "Skywork--Matrix-Game", "matrixgame"]
for pattern in matrixgame_patterns:
if pattern.lower() in pipeline_name_or_path.lower():
return MatrixGame2_SamplingParam
if os.path.exists(pipeline_name_or_path):
config = verify_model_config_and_directory(pipeline_name_or_path)
else:
config = maybe_download_model_index(pipeline_name_or_path)
pipeline_name = config["_class_name"]
# If no match, try to use the fallback config
fallback_config = None
# Try to determine pipeline architecture for fallback
for pipeline_type, detector in SAMPLING_PARAM_DETECTOR.items():
if detector(pipeline_name.lower()):
fallback_config = SAMPLING_FALLBACK_PARAM.get(pipeline_type)
break
logger.warning(
"No match found for pipeline %s, using fallback sampling param %s.",
pipeline_name_or_path, fallback_config)
return fallback_config
+3 -1
View File
@@ -8,6 +8,7 @@ from fastvideo.dataset.preprocessing_datasets import VideoCaptionMergedDataset,
from fastvideo.dataset.transform import (CenterCropResizeVideo, Normalize255,
TemporalRandomCrop)
from fastvideo.dataset.validation_dataset import ValidationDataset
from fastvideo.dataset.rl_prompt_dataset import build_rl_prompt_dataloader
def getdataset(args) -> VideoCaptionMergedDataset:
@@ -47,5 +48,6 @@ def gettextdataset(args) -> TextDataset:
__all__ = [
"build_parquet_map_style_dataloader", "ValidationDataset",
"VideoCaptionMergedDataset", "TextDataset"
"VideoCaptionMergedDataset", "TextDataset",
"build_rl_prompt_dataloader"
]
-41
View File
@@ -116,44 +116,3 @@ pyarrow_schema_text_only = pa.schema([
# --- Metadata ---
pa.field("caption", pa.string()),
])
pyarrow_schema_matrixgame = pa.schema([
pa.field("id", pa.string()),
# --- Image/Video VAE latents ---
# Tensors are stored as raw bytes with shape and dtype info for loading
pa.field("vae_latent_bytes", pa.binary()),
# e.g., [C, T, H, W] or [C, H, W]
pa.field("vae_latent_shape", pa.list_(pa.int64())),
# e.g., 'float32'
pa.field("vae_latent_dtype", pa.string()),
#I2V
pa.field("clip_feature_bytes", pa.binary()),
pa.field("clip_feature_shape", pa.list_(pa.int64())),
pa.field("clip_feature_dtype", pa.string()),
pa.field("first_frame_latent_bytes", pa.binary()),
pa.field("first_frame_latent_shape", pa.list_(pa.int64())),
pa.field("first_frame_latent_dtype", pa.string()),
# --- Action ---
pa.field("mouse_cond_bytes", pa.binary()),
pa.field("mouse_cond_shape", pa.list_(pa.int64())), # [T, 2]
pa.field("mouse_cond_dtype", pa.string()),
pa.field("keyboard_cond_bytes", pa.binary()),
pa.field("keyboard_cond_shape", pa.list_(pa.int64())), # [T, 4]
pa.field("keyboard_cond_dtype", pa.string()),
# I2V Validation
pa.field("pil_image_bytes", pa.binary()),
pa.field("pil_image_shape", pa.list_(pa.int64())),
pa.field("pil_image_dtype", pa.string()),
# --- Metadata ---
pa.field("file_name", pa.string()),
pa.field("caption", pa.string()),
pa.field("media_type", pa.string()), # 'image' or 'video'
pa.field("width", pa.int64()),
pa.field("height", pa.int64()),
# -- Video-specific (can be null/default for images) ---
# Number of frames processed (e.g., 1 for image, N for video)
pa.field("num_frames", pa.int64()),
pa.field("duration_sec", pa.float64()),
pa.field("fps", pa.float64()),
])
+83 -33
View File
@@ -40,7 +40,6 @@ class PreprocessBatch:
num_frames: int | None = None
sample_frame_index: list[int] | None = None
sample_num_frames: int | None = None
action_path: str | None = None
# Processed data
pixel_values: torch.Tensor | None = None
@@ -148,6 +147,62 @@ class DataValidationStage(DatasetFilterStage):
return batch
class ResolutionFilterStage(DatasetFilterStage):
"""Stage for filtering data items based on resolution constraints."""
def __init__(self,
max_h_div_w_ratio: float = 17 / 16,
min_h_div_w_ratio: float = 8 / 16,
max_height: int = 1024,
max_width: int = 1024):
self.max_h_div_w_ratio = max_h_div_w_ratio
self.min_h_div_w_ratio = min_h_div_w_ratio
self.max_height = max_height
self.max_width = max_width
def should_keep(self, batch: PreprocessBatch, **kwargs) -> bool:
"""
Check if data item passes resolution filtering.
Args:
batch: Dataset batch with resolution information
Returns:
True if passes filter, False otherwise
"""
# Only apply to videos
if not batch.is_video:
return True
if batch.resolution is None:
return False
height = batch.resolution.get("height", None)
width = batch.resolution.get("width", None)
if height is None or width is None:
return False
# Check aspect ratio
aspect = self.max_height / self.max_width
hw_aspect_thr = 1.5
return self.filter_resolution(
height,
width,
max_h_div_w_ratio=hw_aspect_thr * aspect,
min_h_div_w_ratio=1 / hw_aspect_thr * aspect,
)
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
"""Process does nothing for resolution filtering - filtering is handled by should_keep."""
return batch
def filter_resolution(self, h: int, w: int, max_h_div_w_ratio: float,
min_h_div_w_ratio: float) -> bool:
"""Filter based on height/width ratio."""
return h / w <= max_h_div_w_ratio and h / w >= min_h_div_w_ratio
class FrameSamplingStage(DatasetFilterStage):
"""Stage for temporal frame sampling and indexing."""
@@ -272,6 +327,11 @@ class VideoTransformStage(DatasetStage):
video = rearrange(video, "t c h w -> c t h w")
video = video.to(torch.uint8)
h, w = video.shape[-2:]
assert (
h / w <= 17 / 16 and h / w >= 8 / 16
), f"Only videos with a ratio (h/w) less than 17/16 and more than 8/16 are supported. But video ({batch.path}) found ratio is {round(h / w, 2)} with the shape of {video.shape}"
video = video.float() / 127.5 - 1.0
batch.pixel_values = video
return batch
@@ -305,11 +365,9 @@ class ImageTransformStage(DatasetStage):
image = self.transform_topcrop(image)
elif self.transform is not None:
image = self.transform(image)
image = image.float() / 127.5 - 1.0
else:
image = image.float() / 127.5 - 1.0
image = image.transpose(0, 1) # [1 C H W] -> [C 1 H W]
image = image.float() / 127.5 - 1.0
batch.pixel_values = image
return batch
@@ -412,11 +470,8 @@ class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
# Initialize tokenizer
tokenizer_path = os.path.join(args.model_path, "tokenizer")
if os.path.exists(tokenizer_path):
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path,
cache_dir=args.cache_dir)
else:
tokenizer = None
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path,
cache_dir=args.cache_dir)
# Initialize processing stages
self._init_stages(args, transform, transform_topcrop, tokenizer)
@@ -428,6 +483,8 @@ class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
tokenizer) -> None:
"""Initialize all processing stages."""
self.validation_stage = DataValidationStage()
self.resolution_filter_stage = ResolutionFilterStage(
max_height=args.max_height, max_width=args.max_width)
self.frame_sampling_stage = FrameSamplingStage(
num_frames=args.num_frames,
train_fps=args.train_fps,
@@ -438,14 +495,11 @@ class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
self.video_transform_stage = VideoTransformStage(transform)
self.image_transform_stage = ImageTransformStage(
transform, transform_topcrop)
if tokenizer is not None:
self.text_encoding_stage = TextEncodingStage(
tokenizer=tokenizer,
text_max_length=args.text_max_length,
cfg_rate=args.training_cfg_rate,
seed=self.seed)
else:
self.text_encoding_stage = None
self.text_encoding_stage = TextEncodingStage(
tokenizer=tokenizer,
text_max_length=args.text_max_length,
cfg_rate=args.training_cfg_rate,
seed=self.seed)
def _load_raw_data(self) -> list[dict]:
"""Load raw data from JSON files."""
@@ -467,8 +521,6 @@ class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
# Update paths with folder prefix
for item in data_items:
item["path"] = opj(folder, item["path"])
if "action_path" in item and item["action_path"]:
item["action_path"] = opj(folder, item["action_path"])
return data_items
@@ -480,6 +532,7 @@ class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
# Initialize counters
filter_counts = {
"validation_failed": 0,
"resolution_failed": 0,
"frame_sampling_failed": 0
}
sample_num_frames: list[int] = []
@@ -489,8 +542,7 @@ class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
cap=item["cap"],
resolution=item.get("resolution"),
fps=item.get("fps"),
duration=item.get("duration"),
action_path=item.get("action_path"))
duration=item.get("duration"))
# Apply filtering stages
if not self._apply_filter_stages(batch, filter_counts):
@@ -515,6 +567,10 @@ class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
filter_counts["validation_failed"] += 1
return False
if not self.resolution_filter_stage.should_keep(batch):
filter_counts["resolution_failed"] += 1
return False
if not self.frame_sampling_stage.should_keep(batch):
filter_counts["frame_sampling_failed"] += 1
return False
@@ -526,9 +582,10 @@ class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
after_count: int):
"""Log filtering statistics."""
logger.info(
"validation_failed: %d, frame_sampling_failed: %d, "
"validation_failed: %d, resolution_failed: %d, frame_sampling_failed: %d, "
"Counter(sample_num_frames): %s, before filter: %d, after filter: %d",
filter_counts['validation_failed'],
filter_counts['resolution_failed'],
filter_counts['frame_sampling_failed'], Counter(sample_num_frames),
before_count, after_count)
@@ -547,27 +604,20 @@ class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
# Apply transformation stages
batch = self.video_transform_stage.process(batch)
batch = self.image_transform_stage.process(batch)
if self.text_encoding_stage is not None:
batch = self.text_encoding_stage.process(batch)
batch = self.text_encoding_stage.process(batch)
# Build result dictionary
result = {
"pixel_values": batch.pixel_values,
"text": batch.text,
"input_ids": batch.input_ids,
"cond_mask": batch.cond_mask,
"path": batch.path,
}
if batch.text is not None:
result["text"] = batch.text
result["input_ids"] = batch.input_ids
result["cond_mask"] = batch.cond_mask
# Add video-specific fields
if batch.is_video:
result.update({"fps": batch.fps, "duration": batch.duration})
# Add action_path
if batch.action_path:
result["action_path"] = batch.action_path
return result
+174
View File
@@ -0,0 +1,174 @@
# SPDX-License-Identifier: Apache-2.0
import torch
from torch.utils.data import Dataset, DataLoader, Sampler
import json
import os
class TextPromptDataset(Dataset):
"""Dataset for loading text prompts from a simple text file (one prompt per line)."""
def __init__(self, dataset, split='train'):
self.file_path = os.path.join(dataset, f'{split}.txt')
with open(self.file_path, 'r') as f:
self.prompts = [line.strip() for line in f.readlines()]
def __len__(self):
return len(self.prompts)
def __getitem__(self, idx):
return {"prompt": self.prompts[idx], "metadata": {}}
@staticmethod
def collate_fn(examples):
prompts = [example["prompt"] for example in examples]
metadatas = [example["metadata"] for example in examples]
return prompts, metadatas
class GenevalPromptDataset(Dataset):
"""Dataset for loading prompts with metadata from JSONL files (e.g., GenEval format)."""
def __init__(self, dataset, split='train'):
self.file_path = os.path.join(dataset, f'{split}_metadata.jsonl')
with open(self.file_path, 'r', encoding='utf-8') as f:
self.metadatas = [json.loads(line) for line in f]
self.prompts = [item['prompt'] for item in self.metadatas]
def __len__(self):
return len(self.prompts)
def __getitem__(self, idx):
return {"prompt": self.prompts[idx], "metadata": self.metadatas[idx]}
@staticmethod
def collate_fn(examples):
prompts = [example["prompt"] for example in examples]
metadatas = [example["metadata"] for example in examples]
return prompts, metadatas
class KRepeatSampler(Sampler):
"""Sampler that repeats each sample k times, ensuring synchronized random selection. For single-node training, set num_replicas=1 and rank=0."""
def __init__(self, dataset, batch_size, k, num_replicas, rank, seed=0):
self.dataset = dataset
self.batch_size = batch_size # Batch size per GPU/card
self.k = k # Number of repetitions per sample
self.num_replicas = num_replicas # Total number of GPUs/cards
self.rank = rank # Current GPU/card rank
self.seed = seed # Random seed for synchronization
# Calculate the number of unique samples needed for each iteration
self.total_samples = self.num_replicas * self.batch_size
assert self.total_samples % self.k == 0, f"k can not div n*b, k{k}-num_replicas{num_replicas}-batch_size{batch_size}"
self.m = self.total_samples // self.k # different number of samples
self.step = 0
def __iter__(self):
while True:
# Generate a deterministic random sequence to ensure all cards are synchronized
g = torch.Generator()
g.manual_seed(self.seed + self.step)
# Randomly select m unique samples
indices = torch.randperm(len(self.dataset), generator=g)[:self.m].tolist()
# Repeat each sample k times to generate a total of n*b samples
repeated_indices = [idx for idx in indices for _ in range(self.k)]
# Shuffle the order to ensure even distribution
shuffled_indices = torch.randperm(len(repeated_indices), generator=g).tolist()
shuffled_samples = [repeated_indices[i] for i in shuffled_indices]
# Split samples among all cards
per_card_samples = []
for i in range(self.num_replicas):
start = i * self.batch_size
end = start + self.batch_size
per_card_samples.append(shuffled_samples[start:end])
# Return the sample indices for the current card
yield per_card_samples[self.rank]
def __len__(self):
return len(self.dataset) // self.batch_size
def set_step(self, step):
"""Used to synchronize the random state for different epochs."""
self.step = step
def build_rl_prompt_dataloader(
dataset_path: str,
dataset_type: str = "text",
split: str = "train",
train_batch_size: int = 8,
test_batch_size: int = 8,
k: int = 1,
seed: int = 42,
train_num_workers: int = 1,
test_num_workers: int = 8,
num_replicas: int = 1,
rank: int = 0,
) -> tuple[DataLoader, DataLoader]:
"""
Factory function to create train and test dataloaders for RL prompt datasets.
Args:
dataset_path: Path to dataset directory
dataset_type: "text" for TextPromptDataset or "geneval" for GenevalPromptDataset
split: Dataset split ("train" or "test")
train_batch_size: Batch size per GPU for training
test_batch_size: Batch size for testing
k: Number of times to repeat each sample (num_image_per_prompt)
seed: Random seed for sampler synchronization
train_num_workers: Number of workers for training dataloader
test_num_workers: Number of workers for test dataloader
num_replicas: Number of replicas (default 1 for single-node)
rank: Rank of current process (default 0 for single-node)
Returns:
Tuple of (train_dataloader, test_dataloader)
"""
# Create datasets based on type
if dataset_type == "text":
train_dataset = TextPromptDataset(dataset_path, 'train')
test_dataset = TextPromptDataset(dataset_path, 'test')
collate_fn = TextPromptDataset.collate_fn
elif dataset_type == "geneval":
train_dataset = GenevalPromptDataset(dataset_path, 'train')
test_dataset = GenevalPromptDataset(dataset_path, 'test')
collate_fn = GenevalPromptDataset.collate_fn
else:
raise ValueError(f"Unknown dataset_type: {dataset_type}. Must be 'text' or 'geneval'")
# Create infinite-loop training sampler
train_sampler = KRepeatSampler(
dataset=train_dataset,
batch_size=train_batch_size,
k=k,
num_replicas=num_replicas,
rank=rank,
seed=seed
)
# Create training dataloader with batch_sampler (infinite loop)
train_dataloader = DataLoader(
train_dataset,
batch_sampler=train_sampler,
num_workers=train_num_workers,
collate_fn=collate_fn,
)
# Create standard test dataloader
test_dataloader = DataLoader(
test_dataset,
batch_size=test_batch_size,
collate_fn=collate_fn,
shuffle=False,
num_workers=test_num_workers,
)
return train_dataloader, test_dataloader, train_dataset, test_dataset
-1
View File
@@ -18,7 +18,6 @@ __all__ = [
"cleanup_dist_env_and_memory",
"model_parallel_is_initialized",
"maybe_init_distributed_environment_and_model_parallel",
"warmup_sequence_parallel_communication",
# World group
"get_world_group",
+1 -91
View File
@@ -7,17 +7,10 @@ import torch.distributed
from fastvideo.distributed.parallel_state import (get_sp_group,
get_sp_parallel_rank,
get_sp_world_size,
get_tp_group,
model_parallel_is_initialized)
get_tp_group)
from fastvideo.distributed.utils import (unpad_sequence_tensor,
compute_padding_for_sp,
pad_sequence_tensor)
from fastvideo.logger import init_logger
logger = init_logger(__name__)
# Track if SP communication has been warmed up
_sp_warmup_done = False
def tensor_model_parallel_all_reduce(input_: torch.Tensor) -> torch.Tensor:
@@ -108,86 +101,3 @@ def sequence_model_parallel_shard(input_: torch.Tensor,
input_ = input_.movedim(0, dim)
return input_, original_seq_len
def warmup_sequence_parallel_communication(
device: torch.device | None = None) -> None:
"""Warmup NCCL communicators for sequence parallel all-to-all operations.
The first NCCL collective operation is slow due to lazy communicator
initialization. This function runs dummy all-to-all operations to
trigger the initialization upfront, before the first real forward pass.
Args:
device: Device to use for warmup tensors. If None, uses CUDA device 0.
"""
global _sp_warmup_done
if _sp_warmup_done:
return
if not model_parallel_is_initialized():
return
sp_world_size = get_sp_world_size()
if sp_world_size <= 1:
_sp_warmup_done = True
return
if device is None:
device = torch.device("cuda")
logger.info("Warming up sequence parallel communication (SP=%d)...",
sp_world_size)
# Use small but representative tensor shapes for warmup
# Shape: [batch, seq_len, num_heads, head_dim]
# The all-to-all patterns used in attention:
# 1. scatter_dim=2 (heads), gather_dim=1 (seq) - before attention
# 2. scatter_dim=1 (seq), gather_dim=2 (heads) - after attention
batch_size = 1
seq_len_per_rank = 16 # Small sequence per rank
num_heads = sp_world_size * 4 # Must be divisible by sp_world_size
head_dim = 64
# Create dummy tensor for warmup
dummy = torch.zeros(batch_size,
seq_len_per_rank,
num_heads,
head_dim,
device=device,
dtype=torch.bfloat16)
# Warmup pattern 1: scatter heads, gather sequence (before attention)
_ = sequence_model_parallel_all_to_all_4D(dummy,
scatter_dim=2,
gather_dim=1)
# Warmup pattern 2: scatter sequence, gather heads (after attention)
dummy2 = torch.zeros(batch_size,
seq_len_per_rank * sp_world_size,
num_heads // sp_world_size,
head_dim,
device=device,
dtype=torch.bfloat16)
_ = sequence_model_parallel_all_to_all_4D(dummy2,
scatter_dim=1,
gather_dim=2)
# Warmup all-gather (used for replicated tokens)
dummy3 = torch.zeros(batch_size,
8,
num_heads // sp_world_size,
head_dim,
device=device,
dtype=torch.bfloat16)
_ = sequence_model_parallel_all_gather(dummy3, dim=2)
# Synchronize to ensure warmup completes
torch.cuda.synchronize(device)
# Clean up
del dummy, dummy2, dummy3
_sp_warmup_done = True
logger.info("Sequence parallel communication warmup complete.")
@@ -57,9 +57,6 @@ class DistributedAutograd:
ctx.dim = dim
ctx.input_shape = input_.shape
# NCCL all_gather_into_tensor requires contiguous tensors.
if not input_.is_contiguous():
input_ = input_.contiguous()
input_size = input_.size()
output_size = (input_size[0] * world_size, ) + input_size[1:]
output_tensor = torch.empty(output_size,
+2 -141
View File
@@ -9,7 +9,6 @@ diffusion models.
import math
import os
import re
import threading
import time
from copy import deepcopy
from typing import Any
@@ -19,8 +18,6 @@ import numpy as np
import torch
import torchvision
from einops import rearrange
import shutil
import tempfile
from fastvideo.configs.sample import SamplingParam
from fastvideo.fastvideo_args import FastVideoArgs
@@ -32,21 +29,6 @@ from fastvideo.worker.executor import Executor
logger = init_logger(__name__)
def _infer_latent_batch_size(batch: ForwardBatch) -> int:
if isinstance(batch.prompt, list):
latent_batch_size = len(batch.prompt)
elif batch.prompt is not None:
latent_batch_size = 1
elif batch.prompt_embeds is not None and len(batch.prompt_embeds) > 0:
latent_batch_size = batch.prompt_embeds[0].shape[0]
else:
raise ValueError(
"Cannot infer batch size from batch; no prompt or prompt_embeds found"
)
latent_batch_size *= batch.num_videos_per_prompt
return latent_batch_size
class VideoGenerator:
"""
A unified class for generating videos using diffusion models.
@@ -388,31 +370,8 @@ class VideoGenerator:
# Run inference
start_time = time.perf_counter()
# Execute forward pass in a new thread for non-blocking tensor allocation
result_container = {}
def execute_forward_thread():
result_container['output_batch'] = self.executor.execute_forward(
batch, fastvideo_args)
thread = threading.Thread(target=execute_forward_thread)
thread.start()
latent_batch_size = _infer_latent_batch_size(batch)
samples = torch.empty((latent_batch_size, 3, sampling_param.num_frames,
sampling_param.height, sampling_param.width),
device='cpu',
pin_memory=fastvideo_args.pin_cpu_memory)
thread.join()
output_batch = result_container['output_batch']
if output_batch.output.shape == samples.shape:
samples.copy_(output_batch.output)
else:
logger.warning(
"Output shape %s does not match expected shape %s; use slow path",
output_batch.output.shape, samples.shape)
samples = output_batch.output.cpu()
output_batch = self.executor.execute_forward(batch, fastvideo_args)
samples = output_batch.output
logging_info = output_batch.logging_info
gen_time = time.perf_counter() - start_time
@@ -430,11 +389,6 @@ class VideoGenerator:
if batch.save_video:
imageio.mimsave(output_path, frames, fps=batch.fps, format="mp4")
logger.info("Saved video to %s", output_path)
audio = output_batch.extra.get("audio")
audio_sample_rate = output_batch.extra.get("audio_sample_rate")
if (audio is not None and audio_sample_rate is not None and
not self._mux_audio(output_path, audio, audio_sample_rate)):
logger.warning("Audio mux failed; saved video without audio.")
if batch.return_frames:
return frames
@@ -442,7 +396,6 @@ class VideoGenerator:
return {
"samples": samples,
"frames": frames,
"audio": output_batch.extra.get("audio"),
"prompts": prompt,
"size": (target_height, target_width, batch.num_frames),
"generation_time": gen_time,
@@ -452,98 +405,6 @@ class VideoGenerator:
"trajectory_decoded": output_batch.trajectory_decoded,
}
@staticmethod
def _mux_audio(
video_path: str,
audio: torch.Tensor | np.ndarray,
sample_rate: int,
) -> bool:
"""Mux audio into video using PyAV."""
try:
import av
except ImportError:
logger.warning("PyAV not installed; cannot mux audio. "
"Install with: pip install av")
return False
if torch.is_tensor(audio):
audio_np = audio.detach().cpu().float().numpy()
else:
audio_np = np.asarray(audio, dtype=np.float32)
if audio_np.ndim == 1:
audio_np = audio_np[:, None]
elif audio_np.ndim == 2:
if audio_np.shape[0] <= 8 and audio_np.shape[1] > audio_np.shape[0]:
audio_np = audio_np.T
else:
logger.warning("Unexpected audio shape %s; skipping mux.",
audio_np.shape)
return False
audio_np = np.clip(audio_np, -1.0, 1.0)
audio_int16 = (audio_np * 32767.0).astype(np.int16)
num_channels = audio_int16.shape[1]
layout = "stereo" if num_channels == 2 else "mono"
try:
import wave
with tempfile.TemporaryDirectory() as tmpdir:
out_path = os.path.join(tmpdir, "muxed.mp4")
wav_path = os.path.join(tmpdir, "audio.wav")
# Write audio to WAV file
with wave.open(wav_path, "wb") as wav_file:
wav_file.setnchannels(num_channels)
wav_file.setsampwidth(2)
wav_file.setframerate(sample_rate)
wav_file.writeframes(audio_int16.tobytes())
# Open input video and audio
input_video = av.open(video_path)
input_audio = av.open(wav_path)
# Create output with both streams
output = av.open(out_path, mode="w")
# Add video stream (copy codec from input)
in_video_stream = input_video.streams.video[0]
out_video_stream = output.add_stream(
codec_name=in_video_stream.codec_context.name,
rate=in_video_stream.average_rate,
)
out_video_stream.width = in_video_stream.width
out_video_stream.height = in_video_stream.height
out_video_stream.pix_fmt = in_video_stream.pix_fmt
# Add audio stream (AAC)
out_audio_stream = output.add_stream("aac", rate=sample_rate)
out_audio_stream.layout = layout
# Remux video (decode and re-encode to be safe)
for frame in input_video.decode(video=0):
for packet in out_video_stream.encode(frame):
output.mux(packet)
for packet in out_video_stream.encode():
output.mux(packet)
# Encode audio
for frame in input_audio.decode(audio=0):
frame.pts = None # Let encoder assign PTS
for packet in out_audio_stream.encode(frame):
output.mux(packet)
for packet in out_audio_stream.encode():
output.mux(packet)
input_video.close()
input_audio.close()
output.close()
shutil.move(out_path, video_path)
return True
except Exception as e:
logger.warning("Audio mux failed: %s", e)
return False
def set_lora_adapter(self,
lora_nickname: str,
lora_path: str | None = None) -> None:
+328 -89
View File
@@ -10,7 +10,7 @@ from enum import Enum
from typing import Any, TYPE_CHECKING
from fastvideo.configs.configs import PreprocessConfig
from fastvideo.configs.pipelines.base import PipelineConfig
from fastvideo.configs.pipelines.base import PipelineConfig, STA_Mode
from fastvideo.configs.utils import clean_cli_args
from fastvideo.layers.quantization import QUANTIZATION_METHODS, QuantizationMethods
from fastvideo.logger import init_logger
@@ -141,6 +141,8 @@ class FastVideoArgs:
# STA (Sliding Tile Attention) parameters
mask_strategy_file_path: str | None = None
STA_mode: STA_Mode = STA_Mode.STA_INFERENCE
skip_time_steps: int = 15
# Compilation
enable_torch_compile: bool = False
@@ -164,20 +166,11 @@ class FastVideoArgs:
# Prompt text file for batch processing
prompt_txt: str | None = None
# LTX-2 VAE tiling overrides
ltx2_vae_tiling: bool | None = None
ltx2_vae_spatial_tile_size_in_pixels: int | None = None
ltx2_vae_spatial_tile_overlap_in_pixels: int | None = None
ltx2_vae_temporal_tile_size_in_frames: int | None = None
ltx2_vae_temporal_tile_overlap_in_frames: int | None = None
ltx2_initial_latent_path: str | None = None
# model paths for correct deallocation
model_paths: dict[str, str] = field(default_factory=dict)
model_loaded: dict[str, bool] = field(default_factory=lambda: {
"transformer": True,
"vae": True,
"upsampler": True,
})
override_text_encoder_safetensors: str | None = None # path to safetensors file for text encoder override
@@ -210,44 +203,8 @@ class FastVideoArgs:
logger.error("Failed to load V-MoBA config from %s: %s",
self.moba_config_path, e)
raise
self._apply_ltx2_vae_overrides()
self.check_fastvideo_args()
def _apply_ltx2_vae_overrides(self) -> None:
if self.pipeline_config is None:
return
vae_config = self.pipeline_config.vae_config
has_any = any(value is not None for value in (
self.ltx2_vae_spatial_tile_size_in_pixels,
self.ltx2_vae_spatial_tile_overlap_in_pixels,
self.ltx2_vae_temporal_tile_size_in_frames,
self.ltx2_vae_temporal_tile_overlap_in_frames,
))
if self.ltx2_vae_tiling is not None and hasattr(self.pipeline_config,
"vae_tiling"):
self.pipeline_config.vae_tiling = self.ltx2_vae_tiling
elif has_any and hasattr(self.pipeline_config, "vae_tiling"):
self.pipeline_config.vae_tiling = True
if hasattr(vae_config, "ltx2_spatial_tile_size_in_pixels"
) and self.ltx2_vae_spatial_tile_size_in_pixels is not None:
vae_config.ltx2_spatial_tile_size_in_pixels = (
self.ltx2_vae_spatial_tile_size_in_pixels)
if hasattr(
vae_config, "ltx2_spatial_tile_overlap_in_pixels"
) and self.ltx2_vae_spatial_tile_overlap_in_pixels is not None:
vae_config.ltx2_spatial_tile_overlap_in_pixels = (
self.ltx2_vae_spatial_tile_overlap_in_pixels)
if hasattr(vae_config, "ltx2_temporal_tile_size_in_frames"
) and self.ltx2_vae_temporal_tile_size_in_frames is not None:
vae_config.ltx2_temporal_tile_size_in_frames = (
self.ltx2_vae_temporal_tile_size_in_frames)
if hasattr(
vae_config, "ltx2_temporal_tile_overlap_in_frames"
) and self.ltx2_vae_temporal_tile_overlap_in_frames is not None:
vae_config.ltx2_temporal_tile_overlap_in_frames = (
self.ltx2_vae_temporal_tile_overlap_in_frames)
@staticmethod
def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser:
# Model and path configuration
@@ -368,44 +325,6 @@ class FastVideoArgs:
"Path to a text file containing prompts (one per line) for batch processing",
)
# LTX-2 VAE tiling overrides
parser.add_argument(
"--ltx2-vae-tiling",
action=StoreBoolean,
default=FastVideoArgs.ltx2_vae_tiling,
help="Enable LTX-2 VAE tiling overrides.",
)
parser.add_argument(
"--ltx2-vae-spatial-tile-size-in-pixels",
type=int,
default=FastVideoArgs.ltx2_vae_spatial_tile_size_in_pixels,
help="LTX-2 VAE spatial tile size in pixels.",
)
parser.add_argument(
"--ltx2-vae-spatial-tile-overlap-in-pixels",
type=int,
default=FastVideoArgs.ltx2_vae_spatial_tile_overlap_in_pixels,
help="LTX-2 VAE spatial tile overlap in pixels.",
)
parser.add_argument(
"--ltx2-vae-temporal-tile-size-in-frames",
type=int,
default=FastVideoArgs.ltx2_vae_temporal_tile_size_in_frames,
help="LTX-2 VAE temporal tile size in frames.",
)
parser.add_argument(
"--ltx2-vae-temporal-tile-overlap-in-frames",
type=int,
default=FastVideoArgs.ltx2_vae_temporal_tile_overlap_in_frames,
help="LTX-2 VAE temporal tile overlap in frames.",
)
parser.add_argument(
"--ltx2-initial-latent-path",
type=str,
default=FastVideoArgs.ltx2_initial_latent_path,
help="Path to load/save a precomputed LTX-2 initial latent.",
)
# LoRA parameters (inference-time adapter loading)
parser.add_argument(
"--lora-path",
@@ -463,6 +382,20 @@ class FastVideoArgs:
)
# STA (Sliding Tile Attention) parameters
parser.add_argument(
"--STA-mode",
type=str,
default=FastVideoArgs.STA_mode.value,
choices=[mode.value for mode in STA_Mode],
help=
"STA mode contains STA_inference, STA_searching, STA_tuning, STA_tuning_cfg, None",
)
parser.add_argument(
"--skip-time-steps",
type=int,
default=FastVideoArgs.skip_time_steps,
help="Number of time steps to warmup (full attention) for STA",
)
parser.add_argument(
"--mask-strategy-file-path",
type=str,
@@ -807,6 +740,271 @@ def get_current_fastvideo_args() -> FastVideoArgs:
return _current_fastvideo_args
@dataclasses.dataclass
class RLArgs:
"""
Reinforcement Learning (RL) specific arguments
"""
# ============================================================================
# SHARED RL CONFIGURATION
rl_mode: bool = False # Enable RL training mode
rl_algorithm: str = "grpo" # RL algorithm to use: "grpo", "ppo", "dpo"
# Trajectory collection
num_rollouts: int = 4 # Number of rollouts to collect per training step
rollout_steps: str = "20,30" # Random intermediate steps for sampling (comma-separated)
noise_injection_min: int = 10 # Minimum timestep for noise injection
noise_injection_max: int = 40 # Maximum timestep for noise injection
use_sde_sampling: bool = True # Use SDE sampling (Flow-GRPO-Fast)
num_denoising_steps: int = 2 # Number of denoising steps per trajectory (1-2 for fast)
# Advantage estimation
gamma: float = 0.99 # Discount factor for returns
lambda_param: float = 0.95 # GAE lambda parameter
use_gae: bool = True # Use Generalized Advantage Estimation
normalize_advantages: bool = True # Normalize advantages before policy update
# Reward models
reward_models: dict[str, float] = field(default_factory=lambda: {"dummy": 1.0}) # reward models (names, weight)
value_model_path: str = "" # Path to value model (can be empty to train from scratch)
value_model_share_backbone: bool = False # Share transformer backbone between policy and value
# Training schedule
warmup_steps: int = 1000 # Collect SFT-style data before starting RL
collect_on_policy: bool = True # Collect fresh rollouts each step (on-policy)
timestep_fraction: float = 0.99 # Fraction of timesteps to train on
num_inner_epochs: int = 1 # Number of inner epochs per outer epoch
# KL regularization
kl_beta: float = 0.004 # KL loss coefficient (GRPO uses KL loss, DPO uses larger beta)
kl_reward: float = 0.0 # KL reward coefficient (alternative to KL loss, typically 0)
# SFT integration
sft_weight: float = 0.0 # SFT loss weight for supervised learning in RL training
sft_batch_size: int = 3 # Batch size for SFT data
# CFG
guidance_scale = 1.0 # use guidance_scale > 1.0 to enable CFG
# Statistics tracking
global_std: bool = False # Use global std across all samples vs per-group std
per_prompt_stat_tracking: bool = True # Track statistics per prompt
# Training options
use_diffusion_loss: bool = True # Use diffusion loss in training
# ============================================================================
# GRPO-SPECIFIC CONFIGURATION
# Policy optimization
grpo_policy_clip_range: float = 0.001 # PPO-style clipping range for policy ratio
grpo_value_clip_range: float = 0.2 # Value function clipping range
grpo_num_policy_epochs: int = 1 # Number of policy update epochs (GRPO typically uses 1)
grpo_num_value_epochs: int = 1 # Number of value function update epochs
grpo_target_kl: float = 0.01 # Target KL divergence for early stopping
grpo_entropy_coef: float = 0.0 # Entropy coefficient for exploration
grpo_value_loss_coef: float = 0.5 # Value loss coefficient
# GRPO-Guard safety mechanisms
grpo_use_grpo_guard: bool = True # Enable GRPO-Guard safety mechanisms
grpo_ratio_norm_correction: bool = True # RatioNorm: correct importance ratio bias
grpo_gradient_reweighting: bool = True # Reweight gradients across denoising steps
grpo_max_importance_ratio: float = 10.0 # Clip importance ratios above this value
# ============================================================================
# DPO-SPECIFIC CONFIGURATION
dpo_beta: float = 100.0 # DPO regularization parameter (typically much larger than GRPO beta)
dpo_ref_update_step: int = 10000000 # Reference model update frequency for OnlineDPO
dpo_label_smoothing: float = 0.0 # Label smoothing for DPO loss
@staticmethod
def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser:
"""Add RL-specific CLI arguments to the parser."""
# RL (Reinforcement Learning) arguments
parser.add_argument("--rl-mode",
action=StoreBoolean,
help="Enable RL training mode")
parser.add_argument("--rl-algorithm",
type=str,
default=RLArgs.rl_algorithm,
choices=["grpo", "ppo", "dpo"],
help="RL algorithm to use (grpo, ppo, dpo)")
# Trajectory collection (Flow-GRPO-Fast)
parser.add_argument("--rl-num-rollouts",
type=int,
default=RLArgs.num_rollouts,
help="Number of rollouts to collect per training step")
parser.add_argument("--rl-rollout-steps",
type=str,
default=RLArgs.rollout_steps,
help="Random intermediate steps for sampling (comma-separated)")
parser.add_argument("--rl-noise-injection-min",
type=int,
default=RLArgs.noise_injection_min,
help="Minimum timestep for noise injection")
parser.add_argument("--rl-noise-injection-max",
type=int,
default=RLArgs.noise_injection_max,
help="Maximum timestep for noise injection")
parser.add_argument("--rl-use-sde-sampling",
action=StoreBoolean,
help="Use SDE sampling (Flow-GRPO-Fast)")
parser.add_argument("--rl-num-denoising-steps",
type=int,
default=RLArgs.num_denoising_steps,
help="Number of denoising steps per trajectory (1-2 for fast)")
# Advantage estimation
parser.add_argument("--rl-gamma",
type=float,
default=RLArgs.gamma,
help="Discount factor for returns")
parser.add_argument("--rl-lambda",
type=float,
default=RLArgs.lambda_param,
help="GAE lambda parameter")
parser.add_argument("--rl-use-gae",
action=StoreBoolean,
help="Use Generalized Advantage Estimation")
parser.add_argument("--rl-normalize-advantages",
action=StoreBoolean,
help="Normalize advantages before policy update")
# Policy optimization (GRPO/PPO)
parser.add_argument("--rl-policy-clip-range",
type=float,
default=RLArgs.grpo_policy_clip_range,
dest="grpo_policy_clip_range", # Map to RLArgs field name
help="PPO-style clipping range for policy ratio")
parser.add_argument("--rl-value-clip-range",
type=float,
default=RLArgs.grpo_value_clip_range,
help="Value function clipping range")
parser.add_argument("--rl-num-policy-epochs",
type=int,
default=RLArgs.grpo_num_policy_epochs,
help="Number of policy update epochs (GRPO typically uses 1)")
parser.add_argument("--rl-num-value-epochs",
type=int,
default=RLArgs.grpo_num_value_epochs,
help="Number of value function update epochs")
parser.add_argument("--rl-target-kl",
type=float,
default=RLArgs.grpo_target_kl,
help="Target KL divergence for early stopping")
parser.add_argument("--rl-entropy-coef",
type=float,
default=RLArgs.grpo_entropy_coef,
help="Entropy coefficient for exploration")
parser.add_argument("--rl-value-loss-coef",
type=float,
default=RLArgs.grpo_value_loss_coef,
help="Value loss coefficient")
# GRPO-Guard (safety mechanisms)
parser.add_argument("--rl-use-grpo-guard",
action=StoreBoolean,
help="Enable GRPO-Guard safety mechanisms")
parser.add_argument("--rl-ratio-norm-correction",
action=StoreBoolean,
help="RatioNorm: correct importance ratio bias")
parser.add_argument("--rl-gradient-reweighting",
action=StoreBoolean,
help="Reweight gradients across denoising steps")
parser.add_argument("--rl-max-importance-ratio",
type=float,
default=RLArgs.grpo_max_importance_ratio,
help="Clip importance ratios above this value")
# Reward models
parser.add_argument("--reward-models",
type=str,
default='{"dummy": 1.0}',
help="Reward models as JSON dict (e.g., '{\"video_ocr\": 1.0, \"pickscore\": 0.5}')")
parser.add_argument("--value-model-path",
type=str,
default=RLArgs.value_model_path,
help="Path to value model (can be empty to train from scratch)")
parser.add_argument("--value-model-share-backbone",
action=StoreBoolean,
help="Share transformer backbone between policy and value")
# Training schedule
parser.add_argument("--rl-warmup-steps",
type=int,
default=RLArgs.warmup_steps,
help="Collect SFT-style data before starting RL")
parser.add_argument("--rl-collect-on-policy",
action=StoreBoolean,
help="Collect fresh rollouts each step (on-policy)")
parser.add_argument("--rl-timestep-fraction",
type=float,
default=RLArgs.timestep_fraction,
help="Fraction of timesteps to train on")
parser.add_argument("--rl-num-inner-epochs",
type=int,
default=RLArgs.num_inner_epochs,
help="Number of inner epochs per outer epoch")
# KL regularization
parser.add_argument("--rl-kl-beta",
type=float,
default=RLArgs.kl_beta,
dest="kl_beta", # Map CLI arg to RLArgs field name
help="KL loss coefficient (GRPO uses KL loss, DPO uses larger beta)")
parser.add_argument("--rl-kl-reward",
type=float,
default=RLArgs.kl_reward,
help="KL reward coefficient (alternative to KL loss, typically 0)")
# SFT integration
parser.add_argument("--rl-sft-weight",
type=float,
default=RLArgs.sft_weight,
help="SFT loss weight for supervised learning in RL training")
parser.add_argument("--rl-sft-batch-size",
type=int,
default=RLArgs.sft_batch_size,
help="Batch size for SFT data")
# CFG settings
parser.add_argument("--guidance-scale",
type=float,
default=1.0,
help="Guidance scale for CFG")
# Statistics tracking
parser.add_argument("--rl-global-std",
action=StoreBoolean,
help="Use global std across all samples vs per-group std")
parser.add_argument("--rl-per-prompt-stat-tracking",
action=StoreBoolean,
help="Track statistics per prompt")
# Training options
parser.add_argument("--rl-use-diffusion-loss",
action=StoreBoolean,
help="Use diffusion loss in training")
# DPO-specific
parser.add_argument("--dpo-beta",
type=float,
default=RLArgs.dpo_beta,
help="DPO regularization parameter (typically much larger than GRPO beta)")
parser.add_argument("--dpo-ref-update-step",
type=int,
default=RLArgs.dpo_ref_update_step,
help="Reference model update frequency for OnlineDPO")
parser.add_argument("--dpo-label-smoothing",
type=float,
default=RLArgs.dpo_label_smoothing,
help="Label smoothing for DPO loss")
return parser
@dataclasses.dataclass
class TrainingArgs(FastVideoArgs):
"""
@@ -819,6 +1017,11 @@ class TrainingArgs(FastVideoArgs):
num_height: int = 0
num_width: int = 0
num_frames: int = 0
# RL dataset configuration (for RL prompt datasets)
rl_dataset_path: str = "" # Path to RL prompt dataset directory (defaults to data_path if not set)
rl_dataset_type: str = "text" # "text" or "geneval"
rl_num_image_per_prompt: int = 4 # k parameter for KRepeatSampler (num_image_per_prompt)
train_batch_size: int = 0
num_latent_t: int = 0
@@ -916,7 +1119,6 @@ class TrainingArgs(FastVideoArgs):
training_state_checkpointing_steps: int = 0 # for resuming training
weight_only_checkpointing_steps: int = 0 # for inference
log_visualization: bool = False
visualization_steps: int = 0
# simulate generator forward to match inference
simulate_generator_forward: bool = False
warp_denoising_step: bool = False
@@ -930,6 +1132,9 @@ class TrainingArgs(FastVideoArgs):
last_step_only: bool = False # Only use the last timestep for training
context_noise: int = 0 # Context noise level for cache updates
# Nested RL configuration
rl_args: RLArgs = dataclasses.field(default_factory=RLArgs)
@classmethod
def from_cli_args(cls, args: argparse.Namespace) -> "TrainingArgs":
provided_args = clean_cli_args(args)
@@ -954,6 +1159,25 @@ class TrainingArgs(FastVideoArgs):
kwargs[attr] = WorkloadType.from_string(
workload_type_value) if isinstance(
workload_type_value, str) else workload_type_value
elif attr == 'rl_args':
# Construct nested RLArgs from CLI arguments
rl_kwargs = {}
for rl_field in dataclasses.fields(RLArgs):
rl_attr = rl_field.name
if hasattr(args, rl_attr):
value = getattr(args, rl_attr)
# Special handling for reward_models: parse JSON string to dict
if rl_attr == 'reward_models' and isinstance(value, str):
rl_kwargs[rl_attr] = json.loads(value) if value else {}
else:
rl_kwargs[rl_attr] = value
else:
# Use default value from RLArgs
if rl_field.default_factory is not dataclasses.MISSING:
rl_kwargs[rl_attr] = rl_field.default_factory()
elif rl_field.default is not dataclasses.MISSING:
rl_kwargs[rl_attr] = rl_field.default
kwargs[attr] = RLArgs(**rl_kwargs)
# Use getattr with default value from the dataclass for potentially missing attributes
else:
# Get the field to check its default value
@@ -983,11 +1207,26 @@ class TrainingArgs(FastVideoArgs):
parser.add_argument("--data-path",
type=str,
required=True,
help="Path to parquet files")
help="Path to parquet files (or RL prompt dataset directory for RL training)")
parser.add_argument("--dataloader-num-workers",
type=int,
required=True,
help="Number of workers for dataloader")
# RL dataset arguments (optional, defaults to data_path)
parser.add_argument("--rl-dataset-path",
type=str,
default="",
help="Path to RL prompt dataset directory (defaults to --data-path if not set)")
parser.add_argument("--rl-dataset-type",
type=str,
default="text",
choices=["text", "geneval"],
help="RL dataset type: 'text' for TextPromptDataset or 'geneval' for GenevalPromptDataset")
parser.add_argument("--rl-num-image-per-prompt",
type=int,
default=4,
help="Number of times to repeat each prompt (k parameter for KRepeatSampler)")
parser.add_argument("--num-height",
type=int,
required=True,
@@ -1080,9 +1319,6 @@ class TrainingArgs(FastVideoArgs):
parser.add_argument("--log-validation",
action=StoreBoolean,
help="Whether to log validation results")
parser.add_argument("--visualization-steps",
type=int,
help="Number of visualization steps")
parser.add_argument("--tracker-project-name",
type=str,
help="Project name for tracking")
@@ -1355,6 +1591,9 @@ class TrainingArgs(FastVideoArgs):
default=TrainingArgs.context_noise,
help="Context noise level for cache updates")
# RL (Reinforcement Learning) arguments
RLArgs.add_cli_args(parser)
return parser
-43
View File
@@ -193,46 +193,3 @@ class ImageProcessor:
tensor = 2.0 * tensor - 1.0
return tensor
def _preprocess_cosmos25(
self,
image: PIL.Image.Image | np.ndarray | torch.Tensor,
height: int,
width: int,
) -> torch.Tensor:
"""
Cosmos-Predict2.5-style preprocessing for image2world:
- aspect-preserving resize (scale so both dims >= target)
- center crop to (height, width)
- normalize to [-1, 1]
Returns:
torch.Tensor: (1, 3, height, width) in [-1, 1]
"""
if not isinstance(image, PIL.Image.Image):
# Reuse the generic preprocess for non-PIL inputs.
return self.preprocess(image, height=height, width=width)
# Ensure RGB (official uses torchvision.to_tensor on PIL, which yields C=3 for RGB inputs)
image = image.convert("RGB")
# Make sure target dims are multiples of VAE scale factor (FastVideo convention)
height = height - (height % self.vae_scale_factor)
width = width - (width % self.vae_scale_factor)
orig_w, orig_h = image.size
# Scale so that resized image fully covers the target box (like official resize_input()).
scale = max(width / float(orig_w), height / float(orig_h))
resized_w = int(np.ceil(scale * orig_w))
resized_h = int(np.ceil(scale * orig_h))
image = image.resize((resized_w, resized_h),
resample=PIL.Image.Resampling.BILINEAR)
# Center crop
left = max(0, (resized_w - width) // 2)
top = max(0, (resized_h - height) // 2)
image = image.crop((left, top, left + width, top + height))
image_np = np.array(image, dtype=np.float32) / 255.0 # [0,1]
return self._normalize_to_tensor(image_np) # -> [-1,1], (1,3,H,W)
+5 -23
View File
@@ -168,15 +168,9 @@ class ScaleResidualLayerNormScaleShift(nn.Module):
else:
raise NotImplementedError(f"Norm type {norm_type} not implemented")
def forward(
self,
residual: torch.Tensor,
x: torch.Tensor,
gate: torch.Tensor | int,
shift: torch.Tensor,
scale: torch.Tensor,
convert_modulation_dtype: bool = False
) -> tuple[torch.Tensor, torch.Tensor]:
def forward(self, residual: torch.Tensor, x: torch.Tensor,
gate: torch.Tensor | int, shift: torch.Tensor,
scale: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
"""
Apply gated residual connection, followed by layernorm and
scale/shift in a single fused operation.
@@ -211,11 +205,6 @@ class ScaleResidualLayerNormScaleShift(nn.Module):
# Apply normalization
normalized = self.norm(residual_output)
if convert_modulation_dtype:
scale = scale.to(normalized.dtype)
shift = shift.to(normalized.dtype)
# Apply scale and shift
if isinstance(scale, torch.Tensor) and scale.dim() == 4:
# scale.shape: [batch_size, num_frames, 1, inner_dim]
@@ -265,21 +254,14 @@ class LayerNormScaleShift(nn.Module):
else:
raise NotImplementedError(f"Norm type {norm_type} not implemented")
def forward(self,
x: torch.Tensor,
shift: torch.Tensor,
scale: torch.Tensor,
convert_modulation_dtype: bool = False) -> torch.Tensor:
def forward(self, x: torch.Tensor, shift: torch.Tensor,
scale: torch.Tensor) -> torch.Tensor:
"""Apply ln followed by scale and shift in a single fused operation."""
# x.shape: [batch_size, seq_len, inner_dim]
normalized = self.norm(x)
if self.compute_dtype == torch.float32:
normalized = normalized.float()
if convert_modulation_dtype:
scale = scale.to(normalized.dtype)
shift = shift.to(normalized.dtype)
if scale.dim() == 4:
# scale.shape: [batch_size, num_frames, 1, inner_dim]
num_frames = scale.shape[1]
+15 -18
View File
@@ -27,7 +27,7 @@ RESET = '\033[0;0m'
_warned_local_main_process = False
_warned_main_process = False
_FORMAT = (f"{FASTVIDEO_LOGGING_PREFIX}%(levelname)s %(asctime)s.%(msecs)03d "
_FORMAT = (f"{FASTVIDEO_LOGGING_PREFIX}%(levelname)s %(asctime)s "
"[%(filename)s:%(lineno)d] %(message)s")
_DATE_FORMAT = "%m-%d %H:%M:%S"
@@ -102,7 +102,6 @@ def _info(logger: Logger,
- When both are False, the message will be logged from all processes
- By default, only logs from processes with LOCAL_RANK=0
"""
is_distributed = int(os.environ.get("WORLD_SIZE", 1)) > 1
try:
local_rank = int(os.environ["LOCAL_RANK"])
rank = int(os.environ["RANK"])
@@ -119,22 +118,20 @@ def _info(logger: Logger,
global _warned_local_main_process, _warned_main_process
# Only show process-awareness warnings when actually running distributed
if is_distributed:
if not _warned_local_main_process and local_main_process_only:
logger.warning(
'%s By default, logger.info(..) will only log from the local main process. Set logger.info(..., is_local_main_process=False) to log from all processes.%s',
GREEN,
RESET,
)
_warned_local_main_process = True
if not _warned_main_process and main_process_only:
logger.warning(
'%s is_main_process_only is set to True, logging only from the main (RANK==0) process.%s',
GREEN,
RESET,
)
_warned_main_process = True
if not _warned_local_main_process and local_main_process_only:
logger.warning(
'%s By default, logger.info(..) will only log from the local main process. Set logger.info(..., is_local_main_process=False) to log from all processes.%s',
GREEN,
RESET,
)
_warned_local_main_process = True
if not _warned_main_process and main_process_only:
logger.warning(
'%s is_main_process_only is set to True, logging only from the main (RANK==0) process.%s',
GREEN,
RESET,
)
_warned_main_process = True
if not main_process_only and not local_main_process_only:
logger.log(logging.INFO, msg, *args, stacklevel=2, **kwargs)
-9
View File
@@ -1,9 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from fastvideo.models.audio.ltx2_audio_vae import (
LTX2AudioDecoder,
LTX2AudioEncoder,
LTX2Vocoder,
)
__all__ = ["LTX2AudioEncoder", "LTX2AudioDecoder", "LTX2Vocoder"]
File diff suppressed because it is too large Load Diff
+3 -7
View File
@@ -33,7 +33,7 @@ from fastvideo.layers.visual_embedding import (PatchEmbed)
from fastvideo.logger import init_logger
from fastvideo.models.dits.base import BaseDiT
from fastvideo.models.dits.wanvideo import WanT2VCrossAttention, WanTimeTextImageEmbedding
from fastvideo.platforms import AttentionBackendEnum, current_platform
from fastvideo.platforms import AttentionBackendEnum
logger = init_logger(__name__)
class CausalWanSelfAttention(nn.Module):
@@ -286,8 +286,6 @@ class CausalWanTransformerBlock(nn.Module):
null_shift = null_scale = torch.tensor([0], device=hidden_states.device)
norm_hidden_states, hidden_states = self.self_attn_residual_norm(
hidden_states, attn_output, gate_msa, null_shift, null_scale)
norm_hidden_states, hidden_states = norm_hidden_states.to(
orig_dtype), hidden_states.to(orig_dtype)
# 2. Cross-attention
attn_output = self.attn2(norm_hidden_states,
@@ -454,6 +452,7 @@ class CausalWanTransformer3DModel(BaseDiT):
This function will be run for num_frame times.
Process the latent frames one by one (1560 tokens each)
"""
from fastvideo.platforms import current_platform
orig_dtype = hidden_states.dtype
if not isinstance(encoder_hidden_states, torch.Tensor):
@@ -677,7 +676,4 @@ class CausalWanTransformer3DModel(BaseDiT):
# u = torch.einsum('fhwpqrc->cfphqwr', u.contiguous())
u = u.reshape(c, *[i * j for i, j in zip(v, self.patch_size)])
out.append(u)
return out
# Entry point for model registry
EntryClass = CausalWanTransformer3DModel
return out
+1 -4
View File
@@ -723,7 +723,4 @@ class CosmosTransformer3DModel(BaseDiT):
hidden_states = hidden_states.permute(0, 7, 1, 6, 2, 4, 3, 5)
hidden_states = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
return hidden_states
# Entry point for model registry
EntryClass = CosmosTransformer3DModel
return hidden_states
-2
View File
@@ -959,5 +959,3 @@ class Cosmos25Transformer3DModel(BaseDiT):
return hidden_states
# Entry point for model registry
EntryClass = Cosmos25Transformer3DModel
-3
View File
@@ -940,6 +940,3 @@ class FinalLayer(nn.Module):
x = self.norm_final(x) * (1.0 + scale.unsqueeze(1)) + shift.unsqueeze(1)
x, _ = self.linear(x)
return x
# Entry point for model registry
EntryClass = HunyuanVideoTransformer3DModel
+1 -6
View File
@@ -626,8 +626,6 @@ class HunyuanVideo15Transformer3DModel(CachableDiT):
)
# Final layer processing
if get_sp_world_size() > 1:
hidden_states = hidden_states.contiguous()
hidden_states = sequence_model_parallel_all_gather_with_unpad(hidden_states, original_seq_len, dim=1)
hidden_states = self.final_layer(hidden_states, temb)
# Unpatchify to get original shape
@@ -852,7 +850,4 @@ class FinalLayer(nn.Module):
scale, shift = self.adaLN_modulation(c).chunk(2, dim=-1)
x = self.norm_final(x) * (1.0 + scale.unsqueeze(1)) + shift.unsqueeze(1)
x, _ = self.linear(x)
return x
# Entry point for model registry
EntryClass = HunyuanVideo15Transformer3DModel
return x
-27
View File
@@ -1,27 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""
HYWorld (HY-WorldPlay) model components for FastVideo.
This module provides:
- HYWorldTransformer3DModel: The main transformer model with ProPE and action conditioning
- HYWorldVideoGenerator: Extended VideoGenerator for HYWorld inference
- Utilities for pose processing and camera trajectory generation
"""
from .hyworld import HYWorldTransformer3DModel, HYWorldDoubleStreamBlock
# Inference utilities (used by examples)
from .resolution_utils import (
get_resolution_from_image,
)
__all__ = [
# Model (used by model registry)
"HYWorldTransformer3DModel",
"HYWorldDoubleStreamBlock",
# Inference utilities (used by examples)
"get_resolution_from_image",
]
# Entry point for model registry
EntryClass = HYWorldTransformer3DModel
@@ -1,261 +0,0 @@
# HY-WorldPlay/hyvideo/prope/camera_rope.py
# MIT License
#
# Copyright (c) Authors of
# "PRoPE: Projective Positional Encoding for Multiview Transformers"
#
# 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.
#
# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
# SOFTWARE.
# How to use PRoPE attention for self-attention:
#
# 1. Easiest way (fast):
# attn = PropeDotProductAttention(...)
# o = attn(q, k, v, viewmats, Ks)
#
# 2. More flexible way (fast):
# attn = PropeDotProductAttention(...)
# attn._precompute_and_cache_apply_fns(viewmats, Ks)
# q = attn._apply_to_q(q)
# k = attn._apply_to_kv(k)
# v = attn._apply_to_kv(v)
# o = F.scaled_dot_product_attention(q, k, v, **kwargs)
# o = attn._apply_to_o(o)
#
# 3. The most flexible way (but slower because repeated computation of RoPE coefficients):
# o = prope_dot_product_attention(q, k, v, ...)
#
# How to use PRoPE attention for cross-attention:
#
# attn_src = PropeDotProductAttention(...)
# attn_tgt = PropeDotProductAttention(...)
# attn_src._precompute_and_cache_apply_fns(viewmats_src, Ks_src)
# attn_tgt._precompute_and_cache_apply_fns(viewmats_tgt, Ks_tgt)
# q_src = attn_src._apply_to_q(q_src)
# k_tgt = attn_tgt._apply_to_kv(k_tgt)
# v_tgt = attn_tgt._apply_to_kv(v_tgt)
# o_src = F.scaled_dot_product_attention(q_src, k_tgt, v_tgt, **kwargs)
# o_src = attn_src._apply_to_o(o_src)
from functools import partial
from typing import Callable, Optional, Tuple, List
import torch
import torch.nn.functional as F
def prope_qkv(
q: torch.Tensor, # (batch, num_heads, seqlen, head_dim)
k: torch.Tensor, # (batch, num_heads, seqlen, head_dim)
v: torch.Tensor, # (batch, num_heads, seqlen, head_dim)
*,
viewmats: torch.Tensor, # (batch, cameras, 4, 4)
Ks: Optional[torch.Tensor], # (batch, cameras, 3, 3)
patches_x: int = None, # How many patches wide is each image?
patches_y: int = None, # How many patches tall is each image?
image_width: int = None, # Width of the image. Used to normalize intrinsics.
image_height: int = None, # Height of the image. Used to normalize intrinsics.
coeffs_x: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
coeffs_y: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
mask: Optional[torch.Tensor] = None,
kv_cache=None,
is_cache: bool = False,
**kwargs,
) -> torch.Tensor:
"""Similar to torch.nn.functional.scaled_dot_product_attention, but applies PRoPE-style
positional encoding.
Currently, we assume that the sequence length is equal to:
cameras * patches_x * patches_y
And token ordering allows the `(seqlen,)` axis to be reshaped into
`(cameras, patches_x, patches_y)`.
"""
# We're going to assume self-attention: all inputs are the same shape.
(batch, num_heads, seqlen, head_dim) = q.shape
cameras = viewmats.shape[1]
assert q.shape == k.shape == v.shape
assert viewmats.shape == (batch, cameras, 4, 4)
assert Ks is None or Ks.shape == (batch, cameras, 3, 3)
# assert seqlen == cameras * patches_x * patches_y
apply_fn_q, apply_fn_kv, apply_fn_o = _prepare_apply_fns_all_dim(
head_dim=head_dim,
viewmats=viewmats,
Ks=Ks,
patches_x=patches_x,
patches_y=patches_y,
image_width=image_width,
image_height=image_height,
coeffs_x=coeffs_x,
coeffs_y=coeffs_y,
)
query = apply_fn_q(q)
key = apply_fn_kv(k)
value = apply_fn_kv(v)
return query, key, value, apply_fn_o
def _prepare_apply_fns_all_dim(
head_dim: int, # Q/K/V will have this last dimension
viewmats: torch.Tensor, # (batch, cameras, 4, 4)
Ks: Optional[torch.Tensor], # (batch, cameras, 3, 3)
patches_x: int, # How many patches wide is each image?
patches_y: int, # How many patches tall is each image?
image_width: int, # Width of the image. Used to normalize intrinsics.
image_height: int, # Height of the image. Used to normalize intrinsics.
coeffs_x: Optional[torch.Tensor] = None,
coeffs_y: Optional[torch.Tensor] = None,
) -> Tuple[
Callable[[torch.Tensor], torch.Tensor],
Callable[[torch.Tensor], torch.Tensor],
Callable[[torch.Tensor], torch.Tensor],
]:
"""Prepare transforms for PRoPE-style positional encoding."""
device = viewmats.device
(batch, cameras, _, _) = viewmats.shape
# Normalize camera intrinsics.
if Ks is not None:
Ks_norm = torch.zeros_like(Ks)
Ks_norm[..., 0, 0] = Ks[..., 0, 0]
Ks_norm[..., 1, 1] = Ks[..., 1, 1]
Ks_norm[..., 0, 2] = 0
Ks_norm[..., 1, 2] = 0
Ks_norm[..., 2, 2] = 1.0
Ks_norm = Ks_norm.to(dtype=Ks.dtype)
del Ks
# Compute the camera projection matrices we use in PRoPE.
# - K is an `image<-camera` transform.
# - viewmats is a `camera<-world` transform.
# - P = lift(K) @ viewmats is an `image<-world` transform.
P = torch.einsum("...ij,...jk->...ik", _lift_K(Ks_norm), viewmats)
P_T = P.transpose(-1, -2).to(dtype=viewmats.dtype)
P_inv = torch.einsum(
"...ij,...jk->...ik",
_invert_SE3(viewmats),
_lift_K(_invert_K(Ks_norm)),
).to(dtype=viewmats.dtype)
else:
# GTA formula. P is `camera<-world` transform.
P = viewmats
P_T = P.transpose(-1, -2)
P_inv = _invert_SE3(viewmats)
assert P.shape == P_inv.shape == (batch, cameras, 4, 4)
# Block-diagonal transforms to the inputs and outputs of the attention operator.
assert head_dim % 4 == 0
transforms_q = [
(partial(_apply_tiled_projmat, matrix=P_T), head_dim),
]
transforms_kv = [
(partial(_apply_tiled_projmat, matrix=P_inv), head_dim),
]
transforms_o = [
(partial(_apply_tiled_projmat, matrix=P), head_dim),
]
apply_fn_q = partial(_apply_block_diagonal, func_size_pairs=transforms_q)
apply_fn_kv = partial(_apply_block_diagonal, func_size_pairs=transforms_kv)
apply_fn_o = partial(_apply_block_diagonal, func_size_pairs=transforms_o)
return apply_fn_q, apply_fn_kv, apply_fn_o
def _apply_tiled_projmat(
feats: torch.Tensor, # (batch, num_heads, seqlen, feat_dim)
matrix: torch.Tensor, # (batch, cameras, D, D)
) -> torch.Tensor:
"""Apply projection matrix to features."""
# - seqlen => (cameras, patches_x * patches_y)
# - feat_dim => (feat_dim // 4, 4)
(batch, num_heads, seqlen, feat_dim) = feats.shape
cameras = matrix.shape[1]
assert seqlen >= cameras and seqlen % cameras == 0
D = matrix.shape[-1]
assert matrix.shape == (batch, cameras, D, D)
assert feat_dim % D == 0
return torch.einsum(
"bcij,bncpkj->bncpki",
matrix,
feats.reshape((batch, num_heads, cameras, -1, feat_dim // D, D)),
).reshape(feats.shape)
def _apply_block_diagonal(
feats: torch.Tensor, # (..., dim)
func_size_pairs: List[Tuple[Callable[[torch.Tensor], torch.Tensor], int]],
) -> torch.Tensor:
"""Apply a block-diagonal function to an input array.
Each function is specified as a tuple with form:
((Tensor) -> Tensor, int)
Where the integer is the size of the input to the function.
"""
funcs, block_sizes = zip(*func_size_pairs)
assert feats.shape[-1] == sum(block_sizes)
x_blocks = torch.split(feats, block_sizes, dim=-1)
out = torch.cat(
[f(x_block) for f, x_block in zip(funcs, x_blocks)],
dim=-1,
)
assert out.shape == feats.shape, "Input/output shapes should match."
return out
def _invert_SE3(transforms: torch.Tensor) -> torch.Tensor:
"""Invert a 4x4 SE(3) matrix."""
assert transforms.shape[-2:] == (4, 4)
Rinv = transforms[..., :3, :3].transpose(-1, -2)
out = torch.zeros_like(transforms)
out[..., :3, :3] = Rinv
out[..., :3, 3] = -torch.einsum("...ij,...j->...i", Rinv, transforms[..., :3, 3])
out[..., 3, 3] = 1.0
out = out.to(dtype=transforms.dtype)
return out
def _lift_K(Ks: torch.Tensor) -> torch.Tensor:
"""Lift 3x3 matrices to homogeneous 4x4 matrices."""
assert Ks.shape[-2:] == (3, 3)
out = torch.zeros(Ks.shape[:-2] + (4, 4), device=Ks.device)
out[..., :3, :3] = Ks
out[..., 3, 3] = 1.0
out = out.to(dtype=Ks.dtype)
return out
def _invert_K(Ks: torch.Tensor) -> torch.Tensor:
"""Invert 3x3 intrinsics matrices. Assumes no skew."""
assert Ks.shape[-2:] == (3, 3)
out = torch.zeros_like(Ks)
out[..., 0, 0] = 1.0 / Ks[..., 0, 0]
out[..., 1, 1] = 1.0 / Ks[..., 1, 1]
out[..., 0, 2] = -Ks[..., 0, 2] / Ks[..., 0, 0]
out[..., 1, 2] = -Ks[..., 1, 2] / Ks[..., 1, 1]
out[..., 2, 2] = 1.0
out = out.to(dtype=Ks.dtype)
return out
@@ -1,76 +0,0 @@
# HY-WorldPlay/hyvideo/utils/data_utils.py
# Licensed under the TENCENT HUNYUAN COMMUNITY LICENSE AGREEMENT (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# https://github.com/Tencent-Hunyuan/HunyuanVideo-1.5/blob/main/LICENSE
#
# Unless and only to the extent required by applicable law, the Tencent Hunyuan works and any
# output and results therefrom are provided "AS IS" without any express or implied warranties of
# any kind including any warranties of title, merchantability, noninfringement, course of dealing,
# usage of trade, or fitness for a particular purpose. You are solely responsible for determining the
# appropriateness of using, reproducing, modifying, performing, displaying or distributing any of
# the Tencent Hunyuan works or outputs and assume any and all risks associated with your or a
# third party's use or distribution of any of the Tencent Hunyuan works or outputs and your exercise
# of rights and permissions under this agreement.
# See the License for the specific language governing permissions and limitations under the License.
import numpy as np
from PIL import Image
def resize_and_center_crop(image, target_width, target_height):
if target_height == image.shape[0] and target_width == image.shape[1]:
return image
pil_image = Image.fromarray(image)
original_width, original_height = pil_image.size
scale_factor = max(target_width / original_width, target_height / original_height)
resized_width = int(round(original_width * scale_factor))
resized_height = int(round(original_height * scale_factor))
resized_image = pil_image.resize((resized_width, resized_height), Image.LANCZOS)
left = (resized_width - target_width) / 2
top = (resized_height - target_height) / 2
right = (resized_width + target_width) / 2
bottom = (resized_height + target_height) / 2
cropped_image = resized_image.crop((left, top, right, bottom))
return np.array(cropped_image)
def get_closest_ratio(height: float, width: float, ratios: list, buckets: list):
"""
Get the closest ratio in the buckets.
Args:
height (float): video height
width (float): video width
ratios (list): video aspect ratio
buckets (list): buckets generated by `generate_crop_size_list`
Returns:
the closest size in the buckets and the corresponding ratio
"""
aspect_ratio = float(height) / float(width)
ratios_array = np.array(ratios)
closest_ratio_id = np.abs(ratios_array - aspect_ratio).argmin()
closest_size = buckets[closest_ratio_id]
closest_ratio = ratios_array[closest_ratio_id]
return closest_size, closest_ratio
def generate_crop_size_list(base_size=256, patch_size=16, max_ratio=4.0):
num_patches = round((base_size / patch_size) ** 2)
assert max_ratio >= 1.0
crop_size_list = []
wp, hp = num_patches, 1
while wp > 0:
if max(wp, hp) / min(wp, hp) <= max_ratio:
crop_size_list.append((wp * patch_size, hp * patch_size))
if (hp + 1) * wp <= num_patches:
hp += 1
else:
wp -= 1
return crop_size_list
-569
View File
@@ -1,569 +0,0 @@
# Licensed under the TENCENT HUNYUAN COMMUNITY LICENSE AGREEMENT (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# https://github.com/Tencent-Hunyuan/HunyuanVideo-1.5/blob/main/LICENSE
#
# Unless and only to the extent required by applicable law, the Tencent Hunyuan works and any
# output and results therefrom are provided "AS IS" without any express or implied warranties of
# any kind including any warranties of title, merchantability, noninfringement, course of dealing,
# usage of trade, or fitness for a particular purpose. You are solely responsible for determining the
# appropriateness of using, reproducing, modifying, performing, displaying or distributing any of
# the Tencent Hunyuan works or outputs and assume any and all risks associated with your or a
# third party's use or distribution of any of the Tencent Hunyuan works or outputs and your exercise
# of rights and permissions under this agreement.
# See the License for the specific language governing permissions and limitations under the License.
from typing import Any, Optional
import torch
import torch.nn as nn
from einops import rearrange, repeat
from fastvideo.distributed.communication_op import (
sequence_model_parallel_all_gather_with_unpad,
sequence_model_parallel_shard)
from fastvideo.configs.models.dits import HYWorldConfig
from fastvideo.layers.linear import ReplicatedLinear
from fastvideo.layers.rotary_embedding import get_rotary_pos_embed
from fastvideo.layers.visual_embedding import TimestepEmbedder, unpatchify
from fastvideo.models.dits.hunyuanvideo15 import (
MMDoubleStreamBlock,
HunyuanVideo15Transformer3DModel,
)
from fastvideo.platforms import AttentionBackendEnum
from fastvideo.logger import init_logger
from fastvideo.forward_context import set_forward_context
from fastvideo.distributed.parallel_state import get_sp_world_size
from fastvideo.distributed.utils import create_attention_mask_for_padding
from .camera_rope import prope_qkv
logger = init_logger(__name__)
class HYWorldDoubleStreamBlock(MMDoubleStreamBlock):
"""
Extended MMDoubleStreamBlock with ProPE (Projective Positional Encoding) support
for camera-aware attention in HY-World/WorldPlay models.
"""
def __init__(
self,
hidden_size: int,
num_attention_heads: int,
mlp_ratio: float,
dtype: torch.dtype | None = None,
supported_attention_backends: tuple[AttentionBackendEnum, ...] | None = None,
prefix: str = "",
):
super().__init__(
hidden_size=hidden_size,
num_attention_heads=num_attention_heads,
mlp_ratio=mlp_ratio,
dtype=dtype,
supported_attention_backends=supported_attention_backends,
prefix=prefix,
)
self.hidden_size = hidden_size
# Add ProPE projection layer for camera-aware attention
self.img_attn_prope_proj = ReplicatedLinear(
hidden_size,
hidden_size,
bias=True,
params_dtype=dtype,
prefix=f"{prefix}.img_attn_prope_proj"
)
# Zero-initialize ProPE projection (starts as identity)
nn.init.zeros_(self.img_attn_prope_proj.weight)
if self.img_attn_prope_proj.bias is not None:
nn.init.zeros_(self.img_attn_prope_proj.bias)
def forward(
self,
img: torch.Tensor,
txt: torch.Tensor,
encoder_attention_mask: torch.Tensor,
vec: torch.Tensor,
vec_txt: torch.Tensor,
freqs_cis: tuple,
seq_attention_mask: torch.Tensor,
viewmats: torch.Tensor,
Ks: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
"""
Forward pass with ProPE camera conditioning.
Args:
img: Image/video tokens
txt: Text tokens
encoder_attention_mask: Text attention mask
vec: Modulation vector
freqs_cis: Rotary embedding frequencies
seq_attention_mask: Sequence attention mask
viewmats: Camera view matrices for ProPE [B, T, 4, 4]
Ks: Camera intrinsics for ProPE [B, T, 3, 3]
Returns:
Tuple of (img, txt) output tokens
"""
# Process modulation vectors (inherited from parent)
img_mod_outputs = self.img_mod(vec)
(
img_attn_shift,
img_attn_scale,
img_attn_gate,
img_mlp_shift,
img_mlp_scale,
img_mlp_gate,
) = torch.chunk(img_mod_outputs, 6, dim=-1)
txt_mod_outputs = self.txt_mod(vec_txt)
(
txt_attn_shift,
txt_attn_scale,
txt_attn_gate,
txt_mlp_shift,
txt_mlp_scale,
txt_mlp_gate,
) = torch.chunk(txt_mod_outputs, 6, dim=-1)
# Prepare image for attention using fused operation
img_attn_input = self.img_attn_norm(img, img_attn_shift, img_attn_scale, convert_modulation_dtype=True)
# Get QKV for image
img_qkv, _ = self.img_attn_qkv(img_attn_input)
batch_size, image_seq_len = img_qkv.shape[0], img_qkv.shape[1]
# Split QKV
img_qkv = img_qkv.view(batch_size, image_seq_len, 3,
self.num_attention_heads, -1)
img_q, img_k, img_v = img_qkv[:, :, 0], img_qkv[:, :, 1], img_qkv[:, :,
2]
# Apply QK-Norm if needed
img_q = self.img_attn_q_norm(img_q).to(img_v)
img_k = self.img_attn_k_norm(img_k).to(img_v)
# Prepare text for attention using fused operation
txt_attn_input = self.txt_attn_norm(txt, txt_attn_shift, txt_attn_scale, convert_modulation_dtype=True)
# Get QKV for text
txt_qkv, _ = self.txt_attn_qkv(txt_attn_input)
batch_size, text_seq_len = txt_qkv.shape[0], txt_qkv.shape[1]
# Split QKV
txt_qkv = txt_qkv.view(batch_size, text_seq_len, 3,
self.num_attention_heads, -1)
txt_q, txt_k, txt_v = txt_qkv[:, :, 0], txt_qkv[:, :, 1], txt_qkv[:, :,
2]
# Apply QK-Norm if needed
txt_q = self.txt_attn_q_norm(txt_q).to(txt_q.dtype)
txt_k = self.txt_attn_k_norm(txt_k).to(txt_k.dtype)
# begin hyworld: add camera pose through prope
img_q_prope, img_k_prope, img_v_prope, apply_fn_o = prope_qkv(
img_q.permute(0, 2, 1, 3),
img_k.permute(0, 2, 1, 3),
img_v.permute(0, 2, 1, 3),
viewmats=viewmats,
Ks=Ks,
) # [batch, num_heads, seqlen, head_dim]
img_q_prope = img_q_prope.permute(0, 2, 1, 3) # [batch, seqlen, num_heads, head_dim]
img_k_prope = img_k_prope.permute(0, 2, 1, 3) # [batch, seqlen, num_heads, head_dim]
img_v_prope = img_v_prope.permute(0, 2, 1, 3) # [batch, seqlen, num_heads, head_dim]
# end hyworld
from fastvideo.attention.backends.flash_attn import FlashAttnMetadataBuilder
attn_metadata = FlashAttnMetadataBuilder().build(
current_timestep=0,
attn_mask=encoder_attention_mask,
)
# Run distributed attention
with set_forward_context(current_timestep=0, attn_metadata=attn_metadata):
img_attn, txt_attn = self.attn(img_q, img_k, img_v, txt_q, txt_k, txt_v, freqs_cis=freqs_cis, attention_mask=seq_attention_mask)
# begin hyworld
# attention with prope
from fastvideo.attention.backends.flash_attn import FlashAttnMetadataBuilder
attn_metadata_prope = FlashAttnMetadataBuilder().build(
current_timestep=0,
attn_mask=encoder_attention_mask,
)
# NOTE: Do NOT pass freqs_cis to prope attention - HY-WorldPlay does not apply RoPE to prope
with set_forward_context(current_timestep=0, attn_metadata=attn_metadata_prope):
img_attn_prope, _ = self.attn(
img_q_prope, img_k_prope, img_v_prope, txt_q, txt_k, txt_v,
freqs_cis=None, attention_mask=seq_attention_mask # No RoPE for prope attention
)
img_attn_prope = img_attn_prope.reshape(batch_size, image_seq_len, -1)
img_attn_prope = rearrange(
img_attn_prope, "B L (H D) -> B H L D", H=self.num_attention_heads
)
img_attn_prope = apply_fn_o(img_attn_prope) # [batch, num_heads, seqlen, head_dim]
img_attn_prope = rearrange(img_attn_prope, "B H L D -> B L (H D)")
# add prope to img_attn
img_attn_out, _ = self.img_attn_proj(img_attn.view(batch_size, image_seq_len, -1))
img_attn_prope_out, _ = self.img_attn_prope_proj(img_attn_prope)
img_attn_out = img_attn_out + img_attn_prope_out
# end hyworld
# Use fused operation for residual connection, normalization, and modulation
img_mlp_input, img_residual = self.img_attn_residual_mlp_norm(
img, img_attn_out, img_attn_gate, img_mlp_shift, img_mlp_scale, convert_modulation_dtype=True)
# Process image MLP
img_mlp_out = self.img_mlp(img_mlp_input)
img = self.img_mlp_residual(img_residual, img_mlp_out, img_mlp_gate)
# Process text attention output
txt_attn_out, _ = self.txt_attn_proj(
txt_attn.reshape(batch_size, text_seq_len, -1))
# Use fused operation for residual connection, normalization, and modulation
txt_mlp_input, txt_residual = self.txt_attn_residual_mlp_norm(
txt, txt_attn_out, txt_attn_gate, txt_mlp_shift, txt_mlp_scale, convert_modulation_dtype=True)
# Process text MLP
txt_mlp_out = self.txt_mlp(txt_mlp_input)
txt = self.txt_mlp_residual(txt_residual, txt_mlp_out, txt_mlp_gate)
return img, txt
class HYWorldFinalLayer(nn.Module):
"""
Final layer for HYWorld that uses modulate() to handle per-token conditioning.
This matches HY-WorldPlay's FinalLayer behavior.
"""
def __init__(self,
hidden_size,
patch_size,
out_channels,
dtype=None,
prefix: str = "") -> None:
super().__init__()
from fastvideo.layers.linear import ReplicatedLinear
from fastvideo.layers.visual_embedding import ModulateProjection
from fastvideo.layers.layernorm import LayerNormScaleShift
self.norm_final = LayerNormScaleShift(
hidden_size,
norm_type="layer",
eps=1e-6,
elementwise_affine=False,
dtype=dtype,
prefix=f"{prefix}.norm_final")
output_dim = patch_size[0] * patch_size[1] * patch_size[2] * out_channels
self.linear = ReplicatedLinear(hidden_size,
output_dim,
bias=True,
params_dtype=dtype,
prefix=f"{prefix}.linear")
# Modulation projection to get shift/scale
self.adaLN_modulation = ModulateProjection(
hidden_size,
factor=2,
act_layer="silu",
dtype=dtype,
prefix=f"{prefix}.adaLN_modulation")
def forward(self, x, c):
shift, scale = self.adaLN_modulation(c).chunk(2, dim=-1)
x = self.norm_final(x, shift, scale, convert_modulation_dtype=True)
x, _ = self.linear(x)
return x
class HYWorldTransformer3DModel(HunyuanVideo15Transformer3DModel):
r"""
HY-World Transformer extending HunyuanVideo15 with:
- ProPE (Projective Positional Encoding) for camera-aware attention
- Action conditioning for interactive video generation
"""
# Class attributes for weight loading - use HYWorld-specific mapping
_fsdp_shard_conditions = HYWorldConfig().arch_config._fsdp_shard_conditions
_compile_conditions = HYWorldConfig().arch_config._compile_conditions
param_names_mapping = HYWorldConfig().arch_config.param_names_mapping
reverse_param_names_mapping = HYWorldConfig().arch_config.reverse_param_names_mapping
def __init__(
self,
config: HYWorldConfig,
hf_config: dict[str, Any],
) -> None:
super().__init__(config=config, hf_config=hf_config)
# Replace double_blocks with HY-World version that supports ProPE
self.double_blocks = nn.ModuleList([
HYWorldDoubleStreamBlock(
hidden_size=self.hidden_size,
num_attention_heads=self.num_attention_heads,
mlp_ratio=config.arch_config.mlp_ratio,
dtype=None,
supported_attention_backends=self._supported_attention_backends,
prefix=f"{config.prefix}.double_blocks.{i}"
)
for i in range(config.arch_config.num_layers)
])
# Add action conditioning module
self.action_in = TimestepEmbedder(
self.hidden_size,
act_layer="silu",
dtype=None,
prefix=f"{config.prefix}.action_in"
)
# Zero-initialize action embedding (starts with no effect)
nn.init.zeros_(self.action_in.mlp.fc_out.weight)
if self.action_in.mlp.fc_out.bias is not None:
nn.init.zeros_(self.action_in.mlp.fc_out.bias)
# Override final_layer with HYWorld version that uses per-token modulate()
self.final_layer = HYWorldFinalLayer(
hidden_size=self.hidden_size,
patch_size=self.patch_size,
out_channels=self.out_channels,
dtype=None,
prefix=f"{config.prefix}.final_layer"
)
self.gradient_checkpointing = False
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: list[torch.Tensor],
timestep: torch.LongTensor,
encoder_hidden_states_image: list[torch.Tensor],
encoder_attention_mask: list[torch.Tensor],
action: torch.Tensor,
viewmats: torch.Tensor,
Ks: torch.Tensor,
timestep_txt: torch.LongTensor,
guidance: Optional[torch.Tensor] = None,
timestep_r: Optional[torch.LongTensor] = None,
attention_kwargs: Optional[dict[str, Any]] = None,
):
"""
Forward pass with action and camera conditioning.
Args:
action: Action tensor for action conditioning [B, T] or [B*T]
viewmats: Camera view matrices [B, T, 4, 4]
Ks: Camera intrinsics [B, T, 3, 3]
... (other args same as parent)
"""
encoder_hidden_states_image = encoder_hidden_states_image[0]
encoder_hidden_states, encoder_hidden_states_2 = encoder_hidden_states
encoder_attention_mask, encoder_attention_mask_2 = encoder_attention_mask
batch_size, num_channels, num_frames, height, width = hidden_states.shape
p_t, p_h, p_w = self.config.patch_size_t, self.config.patch_size, self.config.patch_size
post_patch_num_frames = num_frames // p_t
post_patch_height = height // p_h
post_patch_width = width // p_w
# 1. RoPE
# Get rotary embeddings
freqs_cos, freqs_sin = get_rotary_pos_embed(
(post_patch_num_frames, post_patch_height, post_patch_width),
self.hidden_size,
self.num_attention_heads,
self.config.rope_axes_dim,
self.config.rope_theta
)
freqs_cos = freqs_cos.to(hidden_states.device)
freqs_sin = freqs_sin.to(hidden_states.device)
# NOTE: freqs_cis does NOT need sharding because FastVideo's DistributedAttention
# uses all-to-all to gather the full sequence before applying RoPE
freqs_cis = (freqs_cos, freqs_sin) if freqs_cos is not None else None
# 2. Conditional embeddings
temb = self.time_in(timestep, timestep_r=timestep_r)
temb_txt = self.time_in(timestep_txt, timestep_r=timestep_r)
# Add action conditioning if provided
# temb shape: [B*T, C] where T = num_frames
temb = temb + self.action_in(action.reshape(-1))
# Broadcast timestep embedding for transformer blocks (one per spatial token)
# [B*T, C] -> [B, T*H*W, C] -> [B*T*H*W, C]
temb = repeat(temb, "(B T) C -> B (T H W) C", B=batch_size, H=post_patch_height, W=post_patch_width)
hidden_states = self.img_in(hidden_states)
hidden_states, original_seq_len = sequence_model_parallel_shard(hidden_states, dim=1)
current_seq_len = hidden_states.shape[1]
sp_world_size = get_sp_world_size()
padded_seq_len = current_seq_len * sp_world_size
if padded_seq_len > original_seq_len:
seq_attention_mask = create_attention_mask_for_padding(
seq_len=original_seq_len,
padded_seq_len=padded_seq_len,
batch_size=batch_size,
device=hidden_states.device,
)
else:
seq_attention_mask = None
viewmats_seq = repeat(
viewmats, "B T M N->B (T H W) M N",
H=post_patch_height,
W=post_patch_width
)
Ks_seq = repeat(
Ks, "B T M N->B (T H W) M N",
H=post_patch_height,
W=post_patch_width
)
# Shard viewmats, Ks, and temb for sequence parallelism (shard along sequence dim=1)
# Note that temb in HY1.5 does not need sharding because it is per-sample modulation
# In HYWorld, temb is per-token modulation.
if sp_world_size > 1:
viewmats_seq, _ = sequence_model_parallel_shard(viewmats_seq, dim=1)
Ks_seq, _ = sequence_model_parallel_shard(Ks_seq, dim=1)
temb, _ = sequence_model_parallel_shard(temb, dim=1)
# Rearrange temb after sharding to match expected shape
temb = rearrange(temb, "B S C -> (B S) C")
# qwen text embedding
encoder_hidden_states = self.txt_in(encoder_hidden_states, timestep_txt, encoder_attention_mask)
encoder_hidden_states_cond_emb = self.cond_type_embed(
torch.zeros_like(encoder_hidden_states[:, :, 0], dtype=torch.long)
)
encoder_hidden_states = encoder_hidden_states + encoder_hidden_states_cond_emb
# byt5 text embedding
encoder_hidden_states_2 = self.txt_in_2(encoder_hidden_states_2)
encoder_hidden_states_2_cond_emb = self.cond_type_embed(
torch.ones_like(encoder_hidden_states_2[:, :, 0], dtype=torch.long)
)
encoder_hidden_states_2 = encoder_hidden_states_2 + encoder_hidden_states_2_cond_emb
# image embed
encoder_hidden_states_3 = self.image_embedder(encoder_hidden_states_image)
is_t2v = torch.all(encoder_hidden_states_image == 0)
if is_t2v:
encoder_hidden_states_3 = encoder_hidden_states_3 * 0.0
encoder_attention_mask_3 = torch.zeros(
(batch_size, encoder_hidden_states_3.shape[1]),
dtype=encoder_attention_mask.dtype,
device=encoder_attention_mask.device,
)
else:
encoder_attention_mask_3 = torch.ones(
(batch_size, encoder_hidden_states_3.shape[1]),
dtype=encoder_attention_mask.dtype,
device=encoder_attention_mask.device,
)
encoder_hidden_states_3_cond_emb = self.cond_type_embed(
2
* torch.ones_like(
encoder_hidden_states_3[:, :, 0],
dtype=torch.long,
)
)
encoder_hidden_states_3 = encoder_hidden_states_3 + encoder_hidden_states_3_cond_emb
# reorder and combine text tokens: combine valid tokens first, then padding
encoder_attention_mask = encoder_attention_mask.bool()
encoder_attention_mask_2 = encoder_attention_mask_2.bool()
encoder_attention_mask_3 = encoder_attention_mask_3.bool()
new_encoder_hidden_states = []
new_encoder_attention_mask = []
for text, text_mask, text_2, text_mask_2, image, image_mask in zip(
encoder_hidden_states,
encoder_attention_mask,
encoder_hidden_states_2,
encoder_attention_mask_2,
encoder_hidden_states_3,
encoder_attention_mask_3,
):
# Concatenate: [valid_image, valid_byt5, valid_mllm, invalid_image, invalid_byt5, invalid_mllm]
new_encoder_hidden_states.append(
torch.cat(
[
image[image_mask], # valid image
text_2[text_mask_2], # valid byt5
text[text_mask], # valid mllm
image[~image_mask], # invalid image (zeroed)
torch.zeros_like(text_2[~text_mask_2]), # invalid byt5 (zeroed)
torch.zeros_like(text[~text_mask]), # invalid mllm (zeroed)
],
dim=0,
)
)
# Apply same reordering to attention masks
new_encoder_attention_mask.append(
torch.cat(
[
image_mask[image_mask],
text_mask_2[text_mask_2],
text_mask[text_mask],
image_mask[~image_mask],
text_mask_2[~text_mask_2],
text_mask[~text_mask],
],
dim=0,
)
)
encoder_hidden_states = torch.stack(new_encoder_hidden_states)
encoder_attention_mask = torch.stack(new_encoder_attention_mask)
# 4. Transformer blocks
if torch.is_grad_enabled() and self.gradient_checkpointing:
for block in self.double_blocks:
hidden_states, encoder_hidden_states = self._gradient_checkpointing_func(
block,
hidden_states,
encoder_hidden_states,
encoder_attention_mask,
temb,
temb_txt,
freqs_cis,
seq_attention_mask,
viewmats_seq,
Ks_seq,
)
else:
for block in self.double_blocks:
hidden_states, encoder_hidden_states = block(
hidden_states,
encoder_hidden_states,
encoder_attention_mask,
temb,
temb_txt,
freqs_cis,
seq_attention_mask,
viewmats_seq,
Ks_seq,
)
# Final layer processing (per-token conditioning via HYWorldFinalLayer)
# Apply final_layer on sharded data first, then gather
hidden_states = self.final_layer(hidden_states, temb)
# Gather the output from all ranks
if get_sp_world_size() > 1:
hidden_states = hidden_states.contiguous()
hidden_states = sequence_model_parallel_all_gather_with_unpad(hidden_states, original_seq_len, dim=1)
# Unpatchify to get original shape
hidden_states = unpatchify(hidden_states, post_patch_num_frames, post_patch_height, post_patch_width, self.patch_size, self.out_channels)
return hidden_states
-413
View File
@@ -1,413 +0,0 @@
# Some functions from HY-WorldPlay/hyvideo/generate.py
"""
Pose processing utilities for HYWorld video generation.
This module provides functions to convert camera poses to model input tensors,
including viewmats, intrinsics, and action labels.
Adapted from HY-WorldPlay: https://github.com/Tencent-Hunyuan/HY-WorldPlay
"""
import json
import numpy as np
import torch
from scipy.spatial.transform import Rotation as R
from typing import Union, Optional
from .trajectory import generate_camera_trajectory_local
# Mapping from one-hot action encoding to single label
mapping = {
(0, 0, 0, 0): 0,
(1, 0, 0, 0): 1,
(0, 1, 0, 0): 2,
(0, 0, 1, 0): 3,
(0, 0, 0, 1): 4,
(1, 0, 1, 0): 5,
(1, 0, 0, 1): 6,
(0, 1, 1, 0): 7,
(0, 1, 0, 1): 8,
}
# Default camera intrinsic matrix (for 1920x1080 resolution)
DEFAULT_INTRINSIC = [
[969.6969696969696, 0.0, 960.0],
[0.0, 969.6969696969696, 540.0],
[0.0, 0.0, 1.0],
]
# Default movement speeds
DEFAULT_FORWARD_SPEED = 0.08 # units per frame
DEFAULT_YAW_SPEED = np.deg2rad(3) # radians per frame
DEFAULT_PITCH_SPEED = np.deg2rad(3) # radians per frame
def one_hot_to_one_dimension(one_hot: torch.Tensor) -> torch.Tensor:
"""Convert one-hot action encoding to single dimension labels."""
return torch.tensor([mapping[tuple(row.tolist())] for row in one_hot])
def parse_pose_string(
pose_string: str,
forward_speed: float = DEFAULT_FORWARD_SPEED,
yaw_speed: float = DEFAULT_YAW_SPEED,
pitch_speed: float = DEFAULT_PITCH_SPEED,
) -> list[dict]:
"""
Parse pose string to motions list.
Format: "w-3, right-0.5, d-4"
- w: forward movement
- s: backward movement
- a: left movement
- d: right movement
- up: pitch up rotation
- down: pitch down rotation
- left: yaw left rotation
- right: yaw right rotation
- number after dash: duration in frames/latents
Args:
pose_string: Comma-separated pose commands
forward_speed: Movement amount per frame
yaw_speed: Yaw rotation amount per frame (radians)
pitch_speed: Pitch rotation amount per frame (radians)
Returns:
List of motion dictionaries for generate_camera_trajectory_local
"""
motions = []
commands = [cmd.strip() for cmd in pose_string.split(",")]
for cmd in commands:
if not cmd:
continue
parts = cmd.split("-")
if len(parts) != 2:
raise ValueError(
f"Invalid pose command: {cmd}. Expected format: 'action-duration'"
)
action = parts[0].strip()
try:
duration = float(parts[1].strip())
except ValueError:
raise ValueError(f"Invalid duration in command: {cmd}")
num_frames = int(duration)
# Parse action and create motion dicts
if action == "w":
# Forward
for _ in range(num_frames):
motions.append({"forward": forward_speed})
elif action == "s":
# Backward
for _ in range(num_frames):
motions.append({"forward": -forward_speed})
elif action == "a":
# Left
for _ in range(num_frames):
motions.append({"right": -forward_speed})
elif action == "d":
# Right
for _ in range(num_frames):
motions.append({"right": forward_speed})
elif action == "up":
# Pitch up
for _ in range(num_frames):
motions.append({"pitch": pitch_speed})
elif action == "down":
# Pitch down
for _ in range(num_frames):
motions.append({"pitch": -pitch_speed})
elif action == "left":
# Yaw left
for _ in range(num_frames):
motions.append({"yaw": -yaw_speed})
elif action == "right":
# Yaw right
for _ in range(num_frames):
motions.append({"yaw": yaw_speed})
else:
raise ValueError(
f"Unknown action: {action}. "
f"Supported actions: w, s, a, d, up, down, left, right"
)
return motions
def pose_string_to_json(
pose_string: str,
intrinsic: Optional[list[list[float]]] = None,
) -> dict:
"""
Convert pose string to pose JSON format.
Args:
pose_string: Comma-separated pose commands
intrinsic: Camera intrinsic matrix (default: DEFAULT_INTRINSIC from trajectory)
Returns:
Dict with frame indices as keys, containing extrinsic and K (intrinsic) matrices
"""
if intrinsic is None:
intrinsic = DEFAULT_INTRINSIC
motions = parse_pose_string(pose_string)
poses = generate_camera_trajectory_local(motions)
pose_json = {}
for i, p in enumerate(poses):
pose_json[str(i)] = {"extrinsic": p.tolist(), "K": intrinsic}
return pose_json
def pose_to_input(
pose_data: Union[str, dict],
latent_num: int,
tps: bool = False,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""
Convert pose data to model input tensors.
Args:
pose_data: One of:
- str ending with '.json': path to JSON file
- str: pose string (e.g., "w-3, right-0.5, d-4")
- dict: pose JSON data
latent_num: Number of latents (frames in latent space)
tps: Third-person mode flag
Returns:
Tuple of (viewmats, intrinsics, action_labels):
- viewmats: World-to-camera matrices [T, 4, 4]
- intrinsics: Normalized camera intrinsics [T, 3, 3]
- action_labels: Action labels for each frame [T]
"""
# Handle different input types
if isinstance(pose_data, str):
if pose_data.endswith(".json"):
# Load from JSON file
with open(pose_data, "r") as f:
pose_json = json.load(f)
else:
# Parse pose string
pose_json = pose_string_to_json(pose_data)
elif isinstance(pose_data, dict):
pose_json = pose_data
else:
raise ValueError(
f"Invalid pose_data type: {type(pose_data)}. Expected str or dict."
)
pose_keys = list(pose_json.keys())
latent_num_from_pose = len(pose_keys)
assert latent_num_from_pose == latent_num, (
f"pose corresponds to {latent_num_from_pose * 4 - 3} frames, num_frames "
f"must be set to {latent_num_from_pose * 4 - 3} to ensure alignment."
)
intrinsic_list = []
w2c_list = []
for i in range(latent_num):
t_key = pose_keys[i]
c2w = np.array(pose_json[t_key]["extrinsic"])
w2c = np.linalg.inv(c2w)
w2c_list.append(w2c)
# Normalize intrinsics
intrinsic = np.array(pose_json[t_key]["K"])
intrinsic[0, 0] /= intrinsic[0, 2] * 2
intrinsic[1, 1] /= intrinsic[1, 2] * 2
intrinsic[0, 2] = 0.5
intrinsic[1, 2] = 0.5
intrinsic_list.append(intrinsic)
w2c_list = np.array(w2c_list)
intrinsic_list = torch.tensor(np.array(intrinsic_list))
# Compute relative camera-to-world transforms
c2ws = np.linalg.inv(w2c_list)
C_inv = np.linalg.inv(c2ws[:-1])
relative_c2w = np.zeros_like(c2ws)
relative_c2w[0, ...] = c2ws[0, ...]
relative_c2w[1:, ...] = C_inv @ c2ws[1:, ...]
# Initialize one-hot action encodings
trans_one_hot = np.zeros((relative_c2w.shape[0], 4), dtype=np.int32)
rotate_one_hot = np.zeros((relative_c2w.shape[0], 4), dtype=np.int32)
move_norm_valid = 0.0001
for i in range(1, relative_c2w.shape[0]):
move_dirs = relative_c2w[i, :3, 3] # direction vector
move_norms = np.linalg.norm(move_dirs)
if move_norms > move_norm_valid: # threshold for movement
move_norm_dirs = move_dirs / move_norms
angles_rad = np.arccos(move_norm_dirs.clip(-1.0, 1.0))
trans_angles_deg = angles_rad * (180.0 / np.pi) # convert to degrees
else:
trans_angles_deg = np.zeros(3)
R_rel = relative_c2w[i, :3, :3]
r = R.from_matrix(R_rel)
rot_angles_deg = r.as_euler("xyz", degrees=True)
# Determine movement and rotation actions
if move_norms > move_norm_valid: # threshold for movement
if (not tps) or (
tps and abs(rot_angles_deg[1]) < 5e-2 and abs(rot_angles_deg[0]) < 5e-2
):
if trans_angles_deg[2] < 60:
trans_one_hot[i, 0] = 1 # forward
elif trans_angles_deg[2] > 120:
trans_one_hot[i, 1] = 1 # backward
if trans_angles_deg[0] < 60:
trans_one_hot[i, 2] = 1 # right
elif trans_angles_deg[0] > 120:
trans_one_hot[i, 3] = 1 # left
if rot_angles_deg[1] > 5e-2:
rotate_one_hot[i, 0] = 1 # right
elif rot_angles_deg[1] < -5e-2:
rotate_one_hot[i, 1] = 1 # left
if rot_angles_deg[0] > 5e-2:
rotate_one_hot[i, 2] = 1 # up
elif rot_angles_deg[0] < -5e-2:
rotate_one_hot[i, 3] = 1 # down
trans_one_hot = torch.tensor(trans_one_hot)
rotate_one_hot = torch.tensor(rotate_one_hot)
# Convert one-hot to single-dimension labels
trans_one_label = one_hot_to_one_dimension(trans_one_hot)
rotate_one_label = one_hot_to_one_dimension(rotate_one_hot)
action_one_label = trans_one_label * 9 + rotate_one_label
return (
torch.as_tensor(w2c_list),
torch.as_tensor(intrinsic_list),
action_one_label,
)
def camera_center_normalization(w2c: np.ndarray) -> np.ndarray:
"""Normalize camera centers relative to the first camera."""
c2w = np.linalg.inv(w2c)
C0_inv = np.linalg.inv(c2w[0])
c2w_aligned = np.array([C0_inv @ C for C in c2w])
return np.linalg.inv(c2w_aligned)
def parse_pose_string_to_actions(pose_string: str, fps: int = 24) -> list[dict]:
"""
Parse pose string to frame-level action timeline.
Format: pose string uses latent counts, where:
- 1 latent = 4 frames
- Special rule: first frame of entire video is extra (frame 0)
- Example: "w-4,d-4" means:
- w-4: forward for frames 0-16 (17 frames total: 1 extra + 4*4)
- d-4: right for frames 17-32 (16 frames total: 4*4)
Args:
pose_string: Comma-separated pose commands (e.g., "w-4,d-4")
fps: Frames per second for video (default: 24)
Returns:
List of dicts with action values for each frame
"""
commands = [cmd.strip() for cmd in pose_string.split(",")]
frame_actions = []
is_first_command = True
for cmd in commands:
if not cmd:
continue
parts = cmd.split("-")
if len(parts) != 2:
raise ValueError(
f"Invalid pose command: {cmd}. Expected format: 'action-duration'"
)
action = parts[0].strip()
try:
num_latents = int(parts[1].strip())
except ValueError:
raise ValueError(f"Invalid duration in command: {cmd}")
# Convert latents to frames
# First command gets 1 extra frame (the special frame 0)
if is_first_command:
num_frames = 1 + num_latents * 4
is_first_command = False
else:
num_frames = num_latents * 4
# Map action to action values
action_values = {"forward": 0, "left": 0, "yaw": 0, "pitch": 0}
if action == "w":
action_values["forward"] = 1
elif action == "s":
action_values["forward"] = -1
elif action == "a":
action_values["left"] = 1
elif action == "d":
action_values["left"] = -1
elif action == "up":
action_values["pitch"] = 1
elif action == "down":
action_values["pitch"] = -1
elif action == "left":
action_values["yaw"] = -1
elif action == "right":
action_values["yaw"] = 1
else:
raise ValueError(f"Unknown action: {action}")
# Add frame-level actions
for _ in range(num_frames):
frame_actions.append(action_values.copy())
return frame_actions
def compute_latent_num(num_frames: int) -> int:
"""
Compute the number of latents from number of frames.
Formula: num_frames = (latent_num - 1) * 4 + 1
So: latent_num = (num_frames - 1) // 4 + 1
Args:
num_frames: Number of video frames
Returns:
Number of latents
"""
return (num_frames - 1) // 4 + 1
def compute_num_frames(latent_num: int) -> int:
"""
Compute the number of frames from number of latents.
Formula: num_frames = (latent_num - 1) * 4 + 1
Args:
latent_num: Number of latents
Returns:
Number of video frames
"""
return (latent_num - 1) * 4 + 1
@@ -1,64 +0,0 @@
import numpy as np
from PIL import Image
import requests
from io import BytesIO
from fastvideo.models.dits.hyworld.data_utils import generate_crop_size_list
# Target resolution configs (matching HY-WorldPlay)
TARGET_SIZE_CONFIG = {
"360p": {"bucket_hw_base_size": 480, "bucket_hw_bucket_stride": 16},
"480p": {"bucket_hw_base_size": 640, "bucket_hw_bucket_stride": 16},
"720p": {"bucket_hw_base_size": 960, "bucket_hw_bucket_stride": 16},
"1080p": {"bucket_hw_base_size": 1440, "bucket_hw_bucket_stride": 16},
}
def get_closest_resolution(image_height, image_width, target_resolution="480p"):
"""
Get closest supported resolution for given image dimensions.
Args:
image_height: Height of input image
image_width: Width of input image
target_resolution: Target resolution string (e.g., "480p", "720p")
Returns:
tuple[int, int]: (height, width) of closest supported resolution
"""
config = TARGET_SIZE_CONFIG[target_resolution]
bucket_hw_base_size = config["bucket_hw_base_size"]
bucket_hw_bucket_stride = config["bucket_hw_bucket_stride"]
crop_size_list = generate_crop_size_list(bucket_hw_base_size, bucket_hw_bucket_stride)
aspect_ratios = np.array([round(float(h) / float(w), 5) for h, w in crop_size_list])
# Find closest aspect ratio
image_ratio = float(image_height) / float(image_width)
closest_idx = np.abs(aspect_ratios - image_ratio).argmin()
closest_size = crop_size_list[closest_idx]
return closest_size[0], closest_size[1] # (height, width)
def get_resolution_from_image(image_path, target_resolution="480p"):
"""
Automatically determine resolution from input image.
Args:
image_path: Path or URL to input image
target_resolution: Target resolution tier ("480p", "720p", etc.)
Returns:
tuple[int, int]: (height, width) matching HY-WorldPlay's bucket selection
"""
# Handle URL inputs
if isinstance(image_path, str) and image_path.startswith(('http://', 'https://')):
response = requests.get(image_path)
response.raise_for_status()
img = Image.open(BytesIO(response.content))
else:
img = Image.open(image_path)
img_width, img_height = img.size
return get_closest_resolution(img_height, img_width, target_resolution)
@@ -1,316 +0,0 @@
# HY-WorldPlay/hyvideo/utils/retrieval_context.py
# Licensed under the TENCENT HUNYUAN COMMUNITY LICENSE AGREEMENT (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# https://github.com/Tencent-Hunyuan/HunyuanVideo-1.5/blob/main/LICENSE
#
# Unless and only to the extent required by applicable law, the Tencent Hunyuan works and any
# output and results therefrom are provided "AS IS" without any express or implied warranties of
# any kind including any warranties of title, merchantability, noninfringement, course of dealing,
# usage of trade, or fitness for a particular purpose. You are solely responsible for determining the
# appropriateness of using, reproducing, modifying, performing, displaying or distributing any of
# the Tencent Hunyuan works or outputs and assume any and all risks associated with your or a
# third party's use or distribution of any of the Tencent Hunyuan works or outputs and your exercise
# of rights and permissions under this agreement.
# See the License for the specific language governing permissions and limitations under the License.
import torch
import numpy as np
from typing import List, Tuple, Dict
import math
def generate_points_in_sphere(n_points: int, radius: float) -> torch.Tensor:
"""
Uniformly sample points within a sphere of a specified radius.
:param n_points: The number of points to generate.
:param radius: The radius of the sphere.
:return: A tensor of shape (n_points, 3), representing the (x, y, z) coordinates of the points.
"""
samples_r = torch.rand(n_points)
samples_phi = torch.rand(n_points)
samples_u = torch.rand(n_points)
r = radius * torch.pow(samples_r, 1 / 3)
phi = 2 * math.pi * samples_phi
theta = torch.acos(1 - 2 * samples_u)
# transfer the coordinates from spherical to cartesian
x = r * torch.sin(theta) * torch.cos(phi)
y = r * torch.sin(theta) * torch.sin(phi)
z = r * torch.cos(theta)
points = torch.stack((x, y, z), dim=1)
return points
def rotation_matrix_to_angles(R: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
"""
Estimate the Pitch and Yaw angles from a 3x3 rotation matrix R in the camera coordinate system.
Assumed Camera Coordinate System: X=Right, Y=Up, Z=Backward
(or NeRF style: X=Right, Y=Down, Z=Forward).
Here we adopt the common Computer Vision convention: Z-axis is Forward.
Note: The angle calculations here are directly based on the conventions of your `is_inside_fov_3d_hv` function:
- Yaw/Azimuth angle is in the XZ plane (atan2(x, z)).
- Pitch/Elevation angle is relative to the horizontal plane (atan2(y, sqrt(x^2 + z^2))).
For the third column R[:, 2] of the W2C matrix R (the direction of the World Z-axis in the Camera frame),
this typically corresponds to the direction the camera is looking
(the representation of the world Z-axis in the camera frame).
To simplify and match your `is_inside_fov` logic, we directly use the camera's Z-axis vector:
Camera Z-axis direction in World Frame (Forward Vector): fwd = R_w2c_inv @ [0, 0, 1]
More simply, the Z-axis vector of the C2W matrix is the camera's forward vector in the world frame.
C2W = W2C_inv
"""
R_c2w = R.T
fwd = R_c2w[:, 2]
x = fwd[0]
y = fwd[1]
z = fwd[2]
# compute yaw and pitch
yaw_rad = torch.atan2(x, z)
yaw_deg = yaw_rad * (180.0 / math.pi)
pitch_rad = torch.atan2(y, torch.sqrt(x**2 + z**2))
pitch_deg = pitch_rad * (180.0 / math.pi)
return pitch_deg, yaw_deg
def is_inside_fov_3d_hv(
points: torch.Tensor,
center: torch.Tensor,
center_pitch: torch.Tensor,
center_yaw: torch.Tensor,
fov_half_h: torch.Tensor,
fov_half_v: torch.Tensor,
) -> torch.Tensor:
"""
Check whether points are inside a 3D view frustum defined by a center coordinate, pitch angle, and yaw angle.
:param points: Tensor of shape (N, 3) or (N, B, 3) representing the coordinates of the sampled points.
:param center: Tensor of shape (3) or (B, 3) representing the camera center coordinates.
:param center_pitch: Tensor of shape (1) or (B) representing the pitch angle of center view direction.
:param center_yaw: Tensor of shape (1) or (B) representing the yaw angle of the center view direction.
:param fov_half_h: The horizontal half field-of-view angle (in degrees).
:param fov_half_v: The vertical half field-of-view angle (in degrees).
:return: Boolean tensor of shape (N) or (N, B), indicating whether each point is inside the FOV.
"""
if points.ndim == 2: # N, 3
vectors = points - center[None, :]
C = 1
elif points.ndim == 3: # N, B, 3
vectors = points - center[None, ...]
center_pitch = center_pitch[None, :] if center_pitch.ndim == 1 else center_pitch
center_yaw = center_yaw[None, :] if center_yaw.ndim == 1 else center_yaw
else:
raise ValueError("points' shape should be (N, 3) or (N, B, 3)")
x = vectors[..., 0]
y = vectors[..., 1]
z = vectors[..., 2]
# Calculate the horizontal angle (yaw/azimuth), assuming the Z-axis is forward.
azimuth = torch.atan2(x, z) * (180 / math.pi)
# Calculate the vertical angle (pitch/elevation).
elevation = torch.atan2(y, torch.sqrt(x**2 + z**2)) * (180 / math.pi)
# Calculate the angular difference from the center view direction (handling angle wrapping).
diff_azimuth = azimuth - center_yaw
diff_azimuth = torch.remainder(diff_azimuth + 180, 360) - 180
diff_elevation = elevation - center_pitch
diff_elevation = torch.remainder(diff_elevation + 180, 360) - 180
# Check if within FOV
in_fov_h = diff_azimuth.abs() < fov_half_h
in_fov_v = diff_elevation.abs() < fov_half_v
return in_fov_h & in_fov_v
def calculate_fov_overlap_similarity(
w2c_matrix_curr: torch.Tensor,
w2c_matrix_hist: torch.Tensor,
fov_h_deg: float = 105.0,
fov_v_deg: float = 75.0,
device=None,
points_local=None,
) -> float:
"""
Calculate the Field-of-View (FOV) overlap similarity between two W2C poses using Monte Carlo sampling.
Similarity = (Number of points in Curr_FOV ∩ Hist_FOV) / (Number of points in Curr_FOV).
:param w2c_matrix_curr: The (4, 4) W2C matrix for the current frame.
:param w2c_matrix_hist: The (4, 4) W2C matrix for the historical frame.
:param num_samples, radius, fov_h_deg, fov_v_deg: Sampling and FOV parameters.
:return: The overlap ratio (a float between 0.0 and 1.0).
"""
w2c_matrix_curr = torch.tensor(w2c_matrix_curr, device=device)
w2c_matrix_hist = torch.tensor(w2c_matrix_hist, device=device)
c2w_matrix_curr = torch.linalg.inv(w2c_matrix_curr)
c2w_matrix_hist = torch.linalg.inv(w2c_matrix_hist)
C_inv = w2c_matrix_curr
w2c_matrix_curr = torch.linalg.inv(C_inv @ c2w_matrix_curr)
w2c_matrix_hist = torch.linalg.inv(C_inv @ c2w_matrix_hist)
R_curr, t_curr = w2c_matrix_curr[:3, :3], w2c_matrix_curr[:3, 3]
R_hist, t_hist = w2c_matrix_hist[:3, :3], w2c_matrix_hist[:3, 3]
P_w_curr = -R_curr.T @ t_curr
P_w_hist = -R_hist.T @ t_hist
# pitch, yaw
pitch_curr, yaw_curr = rotation_matrix_to_angles(R_curr)
pitch_hist, yaw_hist = rotation_matrix_to_angles(R_hist)
fov_half_h = torch.tensor(fov_h_deg / 2.0, device=device)
fov_half_v = torch.tensor(fov_v_deg / 2.0, device=device)
# move to P_w_curr (N, 3)
points_world = points_local + P_w_curr[None, :]
in_fov_curr = is_inside_fov_3d_hv(
points_world,
P_w_curr[None, :],
pitch_curr[None],
yaw_curr[None],
fov_half_h,
fov_half_v,
)
# compute based on angle
in_fov_hist = is_inside_fov_3d_hv(
points_world,
P_w_hist[None, :],
pitch_hist[None],
yaw_hist[None],
fov_half_h,
fov_half_v,
)
# compute based on distance
dist = torch.norm(points_world - P_w_hist.reshape(1, -1), dim=1) < 8.0
in_fov_hist = in_fov_hist.bool() & dist.reshape(1, -1).bool()
overlap_count = (in_fov_curr.bool() & in_fov_hist.bool()).sum().float()
fov_curr_count = in_fov_curr.sum().float()
if fov_curr_count == 0:
return 0.0
overlap_ratio = overlap_count / fov_curr_count
return overlap_ratio.item()
def select_aligned_memory_frames(
w2c_list: List[np.ndarray],
current_frame_idx: int,
memory_frames: int,
temporal_context_size: int,
pred_latent_size: int,
pos_weight: float = 1.0,
ang_weight: float = 1.0,
device=None,
points_local=None,
) -> List[int]:
"""
Selects memory and context frames for a given frame based on a four-frame segment distance calculation.
:param w2c_list: List of all N 4x4 World-to-Camera (W2C) extrinsic matrices (np.ndarray).
:param current_frame_idx: The index of the current frame to be processed.
:param memory_frames: The total number of memory frames to select.
:param context_size: The total number of context frames to select.
:param pos_weight: The weight applied to the spatial (position) distance component.
:param ang_weight: The weight applied to the angular distance component.
:return: List[int]: A list containing the indices of the selected memory frames and context frames.
"""
if current_frame_idx <= memory_frames:
return list(range(0, current_frame_idx))
num_total_frames = len(w2c_list)
if current_frame_idx >= num_total_frames or current_frame_idx < 3:
raise ValueError(
f"The current frame index must be within the valid range of w2c_list and must be at least 3."
f"{current_frame_idx}, {len(w2c_list)}"
)
start_context_idx = max(0, current_frame_idx - temporal_context_size)
context_frames_indices = list(range(start_context_idx, current_frame_idx))
candidate_distances = []
query_clip_indices = list(
range(
current_frame_idx,
(
current_frame_idx + pred_latent_size
if current_frame_idx + pred_latent_size <= num_total_frames
else num_total_frames
),
)
)
historical_clip_indices = list(
range(4, current_frame_idx - temporal_context_size, 4)
)
memory_frames_indices = [0, 1, 2, 3] # add the first chunk as context
memory_frames = memory_frames - temporal_context_size
for hist_idx in historical_clip_indices:
total_dist = 0
hist_w2c_1 = w2c_list[hist_idx]
hist_w2c_2 = w2c_list[hist_idx + 2]
for query_idx in query_clip_indices:
dist_1_for_query_idx = 1.0 - calculate_fov_overlap_similarity(
w2c_list[query_idx],
hist_w2c_1,
fov_h_deg=60.0,
fov_v_deg=35.0,
device=device,
points_local=points_local,
)
dist_2_for_query_idx = 1.0 - calculate_fov_overlap_similarity(
w2c_list[query_idx],
hist_w2c_2,
fov_h_deg=60.0,
fov_v_deg=35.0,
device=device,
points_local=points_local,
)
dist_for_query_idx = (dist_1_for_query_idx + dist_2_for_query_idx) / 2.0
total_dist += dist_for_query_idx
final_clip_distance = total_dist / len(query_clip_indices)
candidate_distances.append((hist_idx, final_clip_distance))
candidate_distances.sort(key=lambda x: x[1])
for start_idx, _ in candidate_distances:
# check the memory frame number
if len(memory_frames_indices) >= memory_frames:
break
if start_idx not in memory_frames_indices:
memory_frames_indices.extend(range(start_idx, start_idx + 4))
# exclude the repeated frames
selected_frames_set = set(context_frames_indices)
selected_frames_set.update(memory_frames_indices)
final_selected_frames = sorted(list(selected_frames_set))
return final_selected_frames
-112
View File
@@ -1,112 +0,0 @@
# HY-WorldPlay/hyvideo/generate_custom_trajectory.py
import numpy as np
import json
def rot_x(theta):
c, s = np.cos(theta), np.sin(theta)
return np.array([[1, 0, 0], [0, c, -s], [0, s, c]])
def rot_y(theta):
c, s = np.cos(theta), np.sin(theta)
return np.array([[c, 0, s], [0, 1, 0], [-s, 0, c]])
def rot_z(theta):
c, s = np.cos(theta), np.sin(theta)
return np.array([[c, -s, 0], [s, c, 0], [0, 0, 1]])
def generate_camera_trajectory_local(motions):
"""
motions: list of dict
{"forward": 1.0}, {"yaw": np.pi/2}, {"pitch": np.pi/6}, {"right": 1.0}
- forward: Translation (Forward or Backward)
- yaw: Rotate (Left or Right)
- pitch: Rotate (Up or Down)
- right: Translation (Right or Left)
- third_yaw: Third Perspective Rotate (Left or Right)
"""
poses = []
T = np.eye(4)
poses.append(T.copy())
for move in motions:
# Rotate (Left or Right)
if "yaw" in move:
R = rot_y(move["yaw"])
T[:3, :3] = T[:3, :3] @ R
# Rotate (Up or Down)
if "pitch" in move:
R = rot_x(move["pitch"])
T[:3, :3] = T[:3, :3] @ R
# Translation (Z-direction of the camera's local coordinate system)
forward = move.get("forward", 0.0)
if forward != 0:
local_t = np.array([0, 0, forward])
world_t = T[:3, :3] @ local_t
T[:3, 3] += world_t
# Translation (Z-direction of the camera's local coordinate system)
right = move.get("right", 0.0)
if right != 0:
local_t = np.array([right, 0, 0])
world_t = T[:3, :3] @ local_t
T[:3, 3] += world_t
# Third Perspective Rotate (Left or Right)
third_yaw = move.get("third_yaw", 0.0)
if third_yaw != 0:
theta = -third_yaw
C = np.array([[1, 0.0, 0, 0], [0, 1, 0, 0], [0, 0, 1, -1.0], [0, 0, 0, 1]])
c_origin = C.copy()
# Rotation around the Y-axis
R_y = np.array(
[
[np.cos(theta), 0, np.sin(theta)],
[0, 1, 0],
[-np.sin(theta), 0, np.cos(theta)],
]
)
# Translation
C[:3, :3] = C[:3, :3] @ R_y
C[:3, 3] = R_y @ C[:3, 3]
c_inv = np.linalg.inv(c_origin)
c_relative = c_inv @ C
T = T @ c_relative
poses.append(T.copy())
return poses
if __name__ == "__main__":
# Examples: Forward 0.08 * 16 -> Right Rotate 3 degree * 16
motions = []
for i in range(15):
motions.append({"forward": 0.08})
for i in range(16):
motions.append({"yaw": np.deg2rad(3)})
intrinsic = [
[969.6969696969696, 0.0, 960.0],
[0.0, 969.6969696969696, 540.0],
[0.0, 0.0, 1.0],
]
poses = generate_camera_trajectory_local(motions)
custom_c2w = {}
for i, p in enumerate(poses):
custom_c2w[str(i)] = {"extrinsic": p.tolist(), "K": intrinsic}
json.dump(
custom_c2w,
open("./assets/pose/pose.json", "w"),
indent=4,
ensure_ascii=False,
)
-2
View File
@@ -1134,5 +1134,3 @@ class LongCatTransformer3DModel(CachableDiT):
)
return x
# Entry point for model registry
EntryClass = LongCatTransformer3DModel
File diff suppressed because it is too large Load Diff
@@ -9,6 +9,3 @@ __all__ = [
"CausalMatrixGameTransformerBlock",
"ActionModule",
]
# Entry point for model registry
EntryClass = [MatrixGameWanModel, CausalMatrixGameWanModel]

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