Compare commits

...
Author SHA1 Message Date
SolitaryThinker a8dddfaa16 missing file 2026-02-10 08:33:00 +00:00
Will Lin 7210c68f1b lint 2026-02-10 00:30:03 -08:00
SolitaryThinker 5602dc1bad revert 2026-02-10 08:22:05 +00:00
SolitaryThinker ac4bc4ab84 uipdate 2026-02-10 07:58:43 +00:00
Matthew Noto bee27f9f74 Merge branch 'main' into ltx-base 2026-02-09 17:51:37 -08:00
ad58f802f3 [Feat] Port LTX2 trainer (#1074)
Co-authored-by: Davids048 <jundasu@ucsd.edu>
Co-authored-by: Matthew Noto <99706358+RandNMR73@users.noreply.github.com>
2026-02-09 17:32:57 -08:00
Wei Zhou 04fa356ee3 [Misc] [Training] Fixed a bunch of bugs in current training pipeline (#1084) 2026-02-09 16:01:05 -08:00
Matthew Notoandgemini-code-assist[bot] f9c076fe2b [misc] add AGENTS.md file (#1085)
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
2026-02-09 01:03:49 -08:00
Will Lin dff0ea401a update 2026-02-08 00:22:21 -08:00
XOR-op f76efe798e [bugfix]: _compile_conditions regression (#1077) 2026-02-06 18:57:43 -08:00
Zhang Peiyuan 09f455233e [misc] readme small fix (#1076) 2026-02-06 16:57:48 -08:00
XOR-op aea300f690 [perf]: use CUDA IPC in multiproc executor to avoid serialization overhead (#1061) 2026-02-06 19:43:18 -05:00
Hao Zhang b92219f6a6 more fix and relocate STA arguments to pipeline config (#1073) 2026-02-06 13:38:35 -08:00
Jinzhe Pan a321b95a8a [Fix] remove video ratio limitation (#1069) 2026-02-05 20:33:08 -08:00
Davids048 becd379f58 Split LTX2 mappings and add registry coverage.
Detail:

- Merged LTX2 sampling behavior into the global registry and removed the ambiguous local converted-path default.
- Replaced single LTX2 sampling mapping with explicit model-ID mappings:
    - Lightricks/LTX-2 -> LTX2BaseSamplingParam
    - FastVideo/LTX2-base -> LTX2BaseSamplingParam
    - FastVideo/LTX2-Distilled-Diffusers -> LTX2DistilledSamplingParam
- Kept LTX2T2VConfig as the pipeline config for all explicitly mapped LTX2 IDs.
- Removed implicit mapping for converted/ltx2_diffusers to avoid guessing base vs distilled for user-local
  conversions.
- Added focused local tests at tests/local_tests/test_ltx2_registry.py for:
    - exact base/distilled sampling resolution,
    - pipeline config resolution,
    - no fallback behavior for ambiguous local converted paths.

Assumptions:

- Canonical rename/ID intent:
    - “base” names map to LTX2BaseSamplingParam.
    - “Distilled” names map to LTX2DistilledSamplingParam.
- converted/ltx2_diffusers is intentionally ambiguous across users and must not be auto-assigned.
- Unknown/non-canonical names containing “LTX”/“LTX2” (but not matching explicit registered IDs) should not auto-
  resolve to base or distilled.
    - Sampling resolver returns None (caller falls back to generic defaults/user overrides).
    - Pipeline config lookup raises a “No match found” error.

Notes:

- This change prioritizes explicitness over convenience: only predetermined, canonical model IDs get LTX2-specific
  defaults; everything else requires user intent.
2026-02-05 20:28:46 -08:00
Davids048 1ed7d7e1b0 Add gemma tokenizer to LTX2 conversion script.
- Also clean up for PR.
2026-02-05 18:53:22 -08:00
Davids048 d9fabcc5ef Add some annotations. 2026-02-05 18:53:22 -08:00
Davids048 9db48498de Update quality test script. 2026-02-05 18:53:22 -08:00
Davids048 1cd7038315 Add LTX2 base model. 2026-02-05 18:53:12 -08:00
Wei Zhou 98308db7e0 [Feature] [Hy1.5] Support HY1.5 super-resolution pipeline for 1080p videos (#1046) 2026-02-05 16:39:20 -08:00
Hao Zhang c1e18f6722 Some minor fixes (#1068) 2026-02-05 16:28:00 -08:00
William Lin 75e193a2c9 [core] Refactor and centralize our registry for models, pipelines, and sampling params (#1066) 2026-02-05 14:40:30 -08:00
XOR-op d6e0a7d0dd [refactor] Action module (#1065) 2026-02-05 13:56:35 -08:00
William Linandgemini-code-assist[bot] 7fc5f241da [misc] Fix naming instruction in runpod.md (#1067)
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
2026-02-05 13:45:49 -08:00
Mingjia Huo aae48a7e90 [feat] HYworld VAE with cache (#1057) 2026-02-05 04:18:17 -08:00
William Lin 88f38eb0f4 [misc] upgrade torch to 2.10 (#1048) 2026-02-05 04:15:22 -08:00
Shao Duan d750b463dc Added Sequence Parallelism for LTX-2 Distilled (#1036) 2026-02-04 15:25:30 -08:00
KyleShaoandWill Lin e10b26a3d8 [feat] Add Cosmos 2.5 I2W/V2W support (staged pipeline + examples) (#1021)
Co-authored-by: Will Lin <wlsaidhi@gmail.com>
2026-02-04 14:19:30 -08:00
XOR-op 74636ba246 [chore]: use higher precision timestamp in logging (#1062) 2026-02-04 13:22:59 -08:00
Wei Zhou 7e2f3f14e7 [Bugfix] [Wan I2V] Fix CLIP Image encoder config (#1063) 2026-02-04 13:20:19 -08:00
145 changed files with 8837 additions and 2215 deletions
+42
View File
@@ -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.
+1
View File
@@ -0,0 +1 @@
@AGENTS.md
+30 -47
View File
@@ -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. [![Star](https://img.shields.io/github/stars/sgl-project/sglang.svg?style=social&label=Star)](https://github.com/sgl-project/sglang)
- [DanceGRPO](https://github.com/XueZeyue/DanceGRPO): A unified framework to adapt Group Relative Policy Optimization (GRPO) to visual generation paradigms. Code based on FastVideo. [![Star](https://img.shields.io/github/stars/XueZeyue/DanceGRPO.svg?style=social&label=Star)](https://github.com/XueZeyue/DanceGRPO)
- [SRPO](https://github.com/Tencent-Hunyuan/SRPO): A method to directly align the full diffusion trajectory with fine-grained human preference. Code based on FastVideo. [![Star](https://img.shields.io/github/stars/Tencent-Hunyuan/SRPO.svg?style=social&label=Star)](https://github.com/Tencent-Hunyuan/SRPO)
- [DCM](https://github.com/Vchitect/DCM): Dual-expert consistency model for efficient and high-quality video generation. Code based on FastVideo. [![Star](https://img.shields.io/github/stars/Vchitect/DCM.svg?style=social&label=Star)](https://github.com/Vchitect/DCM)
- [Hunyuan Video 1.5](https://github.com/Tencent-Hunyuan/HunyuanVideo-1.5): A leading lightweight video generation model, where they proposed SSTA based on Sliding Tile Attention. [![Star](https://img.shields.io/github/stars/Tencent-Hunyuan/HunyuanVideo-1.5.svg?style=social&label=Star)](https://github.com/Tencent-Hunyuan/HunyuanVideo-1.5)
- [Kandinsky-5.0](https://github.com/kandinskylab/kandinsky-5): A family of diffusion models for video & image generation, where their NABLA attention includes a Sliding Tile Attention branch. [![Star](https://img.shields.io/github/stars/kandinskylab/kandinsky-5.svg?style=social&label=Star)](https://github.com/kandinskylab/kandinsky-5)
- [LongCat Video](https://github.com/meituan-longcat/LongCat-Video): A foundational video generation model with 13.6B parameters with block-sparse attention similar to Video Sparse Attention. [![Star](https://img.shields.io/github/stars/meituan-longcat/LongCat-Video.svg?style=social&label=Star)](https://github.com/meituan-longcat/LongCat-Video)
- [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},
+1 -1
View File
@@ -43,7 +43,7 @@ RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
uv pip install --no-cache-dir --upgrade pip && \
uv pip install --no-cache-dir .[dev] && \
uv pip install --no-cache-dir https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.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 . .
+1 -1
View File
@@ -43,7 +43,7 @@ RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
uv pip install --no-cache-dir --upgrade pip && \
uv pip install --no-cache-dir .[dev] && \
uv pip install --no-cache-dir https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.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 . .
+1 -1
View File
@@ -43,7 +43,7 @@ RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
uv pip install --no-cache-dir --upgrade pip && \
uv pip install --no-cache-dir .[dev] && \
uv pip install --no-cache-dir https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.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 . .
+9
View File
@@ -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
+2 -2
View File
@@ -324,8 +324,8 @@ Purpose:
Action:
- Add a new pipeline config + sampling params.
- Register them in `fastvideo/configs/pipelines/registry.py` and
`fastvideo/configs/sample/registry.py`.
- Register them in `fastvideo/registry.py` using explicit
`register_configs(...)` blocks (this file is the single source of truth now).
### 5) Wire pipeline stages
+1 -1
View File
@@ -17,7 +17,7 @@ You can easily use the FastVideo Pod Template on [RunPod](https://www.runpod.io)
- Select the "FastVideo" or "fastvideo-dev" Pod Template.
![RunPod Pod Template Selection](../../assets/images/runpod_create.png)
- Set the Pod name to "<name>-<FastVideo>-<date>".
- Set the Pod name to "`<name>-<FastVideo>-<date>`".
- Finally, once the pod is deployed (will take a few minutes as the image is being pulled), you can SSH into it using "SSH exposed over TCP". You'll need to use the matching private ssh key you provided.
![RunPod SSH](../../assets/images/runpod_ssh.png)
+3 -2
View File
@@ -51,8 +51,9 @@ runtime parameters consistent:
- `fastvideo/configs/pipelines/`: pipeline wiring and required components.
- `fastvideo/configs/sample/`: default sampling parameters (steps, frames,
guidance scale, resolution, fps).
- `fastvideo/configs/registry.py`: pipeline registry.
- `fastvideo/configs/sample/registry.py`: sampling param registry.
- `fastvideo/registry.py`: unified registry for pipeline config + sampling
defaults and model metadata resolution, defined via explicit
`register_configs(...)` blocks (no separate dict registries).
`FastVideoArgs` (in `fastvideo/fastvideo_args.py`) provides runtime settings and
is passed into pipeline construction and stages.
+32 -8
View File
@@ -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,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()
+1 -1
View File
@@ -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()
+7 -165
View File
@@ -1,165 +1,11 @@
# SPDX-License-Identifier: Apache-2.0
"""
Basic example for HYWorld (HY-WorldPlay) video generation using FastVideo.
This example replicates the same functionality as HY-WorldPlay/run.sh,
demonstrating image-to-video generation with camera trajectory control.
"""
import time
import math
import numpy as np
import imageio
import torchvision
from einops import rearrange
from fastvideo import VideoGenerator
from fastvideo.pipelines import ForwardBatch
from fastvideo.utils import shallow_asdict, align_to
from fastvideo.logger import init_logger
from fastvideo.models.dits.hyworld.pose import pose_to_input, compute_latent_num
from fastvideo.models.dits.hyworld.resolution_utils import get_resolution_from_image
logger = init_logger(__name__)
class HYWorldVideoGenerator(VideoGenerator):
"""Extended VideoGenerator that adds HYWorld-specific parameters to batch.extra."""
def _generate_single_video(self, prompt: str, sampling_param=None, **kwargs):
"""Override to add viewmats, Ks, and action to batch.extra."""
fastvideo_args = self.fastvideo_args
pipeline_config = fastvideo_args.pipeline_config
if sampling_param is None:
from fastvideo.configs.sample import SamplingParam
sampling_param = SamplingParam.from_pretrained(fastvideo_args.model_path)
# Update sampling param with kwargs
if kwargs:
for key, value in kwargs.items():
if hasattr(sampling_param, key):
setattr(sampling_param, key, value)
# Get pose string from sampling_param or kwargs
pose = kwargs.get('pose', getattr(sampling_param, 'POSE', 'w-31'))
num_frames = kwargs.get('num_frames', getattr(sampling_param, 'num_frames', 125))
# Calculate number of latents
latent_num = compute_latent_num(num_frames)
# Convert pose to viewmats, Ks, and action
viewmats, Ks, action = pose_to_input(pose, latent_num)
# Convert to tensors and add batch dimension
viewmats = viewmats.unsqueeze(0) # (1, T, 4, 4)
Ks = Ks.unsqueeze(0) # (1, T, 3, 3)
action = action.unsqueeze(0) # (1, T)
# Validate inputs
prompt = prompt.strip()
sampling_param = sampling_param.__class__(**shallow_asdict(sampling_param))
output_path = kwargs.get("output_path", sampling_param.output_path)
sampling_param.prompt = prompt
if sampling_param.negative_prompt is not None:
sampling_param.negative_prompt = sampling_param.negative_prompt.strip()
# Validate dimensions
if (sampling_param.height <= 0 or sampling_param.width <= 0 or
sampling_param.num_frames <= 0):
raise ValueError(
f"Height, width, and num_frames must be positive integers")
temporal_scale_factor = pipeline_config.vae_config.arch_config.temporal_compression_ratio
num_frames = sampling_param.num_frames
num_gpus = fastvideo_args.num_gpus
use_temporal_scaling_frames = pipeline_config.vae_config.use_temporal_scaling_frames
# Adjust number of frames based on number of GPUs
if use_temporal_scaling_frames:
orig_latent_num_frames = (num_frames - 1) // temporal_scale_factor + 1
else:
orig_latent_num_frames = sampling_param.num_frames // 17 * 3
if orig_latent_num_frames % fastvideo_args.num_gpus != 0:
if use_temporal_scaling_frames:
new_num_frames = (orig_latent_num_frames - 1) * temporal_scale_factor + 1
else:
divisor = math.lcm(3, num_gpus)
orig_latent_num_frames = (
(orig_latent_num_frames + divisor - 1) // divisor) * divisor
new_num_frames = orig_latent_num_frames // 3 * 17
logger.info(
"Adjusting number of frames from %s to %s based on number of GPUs (%s)",
sampling_param.num_frames, new_num_frames, fastvideo_args.num_gpus)
sampling_param.num_frames = new_num_frames
# Calculate sizes
target_height = align_to(sampling_param.height, 16)
target_width = align_to(sampling_param.width, 16)
# Calculate latent sizes
latents_size = [(sampling_param.num_frames - 1) // 4 + 1,
sampling_param.height // 8, sampling_param.width // 8]
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
# Prepare batch
batch = ForwardBatch(
**shallow_asdict(sampling_param),
eta=0.0,
n_tokens=n_tokens,
VSA_sparsity=fastvideo_args.VSA_sparsity,
)
# Add HYWorld-specific parameters to batch.extra
batch.extra['viewmats'] = viewmats
batch.extra['Ks'] = Ks
batch.extra['action'] = action
batch.extra['chunk_latent_frames'] = 16 # For bidirectional model
# Run inference
start_time = time.perf_counter()
output_batch = self.executor.execute_forward(batch, fastvideo_args)
samples = output_batch.output
logging_info = output_batch.logging_info
gen_time = time.perf_counter() - start_time
logger.info("Generated successfully in %.2f seconds", gen_time)
# Process outputs
videos = rearrange(samples, "b c t h w -> t b c h w")
frames = []
for x in videos:
x = torchvision.utils.make_grid(x, nrow=6)
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
frames.append((x * 255).numpy().astype(np.uint8))
# Save video if requested
if batch.save_video:
imageio.mimsave(output_path, frames, fps=batch.fps, format="mp4")
logger.info("Saved video to %s", output_path)
if batch.return_frames:
return frames
else:
return {
"samples": samples,
"frames": frames,
"prompts": prompt,
"size": (target_height, target_width, batch.num_frames),
"generation_time": gen_time,
"logging_info": logging_info,
"trajectory": output_batch.trajectory_latents,
"trajectory_timesteps": output_batch.trajectory_timesteps,
"trajectory_decoded": output_batch.trajectory_decoded,
}
# 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
@@ -169,7 +15,7 @@ def main():
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='video_samples_hyworld', help="Output video path")
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")
@@ -185,8 +31,7 @@ def main():
# Initialize generator
print("\nInitializing VideoGenerator for HYWorld...")
generator = HYWorldVideoGenerator.from_pretrained(
generator = VideoGenerator.from_pretrained(
"FastVideo/HY-WorldPlay-Bidirectional-Diffusers",
num_gpus=1,
use_fsdp_inference=True,
@@ -198,11 +43,12 @@ def main():
)
# Generate video
# The pose string is automatically converted to camera matrices by the pipeline
print("\nGenerating video...")
start_time = time.time()
video = generator.generate_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="",
@@ -211,13 +57,9 @@ def main():
height=HEIGHT,
width=WIDTH,
seed=args.seed,
pose=args.pose,
)
elapsed = time.time() - start_time
print(f"\nVideo generated successfully!")
print(f"Saved to: {args.output_path}")
print(f"Time: {elapsed:.2f}s")
print(f"\nVideo saved to: {args.output_path}")
if __name__ == "__main__":
+7 -2
View File
@@ -1,5 +1,6 @@
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 "
@@ -16,16 +17,20 @@ PROMPT = (
def main() -> None:
# Uses FastVideo default sampling settings for LTX2 base.
generator = VideoGenerator.from_pretrained(
"FastVideo/LTX2-Distilled-Diffusers",
"Davids048/LTX2-Base-Diffusers",
num_gpus=1,
)
output_path = "outputs_video/ltx2_basic/output_ltx2_distilled_t2v.mp4"
output_path = "outputs_video/ltx2_basic/output_ltx2_base_t2v_1088_1920_1.1.mp4"
generator.generate_video(
prompt=PROMPT,
output_path=output_path,
save_video=True,
num_frames=121,
height=1088,
width=1920,
)
generator.shutdown()
@@ -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()
@@ -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,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
}
]
}
+5
View File
@@ -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
View File
@@ -5,6 +5,7 @@ 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",
@@ -14,4 +15,5 @@ __all__ = [
"LTX2AudioEncoderConfig",
"LTX2AudioDecoderConfig",
"LTX2VocoderConfig",
"UpsamplerConfig",
]
@@ -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\.(.*)$":
+5 -10
View File
@@ -8,12 +8,8 @@ def is_double_block(n: str, m) -> bool:
return "double" in n and str.isdigit(n.split(".")[-1])
def is_single_block(n: str, m) -> bool:
return "single" in n and str.isdigit(n.split(".")[-1])
# def is_refiner_block(n: str, m) -> bool:
# return "refiner" in n and str.isdigit(n.split(".")[-1])
def is_refiner_block(n: str, m) -> bool:
return "refiner" in n and str.isdigit(n.split(".")[-1])
def is_txt_in(n: str, m) -> bool:
@@ -23,10 +19,10 @@ def is_txt_in(n: str, m) -> bool:
@dataclass
class HYWorldArchConfig(DiTArchConfig):
_fsdp_shard_conditions: list = field(
default_factory=lambda: [is_double_block, is_single_block])
default_factory=lambda: [is_double_block, is_refiner_block])
_compile_conditions: list = field(
default_factory=lambda: [is_double_block, is_single_block, is_txt_in])
default_factory=lambda: [is_double_block, is_refiner_block, is_txt_in])
param_names_mapping: dict = field(
default_factory=lambda: {
@@ -54,8 +50,7 @@ class HYWorldArchConfig(DiTArchConfig):
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 (HYWorld uses TimestepEmbedder directly,
# but FastVideo model inherits HunyuanVideo15TimeEmbedding with timestep_embedder):
# 2. time_in mappings:
r"^time_in\.mlp\.0\.(.*)$":
r"time_in.timestep_embedder.mlp.fc_in.\1",
r"^time_in\.mlp\.2\.(.*)$":
+4 -2
View File
@@ -7,10 +7,12 @@ from dataclasses import dataclass, field
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
import re
def is_ltx2_blocks(name: str, _module) -> bool:
"""FSDP shard condition for LTX-2 transformer blocks."""
return "transformer_blocks" in name
res = re.search(r"(?:^|\.)transformer_blocks\.\d+$", name) is not None
return res
@dataclass
@@ -1,4 +1,6 @@
from dataclasses import dataclass, field
import torch
from fastvideo.configs.models.dits.wanvideo import WanVideoArchConfig, WanVideoConfig
@@ -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])
+3 -3
View File
@@ -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,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
+1 -2
View File
@@ -6,8 +6,7 @@ from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
from fastvideo.configs.pipelines.hunyuan15 import Hunyuan15T2V480PConfig, Hunyuan15T2V720PConfig
from fastvideo.configs.pipelines.hyworld import HYWorldConfig
from fastvideo.configs.pipelines.ltx2 import LTX2T2VConfig
from fastvideo.configs.pipelines.registry import (
get_pipeline_config_cls_from_name)
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,
+30 -5
View File
@@ -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:
+28 -1
View File
@@ -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"
-214
View File
@@ -1,214 +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.hyworld import HYWorldConfig
from fastvideo.configs.pipelines.ltx2 import LTX2T2VConfig
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,
"FastVideo/HY-WorldPlay-Bidirectional-Diffusers": HYWorldConfig,
"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,
# LTX-2 models
"Lightricks/LTX-2": LTX2T2VConfig,
"converted/ltx2_diffusers": LTX2T2VConfig,
# 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(),
"hyworld":
lambda id: "hyworld" 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(),
"ltx2":
lambda id: "ltx2" in id.lower() or "ltx-2" 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
"hyworld":
HYWorldConfig, # HYWorld-specific config as fallback for any HYWorld 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,
"ltx2": LTX2T2VConfig,
# 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
+7 -2
View File
@@ -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()
+12 -8
View File
@@ -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
+26
View File
@@ -28,7 +28,33 @@ class Hunyuan15_480P_SamplingParam(SamplingParam):
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
+40 -2
View File
@@ -5,10 +5,44 @@ from fastvideo.configs.sample.base import SamplingParam
@dataclass
class LTX2SamplingParam(SamplingParam):
"""Default sampling parameters for LTX-2 distilled T2V.
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
@@ -18,3 +52,7 @@ class LTX2SamplingParam(SamplingParam):
guidance_scale: float = 1.0
# No default negative_prompt for distilled models
negative_prompt: str = ""
# Backward compatibility alias.
LTX2SamplingParam = LTX2DistilledSamplingParam
-204
View File
@@ -1,204 +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.hyworld import HYWorld_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
from fastvideo.configs.sample.ltx2 import LTX2SamplingParam
# 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/HY-WorldPlay-Bidirectional-Diffusers": HYWorld_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,
# LTX-2 models
"Lightricks/LTX-2": LTX2SamplingParam,
"FastVideo/LTX2-Distilled-Diffusers": LTX2SamplingParam,
# 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(),
"hyworld":
lambda id: "hyworld" 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(),
"ltx2":
lambda id: "ltx2" in id.lower() or "ltx-2" 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
"hyworld":
HYWorld_SamplingParam, # HYWorld-specific config as fallback for any HYWorld 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,
"ltx2": LTX2SamplingParam,
# 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
+8 -2
View File
@@ -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",
]
@@ -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
+1 -70
View File
@@ -148,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."""
@@ -328,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
@@ -489,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,
@@ -543,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] = []
@@ -579,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
@@ -594,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)
+1
View File
@@ -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",
+91 -1
View File
@@ -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.")
+41 -2
View File
@@ -9,6 +9,7 @@ diffusion models.
import math
import os
import re
import threading
import time
from copy import deepcopy
from typing import Any
@@ -31,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.
@@ -372,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
+14 -17
View File
@@ -10,7 +10,7 @@ from enum import Enum
from typing import Any, TYPE_CHECKING
from fastvideo.configs.configs import PreprocessConfig
from fastvideo.configs.pipelines.base import PipelineConfig, 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
@@ -179,6 +177,7 @@ class FastVideoArgs:
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
@@ -464,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,
@@ -918,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
@@ -931,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
@@ -1094,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")
@@ -1268,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(
+43
View File
@@ -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)
+18 -15
View File
@@ -27,7 +27,7 @@ RESET = '\033[0;0m'
_warned_local_main_process = False
_warned_main_process = False
_FORMAT = (f"{FASTVIDEO_LOGGING_PREFIX}%(levelname)s %(asctime)s "
_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,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()
+3
View File
@@ -1325,3 +1325,6 @@ def decode_audio(
decoded_audio = audio_decoder(latent)
decoded_audio = vocoder(decoded_audio).squeeze(0).float()
return decoded_audio
# Entry point for model registry
EntryClass = [LTX2AudioEncoder, LTX2AudioDecoder, LTX2Vocoder]
+7 -3
View File
@@ -33,7 +33,7 @@ from fastvideo.layers.visual_embedding import (PatchEmbed)
from fastvideo.logger import init_logger
from fastvideo.models.dits.base import BaseDiT
from fastvideo.models.dits.wanvideo import WanT2VCrossAttention, WanTimeTextImageEmbedding
from fastvideo.platforms import AttentionBackendEnum
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
+4 -1
View File
@@ -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
+2
View File
@@ -959,3 +959,5 @@ class Cosmos25Transformer3DModel(BaseDiT):
return hidden_states
# Entry point for model registry
EntryClass = Cosmos25Transformer3DModel
+3
View File
@@ -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
+4 -1
View File
@@ -852,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
@@ -22,3 +22,6 @@ __all__ = [
# Inference utilities (used by examples)
"get_resolution_from_image",
]
# Entry point for model registry
EntryClass = HYWorldTransformer3DModel
+4 -4
View File
@@ -536,8 +536,8 @@ class HYWorldTransformer3DModel(HunyuanVideo15Transformer3DModel):
temb_txt,
freqs_cis,
seq_attention_mask,
viewmats_seq, # hyworld
Ks_seq, # hyworld
viewmats_seq,
Ks_seq,
)
else:
for block in self.double_blocks:
@@ -549,8 +549,8 @@ class HYWorldTransformer3DModel(HunyuanVideo15Transformer3DModel):
temb_txt,
freqs_cis,
seq_attention_mask,
viewmats=viewmats_seq, # hyworld
Ks=Ks_seq, # hyworld
viewmats_seq,
Ks_seq,
)
+2
View File
@@ -1134,3 +1134,5 @@ class LongCatTransformer3DModel(CachableDiT):
)
return x
# Entry point for model registry
EntryClass = LongCatTransformer3DModel
File diff suppressed because it is too large Load Diff
@@ -9,3 +9,6 @@ __all__ = [
"CausalMatrixGameTransformerBlock",
"ActionModule",
]
# Entry point for model registry
EntryClass = [MatrixGameWanModel, CausalMatrixGameWanModel]
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+242 -142
View File
@@ -9,20 +9,32 @@ import torch.nn as nn
from fastvideo.attention import DistributedAttention
from fastvideo.configs.models.dits.matrixgame import MatrixGameWanVideoConfig
from fastvideo.distributed.parallel_state import get_sp_world_size
from fastvideo.layers.layernorm import (FP32LayerNorm, LayerNormScaleShift,
RMSNorm, ScaleResidual,
ScaleResidualLayerNormScaleShift)
from fastvideo.layers.layernorm import (
FP32LayerNorm,
LayerNormScaleShift,
RMSNorm,
ScaleResidual,
ScaleResidualLayerNormScaleShift,
)
from fastvideo.layers.linear import ReplicatedLinear
from fastvideo.layers.mlp import MLP
from fastvideo.layers.rotary_embedding import (_apply_rotary_emb,
get_rotary_pos_embed)
from fastvideo.layers.visual_embedding import PatchEmbed, TimestepEmbedder, ModulateProjection
from fastvideo.layers.rotary_embedding import (
_apply_rotary_emb,
get_rotary_pos_embed,
)
from fastvideo.layers.visual_embedding import (
PatchEmbed,
TimestepEmbedder,
ModulateProjection,
)
from fastvideo.logger import init_logger
from fastvideo.models.dits.base import BaseDiT
from fastvideo.models.dits.wanvideo import (WanSelfAttention,
WanI2VCrossAttention,
WanT2VCrossAttention,
WanImageEmbedding)
from fastvideo.models.dits.wanvideo import (
WanSelfAttention,
WanI2VCrossAttention,
WanT2VCrossAttention,
WanImageEmbedding,
)
from fastvideo.platforms import AttentionBackendEnum, current_platform
# Import ActionModule
@@ -41,11 +53,12 @@ class MatrixGameTimeImageEmbedding(nn.Module):
super().__init__()
self.time_embedder = TimestepEmbedder(
dim, frequency_embedding_size=time_freq_dim, act_layer="silu")
self.time_modulation = ModulateProjection(dim,
factor=6,
act_layer="silu")
dim, frequency_embedding_size=time_freq_dim, act_layer="silu"
)
self.time_modulation = ModulateProjection(
dim, factor=6, act_layer="silu"
)
self.image_embedder = None
if image_embed_dim is not None:
self.image_embedder = WanImageEmbedding(image_embed_dim, dim)
@@ -53,7 +66,7 @@ class MatrixGameTimeImageEmbedding(nn.Module):
def forward(
self,
timestep: torch.Tensor,
encoder_hidden_states: torch.Tensor, # Kept for interface compatibility
encoder_hidden_states: torch.Tensor, # Kept for interface compatibility
encoder_hidden_states_image: torch.Tensor | None = None,
timestep_seq_len: int | None = None,
):
@@ -61,17 +74,25 @@ class MatrixGameTimeImageEmbedding(nn.Module):
timestep_proj = self.time_modulation(temb)
# MatrixGame does not use text embeddings, so we ignore encoder_hidden_states
if encoder_hidden_states_image is not None:
assert self.image_embedder is not None
encoder_hidden_states_image = self.image_embedder(
encoder_hidden_states_image)
encoder_hidden_states_image
)
encoder_hidden_states = torch.zeros((timestep.shape[0], 0, temb.shape[-1]),
device=temb.device,
dtype=temb.dtype)
encoder_hidden_states = torch.zeros(
(timestep.shape[0], 0, temb.shape[-1]),
device=temb.device,
dtype=temb.dtype,
)
return temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image
return (
temb,
timestep_proj,
encoder_hidden_states,
encoder_hidden_states_image,
)
class MatrixGameCrossAttention(WanSelfAttention):
@@ -87,7 +108,7 @@ class MatrixGameCrossAttention(WanSelfAttention):
# compute query, key, value
q = self.norm_q(self.to_q(x)[0]).view(b, -1, n, d)
if crossattn_cache is not None:
if not crossattn_cache["is_init"]:
crossattn_cache["is_init"] = True
@@ -101,7 +122,7 @@ class MatrixGameCrossAttention(WanSelfAttention):
else:
k = self.norm_k(self.to_k(context)[0]).view(b, -1, n, d)
v = self.to_v(context)[0].view(b, -1, n, d)
# compute attention
x = self.attn(q, k, v)
@@ -112,19 +133,20 @@ class MatrixGameCrossAttention(WanSelfAttention):
class MatrixGameTransformerBlock(nn.Module):
def __init__(self,
dim: int,
ffn_dim: int,
num_heads: int,
qk_norm: str = "rms_norm_across_heads",
cross_attn_norm: bool = False,
eps: float = 1e-6,
added_kv_proj_dim: int | None = None,
supported_attention_backends: tuple[AttentionBackendEnum, ...]
| None = None,
prefix: str = "",
action_config: dict | None = None):
def __init__(
self,
dim: int,
ffn_dim: int,
num_heads: int,
qk_norm: str = "rms_norm_across_heads",
cross_attn_norm: bool = False,
eps: float = 1e-6,
added_kv_proj_dim: int | None = None,
supported_attention_backends: tuple[AttentionBackendEnum, ...]
| None = None,
prefix: str = "",
action_config: dict | None = None,
):
super().__init__()
action_config = action_config or {}
@@ -140,7 +162,8 @@ class MatrixGameTransformerBlock(nn.Module):
head_size=dim // num_heads,
causal=False,
supported_attention_backends=supported_attention_backends,
prefix=f"{prefix}.attn1")
prefix=f"{prefix}.attn1",
)
self.hidden_dim = dim
self.num_attention_heads = num_heads
dim_head = dim // num_heads
@@ -161,28 +184,28 @@ class MatrixGameTransformerBlock(nn.Module):
eps=eps,
elementwise_affine=True,
dtype=torch.float32,
compute_dtype=torch.float32)
compute_dtype=torch.float32,
)
# 2. Cross-attention
if added_kv_proj_dim is not None:
# I2V
self.attn2 = WanI2VCrossAttention(dim,
num_heads,
qk_norm=qk_norm,
eps=eps)
self.attn2 = WanI2VCrossAttention(
dim, num_heads, qk_norm=qk_norm, eps=eps
)
else:
# T2V
self.attn2 = WanT2VCrossAttention(dim,
num_heads,
qk_norm=qk_norm,
eps=eps)
self.attn2 = WanT2VCrossAttention(
dim, num_heads, qk_norm=qk_norm, eps=eps
)
self.cross_attn_residual_norm = ScaleResidualLayerNormScaleShift(
dim,
norm_type="layer",
eps=eps,
elementwise_affine=False,
dtype=torch.float32,
compute_dtype=torch.float32)
compute_dtype=torch.float32,
)
# 2.1. Action Module Integration
self.use_action_module = len(action_config) > 0
@@ -204,7 +227,7 @@ class MatrixGameTransformerBlock(nn.Module):
temb: torch.Tensor,
freqs_cis: tuple[torch.Tensor, torch.Tensor],
# Action Module specific args
grid_sizes: torch.Tensor | None = None,
grid_sizes: torch.Tensor,
mouse_cond: torch.Tensor | None = None,
keyboard_cond: torch.Tensor | None = None,
) -> torch.Tensor:
@@ -214,9 +237,16 @@ class MatrixGameTransformerBlock(nn.Module):
if temb.dim() == 4:
# temb: batch_size, seq_len, 6, inner_dim (wan2.2 ti2v)
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = (
self.scale_shift_table.unsqueeze(0) + temb.float()
).chunk(6, dim=2)
(
shift_msa,
scale_msa,
gate_msa,
c_shift_msa,
c_scale_msa,
c_gate_msa,
) = (self.scale_shift_table.unsqueeze(0) + temb.float()).chunk(
6, dim=2
)
# batch_size, seq_len, 1, inner_dim
shift_msa = shift_msa.squeeze(2)
scale_msa = scale_msa.squeeze(2)
@@ -227,13 +257,20 @@ class MatrixGameTransformerBlock(nn.Module):
else:
# temb: batch_size, 6, inner_dim (wan2.1/wan2.2 14B)
e = self.scale_shift_table + temb.float()
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
6, dim=1)
(
shift_msa,
scale_msa,
gate_msa,
c_shift_msa,
c_scale_msa,
c_gate_msa,
) = e.chunk(6, dim=1)
assert shift_msa.dtype == torch.float32
# 1. Self-attention
norm_hidden_states = (self.norm1(hidden_states.float()) *
(1 + scale_msa) + shift_msa).to(orig_dtype)
norm_hidden_states = (
self.norm1(hidden_states.float()) * (1 + scale_msa) + shift_msa
).to(orig_dtype)
query, _ = self.to_q(norm_hidden_states)
key, _ = self.to_k(norm_hidden_states)
value, _ = self.to_v(norm_hidden_states)
@@ -249,9 +286,10 @@ class MatrixGameTransformerBlock(nn.Module):
# Apply rotary embeddings
cos, sin = freqs_cis
query, key = _apply_rotary_emb(query, cos, sin,
is_neox_style=False), _apply_rotary_emb(
key, cos, sin, is_neox_style=False)
query, key = (
_apply_rotary_emb(query, cos, sin, is_neox_style=False),
_apply_rotary_emb(key, cos, sin, is_neox_style=False),
)
attn_output, _ = self.attn1(query, key, value)
attn_output = attn_output.flatten(2)
@@ -260,18 +298,24 @@ class MatrixGameTransformerBlock(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)
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,
context=encoder_hidden_states,
context_lens=None)
attn_output = self.attn2(
norm_hidden_states, context=encoder_hidden_states, context_lens=None
)
norm_hidden_states, hidden_states = self.cross_attn_residual_norm(
hidden_states, attn_output, 1, c_shift_msa, c_scale_msa)
norm_hidden_states, hidden_states = norm_hidden_states.to(
orig_dtype), hidden_states.to(orig_dtype)
hidden_states, attn_output, 1, c_shift_msa, c_scale_msa
)
norm_hidden_states, hidden_states = (
norm_hidden_states.to(orig_dtype),
hidden_states.to(orig_dtype),
)
# ================= Action Module =================
if self.action_model is not None:
@@ -280,9 +324,12 @@ class MatrixGameTransformerBlock(nn.Module):
# ActionModule implementation takes hidden_states directly
hidden_states = self.action_model(
hidden_states,
grid_sizes[0], grid_sizes[1], grid_sizes[2],
mouse_cond, keyboard_cond,
num_frame_per_block=grid_sizes[0],
int(grid_sizes[0]),
int(grid_sizes[1]),
int(grid_sizes[2]),
mouse_cond,
keyboard_cond,
num_frame_per_block=int(grid_sizes[0]),
)
# =================================================
@@ -293,6 +340,7 @@ class MatrixGameTransformerBlock(nn.Module):
return hidden_states
_DEFAULT_MATRIXGAME_CONFIG = MatrixGameWanVideoConfig()
@@ -302,15 +350,23 @@ class MatrixGameWanModel(BaseDiT):
_fsdp_shard_conditions = _DEFAULT_MATRIXGAME_CONFIG._fsdp_shard_conditions
_compile_conditions = _DEFAULT_MATRIXGAME_CONFIG._compile_conditions
_supported_attention_backends = _DEFAULT_MATRIXGAME_CONFIG._supported_attention_backends
_supported_attention_backends = (
_DEFAULT_MATRIXGAME_CONFIG._supported_attention_backends
)
param_names_mapping = _DEFAULT_MATRIXGAME_CONFIG.param_names_mapping
reverse_param_names_mapping = _DEFAULT_MATRIXGAME_CONFIG.reverse_param_names_mapping
lora_param_names_mapping = _DEFAULT_MATRIXGAME_CONFIG.lora_param_names_mapping
reverse_param_names_mapping = (
_DEFAULT_MATRIXGAME_CONFIG.reverse_param_names_mapping
)
lora_param_names_mapping = (
_DEFAULT_MATRIXGAME_CONFIG.lora_param_names_mapping
)
def __init__(self,
config: MatrixGameWanVideoConfig,
hf_config: dict[str, Any],
**kwargs) -> None:
def __init__(
self,
config: MatrixGameWanVideoConfig,
hf_config: dict[str, Any],
**kwargs,
) -> None:
super().__init__(config=config, hf_config=hf_config)
inner_dim = config.num_attention_heads * config.attention_head_dim
@@ -320,10 +376,12 @@ class MatrixGameWanModel(BaseDiT):
self.patch_size = config.patch_size
# 1. Patch & position embedding
self.patch_embedding = PatchEmbed(in_chans=config.in_channels,
embed_dim=inner_dim,
patch_size=config.patch_size,
flatten=False)
self.patch_embedding = PatchEmbed(
in_chans=config.in_channels,
embed_dim=inner_dim,
patch_size=config.patch_size,
flatten=False,
)
# 2. Condition embeddings
self.condition_embedder = MatrixGameTimeImageEmbedding(
@@ -333,57 +391,73 @@ class MatrixGameWanModel(BaseDiT):
)
# 2.1. Get action config
self.action_config = getattr(config, 'action_config', {})
self.action_config = getattr(config, "action_config", {})
# 3. Transformer blocks
self.blocks = nn.ModuleList([
MatrixGameTransformerBlock(
inner_dim,
config.ffn_dim,
config.num_attention_heads,
config.qk_norm,
config.cross_attn_norm,
config.eps,
config.added_kv_proj_dim,
self._supported_attention_backends,
prefix=f"{getattr(config, 'prefix', 'Wan')}.blocks.{i}",
action_config=self.action_config)
for i in range(config.num_layers)
])
self.blocks = nn.ModuleList(
[
MatrixGameTransformerBlock(
inner_dim,
config.ffn_dim,
config.num_attention_heads,
config.qk_norm,
config.cross_attn_norm,
config.eps,
config.added_kv_proj_dim,
self._supported_attention_backends,
prefix=f"{getattr(config, 'prefix', 'Wan')}.blocks.{i}",
action_config=self.action_config,
)
for i in range(config.num_layers)
]
)
# 4. Output norm & projection
self.norm_out = LayerNormScaleShift(inner_dim,
norm_type="layer",
eps=config.eps,
elementwise_affine=False,
dtype=torch.float32,
compute_dtype=torch.float32)
self.norm_out = LayerNormScaleShift(
inner_dim,
norm_type="layer",
eps=config.eps,
elementwise_affine=False,
dtype=torch.float32,
compute_dtype=torch.float32,
)
self.proj_out = nn.Linear(
inner_dim, config.out_channels * math.prod(config.patch_size))
inner_dim, config.out_channels * math.prod(config.patch_size)
)
self.scale_shift_table = nn.Parameter(
torch.randn(1, 2, inner_dim) / inner_dim**0.5)
torch.randn(1, 2, inner_dim) / inner_dim**0.5
)
self.gradient_checkpointing = False
def forward(self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor | list[torch.Tensor],
timestep: torch.LongTensor,
encoder_hidden_states_image: torch.Tensor
| list[torch.Tensor] | None = None,
# Action inputs
mouse_cond: torch.Tensor | None = None,
keyboard_cond: torch.Tensor | None = None,
**kwargs) -> torch.Tensor:
if encoder_hidden_states is not None and not isinstance(encoder_hidden_states, torch.Tensor):
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor | list[torch.Tensor],
timestep: torch.LongTensor,
encoder_hidden_states_image: torch.Tensor
| list[torch.Tensor]
| None = None,
# Action inputs
mouse_cond: torch.Tensor | None = None,
keyboard_cond: torch.Tensor | None = None,
**kwargs,
) -> torch.Tensor:
if encoder_hidden_states is not None and not isinstance(
encoder_hidden_states, torch.Tensor
):
encoder_hidden_states = encoder_hidden_states[0]
if isinstance(encoder_hidden_states_image,
list) and len(encoder_hidden_states_image) > 0:
if (
isinstance(encoder_hidden_states_image, list)
and len(encoder_hidden_states_image) > 0
):
encoder_hidden_states_image = encoder_hidden_states_image[0]
else:
encoder_hidden_states_image = None
batch_size, num_channels, num_frames, height, width = hidden_states.shape
batch_size, num_channels, num_frames, height, width = (
hidden_states.shape
)
p_t, p_h, p_w = self.patch_size
post_patch_num_frames = num_frames // p_t
post_patch_height = height // p_h
@@ -393,18 +467,25 @@ class MatrixGameWanModel(BaseDiT):
d = self.hidden_size // self.num_attention_heads
rope_dim_list = [d - 4 * (d // 6), 2 * (d // 6), 2 * (d // 6)]
freqs_cos, freqs_sin = get_rotary_pos_embed(
(post_patch_num_frames * get_sp_world_size(), post_patch_height,
post_patch_width),
(
post_patch_num_frames * get_sp_world_size(),
post_patch_height,
post_patch_width,
),
self.hidden_size,
self.num_attention_heads,
rope_dim_list,
dtype=torch.float32 if current_platform.is_mps() else torch.float64,
rope_theta=10000,
do_sp_sharding=True)
do_sp_sharding=True,
)
freqs_cos = freqs_cos.to(hidden_states.device)
freqs_sin = freqs_sin.to(hidden_states.device)
freqs_cis = (freqs_cos.float(),
freqs_sin.float()) if freqs_cos is not None else None
freqs_cis = (
(freqs_cos.float(), freqs_sin.float())
if freqs_cos is not None
else None
)
hidden_states = self.patch_embedding(hidden_states)
hidden_states = hidden_states.flatten(2).transpose(1, 2)
@@ -412,8 +493,14 @@ class MatrixGameWanModel(BaseDiT):
if timestep.dim() == 2:
timestep = timestep.flatten()
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
timestep, encoder_hidden_states, encoder_hidden_states_image)
(
temb,
timestep_proj,
encoder_hidden_states,
encoder_hidden_states_image,
) = self.condition_embedder(
timestep, encoder_hidden_states, encoder_hidden_states_image
)
timestep_proj = timestep_proj.unflatten(1, (6, -1))
if encoder_hidden_states is not None:
@@ -429,25 +516,30 @@ class MatrixGameWanModel(BaseDiT):
if encoder_hidden_states_image is not None:
if encoder_hidden_states is not None:
encoder_hidden_states = torch.concat(
[encoder_hidden_states_image, encoder_hidden_states], dim=1)
[encoder_hidden_states_image, encoder_hidden_states], dim=1
)
else:
encoder_hidden_states = encoder_hidden_states_image
# This is [F, H, W] for the ActionModule
grid_sizes = torch.tensor([
post_patch_num_frames, post_patch_height, post_patch_width
],
device=hidden_states.device)
grid_sizes = torch.tensor(
[post_patch_num_frames, post_patch_height, post_patch_width],
device=hidden_states.device,
)
# Blocks
for block in self.blocks:
if torch.is_grad_enabled() and self.gradient_checkpointing:
hidden_states = self._gradient_checkpointing_func(
block, hidden_states, encoder_hidden_states, timestep_proj,
block,
hidden_states,
encoder_hidden_states,
timestep_proj,
freqs_cis,
grid_sizes=grid_sizes,
mouse_cond=mouse_cond,
keyboard_cond=keyboard_cond)
keyboard_cond=keyboard_cond,
)
else:
hidden_states = block(
hidden_states,
@@ -456,19 +548,27 @@ class MatrixGameWanModel(BaseDiT):
freqs_cis,
grid_sizes=grid_sizes,
mouse_cond=mouse_cond,
keyboard_cond=keyboard_cond)
keyboard_cond=keyboard_cond,
)
# Output
shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2,
dim=1)
shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(
2, dim=1
)
hidden_states = self.norm_out(hidden_states, shift, scale)
hidden_states = self.proj_out(hidden_states)
hidden_states = hidden_states.reshape(batch_size, post_patch_num_frames,
post_patch_height,
post_patch_width, p_t, p_h, p_w,
-1)
hidden_states = hidden_states.reshape(
batch_size,
post_patch_num_frames,
post_patch_height,
post_patch_width,
p_t,
p_h,
p_w,
-1,
)
hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6)
output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
return output
return output
+3
View File
@@ -685,3 +685,6 @@ class StepVideoModel(BaseDiT):
output = rearrange(output, '(b f) c h w -> b c f h w', f=frame)
return output
# Entry point for model registry
EntryClass = StepVideoModel
+3 -1
View File
@@ -871,4 +871,6 @@ class WanTransformer3DModel(CachableDiT):
return hidden_states + self.previous_residual_even
else:
return hidden_states + self.previous_residual_odd
# Entry point for model registry
EntryClass = WanTransformer3DModel
+3
View File
@@ -631,3 +631,6 @@ class CLIPVisionModel(ImageEncoder):
weight_loader(param, loaded_weight)
loaded_params.add(name)
return loaded_params
# Entry point for model registry
EntryClass = [CLIPTextModel, CLIPVisionModel]
+59 -2
View File
@@ -3,11 +3,11 @@ from __future__ import annotations
from dataclasses import dataclass
import os
from typing import Iterable
from typing import Any, Iterable
import torch
from torch import nn
from transformers import Gemma3ForConditionalGeneration
from transformers import AutoTokenizer, Gemma3ForConditionalGeneration
from fastvideo.configs.models.encoders import BaseEncoderOutput, TextEncoderConfig
from fastvideo.models.encoders.base import TextEncoder
@@ -447,6 +447,60 @@ class LTX2GemmaTextEncoderModel(TextEncoder):
return encoded, encoded_for_audio, attention_mask.squeeze(-1)
@torch.no_grad()
def preprocess_text_embeddings(
self,
prompts: str | list[str],
tokenizer: AutoTokenizer,
tokenizer_kwargs: dict[str, Any] | None = None,
padding_side: str | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Compute pre-connector text embeddings for LTX-2 training preprocessing."""
if isinstance(prompts, str):
prompts = [prompts]
model = self.gemma_model
kwargs: dict[str, Any] = {
"padding": "max_length",
"truncation": True,
"return_tensors": "pt",
}
if tokenizer_kwargs is not None:
kwargs.update(tokenizer_kwargs)
if "max_length" not in kwargs:
kwargs["max_length"] = self.config.arch_config.text_len
original_padding_side = tokenizer.padding_side
target_padding_side = padding_side or self.padding_side
tokenizer.padding_side = target_padding_side
try:
text_inputs = tokenizer(prompts, **kwargs)
finally:
tokenizer.padding_side = original_padding_side
input_ids = text_inputs["input_ids"].to(device=model.device)
attention_mask = text_inputs["attention_mask"].to(device=model.device)
outputs = model(
input_ids=input_ids,
attention_mask=attention_mask,
output_hidden_states=True,
return_dict=True,
)
prompt_embeds = self._run_feature_extractor(
outputs.hidden_states,
attention_mask,
padding_side=target_padding_side,
)
return prompt_embeds, attention_mask
def run_connectors(
self,
encoded_input: torch.Tensor,
attention_mask: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""Apply embedding connectors to precomputed Gemma features."""
return self._run_connectors(encoded_input, attention_mask)
def forward(
self,
input_ids: torch.Tensor | None,
@@ -561,3 +615,6 @@ def _norm_and_concat_padded_batch(
mask_flattened = mask.reshape(b, t, 1).expand(-1, -1, d * l)
normed = normed.masked_fill(~mask_flattened, 0.0)
return normed
# Entry point for model registry
EntryClass = LTX2GemmaTextEncoderModel
+3
View File
@@ -429,3 +429,6 @@ class LlamaModel(TextEncoder):
weight_loader(param, loaded_weight)
loaded_params.add(name)
return loaded_params
# Entry point for model registry
EntryClass = LlamaModel
+3
View File
@@ -385,3 +385,6 @@ class Qwen2_5_VLTextModel(TextEncoder):
weight_loader(param, loaded_weight)
loaded_params.add(name)
return loaded_params
# Entry point for model registry
EntryClass = Qwen2_5_VLTextModel
+2
View File
@@ -351,3 +351,5 @@ class Reason1TextEncoder(TextEncoder):
tensor.std(dim=-1, keepdim=True) + 1e-8
)
# Entry point for model registry
EntryClass = Reason1TextEncoder
+3
View File
@@ -417,3 +417,6 @@ class SiglipVisionModel(ImageEncoder):
loaded_params.add(name)
return loaded_params
# Entry point for model registry
EntryClass = SiglipVisionModel
+3
View File
@@ -583,3 +583,6 @@ class STEP1TextEncoder(torch.nn.Module):
self.device) if with_mask else None)
y_mask = txt_tokens.attention_mask
return y.transpose(0, 1), y_mask
# Entry point for model registry
EntryClass = STEP1TextEncoder
+3
View File
@@ -719,3 +719,6 @@ class UMT5EncoderModel(TextEncoder):
weight_loader(param, loaded_weight)
loaded_params.add(name)
return loaded_params
# Entry point for model registry
EntryClass = [UMT5EncoderModel, T5EncoderModel]
+51 -2
View File
@@ -79,6 +79,7 @@ class ComponentLoader(ABC):
"scheduler": (SchedulerLoader, "diffusers"),
"transformer": (TransformerLoader, "diffusers"),
"transformer_2": (TransformerLoader, "diffusers"),
"transformer_3": (TransformerLoader, "diffusers"),
"vae": (VAELoader, "diffusers"),
"audio_vae": (AudioDecoderLoader, "diffusers"),
"audio_decoder": (AudioDecoderLoader, "diffusers"),
@@ -90,6 +91,8 @@ class ComponentLoader(ABC):
"image_processor": (ImageProcessorLoader, "transformers"),
"feature_extractor": (ImageProcessorLoader, "transformers"),
"image_encoder": (ImageEncoderLoader, "transformers"),
"upsampler": (UpsamplerLoader, "diffusers"),
"upsampler_2": (UpsamplerLoader, "diffusers"),
}
if module_type in module_loaders:
@@ -558,7 +561,8 @@ class VAELoader(ComponentLoader):
def load(self, model_path: str, fastvideo_args: FastVideoArgs):
"""Load the VAE based on the model path, and inference args."""
config = get_diffusers_config(model=model_path)
class_name = config.get("_class_name")
class_name = config.pop("_class_name")
config.pop("_name_or_path", None)
assert class_name is not None, (
"Model config does not contain a _class_name attribute. Only diffusers format is supported."
)
@@ -731,6 +735,7 @@ class TransformerLoader(ComponentLoader):
config = get_diffusers_config(model=model_path)
hf_config = deepcopy(config)
cls_name = config.pop("_class_name")
config.pop("_name_or_path", None)
if cls_name is None:
raise ValueError(
"Model config does not contain a _class_name attribute. "
@@ -745,7 +750,7 @@ class TransformerLoader(ComponentLoader):
fastvideo_args.model_paths["transformer"] = model_path
# Config from Diffusers supersedes fastvideo's model config
dit_config = fastvideo_args.pipeline_config.dit_config
dit_config = deepcopy(fastvideo_args.pipeline_config.dit_config)
dit_config.update_model_arch(config)
model_cls, _ = ModelRegistry.resolve_model_cls(cls_name)
@@ -877,6 +882,50 @@ class SchedulerLoader(ComponentLoader):
return scheduler
class UpsamplerLoader(ComponentLoader):
"""Loader for upsamplers."""
def load(self, model_path: str, fastvideo_args: FastVideoArgs):
"""Load the upsampler based on the model path, and inference args."""
config_dict = get_diffusers_config(model=model_path)
class_name = config_dict.pop("_class_name", None)
if class_name is None:
raise ValueError(
"Model config does not contain a _class_name attribute. "
"Only diffusers format is supported."
)
try:
upsampler_cfg = deepcopy(fastvideo_args.pipeline_config.upsampler_config[0])
upsampler_cfg.update_model_config(config_dict)
except Exception as e:
upsampler_cfg = deepcopy(fastvideo_args.pipeline_config.upsampler_config[1])
upsampler_cfg.update_model_config(config_dict)
model_cls, _ = ModelRegistry.resolve_model_cls(class_name)
model = model_cls(upsampler_cfg)
target_device = get_local_torch_device()
model = model.to(target_device, dtype=PRECISION_TO_TYPE[fastvideo_args.pipeline_config.upsampler_precision])
# Find all safetensors files
safetensors_list = glob.glob(
os.path.join(str(model_path), "*.safetensors"))
if not safetensors_list:
raise ValueError(f"No safetensors files found in {model_path}")
if len(safetensors_list) == 1:
loaded = safetensors_load_file(safetensors_list[0])
else:
loaded = {}
for sf_file in safetensors_list:
loaded.update(safetensors_load_file(sf_file))
model.load_state_dict(loaded, strict=True)
return model.eval()
class GenericComponentLoader(ComponentLoader):
"""Generic loader for components that don't have a specific loader."""
+103 -2
View File
@@ -1,6 +1,7 @@
# SPDX-License-Identifier: Apache-2.0
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/models/registry.py
import ast
import importlib
import os
import pickle
@@ -98,7 +99,12 @@ _SCHEDULERS = {
("schedulers", "scheduling_rcm", "RCMScheduler"),
}
_FAST_VIDEO_MODELS = {
_UPSAMPLERS = {
"SRTo720pUpsampler": ("upsamplers", "hunyuan15", "SRTo720pUpsampler"),
"SRTo1080pUpsampler": ("upsamplers", "hunyuan15", "SRTo1080pUpsampler"),
}
_LEGACY_FAST_VIDEO_MODELS = {
**_TEXT_TO_VIDEO_DIT_MODELS,
**_IMAGE_TO_VIDEO_DIT_MODELS,
**_TEXT_ENCODER_MODELS,
@@ -106,8 +112,102 @@ _FAST_VIDEO_MODELS = {
**_VAE_MODELS,
**_AUDIO_MODELS,
**_SCHEDULERS,
**_UPSAMPLERS,
}
MODELS_PATH = os.path.dirname(__file__)
@lru_cache(maxsize=None)
def _discover_and_register_models() -> dict[str, tuple[str, str, str]]:
discovered_models: dict[str, tuple[str, str, str]] = {}
for root, dirs, files in os.walk(MODELS_PATH):
dirs[:] = [
d for d in dirs
if not d.startswith(".") and d != "__pycache__"
]
for filename in files:
if not filename.endswith(".py"):
continue
filepath = os.path.join(root, filename)
try:
with open(filepath, "r", encoding="utf-8") as f:
source = f.read()
tree = ast.parse(source, filename=filename)
entry_class_node = None
first_class_def = None
for node in ast.walk(tree):
if isinstance(node, ast.Assign):
for target in node.targets:
if isinstance(target, ast.Name) and target.id == "EntryClass":
entry_class_node = node
break
if first_class_def is None and isinstance(node, ast.ClassDef):
first_class_def = node
if not entry_class_node or not first_class_def:
continue
model_cls_name_list: list[str] = []
value_node = entry_class_node.value
if isinstance(value_node, ast.Name):
model_cls_name_list.append(value_node.id)
elif isinstance(value_node, (ast.List, ast.Tuple)):
for elt in value_node.elts:
if isinstance(elt, ast.Constant) and isinstance(
elt.value, str):
model_cls_name_list.append(elt.value)
elif isinstance(elt, ast.Name):
model_cls_name_list.append(elt.id)
if not model_cls_name_list:
continue
rel_dir = os.path.relpath(root, MODELS_PATH)
if rel_dir == ".":
continue
rel_parts = rel_dir.split(os.sep)
component_name = rel_parts[0]
sub_parts = rel_parts[1:]
if filename == "__init__.py":
mod_relname = ".".join(sub_parts)
else:
mod_base = filename[:-3]
mod_relname = ".".join(sub_parts +
[mod_base]) if sub_parts else mod_base
for model_cls_str in model_cls_name_list:
if model_cls_str in discovered_models:
logger.warning(
"Duplicate architecture found: %s. Overwriting.",
model_cls_str)
discovered_models[model_cls_str] = (
component_name,
mod_relname,
model_cls_str,
)
except Exception as e:
logger.warning("Could not parse %s to find models: %s",
filepath, e)
return discovered_models
_DISCOVERED_MODELS = _discover_and_register_models()
_FAST_VIDEO_MODELS = dict(_DISCOVERED_MODELS)
for model_arch, spec in _LEGACY_FAST_VIDEO_MODELS.items():
if model_arch in _FAST_VIDEO_MODELS:
continue
_FAST_VIDEO_MODELS[model_arch] = spec
_SUBPROCESS_COMMAND = [sys.executable, "-m", "fastvideo.models.dits.registry"]
_T = TypeVar("_T")
@@ -339,7 +439,8 @@ class _ModelRegistry:
ModelRegistry = _ModelRegistry({
model_arch:
_LazyRegisteredModel(
module_name=f"fastvideo.models.{component_name}.{mod_relname}",
module_name=(f"fastvideo.models.{component_name}.{mod_relname}"
if mod_relname else f"fastvideo.models.{component_name}"),
component_name=component_name,
class_name=cls_name,
)
@@ -685,3 +685,6 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
def __len__(self) -> int:
return self.config.num_train_timesteps
# Entry point for model registry
EntryClass = FlowMatchEulerDiscreteScheduler
@@ -853,3 +853,6 @@ class FlowUniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
def __len__(self):
return self.config.num_train_timesteps
# Entry point for model registry
EntryClass = FlowUniPCMultistepScheduler
@@ -321,3 +321,6 @@ class RCMScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
def __len__(self) -> int:
return self.config.num_train_timesteps
# Entry point for model registry
EntryClass = RCMScheduler
@@ -160,4 +160,6 @@ class SelfForcingFlowMatchScheduler(BaseScheduler, ConfigMixin, SchedulerMixin):
def set_shift(self, shift: float) -> None:
self.shift = shift
# Entry point for model registry
EntryClass = SelfForcingFlowMatchScheduler
@@ -1094,4 +1094,7 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
return noisy_samples
def __len__(self):
return self.config.num_train_timesteps
return self.config.num_train_timesteps
# Entry point for model registry
EntryClass = UniPCMultistepScheduler
+170
View File
@@ -0,0 +1,170 @@
# Licensed under the TENCENT HUNYUAN COMMUNITY LICENSE AGREEMENT (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# https://github.com/Tencent-Hunyuan/HunyuanVideo-1.5/blob/main/LICENSE
#
# Unless and only to the extent required by applicable law, the Tencent Hunyuan works and any
# output and results therefrom are provided "AS IS" without any express or implied warranties of
# any kind including any warranties of title, merchantability, noninfringement, course of dealing,
# usage of trade, or fitness for a particular purpose. You are solely responsible for determining the
# appropriateness of using, reproducing, modifying, performing, displaying or distributing any of
# the Tencent Hunyuan works or outputs and assume any and all risks associated with your or a
# third party's use or distribution of any of the Tencent Hunyuan works or outputs and your exercise
# of rights and permissions under this agreement.
# See the License for the specific language governing permissions and limitations under the License.
from collections.abc import Sequence
from dataclasses import dataclass
from enum import Enum
import torch
import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange
from torch import Tensor
from fastvideo.models.vaes.hunyuan15vae import (
HunyuanVideo15CausalConv3d,
HunyuanVideo15RMS_norm,
)
from fastvideo.layers.activation import get_act_fn
from fastvideo.configs.models.upsamplers import SRTo720pUpsamplerConfig, SRTo1080pUpsamplerConfig
class HunyuanVideo15ResnetBlock(nn.Module):
def __init__(
self,
in_channels: int,
out_channels: int | None = None,
non_linearity: str = "swish",
) -> None:
super().__init__()
out_channels = out_channels or in_channels
self.nonlinearity = get_act_fn(non_linearity)
self.norm1 = HunyuanVideo15RMS_norm(in_channels, images=False)
self.conv1 = HunyuanVideo15CausalConv3d(in_channels, out_channels, kernel_size=3)
self.norm2 = HunyuanVideo15RMS_norm(out_channels, images=False)
self.conv2 = HunyuanVideo15CausalConv3d(out_channels, out_channels, kernel_size=3)
self.nin_shortcut = None
if in_channels != out_channels:
self.nin_shortcut = nn.Conv3d(in_channels, out_channels, kernel_size=1, stride=1, padding=0)
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
residual = hidden_states
hidden_states = self.norm1(hidden_states)
hidden_states = self.nonlinearity(hidden_states)
hidden_states = self.conv1(hidden_states)
hidden_states = self.norm2(hidden_states)
hidden_states = self.nonlinearity(hidden_states)
hidden_states = self.conv2(hidden_states)
if self.nin_shortcut is not None:
residual = self.nin_shortcut(residual)
return hidden_states + residual
class SRResidualCausalBlock3D(nn.Module):
def __init__(self, channels: int):
super().__init__()
self.block = nn.Sequential(
HunyuanVideo15CausalConv3d(channels, channels, kernel_size=3),
nn.SiLU(inplace=True),
HunyuanVideo15CausalConv3d(channels, channels, kernel_size=3),
nn.SiLU(inplace=True),
HunyuanVideo15CausalConv3d(channels, channels, kernel_size=3),
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return x + self.block(x)
class SRTo720pUpsampler(nn.Module):
def __init__(
self,
config: SRTo720pUpsamplerConfig,
):
super().__init__()
self.in_conv = HunyuanVideo15CausalConv3d(config.in_channels, config.hidden_channels, kernel_size=3)
self.blocks = nn.ModuleList([SRResidualCausalBlock3D(config.hidden_channels) for _ in range(config.num_blocks)])
self.out_conv = HunyuanVideo15CausalConv3d(config.hidden_channels, config.out_channels, kernel_size=3)
self.global_residual = bool(config.global_residual)
def forward(self, x: torch.Tensor) -> torch.Tensor:
residual = x
y = self.in_conv(x)
for blk in self.blocks:
y = blk(y)
y = self.out_conv(y)
if self.global_residual and (y.shape == residual.shape):
y = y + residual
return y
class SRTo1080pUpsampler(nn.Module):
def __init__(
self,
config: SRTo1080pUpsamplerConfig,
):
super().__init__()
self.num_res_blocks = config.num_res_blocks
self.block_out_channels = config.block_out_channels
self.z_channels = config.z_channels
block_in = config.block_out_channels[0]
self.conv_in = HunyuanVideo15CausalConv3d(config.z_channels, block_in, kernel_size=3)
self.up = nn.ModuleList()
for i_level, ch in enumerate(config.block_out_channels):
block = nn.ModuleList()
block_out = ch
for _ in range(self.num_res_blocks + 1):
block.append(HunyuanVideo15ResnetBlock(in_channels=block_in, out_channels=block_out))
block_in = block_out
up = nn.Module()
up.block = block
self.up.append(up)
self.norm_out = HunyuanVideo15RMS_norm(block_in, images=False)
self.conv_out = HunyuanVideo15CausalConv3d(block_in, config.out_channels, kernel_size=3)
self.gradient_checkpointing = False
self.is_residual = config.is_residual
def forward(self, z: Tensor, target_shape: Sequence[int] = None) -> Tensor:
"""
Args:
z: (B, C, T, H, W)
target_shape: (H, W)
"""
if target_shape is not None and z.shape[-2:] != target_shape:
bsz = z.shape[0]
z = rearrange(z, "b c f h w -> (b f) c h w")
z = F.interpolate(z, size=target_shape, mode="bilinear", align_corners=False)
z = rearrange(z, "(b f) c h w -> b c f h w", b=bsz)
# z to block_in
repeats = self.block_out_channels[0] // (self.z_channels)
h = self.conv_in(z) + z.repeat_interleave(repeats=repeats, dim=1)
# upsampling
for i_level in range(len(self.block_out_channels)):
for i_block in range(self.num_res_blocks + 1):
h = self.up[i_level].block[i_block](h)
if hasattr(self.up[i_level], "upsample"):
h = self.up[i_level].upsample(h)
# end
h = self.norm_out(h)
h = get_act_fn("swish")(h)
h = self.conv_out(h)
return h
+4 -1
View File
@@ -700,4 +700,7 @@ class AutoencoderKLHunyuanVideo15(nn.Module, ParallelTiledVAE):
else:
z = posterior.mode()
dec = self.decode(z)
return dec
return dec
# Entry point for model registry
EntryClass = AutoencoderKLHunyuanVideo15
+3
View File
@@ -850,3 +850,6 @@ class AutoencoderKLHunyuanVideo(nn.Module, ParallelTiledVAE):
z = posterior.mode()
dec = self.decode(z)
return dec
# Entry point for model registry
EntryClass = AutoencoderKLHunyuanVideo
+968 -6
View File
@@ -15,15 +15,977 @@
# See the License for the specific language governing permissions and
# limitations under the License.
from fastvideo.models.vaes.hunyuan15vae import AutoencoderKLHunyuanVideo15
from fastvideo.configs.models.vaes import Hunyuan15VAEConfig
from typing import List, Optional, Tuple, Union
class AutoencoderKLHYWorld(AutoencoderKLHunyuanVideo15):
# TODO(mingjia): add temporal caching support for HYWorld VAE
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.utils.checkpoint
from einops import rearrange
from fastvideo.layers.activation import get_act_fn
from fastvideo.configs.models.vaes import Hunyuan15VAEConfig
from fastvideo.models.vaes.common import ParallelTiledVAE
from fastvideo.models.vaes.hunyuan15vae import (
HunyuanVideo15RMS_norm as HYWorldRMS_norm,
HunyuanVideo15AttnBlock as HYWorldAttnBlock,
)
# Cache size for temporal feature caching (number of frames to cache)
CACHE_T = 2
class HYWorldCausalConv3d(nn.Module):
"""Causal Conv3d with optional cache support for temporal feature caching."""
def __init__(
self,
in_channels: int,
out_channels: int,
kernel_size: Union[int, Tuple[int, int, int]] = 3,
stride: Union[int, Tuple[int, int, int]] = 1,
padding: Union[int, Tuple[int, int, int]] = 0,
dilation: Union[int, Tuple[int, int, int]] = 1,
bias: bool = True,
pad_mode: str = "replicate",
) -> None:
super().__init__()
kernel_size = (kernel_size, kernel_size, kernel_size) if isinstance(kernel_size, int) else kernel_size
self.pad_mode = pad_mode
# Padding format: (W_left, W_right, H_left, H_right, T_left, T_right)
self.time_causal_padding = (
kernel_size[0] // 2, # W_left (spatial)
kernel_size[0] // 2, # W_right (spatial)
kernel_size[1] // 2, # H_left (spatial)
kernel_size[1] // 2, # H_right (spatial)
kernel_size[2] - 1, # T_left (temporal causal padding)
0, # T_right (no future padding for causal)
)
self.conv = nn.Conv3d(in_channels, out_channels, kernel_size, stride, padding, dilation, bias=bias)
def forward(self, hidden_states: torch.Tensor, cache_x: Optional[torch.Tensor] = None) -> torch.Tensor:
"""
Forward pass with optional temporal caching.
Args:
hidden_states: Input tensor of shape (B, C, T, H, W)
cache_x: Optional cached frames from previous chunk, shape (B, C, CACHE_T, H, W)
When provided, uses cached frames instead of padding for temporal dimension.
"""
padding = list(self.time_causal_padding)
if cache_x is not None and self.time_causal_padding[4] > 0: # Has temporal padding and cache
cache_x = cache_x.to(hidden_states.device)
# Concatenate cached frames with current input on temporal dimension
hidden_states = torch.cat([cache_x, hidden_states], dim=2)
# Reduce temporal padding since we have cached frames
padding[4] -= cache_x.shape[2]
hidden_states = F.pad(hidden_states, padding, mode=self.pad_mode)
return self.conv(hidden_states)
class HYWorldUpsample(nn.Module):
"""Hierarchical upsampling with temporal/spatial support and optional caching."""
def __init__(self, in_channels: int, out_channels: int, add_temporal_upsample: bool = True):
super().__init__()
factor = 2 * 2 * 2 if add_temporal_upsample else 1 * 2 * 2
self.conv = HYWorldCausalConv3d(in_channels, out_channels * factor, kernel_size=3)
self.add_temporal_upsample = add_temporal_upsample
self.repeats = factor * out_channels // in_channels
def forward(
self,
x: torch.Tensor,
feat_cache: Optional[List[Optional[torch.Tensor]]] = None,
feat_idx: Optional[List[int]] = None,
first_chunk: bool = False,
):
"""
Forward pass with optional temporal caching.
Args:
x: Input tensor of shape (B, C, T, H, W)
feat_cache: List of cached features for each conv layer
feat_idx: List containing current cache index [idx]
first_chunk: Whether this is the first chunk (affects temporal upsample behavior)
"""
r1 = 2 if self.add_temporal_upsample else 1
if feat_cache is not None and feat_idx is not None:
idx = feat_idx[0]
cache_x = x[:, :, -CACHE_T:, :, :].clone()
if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
cache_x = torch.cat(
[feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), cache_x],
dim=2,
)
h = self.conv(x, feat_cache[idx])
feat_cache[idx] = cache_x
feat_idx[0] += 1
else:
h = self.conv(x)
if self.add_temporal_upsample:
if first_chunk:
# First chunk: only spatial upsample
h = rearrange(h, "b (r2 r3 c) f h w -> b c f (h r2) (w r3)", r2=2, r3=2)
h = h[:, : h.shape[1] // 2]
# Compute the shortcut part
shortcut = rearrange(x, "b (r2 r3 c) f h w -> b c f (h r2) (w r3)", r2=2, r3=2)
shortcut = shortcut.repeat_interleave(repeats=self.repeats // 2, dim=1)
elif feat_cache is None and x.shape[2] > 1:
# No cache and multiple frames: separate first frame and rest
h_first = h[:, :, :1, :, :]
h_rest = h[:, :, 1:, :, :]
x_first = x[:, :, :1, :, :]
x_rest = x[:, :, 1:, :, :]
# First frame: only spatial upsample
h_first = rearrange(h_first, "b (r2 r3 c) f h w -> b c f (h r2) (w r3)", r2=2, r3=2)
h_first = h_first[:, : h_first.shape[1] // 2]
shortcut_first = rearrange(x_first, "b (r2 r3 c) f h w -> b c f (h r2) (w r3)", r2=2, r3=2)
shortcut_first = shortcut_first.repeat_interleave(repeats=self.repeats // 2, dim=1)
out_first = h_first + shortcut_first
# Remaining frames: spatio-temporal upsample
h_rest = rearrange(h_rest, "b (r1 r2 r3 c) f h w -> b c (f r1) (h r2) (w r3)", r1=r1, r2=2, r3=2)
shortcut_rest = rearrange(x_rest, "b (r1 r2 r3 c) f h w -> b c (f r1) (h r2) (w r3)", r1=r1, r2=2, r3=2)
shortcut_rest = shortcut_rest.repeat_interleave(repeats=self.repeats, dim=1)
out_rest = h_rest + shortcut_rest
return torch.cat([out_first, out_rest], dim=2)
else:
# Subsequent chunks with cache: spatio-temporal upsample
h = rearrange(h, "b (r1 r2 r3 c) f h w -> b c (f r1) (h r2) (w r3)", r1=r1, r2=2, r3=2)
shortcut = rearrange(x, "b (r1 r2 r3 c) f h w -> b c (f r1) (h r2) (w r3)", r1=r1, r2=2, r3=2)
shortcut = shortcut.repeat_interleave(repeats=self.repeats, dim=1)
else:
h = rearrange(h, "b (r1 r2 r3 c) f h w -> b c (f r1) (h r2) (w r3)", r1=r1, r2=2, r3=2)
shortcut = x.repeat_interleave(repeats=self.repeats, dim=1)
shortcut = rearrange(shortcut, "b (r1 r2 r3 c) f h w -> b c (f r1) (h r2) (w r3)", r1=r1, r2=2, r3=2)
return h + shortcut
class HYWorldDownsample(nn.Module):
"""Hierarchical downsampling with temporal/spatial support and optional caching."""
def __init__(self, in_channels: int, out_channels: int, add_temporal_downsample: bool = True):
super().__init__()
factor = 2 * 2 * 2 if add_temporal_downsample else 1 * 2 * 2
self.conv = HYWorldCausalConv3d(in_channels, out_channels // factor, kernel_size=3)
self.add_temporal_downsample = add_temporal_downsample
self.group_size = factor * in_channels // out_channels
def forward(
self,
x: torch.Tensor,
feat_cache: Optional[List[Optional[torch.Tensor]]] = None,
feat_idx: Optional[List[int]] = None,
):
"""
Forward pass with optional temporal caching.
Args:
x: Input tensor of shape (B, C, T, H, W)
feat_cache: List of cached features for each conv layer
feat_idx: List containing current cache index [idx]
"""
r1 = 2 if self.add_temporal_downsample else 1
# Apply conv with caching
if feat_cache is not None and feat_idx is not None:
idx = feat_idx[0]
cache_x = x[:, :, -CACHE_T:, :, :].clone()
if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
cache_x = torch.cat(
[feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), cache_x],
dim=2,
)
h = self.conv(x, feat_cache[idx])
feat_cache[idx] = cache_x
feat_idx[0] += 1
else:
h = self.conv(x) # Change the channel, ready for spatial or temporal downsample
if self.add_temporal_downsample:
if x.shape[2] == 1:
# Single frame: only spatial downsample
h = rearrange(h, "b c f (h r2) (w r3) -> b (r2 r3 c) f h w", r2=2, r3=2)
h = torch.cat([h, h], dim=1)
# Compute the shortcut part
shortcut = rearrange(x, "b c f (h r2) (w r3) -> b (r2 r3 c) f h w", r2=2, r3=2)
B, C, T, H, W = shortcut.shape
shortcut = shortcut.view(B, h.shape[1], self.group_size // 2, T, H, W).mean(dim=2)
else:
# Multiple frames: full spatio-temporal downsample
h = rearrange(h, "b c (f r1) (h r2) (w r3) -> b (r1 r2 r3 c) f h w", r1=r1, r2=2, r3=2)
# Shortcut computation
shortcut = rearrange(x, "b c (f r1) (h r2) (w r3) -> b (r1 r2 r3 c) f h w", r1=r1, r2=2, r3=2)
B, C, T, H, W = shortcut.shape
shortcut = shortcut.view(B, h.shape[1], self.group_size, T, H, W).mean(dim=2)
else:
h = rearrange(h, "b c (f r1) (h r2) (w r3) -> b (r1 r2 r3 c) f h w", r1=r1, r2=2, r3=2)
shortcut = rearrange(x, "b c (f r1) (h r2) (w r3) -> b (r1 r2 r3 c) f h w", r1=r1, r2=2, r3=2)
B, C, T, H, W = shortcut.shape
shortcut = shortcut.view(B, h.shape[1], self.group_size, T, H, W).mean(dim=2)
return h + shortcut
class HYWorldResnetBlock(nn.Module):
def __init__(
self,
in_channels: int,
out_channels: Optional[int] = None,
non_linearity: str = "swish",
) -> None:
super().__init__()
self.in_channels = in_channels
out_channels = in_channels if out_channels is None else out_channels
self.out_channels = out_channels
self.nonlinearity = get_act_fn(non_linearity)
self.norm1 = HYWorldRMS_norm(in_channels, images=False)
self.conv1 = HYWorldCausalConv3d(in_channels, out_channels, kernel_size=3)
self.norm2 = HYWorldRMS_norm(out_channels, images=False)
self.conv2 = HYWorldCausalConv3d(out_channels, out_channels, kernel_size=3)
if in_channels != out_channels:
self.nin_shortcut = nn.Conv3d(in_channels, out_channels, kernel_size=1, stride=1, padding=0)
def forward(
self,
hidden_states: torch.Tensor,
feat_cache: Optional[List[Optional[torch.Tensor]]] = None,
feat_idx: Optional[List[int]] = None,
) -> torch.Tensor:
"""
Forward pass with optional temporal feature caching.
Args:
hidden_states: Input tensor of shape (B, C, T, H, W)
feat_cache: List of cached features for each conv layer
feat_idx: List containing current cache index [idx]
"""
residual = hidden_states
hidden_states = self.norm1(hidden_states)
hidden_states = self.nonlinearity(hidden_states)
# apply the feature cacheing mechanism
if feat_cache is not None and feat_idx is not None:
# Retrieve the current layer index.
idx = feat_idx[0]
# Clone the last CACHE_T frames from the current input to store for the next step.
cache_x = hidden_states[:, :, -CACHE_T:, :, :].clone()
# Handle boundary conditions: if the current temporal chunk is too short (< 2 frames)
# and we have a previous cache, prepend the last frame of the previous cache.
# This ensures sufficient temporal context for the convolution kernel.
if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
# actually this means chunk 1 after chunk 0
cache_x = torch.cat(
[feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), cache_x],
dim=2,
)
# Apply the convolution layer using `hidden_states`and the *previous* cached state.
hidden_states = self.conv1(hidden_states, feat_cache[idx])
# Update the cache for this layer with the newly prepared context ('cache_x').
feat_cache[idx] = cache_x
# Increment the global layer index.
feat_idx[0] += 1
else:
hidden_states = self.conv1(hidden_states)
hidden_states = self.norm2(hidden_states)
hidden_states = self.nonlinearity(hidden_states)
# Second conv with caching
if feat_cache is not None and feat_idx is not None:
idx = feat_idx[0]
cache_x = hidden_states[:, :, -CACHE_T:, :, :].clone()
if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
cache_x = torch.cat(
[feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), cache_x],
dim=2,
)
hidden_states = self.conv2(hidden_states, feat_cache[idx])
feat_cache[idx] = cache_x
feat_idx[0] += 1
else:
hidden_states = self.conv2(hidden_states)
if self.in_channels != self.out_channels:
residual = self.nin_shortcut(residual)
return hidden_states + residual
class HYWorldMidBlock(nn.Module):
"""Mid block with attention and resnet blocks, with optional caching support."""
def __init__(
self,
in_channels: int,
num_layers: int = 1,
add_attention: bool = True,
) -> None:
super().__init__()
self.add_attention = add_attention
# There is always at least one resnet
resnets = [
HYWorldResnetBlock(
in_channels=in_channels,
out_channels=in_channels,
)
]
attentions = []
for _ in range(num_layers):
if self.add_attention:
attentions.append(HYWorldAttnBlock(in_channels))
else:
attentions.append(None)
resnets.append(
HYWorldResnetBlock(
in_channels=in_channels,
out_channels=in_channels,
)
)
self.attentions = nn.ModuleList(attentions)
self.resnets = nn.ModuleList(resnets)
self.gradient_checkpointing = False
def forward(
self,
hidden_states: torch.Tensor,
feat_cache: Optional[List[Optional[torch.Tensor]]] = None,
feat_idx: Optional[List[int]] = None,
) -> torch.Tensor:
"""
Forward pass with optional temporal caching.
Args:
hidden_states: Input tensor of shape (B, C, T, H, W)
feat_cache: List of cached features for each conv layer
feat_idx: List containing current cache index [idx]
"""
hidden_states = self.resnets[0](hidden_states, feat_cache, feat_idx)
for attn, resnet in zip(self.attentions, self.resnets[1:]):
if attn is not None:
hidden_states = attn(hidden_states)
hidden_states = resnet(hidden_states, feat_cache, feat_idx)
return hidden_states
class HYWorldDownBlock3D(nn.Module):
"""Down block with resnet blocks and optional downsampling, with caching support."""
def __init__(
self,
in_channels: int,
out_channels: int,
num_layers: int = 1,
downsample_out_channels: Optional[int] = None,
add_temporal_downsample: int = True,
) -> None:
super().__init__()
resnets = []
for i in range(num_layers):
in_channels = in_channels if i == 0 else out_channels
resnets.append(
HYWorldResnetBlock(
in_channels=in_channels,
out_channels=out_channels,
)
)
self.resnets = nn.ModuleList(resnets)
if downsample_out_channels is not None:
self.downsamplers = nn.ModuleList(
[
HYWorldDownsample(
out_channels,
out_channels=downsample_out_channels,
add_temporal_downsample=add_temporal_downsample,
)
]
)
else:
self.downsamplers = None
self.gradient_checkpointing = False
def forward(
self,
hidden_states: torch.Tensor,
feat_cache: Optional[List[Optional[torch.Tensor]]] = None,
feat_idx: Optional[List[int]] = None,
) -> torch.Tensor:
"""
Forward pass with optional temporal caching.
Args:
hidden_states: Input tensor of shape (B, C, T, H, W)
feat_cache: List of cached features for each conv layer
feat_idx: List containing current cache index [idx]
"""
for resnet in self.resnets:
hidden_states = resnet(hidden_states, feat_cache, feat_idx)
if self.downsamplers is not None:
for downsampler in self.downsamplers:
hidden_states = downsampler(hidden_states, feat_cache, feat_idx)
return hidden_states
class HYWorldUpBlock3D(nn.Module):
"""Up block with resnet blocks and optional upsampling, with caching support."""
def __init__(
self,
in_channels: int,
out_channels: int,
num_layers: int = 1,
upsample_out_channels: Optional[int] = None,
add_temporal_upsample: bool = True,
) -> None:
super().__init__()
resnets = []
for i in range(num_layers):
input_channels = in_channels if i == 0 else out_channels
resnets.append(
HYWorldResnetBlock(
in_channels=input_channels,
out_channels=out_channels,
)
)
self.resnets = nn.ModuleList(resnets)
if upsample_out_channels is not None:
self.upsamplers = nn.ModuleList(
[
HYWorldUpsample(
out_channels,
out_channels=upsample_out_channels,
add_temporal_upsample=add_temporal_upsample,
)
]
)
else:
self.upsamplers = None
self.gradient_checkpointing = False
def forward(
self,
hidden_states: torch.Tensor,
feat_cache: Optional[List[Optional[torch.Tensor]]] = None,
feat_idx: Optional[List[int]] = None,
first_chunk: bool = False,
) -> torch.Tensor:
"""
Forward pass with optional temporal caching.
Args:
hidden_states: Input tensor of shape (B, C, T, H, W)
feat_cache: List of cached features for each conv layer
feat_idx: List containing current cache index [idx]
first_chunk: Whether this is the first chunk (for upsampling behavior)
"""
if torch.is_grad_enabled() and self.gradient_checkpointing:
for resnet in self.resnets:
hidden_states = self._gradient_checkpointing_func(resnet, hidden_states)
else:
for resnet in self.resnets:
hidden_states = resnet(hidden_states, feat_cache, feat_idx)
if self.upsamplers is not None:
for upsampler in self.upsamplers:
hidden_states = upsampler(hidden_states, feat_cache, feat_idx, first_chunk)
return hidden_states
class HYWorldEncoder3D(nn.Module):
r"""
3D vae encoder for HunyuanImageRefiner with optional temporal caching.
"""
def __init__(
self,
config: Hunyuan15VAEConfig,
in_channels: int = 3,
out_channels: int = 64,
block_out_channels: Tuple[int, ...] = (128, 256, 512, 1024, 1024),
layers_per_block: int = 2,
temporal_compression_ratio: int = 4,
spatial_compression_ratio: int = 16,
downsample_match_channel: bool = True,
) -> None:
AutoencoderKLHunyuanVideo15.__init__(self, config)
super().__init__()
self.in_channels = in_channels
self.out_channels = out_channels
self.group_size = block_out_channels[-1] // self.out_channels
self.block_out_channels = block_out_channels
self.conv_in = HYWorldCausalConv3d(in_channels, block_out_channels[0], kernel_size=3)
self.mid_block = None
self.down_blocks = nn.ModuleList([])
input_channel = block_out_channels[0]
for i in range(len(block_out_channels)):
add_spatial_downsample = i < np.log2(spatial_compression_ratio)
output_channel = block_out_channels[i]
if not add_spatial_downsample:
down_block = HYWorldDownBlock3D(
num_layers=layers_per_block,
in_channels=input_channel,
out_channels=output_channel,
downsample_out_channels=None,
add_temporal_downsample=False,
)
input_channel = output_channel
else:
add_temporal_downsample = i >= np.log2(spatial_compression_ratio // temporal_compression_ratio)
downsample_out_channels = block_out_channels[i + 1] if downsample_match_channel else output_channel
down_block = HYWorldDownBlock3D(
num_layers=layers_per_block,
in_channels=input_channel,
out_channels=output_channel,
downsample_out_channels=downsample_out_channels,
add_temporal_downsample=add_temporal_downsample,
)
input_channel = downsample_out_channels
self.down_blocks.append(down_block)
self.mid_block = HYWorldMidBlock(in_channels=block_out_channels[-1])
self.norm_out = HYWorldRMS_norm(block_out_channels[-1], images=False)
self.conv_act = nn.SiLU()
self.conv_out = HYWorldCausalConv3d(block_out_channels[-1], out_channels, kernel_size=3)
self.gradient_checkpointing = False
def forward(
self,
hidden_states: torch.Tensor,
feat_cache: Optional[List[Optional[torch.Tensor]]] = None,
feat_idx: Optional[List[int]] = None,
) -> torch.Tensor:
"""
Forward pass with optional temporal caching.
Args:
hidden_states: Input tensor of shape (B, C, T, H, W)
feat_cache: List of cached features for each conv layer
feat_idx: List containing current cache index [idx]
"""
# conv_in with caching
if feat_cache is not None and feat_idx is not None:
idx = feat_idx[0]
cache_x = hidden_states[:, :, -CACHE_T:, :, :].clone()
if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
cache_x = torch.cat(
[feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), cache_x],
dim=2,
)
hidden_states = self.conv_in(hidden_states, feat_cache[idx])
feat_cache[idx] = cache_x
feat_idx[0] += 1
else:
hidden_states = self.conv_in(hidden_states)
if torch.is_grad_enabled() and self.gradient_checkpointing:
for down_block in self.down_blocks:
hidden_states = self._gradient_checkpointing_func(down_block, hidden_states)
hidden_states = self._gradient_checkpointing_func(self.mid_block, hidden_states)
else:
for down_block in self.down_blocks:
hidden_states = down_block(hidden_states, feat_cache, feat_idx)
hidden_states = self.mid_block(hidden_states, feat_cache, feat_idx)
batch_size, _, frame, height, width = hidden_states.shape
short_cut = hidden_states.view(batch_size, -1, self.group_size, frame, height, width).mean(dim=2)
hidden_states = self.norm_out(hidden_states)
hidden_states = self.conv_act(hidden_states)
# conv_out with caching
if feat_cache is not None and feat_idx is not None:
idx = feat_idx[0]
cache_x = hidden_states[:, :, -CACHE_T:, :, :].clone()
if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
cache_x = torch.cat(
[feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), cache_x],
dim=2,
)
hidden_states = self.conv_out(hidden_states, feat_cache[idx])
feat_cache[idx] = cache_x
feat_idx[0] += 1
else:
hidden_states = self.conv_out(hidden_states)
hidden_states += short_cut
return hidden_states
class HYWorldDecoder3D(nn.Module):
r"""
Causal decoder for 3D video-like data used for HunyuanImage-1.5 Refiner with optional temporal caching.
"""
def __init__(
self,
in_channels: int = 32,
out_channels: int = 3,
block_out_channels: Tuple[int, ...] = (1024, 1024, 512, 256, 128),
layers_per_block: int = 2,
spatial_compression_ratio: int = 16,
temporal_compression_ratio: int = 4,
upsample_match_channel: bool = True,
):
super().__init__()
self.layers_per_block = layers_per_block
self.in_channels = in_channels
self.out_channels = out_channels
self.block_out_channels = block_out_channels
self.repeat = block_out_channels[0] // self.in_channels
self.conv_in = HYWorldCausalConv3d(self.in_channels, block_out_channels[0], kernel_size=3)
self.up_blocks = nn.ModuleList([])
# mid
self.mid_block = HYWorldMidBlock(in_channels=block_out_channels[0])
# up
input_channel = block_out_channels[0]
for i in range(len(block_out_channels)):
output_channel = block_out_channels[i]
add_spatial_upsample = i < np.log2(spatial_compression_ratio)
add_temporal_upsample = i < np.log2(temporal_compression_ratio)
if add_spatial_upsample or add_temporal_upsample:
upsample_out_channels = block_out_channels[i + 1] if upsample_match_channel else output_channel
up_block = HYWorldUpBlock3D(
num_layers=self.layers_per_block + 1,
in_channels=input_channel,
out_channels=output_channel,
upsample_out_channels=upsample_out_channels,
add_temporal_upsample=add_temporal_upsample,
)
input_channel = upsample_out_channels
else:
up_block = HYWorldUpBlock3D(
num_layers=self.layers_per_block + 1,
in_channels=input_channel,
out_channels=output_channel,
upsample_out_channels=None,
add_temporal_upsample=False,
)
input_channel = output_channel
self.up_blocks.append(up_block)
# out
self.norm_out = HYWorldRMS_norm(block_out_channels[-1], images=False)
self.conv_act = nn.SiLU()
self.conv_out = HYWorldCausalConv3d(block_out_channels[-1], out_channels, kernel_size=3)
self.gradient_checkpointing = False
def forward(
self,
hidden_states: torch.Tensor,
feat_cache: Optional[List[Optional[torch.Tensor]]] = None,
feat_idx: Optional[List[int]] = None,
first_chunk: bool = False,
) -> torch.Tensor:
"""
Forward pass with optional temporal caching.
Args:
hidden_states: Input tensor of shape (B, C, T, H, W)
feat_cache: List of cached features for each conv layer
feat_idx: List containing current cache index [idx]
first_chunk: Whether this is the first chunk (for upsampling behavior)
"""
# conv_in with caching
if feat_cache is not None and feat_idx is not None:
idx = feat_idx[0]
cache_x = hidden_states[:, :, -CACHE_T:, :, :].clone()
if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
cache_x = torch.cat(
[feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), cache_x],
dim=2,
)
hidden_states = self.conv_in(hidden_states, feat_cache[idx]) + hidden_states.repeat_interleave(
repeats=self.repeat, dim=1
)
feat_cache[idx] = cache_x
feat_idx[0] += 1
else:
hidden_states = self.conv_in(hidden_states) + hidden_states.repeat_interleave(repeats=self.repeat, dim=1)
if torch.is_grad_enabled() and self.gradient_checkpointing:
hidden_states = self._gradient_checkpointing_func(self.mid_block, hidden_states)
for up_block in self.up_blocks:
hidden_states = self._gradient_checkpointing_func(up_block, hidden_states)
else:
hidden_states = self.mid_block(hidden_states, feat_cache, feat_idx)
for up_block in self.up_blocks:
hidden_states = up_block(hidden_states, feat_cache, feat_idx, first_chunk)
# post-process
hidden_states = self.norm_out(hidden_states)
hidden_states = self.conv_act(hidden_states)
# conv_out with caching
if feat_cache is not None and feat_idx is not None:
idx = feat_idx[0]
cache_x = hidden_states[:, :, -CACHE_T:, :, :].clone()
if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
cache_x = torch.cat(
[feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), cache_x],
dim=2,
)
hidden_states = self.conv_out(hidden_states, feat_cache[idx])
feat_cache[idx] = cache_x
feat_idx[0] += 1
else:
hidden_states = self.conv_out(hidden_states)
return hidden_states
class AutoencoderKLHYWorld(nn.Module, ParallelTiledVAE):
r"""
A VAE model with KL loss for encoding videos into latents and decoding latent representations into videos.
Revised from HunyuanVideo-1.5 with temporal caching support for HY-WorldPlay integration.
"""
_supports_gradient_checkpointing = True
def __init__(
self,
config: Hunyuan15VAEConfig, # HYWorld use the same VAE architecture as HunyuanVideo-1.5
) -> None:
nn.Module.__init__(self)
ParallelTiledVAE.__init__(self, config)
self.encoder: Optional[HYWorldEncoder3D] = None
self.decoder: Optional[HYWorldDecoder3D] = None
if config.load_encoder:
self.encoder = HYWorldEncoder3D(
in_channels=config.in_channels,
out_channels=config.latent_channels * 2,
block_out_channels=config.block_out_channels,
layers_per_block=config.layers_per_block,
temporal_compression_ratio=config.temporal_compression_ratio,
spatial_compression_ratio=config.spatial_compression_ratio,
downsample_match_channel=config.downsample_match_channel,
)
if config.load_decoder:
self.decoder = HYWorldDecoder3D(
in_channels=config.latent_channels,
out_channels=config.out_channels,
block_out_channels=list(reversed(config.block_out_channels)),
layers_per_block=config.layers_per_block,
temporal_compression_ratio=config.temporal_compression_ratio,
spatial_compression_ratio=config.spatial_compression_ratio,
upsample_match_channel=config.upsample_match_channel,
)
# TODO: Add spatial tiling.
self.use_tiling = False
# The minimal tile height and width for spatial tiling to be used
self.tile_sample_min_height = 256
self.tile_sample_min_width = 256
self.tile_sample_min_num_frames = 2000 # Fill in a random large number, as hy1.5 vae does not use temporal tiling
# Cache-related attributes (initialized in clear_cache)
self._conv_num: int = 0
self._conv_idx: List[int] = [0]
self._feat_map: List[Optional[torch.Tensor]] = []
self._enc_conv_num: int = 0
self._enc_conv_idx: List[int] = [0]
self._enc_feat_map: List[Optional[torch.Tensor]] = []
# Precompute and cache conv counts for encoder and decoder for clear_cache speedup
self._cached_conv_counts = {
"decoder": (
sum(1 for m in self.decoder.modules() if isinstance(m, HYWorldCausalConv3d))
if self.decoder is not None
else 0
),
"encoder": (
sum(1 for m in self.encoder.modules() if isinstance(m, HYWorldCausalConv3d))
if self.encoder is not None
else 0
),
}
def clear_cache(self) -> None:
"""
Initialize/clear the feature cache for chunk-based encoding/decoding.
This should be called before starting a new encode/decode sequence to ensure
the cache is properly initialized.
"""
# Cache for decoder
self._conv_num = self._cached_conv_counts["decoder"]
self._conv_idx = [0]
self._feat_map: List[Optional[torch.Tensor]] = [None] * self._conv_num
# Cache for encoder
self._enc_conv_num = self._cached_conv_counts["encoder"]
self._enc_conv_idx = [0]
self._enc_feat_map: List[Optional[torch.Tensor]] = [None] * self._enc_conv_num
def _encode(self, x: torch.Tensor) -> torch.Tensor:
"""
Encode video with temporal caching for chunk-based processing.
This processes video in chunks (first frame separately, then 4 frames at a time)
while maintaining temporal context through caching. This matches the HY-WorldPlay
behavior for memory-efficient long video encoding.
Args:
x: Input video tensor of shape (B, C, T, H, W)
Returns:
Encoded latent tensor
"""
assert self.encoder is not None, "Encoder not loaded"
_, _, num_frame, _, _ = x.shape
self.clear_cache()
# Process in chunks: first frame alone, then groups of 4 frames
iter_ = 1 + (num_frame - 1) // 4
for i in range(iter_):
self._enc_conv_idx = [0]
if i == 0:
# First frame
out = self.encoder(
x[:, :, :1, :, :],
feat_cache=self._enc_feat_map,
feat_idx=self._enc_conv_idx,
)
else:
# Subsequent frames in groups of 4
out_ = self.encoder(
x[:, :, 1 + 4 * (i - 1) : 1 + 4 * i, :, :],
feat_cache=self._enc_feat_map,
feat_idx=self._enc_conv_idx,
)
out = torch.cat([out, out_], dim=2)
self.clear_cache()
return out
def _decode(self, z: torch.Tensor) -> torch.Tensor:
"""
Decode latents with temporal caching for chunk-based processing.
This processes latents one frame at a time while maintaining temporal context
through caching. This matches the HY-WorldPlay behavior for memory-efficient
long video decoding.
Args:
z: Latent tensor of shape (B, C, T, H, W)
Returns:
Decoded video tensor
"""
assert self.decoder is not None, "Decoder not loaded"
_, _, num_frame, _, _ = z.shape
self.clear_cache()
# Process one frame at a time with caching
for i in range(num_frame):
self._conv_idx = [0]
if i == 0:
# First frame with first_chunk=True
out = self.decoder(
z[:, :, i : i + 1, :, :],
feat_cache=self._feat_map,
feat_idx=self._conv_idx,
first_chunk=True,
)
else:
# Subsequent frames
out_ = self.decoder(
z[:, :, i : i + 1, :, :],
feat_cache=self._feat_map,
feat_idx=self._conv_idx,
first_chunk=False,
)
out = torch.cat([out, out_], dim=2)
self.clear_cache()
return out
def forward(
self,
sample: torch.Tensor,
sample_posterior: bool = False,
return_dict: bool = True,
generator: Optional[torch.Generator] = None,
) -> torch.Tensor:
r"""
Args:
sample (`torch.Tensor`): Input sample.
sample_posterior (`bool`, *optional*, defaults to `False`):
Whether to sample from the posterior.
return_dict (`bool`, *optional*, defaults to `True`):
Whether or not to return a [`DecoderOutput`] instead of a plain tuple.
generator (`torch.Generator`, *optional*):
Generator for sampling.
"""
x = sample
# encode() uses temporal caching by default (via _encode)
posterior = self.encode(x)
if sample_posterior:
z = posterior.sample(generator=generator)
else:
z = posterior.mode()
# decode() uses temporal caching by default (via _decode)
dec = self.decode(z)
return dec
# Entry point for model registry
EntryClass = AutoencoderKLHYWorld
+3
View File
@@ -1802,3 +1802,6 @@ class LTX2CausalVideoAutoencoder(nn.Module):
if previous_chunk is not None:
previous_weights = previous_weights.clamp(min=1e-8)
yield previous_chunk / previous_weights
# Entry point for model registry
EntryClass = LTX2CausalVideoAutoencoder
+3
View File
@@ -1138,3 +1138,6 @@ class AutoencoderKLStepvideo(nn.Module, ParallelTiledVAE):
z = posterior.mode()
dec = self.decode(z)
return dec
# Entry point for model registry
EntryClass = AutoencoderKLStepvideo
+2
View File
@@ -1319,3 +1319,5 @@ class AutoencoderKLWan(nn.Module, ParallelTiledVAE):
dec = self.decode(z)
return dec
# Entry point for model registry
EntryClass = AutoencoderKLWan
+1 -1
View File
@@ -277,7 +277,7 @@ def load_video(
if convert_method is not None:
pil_images = convert_method(pil_images)
return pil_images, original_fps if return_fps else pil_images
return (pil_images, original_fps) if return_fps else pil_images
def get_default_height_width(
+10 -24
View File
@@ -12,10 +12,9 @@ from fastvideo.logger import init_logger
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.pipelines.lora_pipeline import LoRAPipeline
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch, TrainingBatch
from fastvideo.pipelines.pipeline_registry import (PipelineType,
get_pipeline_registry)
from fastvideo.utils import (maybe_download_model,
verify_model_config_and_directory)
from fastvideo.pipelines.pipeline_registry import PipelineType
from fastvideo.registry import get_model_info
from fastvideo.utils import maybe_download_model
logger = init_logger(__name__)
@@ -42,30 +41,17 @@ def build_pipeline(
# fastvideo_args.downloaded_model_path = model_path
logger.info("Model path: %s", model_path)
config = verify_model_config_and_directory(model_path)
pipeline_name = config.get("_class_name")
if fastvideo_args.override_pipeline_cls_name:
logger.info("Overriding pipeline class name from %s to %s",
pipeline_name, fastvideo_args.override_pipeline_cls_name)
pipeline_name = fastvideo_args.override_pipeline_cls_name
if pipeline_name is None:
raise ValueError(
"Model config does not contain a _class_name attribute. "
"Only diffusers format is supported.")
# Get the appropriate pipeline registry based on pipeline_type
logger.info(
"Building pipeline of type: %s", pipeline_type.value if isinstance(
pipeline_type, PipelineType) else pipeline_type)
pipeline_registry = get_pipeline_registry(pipeline_type)
if isinstance(pipeline_type, str):
pipeline_type = PipelineType.from_string(pipeline_type)
pipeline_cls = pipeline_registry.resolve_pipeline_cls(
pipeline_name, pipeline_type, fastvideo_args.workload_type)
model_info = get_model_info(
model_path=model_path,
pipeline_type=pipeline_type,
workload_type=fastvideo_args.workload_type,
override_pipeline_cls_name=fastvideo_args.override_pipeline_cls_name,
)
pipeline_cls = model_info.pipeline_cls
# instantiate the pipelines
pipeline = pipeline_cls(model_path, fastvideo_args)
@@ -5,8 +5,8 @@ from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.pipelines.stages import (ConditioningStage,
Cosmos25DenoisingStage,
Cosmos25LatentPreparationStage,
Cosmos25AutoDenoisingStage,
Cosmos25AutoLatentPreparationStage,
DecodingStage, InputValidationStage,
Cosmos25TextEncodingStage,
Cosmos25TimestepPreparationStage)
@@ -42,13 +42,13 @@ class Cosmos2_5Pipeline(ComposedPipelineBase):
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="latent_preparation_stage",
stage=Cosmos25LatentPreparationStage(
stage=Cosmos25AutoLatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer"),
vae=self.get_module("vae")))
self.add_stage(stage_name="denoising_stage",
stage=Cosmos25DenoisingStage(
stage=Cosmos25AutoDenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler")))
@@ -0,0 +1,162 @@
# SPDX-License-Identifier: Apache-2.0
"""
Hunyuan video diffusion pipeline implementation.
This module contains an implementation of the Hunyuan video diffusion pipeline
using the modular pipeline architecture.
"""
import torch
import time
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.logger import init_logger
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.pipelines.stages import (
ConditioningStage, DecodingStage, DenoisingStage, InputValidationStage,
LatentPreparationStage, TextEncodingStage, TimestepPreparationStage,
Hy15ImageEncodingStage, SRDenoisingStage)
from fastvideo.distributed.parallel_state import get_local_torch_device
# TODO(will): move PRECISION_TO_TYPE to better place
logger = init_logger(__name__)
class HunyuanVideo152SRPipeline(ComposedPipelineBase):
_required_config_modules = [
"text_encoder", "text_encoder_2", "tokenizer", "tokenizer_2", "vae",
"transformer", "transformer_2", "transformer_3", "scheduler",
"upsampler", "upsampler_2"
]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
self.add_stage(stage_name="input_validation_stage",
stage=InputValidationStage())
self.add_stage(stage_name="prompt_encoding_stage_primary",
stage=TextEncodingStage(
text_encoders=[
self.get_module("text_encoder"),
self.get_module("text_encoder_2")
],
tokenizers=[
self.get_module("tokenizer"),
self.get_module("tokenizer_2")
],
))
self.add_stage(stage_name="conditioning_stage",
stage=ConditioningStage())
self.add_stage(stage_name="timestep_preparation_stage",
stage=TimestepPreparationStage(
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer")))
self.add_stage(stage_name="image_encoding_stage",
stage=Hy15ImageEncodingStage(image_encoder=None,
image_processor=None))
self.add_stage(stage_name="denoising_stage",
stage=DenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="sr_720p_latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer_2")))
self.add_stage(stage_name="sr_720p_denoising_stage",
stage=SRDenoisingStage(
transformer=self.get_module("transformer_2"),
scheduler=self.get_module("scheduler"),
upsampler=self.get_module("upsampler")))
self.add_stage(stage_name="sr_1080p_latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer_3")))
self.add_stage(stage_name="sr_1080p_denoising_stage",
stage=SRDenoisingStage(
transformer=self.get_module("transformer_3"),
scheduler=self.get_module("scheduler"),
upsampler=self.get_module("upsampler_2")))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
@torch.no_grad()
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
"""
Generate a video or image using the pipeline.
Args:
batch: The batch to generate from.
fastvideo_args: The inference arguments.
Returns:
ForwardBatch: The batch with the generated video or image.
"""
if not self.post_init_called:
self.post_init()
self.get_module("transformer").to(get_local_torch_device())
# Execute each stage
logger.info("Running pipeline stages: %s",
self._stage_name_mapping.keys())
# logger.info("Batch: %s", batch)
batch = self.input_validation_stage(batch, fastvideo_args)
batch = self.prompt_encoding_stage_primary(batch, fastvideo_args)
batch = self.conditioning_stage(batch, fastvideo_args)
batch = self.timestep_preparation_stage(batch, fastvideo_args)
batch = self.latent_preparation_stage(batch, fastvideo_args)
batch = self.image_encoding_stage(batch, fastvideo_args)
batch = self.denoising_stage(batch, fastvideo_args)
self.get_module("transformer").to("cpu")
# 720p SR
self.get_module("transformer_2").to(get_local_torch_device())
batch.lq_latents = batch.latents
batch.latents = None
batch.height = 720
batch.width = 1280
batch.num_inference_steps_sr = 6
batch = self.sr_720p_latent_preparation_stage(batch, fastvideo_args)
batch = self.image_encoding_stage(batch, fastvideo_args)
batch = self.sr_720p_denoising_stage(batch, fastvideo_args)
self.get_module("transformer_2").to("cpu")
# 1080p SR
self.get_module("transformer_3").to(get_local_torch_device())
batch.lq_latents = batch.latents
batch.latents = None
batch.height = 1072
batch.width = 1920
batch.num_inference_steps_sr = 8
batch = self.sr_1080p_latent_preparation_stage(batch, fastvideo_args)
batch = self.image_encoding_stage(batch, fastvideo_args)
batch = self.sr_1080p_denoising_stage(batch, fastvideo_args)
self.get_module("transformer_3").to("cpu")
start_time = time.time()
batch = self.decoding_stage(batch, fastvideo_args)
end_time = time.time()
logger.info("Decoding time: %s seconds", end_time - start_time)
# Return the output
return batch
EntryClass = HunyuanVideo152SRPipeline
@@ -0,0 +1,74 @@
# SPDX-License-Identifier: Apache-2.0
"""
Hunyuan video diffusion pipeline implementation.
This module contains an implementation of the Hunyuan video diffusion pipeline
using the modular pipeline architecture.
"""
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.pipelines.stages import (ConditioningStage, DecodingStage,
DenoisingStage, InputValidationStage,
LatentPreparationStage,
TextEncodingStage,
TimestepPreparationStage,
Hy15ImageEncodingStage)
# TODO(will): move PRECISION_TO_TYPE to better place
logger = init_logger(__name__)
class HunyuanVideo15ImageToVideoPipeline(ComposedPipelineBase):
_required_config_modules = [
"text_encoder", "text_encoder_2", "tokenizer", "tokenizer_2", "vae",
"transformer", "scheduler"
]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
self.add_stage(stage_name="input_validation_stage",
stage=InputValidationStage())
self.add_stage(stage_name="prompt_encoding_stage_primary",
stage=TextEncodingStage(
text_encoders=[
self.get_module("text_encoder"),
self.get_module("text_encoder_2")
],
tokenizers=[
self.get_module("tokenizer"),
self.get_module("tokenizer_2")
],
))
self.add_stage(stage_name="conditioning_stage",
stage=ConditioningStage())
self.add_stage(stage_name="timestep_preparation_stage",
stage=TimestepPreparationStage(
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer")))
self.add_stage(stage_name="image_encoding_stage",
stage=Hy15ImageEncodingStage(image_encoder=None,
image_processor=None))
self.add_stage(stage_name="denoising_stage",
stage=DenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
EntryClass = HunyuanVideo15ImageToVideoPipeline
@@ -0,0 +1,136 @@
# SPDX-License-Identifier: Apache-2.0
"""
Hunyuan video diffusion pipeline implementation.
This module contains an implementation of the Hunyuan video diffusion pipeline
using the modular pipeline architecture.
"""
import torch
import time
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.logger import init_logger
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.pipelines.stages import (
ConditioningStage, DecodingStage, DenoisingStage, InputValidationStage,
LatentPreparationStage, TextEncodingStage, TimestepPreparationStage,
Hy15ImageEncodingStage, SRDenoisingStage)
from fastvideo.distributed.parallel_state import get_local_torch_device
# TODO(will): move PRECISION_TO_TYPE to better place
logger = init_logger(__name__)
class HunyuanVideo15SRPipeline(ComposedPipelineBase):
_required_config_modules = [
"text_encoder", "text_encoder_2", "tokenizer", "tokenizer_2", "vae",
"transformer", "transformer_2", "scheduler", "upsampler"
]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
self.add_stage(stage_name="input_validation_stage",
stage=InputValidationStage())
self.add_stage(stage_name="prompt_encoding_stage_primary",
stage=TextEncodingStage(
text_encoders=[
self.get_module("text_encoder"),
self.get_module("text_encoder_2")
],
tokenizers=[
self.get_module("tokenizer"),
self.get_module("tokenizer_2")
],
))
self.add_stage(stage_name="conditioning_stage",
stage=ConditioningStage())
self.add_stage(stage_name="timestep_preparation_stage",
stage=TimestepPreparationStage(
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer")))
self.add_stage(stage_name="image_encoding_stage",
stage=Hy15ImageEncodingStage(image_encoder=None,
image_processor=None))
self.add_stage(stage_name="denoising_stage",
stage=DenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="sr_latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer_2")))
self.add_stage(stage_name="sr_denoising_stage",
stage=SRDenoisingStage(
transformer=self.get_module("transformer_2"),
scheduler=self.get_module("scheduler"),
upsampler=self.get_module("upsampler")))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
@torch.no_grad()
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
"""
Generate a video or image using the pipeline.
Args:
batch: The batch to generate from.
fastvideo_args: The inference arguments.
Returns:
ForwardBatch: The batch with the generated video or image.
"""
if not self.post_init_called:
self.post_init()
# Execute each stage
logger.info("Running pipeline stages: %s",
self._stage_name_mapping.keys())
# logger.info("Batch: %s", batch)
self.get_module("transformer").to(get_local_torch_device())
batch = self.input_validation_stage(batch, fastvideo_args)
batch = self.prompt_encoding_stage_primary(batch, fastvideo_args)
batch = self.conditioning_stage(batch, fastvideo_args)
batch = self.timestep_preparation_stage(batch, fastvideo_args)
batch = self.latent_preparation_stage(batch, fastvideo_args)
batch = self.image_encoding_stage(batch, fastvideo_args)
batch = self.denoising_stage(batch, fastvideo_args)
self.get_module("transformer").to("cpu")
self.get_module("transformer_2").to(get_local_torch_device())
batch.lq_latents = batch.latents
batch.latents = None
batch.height = batch.height_sr
batch.width = batch.width_sr
batch = self.sr_latent_preparation_stage(batch, fastvideo_args)
batch = self.image_encoding_stage(batch, fastvideo_args)
batch = self.sr_denoising_stage(batch, fastvideo_args)
self.get_module("transformer_2").to("cpu")
start_time = time.time()
batch = self.decoding_stage(batch, fastvideo_args)
end_time = time.time()
logger.info("Decoding time: %s seconds", end_time - start_time)
# Return the output
return batch
EntryClass = HunyuanVideo15SRPipeline

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