Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
a8dddfaa16 | ||
|
|
7210c68f1b | ||
|
|
5602dc1bad | ||
|
|
ac4bc4ab84 | ||
|
|
bee27f9f74 | ||
|
|
ad58f802f3 | ||
|
|
04fa356ee3 | ||
|
|
f9c076fe2b | ||
|
|
dff0ea401a | ||
|
|
f76efe798e | ||
|
|
09f455233e | ||
|
|
aea300f690 | ||
|
|
b92219f6a6 | ||
|
|
a321b95a8a | ||
|
|
becd379f58 | ||
|
|
1ed7d7e1b0 | ||
|
|
d9fabcc5ef | ||
|
|
9db48498de | ||
|
|
1cd7038315 | ||
|
|
98308db7e0 | ||
|
|
c1e18f6722 | ||
|
|
75e193a2c9 | ||
|
|
d6e0a7d0dd | ||
|
|
7fc5f241da | ||
|
|
aae48a7e90 | ||
|
|
88f38eb0f4 | ||
|
|
d750b463dc | ||
|
|
e10b26a3d8 | ||
|
|
74636ba246 | ||
|
|
7e2f3f14e7 | ||
|
|
caa1c402ba | ||
|
|
38a6bd93d3 | ||
|
|
b867ef7e7c | ||
|
|
3ae58c277a | ||
|
|
0c6862ca55 | ||
|
|
e8c854bcf1 | ||
|
|
06860e96fe | ||
|
|
1b503554d1 | ||
|
|
351ceb7c59 | ||
|
|
10875e0d7b | ||
|
|
1eaae8a10b | ||
|
|
59e00f6164 | ||
|
|
745cc05b10 | ||
|
|
c5dc244871 | ||
|
|
dbf3917bf4 | ||
|
|
050f189c95 | ||
|
|
029216029f |
@@ -0,0 +1,42 @@
|
||||
# Repository Guidelines
|
||||
|
||||
## Project Structure & Module Organization
|
||||
- Core Python package: `fastvideo/` (models, pipelines, training, distributed runtime, CLI entrypoints).
|
||||
- CUDA/custom kernels: `fastvideo-kernel/` (separate build/test flow).
|
||||
- Tests:
|
||||
- `fastvideo/tests/` for package-level tests (dataset, encoders, inference, training, SSIM, workflow).
|
||||
- `tests/local_tests/` for additional local/component checks.
|
||||
- Docs and guides: `docs/` (MkDocs source), with contributor docs in `docs/contributing/`.
|
||||
- Runnable examples and scripts: `examples/` and `scripts/`.
|
||||
- Static assets: `assets/`, `images/`, `videos/`, and `comfyui/assets/`.
|
||||
|
||||
## Build, Test, and Development Commands
|
||||
- `uv pip install -e .[dev]`: editable install with lint/test extras.
|
||||
- `pre-commit install --hook-type pre-commit --hook-type commit-msg`: enable local hooks.
|
||||
- `pre-commit run --all-files`: run formatter/lint/type/spelling checks.
|
||||
- `pytest tests/`: run top-level test suite.
|
||||
- `pytest fastvideo/tests/ -v`: run package tests.
|
||||
- `pytest fastvideo/tests/ssim/ -vs`: run SSIM regression tests (GPU-heavy).
|
||||
- `cd fastvideo-kernel && ./build.sh`: build kernel extensions.
|
||||
|
||||
## Coding Style & Naming Conventions
|
||||
- Python 3.10+; 4-space indentation; keep code and imports readable and explicit.
|
||||
- Style tools are configured in `pyproject.toml` and `.pre-commit-config.yaml`:
|
||||
- `yapf` (format), `ruff` (lint, auto-fix), `mypy` (typing), `codespell`.
|
||||
- Target line length is 80.
|
||||
- Naming: `snake_case` for functions/files, `PascalCase` for classes, `UPPER_SNAKE_CASE` for constants.
|
||||
|
||||
## Testing Guidelines
|
||||
- Use `pytest` and place tests near relevant domains (e.g., `fastvideo/tests/encoders/`).
|
||||
- Prefer descriptive names like `test_<feature>_<expected_behavior>.py`.
|
||||
- For new pipelines/backends, include at least one regression-oriented test; add SSIM coverage when output quality must be preserved.
|
||||
- Document GPU assumptions in tests that require specific hardware.
|
||||
|
||||
## Commit & Pull Request Guidelines
|
||||
- Follow existing commit style: short subject with optional tag prefix, e.g. `[bugfix]: ...`, `[feat]: ...`, `[misc]: ...`, and include PR reference like `(#1234)` when applicable.
|
||||
- Keep commits focused by concern (feature, refactor, fix).
|
||||
- PRs should include:
|
||||
- clear problem/solution summary,
|
||||
- test evidence (`pytest`/SSIM outputs or rationale if skipped),
|
||||
- linked issue/PR context,
|
||||
- screenshots or sample outputs for UI/demo/docs changes.
|
||||
@@ -9,27 +9,26 @@
|
||||
**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/).
|
||||
|
||||
<details>
|
||||
<summary>More</summary>
|
||||
- `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/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/).
|
||||
### More News
|
||||
|
||||
</details>
|
||||
- `2025/06/14`: Release finetuning and inference code for [VSA](https://arxiv.org/pdf/2505.13389)
|
||||
- `2025/04/24`: [FastVideo V1](https://hao-ai-lab.github.io/blogs/fastvideo/) is released!
|
||||
- `2025/02/18`: Release the inference code for [Sliding Tile Attention](https://hao-ai-lab.github.io/blogs/sta/).
|
||||
|
||||
## Key Features
|
||||
|
||||
FastVideo has the following features:
|
||||
|
||||
- End-to-end post-training support for bidirectional and autoregressive models:
|
||||
- Support full finetuning and LoRA finetuning for state-of-the-art open video DiTs
|
||||
- Data preprocessing pipeline for video, image, and text data
|
||||
- Distribution Matching Distillation (DMD2) stepwise distillation.
|
||||
- Sparse attention with [Video Sparse Attention](https://arxiv.org/pdf/2505.13389)
|
||||
- [Sparse distillation](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/) to achineve >50x denoising speedup
|
||||
- [Sparse distillation](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/) to achieve >50x denoising speedup
|
||||
- Scalable training with FSDP2, sequence parallelism, and selective activation checkpointing.
|
||||
- Causal distillation through Self-Forcing
|
||||
- See this [page](https://hao-ai-lab.github.io/FastVideo/training/overview/) for full list of supported models and recipes.
|
||||
@@ -44,6 +43,7 @@ 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
|
||||
@@ -58,18 +58,20 @@ 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.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) |
|
||||
| Model | Sparse Distillation | Dataset |
|
||||
| ------------------------------------------------------------------------------------- | --------------------------------------------------------------------------------------------------------------- | -------------------------------------------------------------------------------------------------------- |
|
||||
| [FastWan2.1-T2V-1.3B](https://huggingface.co/FastVideo/FastWan2.1-T2V-1.3B-Diffusers) | [Recipe](https://github.com/hao-ai-lab/FastVideo/tree/main/examples/distill/Wan2.1-T2V/Wan-Syn-Data-480P) | [FastVideo Synthetic Wan2.1 480P](https://huggingface.co/datasets/FastVideo/Wan-Syn_77x448x832_600k) |
|
||||
| [FastWan2.2-TI2V-5B](https://huggingface.co/FastVideo/FastWan2.2-TI2V-5B-Diffusers) | [Recipe](https://github.com/hao-ai-lab/FastVideo/tree/main/examples/distill/Wan2.2-TI2V-5B-Diffusers/Data-free) | [FastVideo Synthetic Wan2.2 720P](https://huggingface.co/datasets/FastVideo/Wan2.2-Syn-121x704x1280_32k) |
|
||||
|
||||
## Inference
|
||||
|
||||
### Generating Your First Video
|
||||
|
||||
Here's a minimal example to generate a video using the default settings. Make sure VSA kernels are [installed](https://hao-ai-lab.github.io/FastVideo/video_sparse_attention/installation/). Create a file called `example.py` with the following code:
|
||||
|
||||
```python
|
||||
@@ -108,55 +110,36 @@ 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/).
|
||||
|
||||
### Other docs:
|
||||
## More Guides
|
||||
|
||||
- [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/)
|
||||
<!-- - [Finetuning Guide](https://hao-ai-lab.github.io/FastVideo/training/finetune.html) -->
|
||||
- [Contribution Guide](https://hao-ai-lab.github.io/FastVideo/contributing/overview/)
|
||||
|
||||
## 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. [](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. [](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. [](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. [](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. [](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. [](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. [](https://github.com/meituan-longcat/LongCat-Video)
|
||||
- [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.
|
||||
|
||||
## 🤝 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 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.
|
||||
## Acknowledgement
|
||||
|
||||
We learned the design and reused code from the following projects: [Wan-Video](https://github.com/Wan-Video), [ThunderKittens](https://github.com/HazyResearch/ThunderKittens), [DMD2](https://github.com/tianweiy/DMD2), [diffusers](https://github.com/huggingface/diffusers), [xDiT](https://github.com/xdit-project/xDiT), [vLLM](https://github.com/vllm-project/vllm), [SGLang](https://github.com/sgl-project/sglang). We thank [MBZUAI](https://ifm.mbzuai.ac.ae/), [Anyscale](https://www.anyscale.com/), and [GMI Cloud](https://www.gmicloud.ai/) for their support throughout this project.
|
||||
|
||||
## Citation
|
||||
If you find FastVideo useful, please considering citing our work:
|
||||
|
||||
If you find FastVideo useful, please consider citing our research 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},
|
||||
|
||||
@@ -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.5.4/flash_attn-2.8.3%2Bcu128torch2.9-cp310-cp310-linux_x86_64.whl
|
||||
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
|
||||
|
||||
COPY . .
|
||||
|
||||
|
||||
@@ -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.5.4/flash_attn-2.8.3%2Bcu128torch2.9-cp311-cp311-linux_x86_64.whl
|
||||
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
|
||||
|
||||
COPY . .
|
||||
|
||||
|
||||
@@ -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.5.4/flash_attn-2.8.3%2Bcu128torch2.9-cp312-cp312-linux_x86_64.whl
|
||||
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
|
||||
|
||||
COPY . .
|
||||
|
||||
|
||||
@@ -121,6 +121,15 @@ 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
|
||||
|
||||
|
After Width: | Height: | Size: 211 KiB |
|
After Width: | Height: | Size: 117 KiB |
|
Before Width: | Height: | Size: 18 KiB After Width: | Height: | Size: 461 KiB |
|
After Width: | Height: | Size: 210 KiB |
|
Before Width: | Height: | Size: 27 KiB After Width: | Height: | Size: 437 KiB |
@@ -6,6 +6,8 @@ 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
|
||||
|
||||
|
||||
@@ -0,0 +1,176 @@
|
||||
# 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.
|
||||
@@ -0,0 +1,480 @@
|
||||
# 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`
|
||||
@@ -1,16 +1,55 @@
|
||||
|
||||
# 📦 Developing FastVideo on RunPod
|
||||
|
||||
You can easily use the FastVideo Docker image as a custom container on [RunPod](https://www.runpod.io) for development or experimentation.
|
||||
You can easily use the FastVideo Pod Template on [RunPod](https://www.runpod.io) for development or experimentation.
|
||||
|
||||
## Creating a new pod
|
||||
|
||||
Choose a GPU that supports CUDA 12.8
|
||||
|
||||
Pick 1 or 2 L40S GPU(s)
|
||||
- Make sure you are using the correct RunPod account.
|
||||

|
||||
|
||||
- Use "Additional Filters" to select CUDA 12.8.
|
||||

|
||||
|
||||
- Click "Deploy" and Pick a single A40 or RTX 4090 GPU.
|
||||

|
||||
|
||||
- Select the "FastVideo" or "fastvideo-dev" Pod Template.
|
||||

|
||||
|
||||
- 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.
|
||||

|
||||
|
||||
## 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:
|
||||
|
||||
```
|
||||
@@ -24,30 +63,3 @@ bash -c "apt update;DEBIAN_FRONTEND=noninteractive apt-get install openssh-serve
|
||||
```
|
||||
|
||||

|
||||
|
||||
After deploying, the pod will take a few minutes to pull the image and start the SSH service.
|
||||
|
||||

|
||||
|
||||
## 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/
|
||||
```
|
||||
|
||||
@@ -1,76 +1,84 @@
|
||||
|
||||
# 🛠️ Contributing to FastVideo
|
||||
|
||||
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!
|
||||
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.
|
||||
|
||||
Our community is open to everyone and welcomes any contributions no matter how large or small.
|
||||
## Quick prerequisites
|
||||
|
||||
# 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.
|
||||
- **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).
|
||||
|
||||
We recommend using a fresh Python 3.10 Conda environment to develop FastVideo:
|
||||
For a full install checklist, see `docs/getting_started/installation/gpu.md`.
|
||||
|
||||
## Local development (Conda + editable install)
|
||||
|
||||
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 for FastVideo:
|
||||
Create and activate a Conda environment:
|
||||
|
||||
```
|
||||
```bash
|
||||
conda create -n fastvideo python=3.12 -y
|
||||
conda activate fastvideo
|
||||
```
|
||||
|
||||
Install `uv` (optional, but recommended):
|
||||
|
||||
From instructions on [uv](https://astral.sh/uv/):
|
||||
|
||||
```
|
||||
```bash
|
||||
curl -LsSf https://astral.sh/uv/install.sh | sh
|
||||
# or
|
||||
# or
|
||||
wget -qO- https://astral.sh/uv/install.sh | sh
|
||||
```
|
||||
|
||||
Clone the FastVideo repository and go to the FastVideo directory:
|
||||
Clone the repo:
|
||||
|
||||
```
|
||||
```bash
|
||||
git clone https://github.com/hao-ai-lab/FastVideo.git && cd FastVideo
|
||||
|
||||
```
|
||||
|
||||
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.
|
||||
Install FastVideo in editable mode and set up hooks:
|
||||
|
||||
```bash
|
||||
uv pip install -e .[dev]
|
||||
|
||||
# Can also install flash-attn (optional)
|
||||
uv pip install flash-attn --no-build-isolation
|
||||
# Optional: FlashAttention (builds native kernels)
|
||||
uv pip install flash-attn --no-build-isolation
|
||||
|
||||
# Linting, formatting and static type checking
|
||||
# Linting, formatting, static typing
|
||||
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, you should also install [FA3](https://github.com/Dao-AILab/flash-attention) for much better performance:
|
||||
If you are on a Hopper GPU, installing FlashAttention 3 can improve
|
||||
performance (see `docs/inference/optimizations.md`).
|
||||
|
||||
```
|
||||
git clone https://github.com/Dao-AILab/flash-attention.git && cd flash-attention/hopper
|
||||
## Docker development (optional)
|
||||
|
||||
# make sure you have ninja installed
|
||||
uv pip install ninja
|
||||
|
||||
python setup.py install
|
||||
```
|
||||
If you prefer a containerized environment, use the dev image documented in
|
||||
`docs/contributing/developer_env/docker.md`.
|
||||
|
||||
## Testing
|
||||
|
||||
Please refer to the [Testing Guide](testing.md) for more information on how to add and run tests in FastVideo.
|
||||
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`.
|
||||
|
||||
@@ -1,424 +1,179 @@
|
||||
# 🔍 FastVideo Overview
|
||||
# FastVideo Architecture Overview
|
||||
|
||||
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.
|
||||
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.
|
||||
|
||||
## Table of Contents - Directory Structure and Files
|
||||
## FastVideo structure at a glance
|
||||
|
||||
- [`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 maps a Diffusers-style repo into a pipeline like this:
|
||||
|
||||
## Core Architecture
|
||||
- `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.
|
||||
|
||||
FastVideo separates model components from execution logic with these principles:
|
||||
Flow:
|
||||
`model_index.json` -> component loaders -> model modules -> pipeline stages ->
|
||||
sampling params.
|
||||
|
||||
- **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:
|
||||
Minimal usage (from `examples/inference/basic/basic.py`):
|
||||
|
||||
```python
|
||||
# Load arguments from command line
|
||||
fastvideo_args = prepare_fastvideo_args(sys.argv[1:])
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
|
||||
# Access parameters
|
||||
model = load_model(fastvideo_args.model_path)
|
||||
model_id = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # or official_weights/<model_name>/
|
||||
generator = VideoGenerator.from_pretrained(model_id, num_gpus=1)
|
||||
|
||||
# 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
|
||||
```
|
||||
|
||||

|
||||
|
||||
### 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.
|
||||
|
||||

|
||||
|
||||
## 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 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
|
||||
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,
|
||||
)
|
||||
```
|
||||
|
||||
### Sequence Parallelism
|
||||
## Configuration system
|
||||
|
||||
Sequence parallelism splits sequences across devices:
|
||||
FastVideo uses typed configs to keep model definitions, pipeline wiring, and
|
||||
runtime parameters consistent:
|
||||
|
||||
- **Implementation**: Through `DistributedAttention` and sequence splitting
|
||||
- **Use cases**: Long video sequences or high-resolution processing. Used by DiT models.
|
||||
- `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).
|
||||
|
||||
```python
|
||||
# Distributed attention for long sequences
|
||||
from fastvideo.attention import DistributedAttention
|
||||
`FastVideoArgs` (in `fastvideo/fastvideo_args.py`) provides runtime settings and
|
||||
is passed into pipeline construction and stages.
|
||||
|
||||
self.attn = DistributedAttention(
|
||||
num_heads=num_heads,
|
||||
head_size=head_dim,
|
||||
causal=False,
|
||||
supported_attention_backends=(_Backend.SLIDING_TILE_ATTN, _Backend.FLASH_ATTN)
|
||||
)
|
||||
## 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
|
||||
```
|
||||
|
||||
### Communication Primitives
|
||||
Key points:
|
||||
|
||||
Efficient distributed operations via AllGather, AllReduce, and synchronization mechanisms.
|
||||
- `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 communication primitives minimize distributed overhead:
|
||||
Note on tensor names:
|
||||
|
||||
- **Sequence-Parallel AllGather**: Collects sequence chunks
|
||||
- **Tensor-Parallel AllReduce**: Combines partial results
|
||||
- **Distributed Synchronization**: Coordinates execution
|
||||
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.
|
||||
|
||||
## Forward Context Management
|
||||
Example HF repo (Wan 2.1 T2V 1.3B Diffusers):
|
||||
|
||||
### 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)
|
||||
```
|
||||
https://huggingface.co/Wan-AI/Wan2.1-T2V-1.3B-Diffusers/tree/main
|
||||
```
|
||||
|
||||
## Executor and Worker System
|
||||
Example `model_index.json` from that repo:
|
||||
|
||||
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
|
||||
```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 platform system is designed to be extensible for future hardware targets.
|
||||
How this maps to FastVideo:
|
||||
|
||||
## Logger
|
||||
- `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`
|
||||
|
||||
See [PR](https://github.com/hao-ai-lab/FastVideo/pull/356)
|
||||
## Pipeline system
|
||||
|
||||
*TODO*: (help wanted) Add an environment variable that disables process-aware logging.
|
||||
- `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.
|
||||
|
||||
## Contributing to FastVideo
|
||||
## Model components
|
||||
|
||||
If you're a new contributor, here are some common areas to explore:
|
||||
- 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/`
|
||||
|
||||
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 and distributed execution
|
||||
|
||||
When adding code, follow these practices:
|
||||
- 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/`.
|
||||
|
||||
- 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
|
||||
## Related docs
|
||||
|
||||
- [Contributing overview](../contributing/overview.md)
|
||||
- [Coding agents workflow](../contributing/coding_agents.md)
|
||||
- [Testing guide](../contributing/testing.md)
|
||||
|
||||
@@ -131,17 +131,41 @@ self.attn = DistributedAttention(
|
||||
|
||||
### Registering Models
|
||||
|
||||
Register implemented modules in the model registry:
|
||||
Register implemented modules for auto‑discovery by adding `EntryClass` in each
|
||||
model module (the registry scans for it):
|
||||
|
||||
```python
|
||||
# In fastvideo/models/registry.py
|
||||
_TEXT_TO_VIDEO_DIT_MODELS = {
|
||||
"YourTransformerModel": ("dits", "yourmodule", "YourTransformerClass"),
|
||||
}
|
||||
# In fastvideo/models/dits/your_module.py
|
||||
class YourTransformerModel(...):
|
||||
...
|
||||
|
||||
_VAE_MODELS = {
|
||||
"YourVAEModel": ("vaes", "yourvae", "YourVAEClass"),
|
||||
}
|
||||
# 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(),
|
||||
],
|
||||
)
|
||||
```
|
||||
|
||||
## Step 2: Directory Structure
|
||||
|
||||
@@ -0,0 +1,140 @@
|
||||
# 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)
|
||||
```
|
||||
@@ -103,7 +103,7 @@ python setup.py install # or pip install -e .
|
||||
|
||||
**`SAGE_ATTN_THREE`**
|
||||
|
||||
[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.
|
||||
[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.
|
||||
|
||||
#### Hardware Requirements
|
||||
|
||||
@@ -113,11 +113,7 @@ 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, 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
|
||||
```
|
||||
To use Sage Attention 3 in FastVideo, follow the `README.md` in the linked repository to install the package from source.
|
||||
|
||||
## Teacache
|
||||
|
||||
|
||||
@@ -0,0 +1,51 @@
|
||||
# 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,4 +1,5 @@
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
|
||||
|
||||
def main():
|
||||
@@ -15,19 +16,27 @@ 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."
|
||||
)
|
||||
|
||||
video = generator.generate_video(
|
||||
generator.generate_video(
|
||||
prompt,
|
||||
negative_prompt="",
|
||||
height=704,
|
||||
width=1280,
|
||||
num_frames=77,
|
||||
num_inference_steps=35,
|
||||
guidance_scale=7.0,
|
||||
fps=24,
|
||||
sampling_param=sampling_param,
|
||||
output_path="outputs_video/cosmos2_5_t2w.mp4",
|
||||
save_video=True,
|
||||
)
|
||||
|
||||
@@ -0,0 +1,54 @@
|
||||
# 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()
|
||||
|
||||
|
||||
@@ -39,4 +39,4 @@ def main():
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
main()
|
||||
@@ -0,0 +1,43 @@
|
||||
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()
|
||||
@@ -0,0 +1,66 @@
|
||||
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()
|
||||
@@ -0,0 +1,39 @@
|
||||
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:
|
||||
# Uses FastVideo default sampling settings for LTX2 base.
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Davids048/LTX2-Base-Diffusers",
|
||||
num_gpus=1,
|
||||
)
|
||||
|
||||
output_path = "outputs_video/ltx2_basic/output_ltx2_base_t2v_1088_1920_1.1.mp4"
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
output_path=output_path,
|
||||
save_video=True,
|
||||
num_frames=121,
|
||||
height=1088,
|
||||
width=1920,
|
||||
)
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,34 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
PROMPT = (
|
||||
"A warm sunny backyard. The camera starts in a tight cinematic close-up "
|
||||
"of a woman and a man in their 30s, facing each other with serious "
|
||||
"expressions. The woman, emotional and dramatic, says softly, \"That's "
|
||||
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
|
||||
"annoyed: \"Stop being so dramatic, Jess.\" A beat. He glances aside, "
|
||||
"then mutters defensively, \"He's just having fun.\" The camera slowly "
|
||||
"pans right, revealing the grandfather in the garden wearing enormous "
|
||||
"butterfly wings, waving his arms in the air like he's trying to take "
|
||||
"off. He shouts, \"Wheeeew!\" as he flaps his wings with full commitment. "
|
||||
"The woman covers her face, on the verge of tears. The tone is deadpan, "
|
||||
"absurd, and quietly tragic."
|
||||
)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/LTX2-Distilled-Diffusers",
|
||||
num_gpus=1,
|
||||
)
|
||||
|
||||
output_path = "outputs_video/ltx2_basic/output_ltx2_distilled_t2v.mp4"
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
output_path=output_path,
|
||||
save_video=True,
|
||||
)
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,6 +1,5 @@
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.configs.pipelines.wan import MatrixGameI2V480PConfig
|
||||
from fastvideo.models.dits.matrix_game.utils import create_action_presets
|
||||
from fastvideo.models.dits.matrixgame.utils import create_action_presets
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
from fastvideo.entrypoints.streaming_generator import StreamingVideoGenerator
|
||||
from fastvideo.models.dits.matrix_game.utils import get_current_action_async, expand_action_to_frames
|
||||
from fastvideo.models.dits.matrixgame.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.matrix_game.utils import expand_action_to_frames
|
||||
from fastvideo.models.dits.matrixgame.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="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
|
||||
DATA_DIR="data/crush-smol_processed_t2v_1_3b_ode_init/"
|
||||
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
|
||||
NUM_GPUS=1
|
||||
@@ -14,7 +14,6 @@ NUM_GPUS=1
|
||||
training_args=(
|
||||
--tracker_project_name "wan_ode_init"
|
||||
--output_dir "wan_ode_init_crush_smol"
|
||||
--override_transformer_cls_name "CausalWanTransformer3DModel"
|
||||
--wandb_run_name "wan_ode_init_crush_smol"
|
||||
--max_train_steps 6000
|
||||
--train_batch_size 1
|
||||
@@ -34,7 +33,7 @@ parallel_args=(
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 1
|
||||
--hsdp_shard_dim 1
|
||||
--hsdp_shard_dim $NUM_GPUS
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
@@ -51,20 +50,17 @@ dataset_args=(
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 50
|
||||
--validation_sampling_steps "50"
|
||||
--validation_guidance_scale "6.0"
|
||||
--log-visualization
|
||||
--visualization-steps 100
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 6e-6
|
||||
--learning_rate 1e-5
|
||||
--mixed_precision "bf16"
|
||||
--weight_only_checkpointing_steps 1000
|
||||
--training_state_checkpointing_steps 1000
|
||||
--weight_decay 1e-4
|
||||
--weight_decay 0.01
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
|
||||
@@ -0,0 +1,93 @@
|
||||
#!/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[@]}"
|
||||
@@ -0,0 +1,27 @@
|
||||
#!/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
|
||||
@@ -0,0 +1,23 @@
|
||||
# LTX-2 Crush-Smol Example
|
||||
# TODO: Update this doc.
|
||||
|
||||
These are e2e example scripts for finetuning LTX-2 on the crush-smol dataset.
|
||||
|
||||
## Execute the following commands from `FastVideo/` to run training:
|
||||
|
||||
### Download crush-smol dataset:
|
||||
|
||||
`bash examples/training/finetune/ltx2/overfit/download_dataset.sh`
|
||||
|
||||
### Preprocess the videos and captions into latents:
|
||||
|
||||
`bash examples/training/finetune/ltx2/overfit/preprocess_ltx2_data_t2v_new.sh`
|
||||
|
||||
### Edit the following file and run finetuning:
|
||||
|
||||
`bash examples/training/finetune/ltx2/overfit/finetune_t2v.sh`
|
||||
|
||||
Notes:
|
||||
- Update `DATASET_PATH` in the preprocess script to point to your merged dataset root (`videos/` + `videos2caption.json`).
|
||||
- `MODEL_PATH` should point to a local LTX-2 diffusers-style directory that contains `model_index.json` and `text_encoder/gemma`.
|
||||
|
||||
@@ -0,0 +1,3 @@
|
||||
# #!/bin/bash
|
||||
#
|
||||
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
|
||||
@@ -0,0 +1,95 @@
|
||||
#!/bin/bash
|
||||
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
|
||||
MODEL_PATH="Davids048/LTX2-Base-Diffusers"
|
||||
# Also can use simple 1 video for overfitting experiments.
|
||||
# DATA_DIR="/home/hal-jundas/codes/FastVideo/data/crush-smol"
|
||||
DATA_DIR="<PATH_TO_PROCESSED_DATASET>"
|
||||
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
|
||||
echo VALIDATION_DATASET_FILE: $VALIDATION_DATASET_FILE
|
||||
NUM_GPUS=4
|
||||
OVERFIT_HEIGHT=480
|
||||
OVERFIT_WIDTH=832
|
||||
OVERFIT_FRAMES=73
|
||||
|
||||
training_args=(
|
||||
--tracker_project_name "ltx2_t2v_finetune"
|
||||
--output_dir "checkpoints/ltx2_t2v_finetune"
|
||||
--max_train_steps 5000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 10
|
||||
--num_height $OVERFIT_HEIGHT
|
||||
--num_width $OVERFIT_WIDTH
|
||||
--num_frames $OVERFIT_FRAMES
|
||||
--ltx2-first-frame-conditioning-p 0.1
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
--mode "finetuning"
|
||||
)
|
||||
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size $NUM_GPUS
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 1
|
||||
--hsdp_shard_dim $NUM_GPUS
|
||||
)
|
||||
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
dataset_args=(
|
||||
--data_path $DATA_DIR
|
||||
--dataloader_num_workers 1
|
||||
)
|
||||
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file $VALIDATION_DATASET_FILE
|
||||
--validation_steps 50
|
||||
--validation_sampling_steps "50"
|
||||
--validation_guidance_scale "3.0"
|
||||
)
|
||||
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-5
|
||||
--mixed_precision "bf16"
|
||||
--weight_only_checkpointing_steps 1000
|
||||
--training_state_checkpointing_steps 1000
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
--lr_scheduler "linear"
|
||||
)
|
||||
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--dit_precision "fp32"
|
||||
--dit_cpu_offload False
|
||||
--dit_layerwise_offload False
|
||||
--text_encoder_cpu_offload False
|
||||
--image_encoder_cpu_offload False
|
||||
--vae_cpu_offload False
|
||||
)
|
||||
|
||||
# NOTE: Setting this environment variable to TORCH_SDPA to avoid the issue of stacking that failed in flash attn.
|
||||
export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
|
||||
torchrun \
|
||||
--nnodes 1 \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
fastvideo/training/ltx2_training_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}"
|
||||
@@ -0,0 +1,80 @@
|
||||
#!/bin/bash
|
||||
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
|
||||
MODEL_PATH="/path/to/LTX-2"
|
||||
DATA_DIR="data/crush-smol"
|
||||
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
|
||||
NUM_GPUS=1
|
||||
|
||||
training_args=(
|
||||
--tracker_project_name "ltx2_t2v_lora_finetune"
|
||||
--output_dir "checkpoints/ltx2_t2v_lora_finetune"
|
||||
--max_train_steps 2000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 8
|
||||
--num_latent_t 10
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 77
|
||||
--ltx2-first-frame-conditioning-p 0.1
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size $NUM_GPUS
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 1
|
||||
--hsdp_shard_dim $NUM_GPUS
|
||||
)
|
||||
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
dataset_args=(
|
||||
--data_path $DATA_DIR
|
||||
--dataloader_num_workers 1
|
||||
)
|
||||
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file $VALIDATION_DATASET_FILE
|
||||
--validation_steps 200
|
||||
--validation_sampling_steps "50"
|
||||
--validation_guidance_scale "3.0"
|
||||
)
|
||||
|
||||
optimizer_args=(
|
||||
--learning_rate 2e-4
|
||||
--mixed_precision "bf16"
|
||||
--weight_only_checkpointing_steps 1000
|
||||
--training_state_checkpointing_steps 1000
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
--lora_training True
|
||||
--lora_rank 16
|
||||
--lora_alpha 16
|
||||
)
|
||||
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
)
|
||||
|
||||
torchrun \
|
||||
--nnodes 1 \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
fastvideo/training/ltx2_training_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}"
|
||||
@@ -0,0 +1,35 @@
|
||||
#!/bin/bash
|
||||
|
||||
GPU_NUM=1
|
||||
MODEL_PATH="Davids048/LTX2-Base-Diffusers"
|
||||
# DATASET_PATH="data/overfit"
|
||||
DATASET_PATH="data/crush-smol"
|
||||
OUTPUT_DIR="$DATASET_PATH"
|
||||
WITH_AUDIO=true
|
||||
|
||||
# Convert one-file overfit metadata into merged format if needed.
|
||||
if [ ! -f "$DATASET_PATH/videos2caption.json" ] && [ -f "$DATASET_PATH/overfit.json" ]; then
|
||||
python scripts/dataset_preparation/convert_to_merged_dataset.py \
|
||||
--items-json "$DATASET_PATH/overfit.json" \
|
||||
--output-dir "$DATASET_PATH"
|
||||
fi
|
||||
|
||||
|
||||
torchrun --nproc_per_node=$GPU_NUM \
|
||||
--master_port=29513 \
|
||||
-m fastvideo.pipelines.preprocess.v1_preprocessing_new \
|
||||
--model_path $MODEL_PATH \
|
||||
--mode preprocess \
|
||||
--workload_type t2v \
|
||||
--preprocess.video_loader_type torchvision \
|
||||
--preprocess.dataset_type merged \
|
||||
--preprocess.dataset_path $DATASET_PATH \
|
||||
--preprocess.dataset_output_dir $OUTPUT_DIR \
|
||||
--preprocess.with_audio $WITH_AUDIO \
|
||||
--preprocess.preprocess_video_batch_size 1 \
|
||||
--preprocess.dataloader_num_workers 0 \
|
||||
--preprocess.max_height 480 \
|
||||
--preprocess.max_width 832 \
|
||||
--preprocess.num_frames 73 \
|
||||
--preprocess.train_fps 16 \
|
||||
--preprocess.video_length_tolerance_range 5
|
||||
@@ -0,0 +1,13 @@
|
||||
{
|
||||
"data": [
|
||||
{
|
||||
"caption": "The camera opens in a calm, sunlit frog yoga studio. Warm morning light washes over the wooden floor as incense smoke drifts lazily in the air. The senior frog instructor sits cross-legged at the center, eyes closed, voice deep and calm. “We are one with the pond.” All the frogs answer softly: “Ommm...” “We are one with the mud.” “Ommm...” He smiles faintly. “We are one with the flies.” A quiet pause. The camera slowly pans to the side — one frog twitches, eyes darting. Suddenly — *thwip!* — its tongue snaps out, catching a fly mid-air and pulling it into its mouth. The master exhales slowly, still serene.",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 50,
|
||||
"height": 1088,
|
||||
"width": 1920,
|
||||
"num_frames": 121
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -0,0 +1,31 @@
|
||||
{
|
||||
"data": [
|
||||
{
|
||||
"caption": "A large metal cylinder is seen pressing down on a pile of Oreo cookies, flattening them as if they were under a hydraulic press.",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 50,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 50,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A large metal cylinder is seen pressing down on a pile of colorful candies, flattening them as if they were under a hydraulic press. The candies are crushed and broken into small pieces, creating a mess on the table.",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 50,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -40,6 +40,17 @@ 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.
|
||||
|
||||
@@ -0,0 +1,166 @@
|
||||
#!/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()
|
||||
|
||||
|
||||
@@ -9,7 +9,7 @@ build-backend = "scikit_build_core.build"
|
||||
|
||||
[project]
|
||||
name = "fastvideo-kernel"
|
||||
version = "0.2.4"
|
||||
version = "0.2.5"
|
||||
description = "Unified CUDA kernels for FastVideo"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.10"
|
||||
|
||||
@@ -30,13 +30,12 @@ def _force_triton() -> bool:
|
||||
return os.environ.get("FASTVIDEO_KERNEL_VSA_FORCE_TRITON", "0") == "1"
|
||||
|
||||
|
||||
def _map_to_index_torch(block_map: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
def _map_to_index(block_map: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
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)
|
||||
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.
|
||||
"""
|
||||
if block_map.dim() == 3:
|
||||
block_map = block_map.unsqueeze(0)
|
||||
@@ -45,20 +44,17 @@ def _map_to_index_torch(block_map: torch.Tensor) -> Tuple[torch.Tensor, torch.Te
|
||||
if block_map.dtype != torch.bool:
|
||||
block_map = block_map.to(torch.bool)
|
||||
|
||||
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)
|
||||
if not block_map.is_cuda:
|
||||
raise RuntimeError("block_map must be a CUDA tensor (Triton map_to_index required).")
|
||||
|
||||
# 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
|
||||
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)
|
||||
|
||||
|
||||
@torch.library.custom_op(
|
||||
@@ -77,7 +73,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_torch(block_map)
|
||||
q2k_idx, q2k_num = _map_to_index(block_map)
|
||||
|
||||
from fastvideo_kernel.triton_kernels.block_sparse_attn_triton import ( # local import
|
||||
triton_block_sparse_attn_forward,
|
||||
@@ -87,6 +83,7 @@ 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,
|
||||
@@ -117,8 +114,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_torch(block_map)
|
||||
k2q_idx, k2q_num = _map_to_index_torch(block_map.transpose(-1, -2).contiguous())
|
||||
q2k_idx, q2k_num = _map_to_index(block_map)
|
||||
k2q_idx, k2q_num = _map_to_index(block_map.transpose(-1, -2).contiguous())
|
||||
|
||||
from fastvideo_kernel.triton_kernels.block_sparse_attn_triton import ( # local import
|
||||
triton_block_sparse_attn_backward,
|
||||
@@ -182,7 +179,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_torch(block_map)
|
||||
q2k_idx, q2k_num = _map_to_index(block_map)
|
||||
|
||||
o_padded, lse_padded = block_sparse_fwd(
|
||||
q_padded, k_padded, v_padded, q2k_idx, q2k_num, variable_block_sizes.int()
|
||||
@@ -224,7 +221,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_torch(block_map.transpose(-1, -2).contiguous())
|
||||
k2q_idx, k2q_num = _map_to_index(block_map.transpose(-1, -2).contiguous())
|
||||
|
||||
dq, dk, dv = block_sparse_bwd(
|
||||
q_padded,
|
||||
|
||||
@@ -1 +1 @@
|
||||
__version__ = "0.2.4"
|
||||
__version__ = "0.2.5"
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import torch
|
||||
from fastvideo.attention.backends.sageattn.api import sageattn_blackwell
|
||||
from sageattn3 import sageattn3_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 = sageattn_blackwell(query, key, value, is_causal=self.causal)
|
||||
output = sageattn3_blackwell(query, key, value, is_causal=self.causal)
|
||||
output = output.transpose(1, 2)
|
||||
return output
|
||||
|
||||
@@ -86,6 +86,7 @@ class PreprocessConfig:
|
||||
|
||||
# Model configuration
|
||||
training_cfg_rate: float = 0.0
|
||||
with_audio: bool = False
|
||||
|
||||
# framework configuration
|
||||
seed: int = 42
|
||||
@@ -190,6 +191,10 @@ class PreprocessConfig:
|
||||
type=float,
|
||||
default=PreprocessConfig.training_cfg_rate,
|
||||
help="Training CFG rate")
|
||||
preprocess_args.add_argument(f"--{prefix_with_dot}with-audio",
|
||||
action=StoreBoolean,
|
||||
default=PreprocessConfig.with_audio,
|
||||
help="Whether to extract and encode audio")
|
||||
preprocess_args.add_argument(f"--{prefix_with_dot}seed",
|
||||
type=int,
|
||||
default=PreprocessConfig.seed,
|
||||
|
||||
@@ -2,5 +2,18 @@ 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"]
|
||||
__all__ = [
|
||||
"ModelConfig",
|
||||
"VAEConfig",
|
||||
"DiTConfig",
|
||||
"EncoderConfig",
|
||||
"LTX2AudioEncoderConfig",
|
||||
"LTX2AudioDecoderConfig",
|
||||
"LTX2VocoderConfig",
|
||||
"UpsamplerConfig",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,13 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from fastvideo.configs.models.audio.ltx2_audio_vae import (
|
||||
LTX2AudioDecoderConfig,
|
||||
LTX2AudioEncoderConfig,
|
||||
LTX2VocoderConfig,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"LTX2AudioEncoderConfig",
|
||||
"LTX2AudioDecoderConfig",
|
||||
"LTX2VocoderConfig",
|
||||
]
|
||||
@@ -0,0 +1,31 @@
|
||||
# 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"]))
|
||||
@@ -3,11 +3,13 @@ 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"
|
||||
"LongCatVideoConfig", "LTX2VideoConfig", "HYWorldConfig"
|
||||
]
|
||||
|
||||
@@ -55,6 +55,8 @@ 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\.(.*)$":
|
||||
|
||||
@@ -0,0 +1,202 @@
|
||||
# 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"
|
||||
@@ -0,0 +1,86 @@
|
||||
# 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
|
||||
|
||||
import re
|
||||
|
||||
|
||||
def is_ltx2_blocks(name: str, _module) -> bool:
|
||||
res = re.search(r"(?:^|\.)transformer_blocks\.\d+$", name) is not None
|
||||
return res
|
||||
|
||||
|
||||
@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"
|
||||
@@ -1,4 +1,6 @@
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
import torch
|
||||
from fastvideo.configs.models.dits.wanvideo import WanVideoArchConfig, WanVideoConfig
|
||||
|
||||
|
||||
@@ -8,8 +10,8 @@ class MatrixGameWanVideoArchConfig(WanVideoArchConfig):
|
||||
# because MatrixGame checkpoints already have patch_embedding.proj format
|
||||
param_names_mapping: dict = field(
|
||||
default_factory=lambda: {
|
||||
# Removed: r"^patch_embedding\.(.*)$": r"patch_embedding.proj.\1"
|
||||
# because checkpoint already has correct format
|
||||
r"^patch_embedding\.(?!proj\.)(.*)$":
|
||||
r"patch_embedding.proj.\1",
|
||||
r"^condition_embedder\.text_embedder\.linear_1\.(.*)$":
|
||||
r"condition_embedder.text_embedder.fc_in.\1",
|
||||
r"^condition_embedder\.text_embedder\.linear_2\.(.*)$":
|
||||
@@ -76,8 +78,14 @@ 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,11 +7,14 @@ 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"
|
||||
"Qwen2_5_VLConfig", "Reason1ArchConfig", "Reason1Config", "LTX2GemmaConfig",
|
||||
"SiglipVisionConfig"
|
||||
]
|
||||
|
||||
@@ -87,10 +87,10 @@ class CLIPVisionConfig(ImageEncoderConfig):
|
||||
arch_config: ImageEncoderArchConfig = field(
|
||||
default_factory=CLIPVisionArchConfig)
|
||||
|
||||
num_hidden_layers_override: int | None = None
|
||||
num_hidden_layers_override: int | None = 31
|
||||
require_post_norm: bool | None = None
|
||||
enable_scale: bool = True
|
||||
is_causal: bool = True
|
||||
enable_scale: bool = False
|
||||
is_causal: bool = False
|
||||
prefix: str = "clip"
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,48 @@
|
||||
# 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"
|
||||
@@ -0,0 +1,53 @@
|
||||
# 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"
|
||||
@@ -0,0 +1,6 @@
|
||||
from fastvideo.configs.models.upsamplers.hunyuan15 import SRTo720pUpsamplerConfig, SRTo1080pUpsamplerConfig
|
||||
from fastvideo.configs.models.upsamplers.base import UpsamplerConfig
|
||||
|
||||
__all__ = [
|
||||
"SRTo720pUpsamplerConfig", "SRTo1080pUpsamplerConfig", "UpsamplerConfig"
|
||||
]
|
||||
@@ -0,0 +1,7 @@
|
||||
from dataclasses import dataclass
|
||||
from fastvideo.configs.models.base import ModelConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class UpsamplerConfig(ModelConfig):
|
||||
pass
|
||||
@@ -0,0 +1,20 @@
|
||||
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,6 +2,7 @@ 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
|
||||
|
||||
@@ -12,4 +13,5 @@ __all__ = [
|
||||
"CosmosVAEConfig",
|
||||
"Cosmos25VAEConfig",
|
||||
"Hunyuan15VAEConfig",
|
||||
"LTX2VAEConfig",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,45 @@
|
||||
# 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
|
||||
@@ -4,8 +4,9 @@ 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.registry import (
|
||||
get_pipeline_config_cls_from_name)
|
||||
from fastvideo.configs.pipelines.hyworld import HYWorldConfig
|
||||
from fastvideo.configs.pipelines.ltx2 import LTX2T2VConfig
|
||||
from fastvideo.registry import get_pipeline_config_cls_from_name
|
||||
from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
|
||||
from fastvideo.configs.pipelines.wan import (SelfForcingWanT2V480PConfig,
|
||||
WanI2V480PConfig, WanI2V720PConfig,
|
||||
@@ -16,5 +17,6 @@ __all__ = [
|
||||
"Hunyuan15T2V480PConfig", "Hunyuan15T2V720PConfig", "SlidingTileAttnConfig",
|
||||
"WanT2V480PConfig", "WanI2V480PConfig", "WanT2V720PConfig",
|
||||
"WanI2V720PConfig", "StepVideoT2VConfig", "SelfForcingWanT2V480PConfig",
|
||||
"CosmosConfig", "Cosmos25Config", "get_pipeline_config_cls_from_name"
|
||||
"CosmosConfig", "Cosmos25Config", "LTX2T2VConfig", "HYWorldConfig",
|
||||
"get_pipeline_config_cls_from_name"
|
||||
]
|
||||
|
||||
@@ -8,7 +8,7 @@ from typing import Any, cast
|
||||
import torch
|
||||
|
||||
from fastvideo.configs.models import (DiTConfig, EncoderConfig, ModelConfig,
|
||||
VAEConfig)
|
||||
VAEConfig, UpsamplerConfig)
|
||||
from fastvideo.configs.models.encoders import BaseEncoderOutput
|
||||
from fastvideo.configs.utils import update_config_from_args
|
||||
from fastvideo.logger import init_logger
|
||||
@@ -44,12 +44,15 @@ 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)
|
||||
@@ -216,6 +219,24 @@ 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")
|
||||
@@ -245,8 +266,7 @@ class PipelineConfig:
|
||||
"""
|
||||
use the pipeline class setting from model_path to match the pipeline config
|
||||
"""
|
||||
from fastvideo.configs.pipelines.registry import (
|
||||
get_pipeline_config_cls_from_name)
|
||||
from fastvideo.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))
|
||||
@@ -260,8 +280,7 @@ class PipelineConfig:
|
||||
kwargs: dictionary of kwargs
|
||||
config_cli_prefix: prefix of CLI arguments for this PipelineConfig instance
|
||||
"""
|
||||
from fastvideo.configs.pipelines.registry import (
|
||||
get_pipeline_config_cls_from_name)
|
||||
from fastvideo.registry import get_pipeline_config_cls_from_name
|
||||
|
||||
prefix_with_dot = f"{config_cli_prefix}." if (config_cli_prefix.strip()
|
||||
!= "") else ""
|
||||
@@ -298,6 +317,12 @@ 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:
|
||||
|
||||
@@ -11,7 +11,8 @@ 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.pipelines.base import PipelineConfig
|
||||
from fastvideo.configs.models.upsamplers import SRTo720pUpsamplerConfig, SRTo1080pUpsamplerConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig, UpsamplerConfig
|
||||
|
||||
PROMPT_TEMPLATE_TOKEN_LENGTH = 108
|
||||
|
||||
@@ -131,9 +132,35 @@ 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"
|
||||
|
||||
@@ -0,0 +1,29 @@
|
||||
# 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
|
||||
@@ -0,0 +1,50 @@
|
||||
# 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
|
||||
@@ -1,201 +0,0 @@
|
||||
# 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
|
||||
@@ -196,6 +196,12 @@ 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)
|
||||
|
||||
@@ -28,6 +28,9 @@ 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)
|
||||
@@ -55,10 +58,13 @@ 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
|
||||
@@ -92,8 +98,7 @@ class SamplingParam:
|
||||
|
||||
@classmethod
|
||||
def from_pretrained(cls, model_path: str) -> "SamplingParam":
|
||||
from fastvideo.configs.sample.registry import (
|
||||
get_sampling_param_cls_for_name)
|
||||
from fastvideo.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()
|
||||
|
||||
@@ -5,15 +5,19 @@ from fastvideo.configs.sample.base import SamplingParam
|
||||
|
||||
|
||||
@dataclass
|
||||
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
|
||||
class Cosmos25SamplingParamBase(SamplingParam):
|
||||
height: int = 704
|
||||
width: int = 1280
|
||||
num_frames: int = 77
|
||||
fps: int = 24
|
||||
seed: int = 0
|
||||
|
||||
guidance_scale: float = 7.0
|
||||
# Official Cosmos2.5 sampling uses empty string as unconditional.
|
||||
negative_prompt: str = ""
|
||||
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.")
|
||||
num_inference_steps: int = 35
|
||||
|
||||
@@ -22,8 +22,39 @@ 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
|
||||
|
||||
@@ -0,0 +1,26 @@
|
||||
# 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 = ""
|
||||
@@ -0,0 +1,58 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
|
||||
|
||||
@dataclass
|
||||
class LTX2BaseSamplingParam(SamplingParam):
|
||||
"""Default sampling parameters for LTX-2 base one-stage T2V.
|
||||
|
||||
Values follow the official LTX-2 one-stage defaults.
|
||||
"""
|
||||
|
||||
seed: int = 10
|
||||
num_frames: int = 121
|
||||
height: int = 512
|
||||
width: int = 768
|
||||
fps: int = 24
|
||||
num_inference_steps: int = 40
|
||||
guidance_scale: float = 3.0
|
||||
# Copied/following official LTX-2 DEFAULT_NEGATIVE_PROMPT.
|
||||
negative_prompt: str = (
|
||||
"blurry, out of focus, overexposed, underexposed, low contrast, "
|
||||
"washed out colors, excessive noise, grainy texture, poor lighting, "
|
||||
"flickering, motion blur, distorted proportions, unnatural skin "
|
||||
"tones, deformed facial features, asymmetrical face, missing facial "
|
||||
"features, extra limbs, disfigured hands, wrong hand count, "
|
||||
"artifacts around text, inconsistent perspective, camera shake, "
|
||||
"incorrect depth of field, background too sharp, background clutter, "
|
||||
"distracting reflections, harsh shadows, inconsistent lighting "
|
||||
"direction, color banding, cartoonish rendering, 3D CGI look, "
|
||||
"unrealistic materials, uncanny valley effect, incorrect ethnicity, "
|
||||
"wrong gender, exaggerated expressions, wrong gaze direction, "
|
||||
"mismatched lip sync, silent or muted audio, distorted voice, "
|
||||
"robotic voice, echo, background noise, off-sync audio, incorrect "
|
||||
"dialogue, added dialogue, repetitive speech, jittery movement, "
|
||||
"awkward pauses, incorrect timing, unnatural transitions, "
|
||||
"inconsistent framing, tilted camera, flat lighting, inconsistent "
|
||||
"tone, cinematic oversaturation, stylized filters, or AI artifacts.")
|
||||
|
||||
|
||||
@dataclass
|
||||
class LTX2DistilledSamplingParam(SamplingParam):
|
||||
"""Default sampling parameters for LTX-2 distilled one-stage T2V."""
|
||||
|
||||
seed: int = 10
|
||||
num_frames: int = 121
|
||||
height: int = 1024
|
||||
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 = ""
|
||||
|
||||
|
||||
# Backward compatibility alias.
|
||||
LTX2SamplingParam = LTX2DistilledSamplingParam
|
||||
@@ -1,206 +0,0 @@
|
||||
# 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
|
||||
@@ -4,6 +4,8 @@ from torchvision.transforms import Lambda
|
||||
|
||||
from fastvideo.dataset.parquet_dataset_map_style import (
|
||||
build_parquet_map_style_dataloader)
|
||||
from fastvideo.dataset.ltx2_precomputed_dataset import (
|
||||
build_ltx2_precomputed_dataloader, LTX2PrecomputedDataset)
|
||||
from fastvideo.dataset.preprocessing_datasets import VideoCaptionMergedDataset, TextDataset
|
||||
from fastvideo.dataset.transform import (CenterCropResizeVideo, Normalize255,
|
||||
TemporalRandomCrop)
|
||||
@@ -46,6 +48,10 @@ def gettextdataset(args) -> TextDataset:
|
||||
|
||||
|
||||
__all__ = [
|
||||
"build_parquet_map_style_dataloader", "ValidationDataset",
|
||||
"VideoCaptionMergedDataset", "TextDataset"
|
||||
"build_parquet_map_style_dataloader",
|
||||
"build_ltx2_precomputed_dataloader",
|
||||
"LTX2PrecomputedDataset",
|
||||
"ValidationDataset",
|
||||
"VideoCaptionMergedDataset",
|
||||
"TextDataset",
|
||||
]
|
||||
|
||||
@@ -116,3 +116,44 @@ 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()),
|
||||
])
|
||||
|
||||
@@ -0,0 +1,210 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Dataset utilities for loading LTX2 precomputed training artifacts.
|
||||
#
|
||||
# Usage:
|
||||
# - Input root can be either `<data_root>/` or `<data_root>/.precomputed/`.
|
||||
# - Required sources are `latents/` and `conditions/` with matching `.pt` files.
|
||||
# - Optional source `audio_latents/` is loaded when provided in `data_sources`.
|
||||
# - `build_ltx2_precomputed_dataloader(...)` is the intended entrypoint used by
|
||||
# `fastvideo/training/ltx2_training_pipeline.py`.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from einops import rearrange
|
||||
from torch.utils.data import Dataset
|
||||
from torchdata.stateful_dataloader import StatefulDataLoader
|
||||
|
||||
from fastvideo.dataset.parquet_dataset_map_style import DP_SP_BatchSampler
|
||||
from fastvideo.distributed import get_sp_world_size, get_world_rank, get_world_size
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
PRECOMPUTED_DIR_NAME = ".precomputed"
|
||||
|
||||
|
||||
class LTX2PrecomputedDataset(Dataset):
|
||||
"""Dataset for LTX-2 precomputed latents and conditions.
|
||||
|
||||
Expected directory structure (data_root):
|
||||
.precomputed/
|
||||
latents/*.pt
|
||||
conditions/*.pt
|
||||
audio_latents/*.pt (optional)
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
data_root: str,
|
||||
data_sources: dict[str, str] | list[str] | None = None,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.data_root = self._setup_data_root(data_root)
|
||||
self.data_sources = self._normalize_data_sources(data_sources)
|
||||
self.source_paths = self._setup_source_paths()
|
||||
self.sample_files = self._discover_samples()
|
||||
self._validate_setup()
|
||||
|
||||
@staticmethod
|
||||
def _setup_data_root(data_root: str) -> Path:
|
||||
data_root_path = Path(data_root).expanduser().resolve()
|
||||
if not data_root_path.exists():
|
||||
raise FileNotFoundError(
|
||||
f"Data root directory does not exist: {data_root_path}")
|
||||
if (data_root_path / PRECOMPUTED_DIR_NAME).exists():
|
||||
data_root_path = data_root_path / PRECOMPUTED_DIR_NAME
|
||||
return data_root_path
|
||||
|
||||
@staticmethod
|
||||
def _normalize_data_sources(
|
||||
data_sources: dict[str, str] | list[str] | None,
|
||||
) -> dict[str, str]:
|
||||
if data_sources is None:
|
||||
return {"latents": "latents", "conditions": "conditions"}
|
||||
if isinstance(data_sources, list):
|
||||
return {source: source for source in data_sources}
|
||||
if isinstance(data_sources, dict):
|
||||
return data_sources.copy()
|
||||
raise TypeError(
|
||||
f"data_sources must be dict, list, or None, got {type(data_sources)}")
|
||||
|
||||
def _setup_source_paths(self) -> dict[str, Path]:
|
||||
source_paths: dict[str, Path] = {}
|
||||
for dir_name in self.data_sources:
|
||||
source_path = self.data_root / dir_name
|
||||
if not source_path.exists():
|
||||
raise FileNotFoundError(
|
||||
f"Required {dir_name} directory does not exist: {source_path}")
|
||||
source_paths[dir_name] = source_path
|
||||
return source_paths
|
||||
|
||||
def _discover_samples(self) -> dict[str, list[Path]]:
|
||||
data_key = ("latents"
|
||||
if "latents" in self.data_sources else next(iter(
|
||||
self.data_sources.keys())))
|
||||
data_path = self.source_paths[data_key]
|
||||
data_files = list(data_path.glob("**/*.pt"))
|
||||
if not data_files:
|
||||
raise ValueError(f"No data files found in {data_path}")
|
||||
|
||||
sample_files = {output_key: [] for output_key in self.data_sources.values()}
|
||||
for data_file in data_files:
|
||||
rel_path = data_file.relative_to(data_path)
|
||||
if self._all_source_files_exist(data_file, rel_path):
|
||||
self._fill_sample_data_files(data_file, rel_path, sample_files)
|
||||
return sample_files
|
||||
|
||||
def _all_source_files_exist(self, data_file: Path, rel_path: Path) -> bool:
|
||||
for dir_name in self.data_sources:
|
||||
expected_path = self._get_expected_file_path(dir_name, data_file,
|
||||
rel_path)
|
||||
if not expected_path.exists():
|
||||
logger.warning(
|
||||
"No matching %s file found for: %s (expected in: %s)",
|
||||
dir_name,
|
||||
data_file.name,
|
||||
expected_path,
|
||||
)
|
||||
return False
|
||||
return True
|
||||
|
||||
def _get_expected_file_path(self, dir_name: str, data_file: Path,
|
||||
rel_path: Path) -> Path:
|
||||
source_path = self.source_paths[dir_name]
|
||||
if dir_name == "conditions" and data_file.name.startswith("latent_"):
|
||||
return source_path / f"condition_{data_file.stem[7:]}.pt"
|
||||
return source_path / rel_path
|
||||
|
||||
def _fill_sample_data_files(self, data_file: Path, rel_path: Path,
|
||||
sample_files: dict[str, list[Path]]) -> None:
|
||||
for dir_name, output_key in self.data_sources.items():
|
||||
expected_path = self._get_expected_file_path(dir_name, data_file,
|
||||
rel_path)
|
||||
sample_files[output_key].append(
|
||||
expected_path.relative_to(self.source_paths[dir_name]))
|
||||
|
||||
def _validate_setup(self) -> None:
|
||||
if not self.sample_files:
|
||||
raise ValueError(
|
||||
"No valid samples found - all data sources must have matching files"
|
||||
)
|
||||
sample_counts = {
|
||||
key: len(files)
|
||||
for key, files in self.sample_files.items()
|
||||
}
|
||||
if len(set(sample_counts.values())) > 1:
|
||||
raise ValueError(
|
||||
f"Mismatched sample counts across sources: {sample_counts}")
|
||||
|
||||
def __len__(self) -> int:
|
||||
first_key = next(iter(self.sample_files.keys()))
|
||||
return len(self.sample_files[first_key])
|
||||
|
||||
def __getitem__(self, index: int) -> dict[str, torch.Tensor]:
|
||||
result: dict[str, Any] = {}
|
||||
for dir_name, output_key in self.data_sources.items():
|
||||
source_path = self.source_paths[dir_name]
|
||||
file_rel_path = self.sample_files[output_key][index]
|
||||
file_path = source_path / file_rel_path
|
||||
try:
|
||||
data = torch.load(file_path, map_location="cpu", weights_only=True)
|
||||
if "latent" in dir_name.lower():
|
||||
data = self._normalize_video_latents(data)
|
||||
result[output_key] = data
|
||||
except Exception as e:
|
||||
raise RuntimeError(
|
||||
f"Failed to load {output_key} from {file_path}: {e}") from e
|
||||
result["idx"] = index
|
||||
return result
|
||||
|
||||
@staticmethod
|
||||
def _normalize_video_latents(data: dict) -> dict:
|
||||
latents = data["latents"]
|
||||
if latents.dim() == 2:
|
||||
num_frames = data["num_frames"]
|
||||
height = data["height"]
|
||||
width = data["width"]
|
||||
latents = rearrange(
|
||||
latents,
|
||||
"(f h w) c -> c f h w",
|
||||
f=num_frames,
|
||||
h=height,
|
||||
w=width,
|
||||
)
|
||||
data = data.copy()
|
||||
data["latents"] = latents
|
||||
return data
|
||||
|
||||
|
||||
def build_ltx2_precomputed_dataloader(
|
||||
path: str,
|
||||
batch_size: int,
|
||||
num_data_workers: int,
|
||||
data_sources: dict[str, str] | list[str] | None = None,
|
||||
drop_last: bool = True,
|
||||
seed: int = 42,
|
||||
) -> tuple[LTX2PrecomputedDataset, StatefulDataLoader]:
|
||||
dataset = LTX2PrecomputedDataset(path, data_sources=data_sources)
|
||||
sampler = DP_SP_BatchSampler(
|
||||
batch_size=batch_size,
|
||||
dataset_size=len(dataset),
|
||||
num_sp_groups=get_world_size() // get_sp_world_size(),
|
||||
sp_world_size=get_sp_world_size(),
|
||||
global_rank=get_world_rank(),
|
||||
drop_last=drop_last,
|
||||
drop_first_row=False,
|
||||
seed=seed,
|
||||
)
|
||||
loader = StatefulDataLoader(
|
||||
dataset,
|
||||
batch_sampler=sampler,
|
||||
collate_fn=None,
|
||||
num_workers=num_data_workers,
|
||||
pin_memory=True,
|
||||
persistent_workers=num_data_workers > 0,
|
||||
)
|
||||
return dataset, loader
|
||||
@@ -40,6 +40,7 @@ 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
|
||||
@@ -147,62 +148,6 @@ 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."""
|
||||
|
||||
@@ -327,11 +272,6 @@ 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
|
||||
@@ -365,9 +305,11 @@ 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
|
||||
|
||||
@@ -470,8 +412,11 @@ class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
|
||||
|
||||
# Initialize tokenizer
|
||||
tokenizer_path = os.path.join(args.model_path, "tokenizer")
|
||||
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path,
|
||||
cache_dir=args.cache_dir)
|
||||
if os.path.exists(tokenizer_path):
|
||||
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path,
|
||||
cache_dir=args.cache_dir)
|
||||
else:
|
||||
tokenizer = None
|
||||
|
||||
# Initialize processing stages
|
||||
self._init_stages(args, transform, transform_topcrop, tokenizer)
|
||||
@@ -483,8 +428,6 @@ 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,
|
||||
@@ -495,11 +438,14 @@ class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
|
||||
self.video_transform_stage = VideoTransformStage(transform)
|
||||
self.image_transform_stage = ImageTransformStage(
|
||||
transform, transform_topcrop)
|
||||
self.text_encoding_stage = TextEncodingStage(
|
||||
tokenizer=tokenizer,
|
||||
text_max_length=args.text_max_length,
|
||||
cfg_rate=args.training_cfg_rate,
|
||||
seed=self.seed)
|
||||
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
|
||||
|
||||
def _load_raw_data(self) -> list[dict]:
|
||||
"""Load raw data from JSON files."""
|
||||
@@ -521,6 +467,8 @@ 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
|
||||
|
||||
@@ -532,7 +480,6 @@ 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] = []
|
||||
@@ -542,7 +489,8 @@ class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
|
||||
cap=item["cap"],
|
||||
resolution=item.get("resolution"),
|
||||
fps=item.get("fps"),
|
||||
duration=item.get("duration"))
|
||||
duration=item.get("duration"),
|
||||
action_path=item.get("action_path"))
|
||||
|
||||
# Apply filtering stages
|
||||
if not self._apply_filter_stages(batch, filter_counts):
|
||||
@@ -567,10 +515,6 @@ 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
|
||||
@@ -582,10 +526,9 @@ class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
|
||||
after_count: int):
|
||||
"""Log filtering statistics."""
|
||||
logger.info(
|
||||
"validation_failed: %d, resolution_failed: %d, frame_sampling_failed: %d, "
|
||||
"validation_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)
|
||||
|
||||
@@ -604,20 +547,27 @@ class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
|
||||
# Apply transformation stages
|
||||
batch = self.video_transform_stage.process(batch)
|
||||
batch = self.image_transform_stage.process(batch)
|
||||
batch = self.text_encoding_stage.process(batch)
|
||||
if self.text_encoding_stage is not None:
|
||||
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
|
||||
|
||||
|
||||
@@ -18,6 +18,7 @@ __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",
|
||||
|
||||
@@ -7,10 +7,17 @@ import torch.distributed
|
||||
from fastvideo.distributed.parallel_state import (get_sp_group,
|
||||
get_sp_parallel_rank,
|
||||
get_sp_world_size,
|
||||
get_tp_group)
|
||||
get_tp_group,
|
||||
model_parallel_is_initialized)
|
||||
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:
|
||||
@@ -101,3 +108,86 @@ 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,6 +57,9 @@ 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,
|
||||
|
||||
@@ -9,6 +9,7 @@ diffusion models.
|
||||
import math
|
||||
import os
|
||||
import re
|
||||
import threading
|
||||
import time
|
||||
from copy import deepcopy
|
||||
from typing import Any
|
||||
@@ -18,6 +19,8 @@ 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
|
||||
@@ -29,6 +32,21 @@ 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.
|
||||
@@ -370,8 +388,31 @@ class VideoGenerator:
|
||||
|
||||
# Run inference
|
||||
start_time = time.perf_counter()
|
||||
output_batch = self.executor.execute_forward(batch, fastvideo_args)
|
||||
samples = output_batch.output
|
||||
|
||||
# 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()
|
||||
logging_info = output_batch.logging_info
|
||||
|
||||
gen_time = time.perf_counter() - start_time
|
||||
@@ -389,6 +430,11 @@ 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
|
||||
@@ -396,6 +442,7 @@ 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,
|
||||
@@ -405,6 +452,98 @@ 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:
|
||||
|
||||
@@ -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, STA_Mode
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
from fastvideo.configs.utils import clean_cli_args
|
||||
from fastvideo.layers.quantization import QUANTIZATION_METHODS, QuantizationMethods
|
||||
from fastvideo.logger import init_logger
|
||||
@@ -141,8 +141,6 @@ 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
|
||||
@@ -166,11 +164,20 @@ 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
|
||||
@@ -203,8 +210,44 @@ 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
|
||||
@@ -325,6 +368,44 @@ 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",
|
||||
@@ -382,20 +463,6 @@ 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,
|
||||
@@ -836,6 +903,7 @@ class TrainingArgs(FastVideoArgs):
|
||||
lora_rank: int | None = None
|
||||
lora_alpha: int | None = None
|
||||
lora_training: bool = False
|
||||
ltx2_first_frame_conditioning_p: float = 0.1
|
||||
|
||||
# distillation args
|
||||
generator_update_interval: int = 5
|
||||
@@ -849,6 +917,7 @@ class TrainingArgs(FastVideoArgs):
|
||||
training_state_checkpointing_steps: int = 0 # for resuming training
|
||||
weight_only_checkpointing_steps: int = 0 # for inference
|
||||
log_visualization: bool = False
|
||||
visualization_steps: int = 0
|
||||
# simulate generator forward to match inference
|
||||
simulate_generator_forward: bool = False
|
||||
warp_denoising_step: bool = False
|
||||
@@ -1012,6 +1081,9 @@ class TrainingArgs(FastVideoArgs):
|
||||
parser.add_argument("--log-validation",
|
||||
action=StoreBoolean,
|
||||
help="Whether to log validation results")
|
||||
parser.add_argument("--visualization-steps",
|
||||
type=int,
|
||||
help="Number of visualization steps")
|
||||
parser.add_argument("--tracker-project-name",
|
||||
type=str,
|
||||
help="Project name for tracking")
|
||||
@@ -1186,6 +1258,13 @@ class TrainingArgs(FastVideoArgs):
|
||||
help="Whether to use LoRA training")
|
||||
parser.add_argument("--lora-rank", type=int, help="LoRA rank")
|
||||
parser.add_argument("--lora-alpha", type=int, help="LoRA alpha")
|
||||
parser.add_argument(
|
||||
"--ltx2-first-frame-conditioning-p",
|
||||
type=float,
|
||||
default=TrainingArgs.ltx2_first_frame_conditioning_p,
|
||||
help=
|
||||
"Probability of conditioning on the first frame during LTX-2 training",
|
||||
)
|
||||
|
||||
# V-MoBA parameters
|
||||
parser.add_argument(
|
||||
|
||||
@@ -193,3 +193,46 @@ 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)
|
||||
|
||||
@@ -168,9 +168,15 @@ 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) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
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]:
|
||||
"""
|
||||
Apply gated residual connection, followed by layernorm and
|
||||
scale/shift in a single fused operation.
|
||||
@@ -205,6 +211,11 @@ 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]
|
||||
@@ -254,14 +265,21 @@ 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) -> torch.Tensor:
|
||||
def forward(self,
|
||||
x: torch.Tensor,
|
||||
shift: torch.Tensor,
|
||||
scale: torch.Tensor,
|
||||
convert_modulation_dtype: bool = False) -> 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]
|
||||
|
||||
@@ -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 "
|
||||
_FORMAT = (f"{FASTVIDEO_LOGGING_PREFIX}%(levelname)s %(asctime)s.%(msecs)03d "
|
||||
"[%(filename)s:%(lineno)d] %(message)s")
|
||||
_DATE_FORMAT = "%m-%d %H:%M:%S"
|
||||
|
||||
@@ -102,6 +102,7 @@ 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"])
|
||||
@@ -118,20 +119,22 @@ def _info(logger: Logger,
|
||||
|
||||
global _warned_local_main_process, _warned_main_process
|
||||
|
||||
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
|
||||
# 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 main_process_only and not local_main_process_only:
|
||||
logger.log(logging.INFO, msg, *args, stacklevel=2, **kwargs)
|
||||
|
||||
@@ -0,0 +1,9 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from fastvideo.models.audio.ltx2_audio_vae import (
|
||||
LTX2AudioDecoder,
|
||||
LTX2AudioEncoder,
|
||||
LTX2Vocoder,
|
||||
)
|
||||
|
||||
__all__ = ["LTX2AudioEncoder", "LTX2AudioDecoder", "LTX2Vocoder"]
|
||||
@@ -0,0 +1,61 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Audio preprocessing helpers for LTX-2 training."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
import torchaudio
|
||||
from torch import nn
|
||||
|
||||
|
||||
class AudioProcessor(nn.Module):
|
||||
"""Converts audio waveforms to log-mel spectrograms with resampling."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
sample_rate: int,
|
||||
mel_bins: int,
|
||||
mel_hop_length: int,
|
||||
n_fft: int,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.sample_rate = sample_rate
|
||||
self.mel_transform = torchaudio.transforms.MelSpectrogram(
|
||||
sample_rate=sample_rate,
|
||||
n_fft=n_fft,
|
||||
win_length=n_fft,
|
||||
hop_length=mel_hop_length,
|
||||
f_min=0.0,
|
||||
f_max=sample_rate / 2.0,
|
||||
n_mels=mel_bins,
|
||||
window_fn=torch.hann_window,
|
||||
center=True,
|
||||
pad_mode="reflect",
|
||||
power=1.0,
|
||||
mel_scale="slaney",
|
||||
norm="slaney",
|
||||
)
|
||||
|
||||
def resample_waveform(
|
||||
self,
|
||||
waveform: torch.Tensor,
|
||||
source_rate: int,
|
||||
target_rate: int,
|
||||
) -> torch.Tensor:
|
||||
if source_rate == target_rate:
|
||||
return waveform
|
||||
resampled = torchaudio.functional.resample(
|
||||
waveform, source_rate, target_rate)
|
||||
return resampled.to(device=waveform.device, dtype=waveform.dtype)
|
||||
|
||||
def waveform_to_mel(
|
||||
self,
|
||||
waveform: torch.Tensor,
|
||||
waveform_sample_rate: int,
|
||||
) -> torch.Tensor:
|
||||
waveform = self.resample_waveform(
|
||||
waveform, waveform_sample_rate, self.sample_rate)
|
||||
mel = self.mel_transform(waveform)
|
||||
mel = torch.log(torch.clamp(mel, min=1e-5))
|
||||
mel = mel.to(device=waveform.device, dtype=waveform.dtype)
|
||||
return mel.permute(0, 1, 3, 2).contiguous()
|
||||
@@ -33,7 +33,7 @@ from fastvideo.layers.visual_embedding import (PatchEmbed)
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.dits.base import BaseDiT
|
||||
from fastvideo.models.dits.wanvideo import WanT2VCrossAttention, WanTimeTextImageEmbedding
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
from fastvideo.platforms import AttentionBackendEnum, current_platform
|
||||
|
||||
logger = init_logger(__name__)
|
||||
class CausalWanSelfAttention(nn.Module):
|
||||
@@ -286,6 +286,8 @@ class CausalWanTransformerBlock(nn.Module):
|
||||
null_shift = null_scale = torch.tensor([0], device=hidden_states.device)
|
||||
norm_hidden_states, hidden_states = self.self_attn_residual_norm(
|
||||
hidden_states, attn_output, gate_msa, null_shift, null_scale)
|
||||
norm_hidden_states, hidden_states = norm_hidden_states.to(
|
||||
orig_dtype), hidden_states.to(orig_dtype)
|
||||
|
||||
# 2. Cross-attention
|
||||
attn_output = self.attn2(norm_hidden_states,
|
||||
@@ -452,7 +454,6 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
This function will be run for num_frame times.
|
||||
Process the latent frames one by one (1560 tokens each)
|
||||
"""
|
||||
from fastvideo.platforms import current_platform
|
||||
|
||||
orig_dtype = hidden_states.dtype
|
||||
if not isinstance(encoder_hidden_states, torch.Tensor):
|
||||
@@ -676,4 +677,7 @@ 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
|
||||
return out
|
||||
|
||||
# Entry point for model registry
|
||||
EntryClass = CausalWanTransformer3DModel
|
||||
|
||||
@@ -723,4 +723,7 @@ 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
|
||||
return hidden_states
|
||||
|
||||
# Entry point for model registry
|
||||
EntryClass = CosmosTransformer3DModel
|
||||
|
||||
@@ -959,3 +959,5 @@ class Cosmos25Transformer3DModel(BaseDiT):
|
||||
|
||||
return hidden_states
|
||||
|
||||
# Entry point for model registry
|
||||
EntryClass = Cosmos25Transformer3DModel
|
||||
|
||||
@@ -940,3 +940,6 @@ 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
|
||||
|
||||
@@ -626,6 +626,8 @@ 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
|
||||
@@ -850,4 +852,7 @@ 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
|
||||
return x
|
||||
|
||||
# Entry point for model registry
|
||||
EntryClass = HunyuanVideo15Transformer3DModel
|
||||
|
||||