Compare commits
19
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
8cb2ae9d27 | ||
|
|
f76efe798e | ||
|
|
09f455233e | ||
|
|
aea300f690 | ||
|
|
b92219f6a6 | ||
|
|
a321b95a8a | ||
|
|
98308db7e0 | ||
|
|
c1e18f6722 | ||
|
|
75e193a2c9 | ||
|
|
d6e0a7d0dd | ||
|
|
7fc5f241da | ||
|
|
aae48a7e90 | ||
|
|
88f38eb0f4 | ||
|
|
d750b463dc | ||
|
|
e10b26a3d8 | ||
|
|
74636ba246 | ||
|
|
7e2f3f14e7 | ||
|
|
caa1c402ba | ||
|
|
38a6bd93d3 |
@@ -1,35 +1,29 @@
|
||||
<div align="center">
|
||||
<img src=assets/logos/logo.svg width="30%"/>
|
||||
</div>
|
||||
|
||||
<p align="center">
|
||||
| <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start/"><b> Quick Start</b></a> | <a href="https://github.com/hao-ai-lab/FastVideo/discussions/982" target="_blank"><b>Weekly Dev Meeting</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-3f4lao1uq-u~Ipx6Lt4J27AlD2y~IdLQ" target="_blank"> <b>Slack</b> </a> | 🟣💬 <a href="https://ibb.co/sv3MMKyv" target="_blank"> <b> WeChat </b> </a> |
|
||||
</p>
|
||||
| **[Documentation](https://hao-ai-lab.github.io/FastVideo)** | **[Quick Start](https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start/)** | **[Weekly Dev Meeting](https://github.com/hao-ai-lab/FastVideo/discussions/982)** | 🟣💬 **[Slack](https://join.slack.com/t/fastvideo/shared_invite/zt-3f4lao1uq-u~Ipx6Lt4J27AlD2y~IdLQ)** |
|
||||
|
||||
**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 +38,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 +53,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 +105,36 @@ python example.py
|
||||
|
||||
For a more detailed guide, please see our [inference quick start](https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start/).
|
||||
|
||||
### Other docs:
|
||||
## More Guides
|
||||
|
||||
- [Design Overview](https://hao-ai-lab.github.io/FastVideo/design/overview/)
|
||||
- [Contribution Guide](https://hao-ai-lab.github.io/FastVideo/getting_started/installation/)
|
||||
|
||||
## Distillation and Finetuning
|
||||
- [Distillation Guide](https://hao-ai-lab.github.io/FastVideo/distillation/dmd/)
|
||||
<!-- - [Finetuning Guide](https://hao-ai-lab.github.io/FastVideo/training/finetune.html) -->
|
||||
- [Contribution Guide](https://hao-ai-lab.github.io/FastVideo/contributing/overview/)
|
||||
|
||||
## Awesome work using FastVideo or our research projects
|
||||
|
||||
- [SGLang](https://github.com/sgl-project/sglang/tree/main/python/sglang/multimodal_gen): SGLang's diffusion inference functionality is based on a fork of FastVideo on Sept. 24, 2025. [](https://github.com/sgl-project/sglang)
|
||||
|
||||
- [DanceGRPO](https://github.com/XueZeyue/DanceGRPO): A unified framework to adapt Group Relative Policy Optimization (GRPO) to visual generation paradigms. Code based on FastVideo. [](https://github.com/XueZeyue/DanceGRPO)
|
||||
- [SRPO](https://github.com/Tencent-Hunyuan/SRPO): A method to directly align the full diffusion trajectory with fine-grained human preference. Code based on FastVideo. [](https://github.com/Tencent-Hunyuan/SRPO)
|
||||
- [DCM](https://github.com/Vchitect/DCM): Dual-expert consistency model for efficient and high-quality video generation. Code based on FastVideo. [](https://github.com/Vchitect/DCM)
|
||||
- [Hunyuan Video 1.5](https://github.com/Tencent-Hunyuan/HunyuanVideo-1.5): A leading lightweight video generation model, where they proposed SSTA based on Sliding Tile Attention. [](https://github.com/Tencent-Hunyuan/HunyuanVideo-1.5)
|
||||
- [Kandinsky-5.0](https://github.com/kandinskylab/kandinsky-5): A family of diffusion models for video & image generation, where their NABLA attention includes a Sliding Tile Attention branch. [](https://github.com/kandinskylab/kandinsky-5)
|
||||
- [LongCat Video](https://github.com/meituan-longcat/LongCat-Video): A foundational video generation model with 13.6B parameters with block-sparse attention similar to Video Sparse Attention. [](https://github.com/meituan-longcat/LongCat-Video)
|
||||
- [SGLang](https://github.com/sgl-project/sglang/tree/main/python/sglang/multimodal_gen): SGLang's diffusion inference functionality is based on a fork of FastVideo on Sept. 24, 2025.
|
||||
- [DanceGRPO](https://github.com/XueZeyue/DanceGRPO): A unified framework to adapt Group Relative Policy Optimization (GRPO) to visual generation paradigms. Code based on FastVideo.
|
||||
- [SRPO](https://github.com/Tencent-Hunyuan/SRPO): A method to directly align the full diffusion trajectory with fine-grained human preference. Code based on FastVideo.
|
||||
- [DCM](https://github.com/Vchitect/DCM): Dual-expert consistency model for efficient and high-quality video generation. Code based on FastVideo.
|
||||
- [Hunyuan Video 1.5](https://github.com/Tencent-Hunyuan/HunyuanVideo-1.5): A leading lightweight video generation model, where they proposed SSTA based on Sliding Tile Attention.
|
||||
- [Kandinsky-5.0](https://github.com/kandinskylab/kandinsky-5): A family of diffusion models for video & image generation, where their NABLA attention includes a Sliding Tile Attention branch.
|
||||
- [LongCat Video](https://github.com/meituan-longcat/LongCat-Video): A foundational video generation model with 13.6B parameters with block-sparse attention similar to Video Sparse Attention.
|
||||
|
||||
## 🤝 Contributing
|
||||
|
||||
We welcome all contributions. Please check out our guide [here](https://hao-ai-lab.github.io/FastVideo/contributing/overview/).
|
||||
See details in [development roadmap](https://github.com/hao-ai-lab/FastVideo/issues/899).
|
||||
## Acknowledgement
|
||||
We learned and reused code from the following projects:
|
||||
- [Wan-Video](https://github.com/Wan-Video)
|
||||
- [ThunderKittens](https://github.com/HazyResearch/ThunderKittens)
|
||||
- [Triton](https://github.com/triton-lang/triton)
|
||||
- [DMD2](https://github.com/tianweiy/DMD2)
|
||||
- [diffusers](https://github.com/huggingface/diffusers)
|
||||
- [xDiT](https://github.com/xdit-project/xDiT)
|
||||
- [vLLM](https://github.com/vllm-project/vllm)
|
||||
- [SGLang](https://github.com/sgl-project/sglang)
|
||||
|
||||
We thank [MBZUAI](https://ifm.mbzuai.ac.ae/), [Anyscale](https://www.anyscale.com/), and [GMI Cloud](https://www.gmicloud.ai/) for their support throughout this project.
|
||||
## Acknowledgement
|
||||
|
||||
We learned the design and reused code from the following projects: [Wan-Video](https://github.com/Wan-Video), [ThunderKittens](https://github.com/HazyResearch/ThunderKittens), [DMD2](https://github.com/tianweiy/DMD2), [diffusers](https://github.com/huggingface/diffusers), [xDiT](https://github.com/xdit-project/xDiT), [vLLM](https://github.com/vllm-project/vllm), [SGLang](https://github.com/sgl-project/sglang). We thank [MBZUAI](https://ifm.mbzuai.ac.ae/), [Anyscale](https://www.anyscale.com/), and [GMI Cloud](https://www.gmicloud.ai/) for their support throughout this project.
|
||||
|
||||
## Citation
|
||||
If you find FastVideo useful, please considering citing our work:
|
||||
|
||||
If you find FastVideo useful, please consider citing our research work:
|
||||
|
||||
```bibtex
|
||||
@software{fastvideo2024,
|
||||
title = {FastVideo: A Unified Framework for Accelerated Video Generation},
|
||||
author = {The FastVideo Team},
|
||||
url = {https://github.com/hao-ai-lab/FastVideo},
|
||||
month = apr,
|
||||
year = {2024},
|
||||
}
|
||||
|
||||
@article{zhang2025vsa,
|
||||
title={Vsa: Faster video diffusion with trainable sparse attention},
|
||||
author={Zhang, Peiyuan and Chen, Yongqi and Huang, Haofeng and Lin, Will and Liu, Zhengzhong and Stoica, Ion and Xing, Eric and Zhang, Hao},
|
||||
|
||||
@@ -43,7 +43,7 @@ RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir --upgrade pip && \
|
||||
uv pip install --no-cache-dir .[dev] && \
|
||||
uv pip install --no-cache-dir https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.5.4/flash_attn-2.8.3%2Bcu128torch2.9-cp310-cp310-linux_x86_64.whl
|
||||
uv pip install --no-cache-dir https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.7.16/flash_attn-2.8.3+cu128torch2.10-cp310-cp310-linux_x86_64.whl
|
||||
|
||||
COPY . .
|
||||
|
||||
|
||||
@@ -43,7 +43,7 @@ RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir --upgrade pip && \
|
||||
uv pip install --no-cache-dir .[dev] && \
|
||||
uv pip install --no-cache-dir https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.5.4/flash_attn-2.8.3%2Bcu128torch2.9-cp311-cp311-linux_x86_64.whl
|
||||
uv pip install --no-cache-dir https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.7.16/flash_attn-2.8.3+cu128torch2.10-cp311-cp311-linux_x86_64.whl
|
||||
|
||||
COPY . .
|
||||
|
||||
|
||||
@@ -43,7 +43,7 @@ RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir --upgrade pip && \
|
||||
uv pip install --no-cache-dir .[dev] && \
|
||||
uv pip install --no-cache-dir https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.5.4/flash_attn-2.8.3%2Bcu128torch2.9-cp312-cp312-linux_x86_64.whl
|
||||
uv pip install --no-cache-dir https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.7.16/flash_attn-2.8.3+cu128torch2.10-cp312-cp312-linux_x86_64.whl
|
||||
|
||||
COPY . .
|
||||
|
||||
|
||||
@@ -121,6 +121,15 @@ This page contains the complete API reference for the FastVideo library.
|
||||
show_root_toc_entry: true
|
||||
heading_level: 4
|
||||
|
||||
## fastvideo.registry
|
||||
|
||||
::: fastvideo.registry
|
||||
options:
|
||||
show_source: true
|
||||
show_root_heading: true
|
||||
show_root_toc_entry: true
|
||||
heading_level: 3
|
||||
|
||||
## fastvideo.pipelines
|
||||
|
||||
::: fastvideo.pipelines
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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.
|
||||

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

|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -103,7 +103,7 @@ python setup.py install # or pip install -e .
|
||||
|
||||
**`SAGE_ATTN_THREE`**
|
||||
|
||||
[SageAttention 3](https://huggingface.co/jt-zhang/SageAttention3) is an advanced attention mechanism that leverages FP4 quantization and Blackwell GPU Tensor Cores for significant performance improvements.
|
||||
[SageAttention 3](https://github.com/thu-ml/SageAttention/tree/main/sageattention3_blackwell) is an advanced attention mechanism that leverages FP4 quantization and Blackwell GPU Tensor Cores for significant performance improvements.
|
||||
|
||||
#### Hardware Requirements
|
||||
|
||||
@@ -113,11 +113,7 @@ python setup.py install # or pip install -e .
|
||||
|
||||
Note that Sage Attention 3 requires `python>=3.13`, `torch>=2.8.0`, `CUDA >=12.8`. If you are using `uv` and using `torch==2.8.0` make sure that `sentencepiece==0.2.1` in the pyproject.toml file.
|
||||
|
||||
To use Sage Attention 3 in FastVideo, first get access to the SageAttention3 code, then move `sageattn/` and `setup.py` to the directory `fastvideo/attention/backends`, then install from using:
|
||||
|
||||
```bash
|
||||
python setup.py install
|
||||
```
|
||||
To use Sage Attention 3 in FastVideo, follow the `README.md` in the linked repository to install the package from source.
|
||||
|
||||
## Teacache
|
||||
|
||||
|
||||
@@ -0,0 +1,51 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
|
||||
|
||||
def main():
|
||||
# Point this to your local diffusers model dir (or replace with a HF model ID).
|
||||
model_path = "KyleShao/Cosmos-Predict2.5-2B-Diffusers"
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
model_path,
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set True if GPU is out of memory
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True,
|
||||
)
|
||||
|
||||
sampling_param = SamplingParam.from_pretrained(model_path)
|
||||
|
||||
# image2world example from official repo
|
||||
image_path = "images/bus_terminal.jpg"
|
||||
|
||||
prompt = (
|
||||
"A nighttime city bus terminal gradually shifts from stillness to subtle movement. "
|
||||
"At first, multiple double-decker buses are parked under the glow of overhead lights, "
|
||||
"with a central bus labeled '87D' facing forward and stationary. "
|
||||
"As the video progresses, the bus in the middle moves ahead slowly, its headlights brightening the surrounding area "
|
||||
"and casting reflections onto adjacent vehicles. "
|
||||
"The motion creates space in the lineup, signaling activity within the otherwise quiet station. "
|
||||
"It then comes to a smooth stop, resuming its position in line. "
|
||||
"Overhead signage in Chinese characters remains illuminated, enhancing the vibrant, urban night scene."
|
||||
)
|
||||
|
||||
generator.generate_video(
|
||||
prompt,
|
||||
sampling_param=sampling_param,
|
||||
image_path=str(image_path),
|
||||
num_cond_frames=1,
|
||||
output_path="outputs_video/cosmos2_5_i2w.mp4",
|
||||
save_video=True,
|
||||
)
|
||||
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
|
||||
|
||||
def main():
|
||||
@@ -15,19 +16,27 @@ def main():
|
||||
pin_cpu_memory=True,
|
||||
)
|
||||
|
||||
# Load default sampling parameters (negative_prompt, resolution, steps, etc.)
|
||||
sampling_param = SamplingParam.from_pretrained(model_path)
|
||||
|
||||
prompt = (
|
||||
"A high-definition video captures the precision of robotic welding in an industrial setting. The first frame showcases a robotic arm, equipped with a welding torch, positioned over a large metal structure. The welding process is in full swing, with bright sparks and intense light illuminating the scene, creating a vivid display of blue and white hues. A significant amount of smoke billows around the welding area, partially obscuring the view but emphasizing the heat and activity. The background reveals parts of the workshop environment, including a ventilation system and various pieces of machinery, indicating a busy and functional industrial workspace. As the video progresses, the robotic arm maintains its steady position, continuing the welding process and moving to its left. The welding torch consistently emits sparks and light, and the smoke continues to rise, diffusing slightly as it moves upward. The metal surface beneath the torch shows ongoing signs of heating and melting. The scene retains its industrial ambiance, with the welding sparks and smoke dominating the visual field, underscoring the ongoing nature of the welding operation."
|
||||
"A high-definition video captures the precision of robotic welding in an industrial setting. "
|
||||
"The first frame showcases a robotic arm, equipped with a welding torch, positioned over a large metal structure. "
|
||||
"The welding process is in full swing, with bright sparks and intense light illuminating the scene, "
|
||||
"creating a vivid display of blue and white hues. "
|
||||
"A significant amount of smoke billows around the welding area, partially obscuring the view but emphasizing the heat and activity. "
|
||||
"The background reveals parts of the workshop environment, including a ventilation system and various pieces of machinery, "
|
||||
"indicating a busy and functional industrial workspace. "
|
||||
"As the video progresses, the robotic arm maintains its steady position, continuing the welding process and moving to its left. "
|
||||
"The welding torch consistently emits sparks and light, and the smoke continues to rise, diffusing slightly as it moves upward. "
|
||||
"The metal surface beneath the torch shows ongoing signs of heating and melting. "
|
||||
"The scene retains its industrial ambiance, with the welding sparks and smoke dominating the visual field, "
|
||||
"underscoring the ongoing nature of the welding operation."
|
||||
)
|
||||
|
||||
video = generator.generate_video(
|
||||
generator.generate_video(
|
||||
prompt,
|
||||
negative_prompt="",
|
||||
height=704,
|
||||
width=1280,
|
||||
num_frames=77,
|
||||
num_inference_steps=35,
|
||||
guidance_scale=7.0,
|
||||
fps=24,
|
||||
sampling_param=sampling_param,
|
||||
output_path="outputs_video/cosmos2_5_t2w.mp4",
|
||||
save_video=True,
|
||||
)
|
||||
|
||||
@@ -0,0 +1,54 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
|
||||
|
||||
def main():
|
||||
# Point this to your local diffusers model dir (or replace with a HF model ID).
|
||||
model_path = "KyleShao/Cosmos-Predict2.5-2B-Diffusers"
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
model_path,
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set True if GPU is out of memory
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True,
|
||||
)
|
||||
|
||||
sampling_param = SamplingParam.from_pretrained(model_path)
|
||||
|
||||
# video2world example from official repo
|
||||
video_path = "videos/robot_pouring.mp4"
|
||||
|
||||
prompt = (
|
||||
"A robotic arm, primarily white with black joints and cables, is shown in a clean, modern indoor setting with a white tabletop. "
|
||||
"The arm, equipped with a gripper holding a small, light green pitcher, is positioned above a clear glass containing a reddish-brown liquid and a spoon. "
|
||||
"The robotic arm is in the process of pouring a transparent liquid into the glass. "
|
||||
"To the left of the pitcher, there is an opened jar with a similar reddish-brown substance visible through its transparent body. "
|
||||
"In the background, a vase with white flowers and a brown couch are partially visible, adding to the contemporary ambiance. "
|
||||
"The lighting is bright, casting soft shadows on the table. "
|
||||
"The robotic arm's movements are smooth and controlled, demonstrating precision in its task. "
|
||||
"As the video progresses, the robotic arm completes the pour, leaving the glass half-filled with the reddish-brown liquid. "
|
||||
"The jar remains untouched throughout the sequence, and the spoon inside the glass remains stationary. "
|
||||
"The other robotic arm on the right side also stays stationary throughout the video. "
|
||||
"The final frame captures the robotic arm with the pitcher finishing the pour, with the glass now filled to a higher level, while the pitcher is slightly tilted but still held securely by the gripper."
|
||||
)
|
||||
|
||||
generator.generate_video(
|
||||
prompt,
|
||||
sampling_param=sampling_param,
|
||||
video_path=str(video_path),
|
||||
num_cond_frames=1,
|
||||
output_path="outputs_video/cosmos2_5_v2w.mp4",
|
||||
save_video=True,
|
||||
)
|
||||
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
|
||||
@@ -39,4 +39,4 @@ def main():
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
main()
|
||||
@@ -0,0 +1,43 @@
|
||||
from fastvideo import VideoGenerator
|
||||
import json
|
||||
# from fastvideo.configs.sample import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples_hy15_1080p"
|
||||
def main():
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
# model.
|
||||
# If a local path is provided, FastVideo will make a best effort
|
||||
# attempt to identify the optimal arguments.
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"weizhou03/HunyuanVideo-1.5-Diffusers-1080p-2SR", # 480p -> 720p -> 1080p
|
||||
# or "weizhou03/HunyuanVideo-1.5-Diffusers-1080p" # 720p -> 1080p
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=True,
|
||||
vae_cpu_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
# image_encoder_cpu_offload=False,
|
||||
)
|
||||
|
||||
prompt = (
|
||||
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
|
||||
)
|
||||
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, negative_prompt="")
|
||||
|
||||
prompt2 = (
|
||||
"A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
|
||||
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True, negative_prompt="")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -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__":
|
||||
|
||||
@@ -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
|
||||
)
|
||||
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import torch
|
||||
from fastvideo.attention.backends.sageattn.api import sageattn_blackwell
|
||||
from sageattn3 import sageattn3_blackwell
|
||||
|
||||
from fastvideo.attention.backends.abstract import (AttentionBackend,
|
||||
AttentionImpl,
|
||||
@@ -67,6 +67,6 @@ class SageAttention3Impl(AttentionImpl):
|
||||
query = query.transpose(1, 2)
|
||||
key = key.transpose(1, 2)
|
||||
value = value.transpose(1, 2)
|
||||
output = sageattn_blackwell(query, key, value, is_causal=self.causal)
|
||||
output = sageattn3_blackwell(query, key, value, is_causal=self.causal)
|
||||
output = output.transpose(1, 2)
|
||||
return output
|
||||
|
||||
@@ -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\.(.*)$":
|
||||
|
||||
@@ -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\.(.*)$":
|
||||
|
||||
@@ -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])
|
||||
|
||||
@@ -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
|
||||
@@ -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,
|
||||
|
||||
@@ -8,7 +8,7 @@ from typing import Any, cast
|
||||
import torch
|
||||
|
||||
from fastvideo.configs.models import (DiTConfig, EncoderConfig, ModelConfig,
|
||||
VAEConfig)
|
||||
VAEConfig, UpsamplerConfig)
|
||||
from fastvideo.configs.models.encoders import BaseEncoderOutput
|
||||
from fastvideo.configs.utils import update_config_from_args
|
||||
from fastvideo.logger import init_logger
|
||||
@@ -44,12 +44,15 @@ class PipelineConfig:
|
||||
# Video generation parameters
|
||||
embedded_cfg_scale: float = 6.0
|
||||
flow_shift: float | None = None
|
||||
flow_shift_sr: float | None = None
|
||||
disable_autocast: bool = False
|
||||
is_causal: bool = False
|
||||
|
||||
# Model configuration
|
||||
dit_config: DiTConfig = field(default_factory=DiTConfig)
|
||||
dit_precision: str = "bf16"
|
||||
upsampler_config: UpsamplerConfig = field(default_factory=UpsamplerConfig)
|
||||
upsampler_precision: str = "fp32"
|
||||
|
||||
# VAE configuration
|
||||
vae_config: VAEConfig = field(default_factory=VAEConfig)
|
||||
@@ -216,6 +219,24 @@ class PipelineConfig:
|
||||
"Comma-separated list of denoising steps (e.g., '1000,757,522')",
|
||||
)
|
||||
|
||||
# STA (Sliding Tile Attention) parameters
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}STA-mode",
|
||||
type=str,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}STA_mode",
|
||||
default=PipelineConfig.STA_mode.value,
|
||||
choices=[mode.value for mode in STA_Mode],
|
||||
help=
|
||||
"STA mode: STA_inference, STA_searching, STA_tuning, STA_tuning_cfg, None",
|
||||
)
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}skip-time-steps",
|
||||
type=int,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}skip_time_steps",
|
||||
default=PipelineConfig.skip_time_steps,
|
||||
help="Number of time steps to warmup (full attention) for STA",
|
||||
)
|
||||
|
||||
# Add VAE configuration arguments
|
||||
from fastvideo.configs.models.vaes.base import VAEConfig
|
||||
VAEConfig.add_cli_args(parser, prefix=f"{prefix_with_dot}vae-config")
|
||||
@@ -245,8 +266,7 @@ class PipelineConfig:
|
||||
"""
|
||||
use the pipeline class setting from model_path to match the pipeline config
|
||||
"""
|
||||
from fastvideo.configs.pipelines.registry import (
|
||||
get_pipeline_config_cls_from_name)
|
||||
from fastvideo.registry import get_pipeline_config_cls_from_name
|
||||
pipeline_config_cls = get_pipeline_config_cls_from_name(model_path)
|
||||
|
||||
return cast(PipelineConfig, pipeline_config_cls(model_path=model_path))
|
||||
@@ -260,8 +280,7 @@ class PipelineConfig:
|
||||
kwargs: dictionary of kwargs
|
||||
config_cli_prefix: prefix of CLI arguments for this PipelineConfig instance
|
||||
"""
|
||||
from fastvideo.configs.pipelines.registry import (
|
||||
get_pipeline_config_cls_from_name)
|
||||
from fastvideo.registry import get_pipeline_config_cls_from_name
|
||||
|
||||
prefix_with_dot = f"{config_cli_prefix}." if (config_cli_prefix.strip()
|
||||
!= "") else ""
|
||||
@@ -298,6 +317,12 @@ class PipelineConfig:
|
||||
# 4. Update PipelineConfig from CLI arguments if provided
|
||||
kwargs[prefix_with_dot + 'model_path'] = model_path
|
||||
pipeline_config.update_config_from_dict(kwargs, config_cli_prefix)
|
||||
|
||||
# Convert STA_mode string to enum if necessary
|
||||
if isinstance(pipeline_config.STA_mode, str) and not isinstance(
|
||||
pipeline_config.STA_mode, STA_Mode):
|
||||
pipeline_config.STA_mode = STA_Mode(pipeline_config.STA_mode)
|
||||
|
||||
return pipeline_config
|
||||
|
||||
def check_pipeline_config(self) -> None:
|
||||
|
||||
@@ -11,7 +11,8 @@ from fastvideo.configs.models.dits import HunyuanVideo15Config
|
||||
from fastvideo.configs.models.encoders import (BaseEncoderOutput,
|
||||
Qwen2_5_VLConfig, T5Config)
|
||||
from fastvideo.configs.models.vaes import Hunyuan15VAEConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
from fastvideo.configs.models.upsamplers import SRTo720pUpsamplerConfig, SRTo1080pUpsamplerConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig, UpsamplerConfig
|
||||
|
||||
PROMPT_TEMPLATE_TOKEN_LENGTH = 108
|
||||
|
||||
@@ -131,9 +132,35 @@ class Hunyuan15T2V480PConfig(PipelineConfig):
|
||||
self.vae_config.load_decoder = True
|
||||
|
||||
|
||||
@dataclass
|
||||
class Hunyuan15I2V480PStepDistilledConfig(Hunyuan15T2V480PConfig):
|
||||
flow_shift: int = 7
|
||||
|
||||
|
||||
@dataclass
|
||||
class Hunyuan15T2V720PConfig(Hunyuan15T2V480PConfig):
|
||||
"""Base configuration for HunYuan pipeline architecture."""
|
||||
|
||||
# HunyuanConfig-specific parameters with defaults
|
||||
flow_shift: int = 9
|
||||
|
||||
|
||||
@dataclass
|
||||
class Hunyuan15I2V720PConfig(Hunyuan15T2V720PConfig):
|
||||
"""Base configuration for HunYuan pipeline architecture."""
|
||||
|
||||
# HunyuanConfig-specific parameters with defaults
|
||||
flow_shift: int = 7
|
||||
|
||||
|
||||
@dataclass
|
||||
class Hunyuan15SR1080PConfig(Hunyuan15T2V720PConfig):
|
||||
"""Base configuration for HunYuan pipeline architecture."""
|
||||
|
||||
# HunyuanConfig-specific parameters with defaults
|
||||
flow_shift: int = 7
|
||||
flow_shift_sr: int = 2
|
||||
upsampler_config: tuple[UpsamplerConfig, ...] = field(
|
||||
default_factory=lambda:
|
||||
(SRTo720pUpsamplerConfig(), SRTo1080pUpsamplerConfig()))
|
||||
upsampler_precision: str = "fp32"
|
||||
|
||||
@@ -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
|
||||
@@ -28,6 +28,9 @@ class SamplingParam:
|
||||
keyboard_cond: Any | None = None # Shape: (B, T, K)
|
||||
grid_sizes: Any | None = None # Shape: (3,) [F,H,W]
|
||||
|
||||
# Camera control inputs (HYWorld)
|
||||
pose: str | None = None # Camera trajectory: pose string (e.g., 'w-31') or JSON file path
|
||||
|
||||
# Refine inputs (LongCat 480p->720p upscaling)
|
||||
# Path-based refine (load stage1 video from disk, e.g. MP4)
|
||||
refine_from: str | None = None # Path to stage1 video (480p output from distill)
|
||||
@@ -55,10 +58,13 @@ class SamplingParam:
|
||||
num_frames_round_down: bool = False # Whether to round down num_frames if it's not divisible by num_gpus
|
||||
height: int = 720
|
||||
width: int = 1280
|
||||
height_sr: int = 1072
|
||||
width_sr: int = 1920
|
||||
fps: int = 24
|
||||
|
||||
# Denoising parameters
|
||||
num_inference_steps: int = 50
|
||||
num_inference_steps_sr: int = 50
|
||||
guidance_scale: float = 1.0
|
||||
guidance_rescale: float = 0.0
|
||||
boundary_ratio: float | None = None
|
||||
@@ -92,8 +98,7 @@ class SamplingParam:
|
||||
|
||||
@classmethod
|
||||
def from_pretrained(cls, model_path: str) -> "SamplingParam":
|
||||
from fastvideo.configs.sample.registry import (
|
||||
get_sampling_param_cls_for_name)
|
||||
from fastvideo.registry import get_sampling_param_cls_for_name
|
||||
sampling_cls = get_sampling_param_cls_for_name(model_path)
|
||||
if sampling_cls is not None:
|
||||
sampling_param: SamplingParam = sampling_cls()
|
||||
|
||||
@@ -5,15 +5,19 @@ from fastvideo.configs.sample.base import SamplingParam
|
||||
|
||||
|
||||
@dataclass
|
||||
class Cosmos_Predict2_5_2B_Diffusers_SamplingParam(SamplingParam):
|
||||
"""Defaults for Cosmos 2.5 (Predict2.5) text-to-video diffusers-format model."""
|
||||
|
||||
height: int = 480
|
||||
width: int = 832
|
||||
num_frames: int = 121
|
||||
class Cosmos25SamplingParamBase(SamplingParam):
|
||||
height: int = 704
|
||||
width: int = 1280
|
||||
num_frames: int = 77
|
||||
fps: int = 24
|
||||
seed: int = 0
|
||||
|
||||
guidance_scale: float = 7.0
|
||||
# Official Cosmos2.5 sampling uses empty string as unconditional.
|
||||
negative_prompt: str = ""
|
||||
negative_prompt: str = (
|
||||
"The video captures a series of frames showing ugly scenes, static with no motion, motion blur, "
|
||||
"over-saturation, shaky footage, low resolution, grainy texture, pixelated images, poorly lit areas, "
|
||||
"underexposed and overexposed scenes, poor color balance, washed out colors, choppy sequences, jerky movements, "
|
||||
"low frame rate, artifacting, color banding, unnatural transitions, outdated special effects, fake elements, "
|
||||
"unconvincing visuals, poorly edited content, jump cuts, visual noise, and flickering. "
|
||||
"Overall, the video is of poor quality.")
|
||||
num_inference_steps: int = 35
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -366,9 +305,11 @@ class ImageTransformStage(DatasetStage):
|
||||
image = self.transform_topcrop(image)
|
||||
elif self.transform is not None:
|
||||
image = self.transform(image)
|
||||
image = image.float() / 127.5 - 1.0
|
||||
else:
|
||||
image = image.float() / 127.5 - 1.0
|
||||
|
||||
image = image.transpose(0, 1) # [1 C H W] -> [C 1 H W]
|
||||
image = image.float() / 127.5 - 1.0
|
||||
batch.pixel_values = image
|
||||
return batch
|
||||
|
||||
@@ -487,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,
|
||||
@@ -541,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] = []
|
||||
@@ -577,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
|
||||
@@ -592,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)
|
||||
|
||||
|
||||
@@ -18,6 +18,7 @@ __all__ = [
|
||||
"cleanup_dist_env_and_memory",
|
||||
"model_parallel_is_initialized",
|
||||
"maybe_init_distributed_environment_and_model_parallel",
|
||||
"warmup_sequence_parallel_communication",
|
||||
|
||||
# World group
|
||||
"get_world_group",
|
||||
|
||||
@@ -7,10 +7,17 @@ import torch.distributed
|
||||
from fastvideo.distributed.parallel_state import (get_sp_group,
|
||||
get_sp_parallel_rank,
|
||||
get_sp_world_size,
|
||||
get_tp_group)
|
||||
get_tp_group,
|
||||
model_parallel_is_initialized)
|
||||
from fastvideo.distributed.utils import (unpad_sequence_tensor,
|
||||
compute_padding_for_sp,
|
||||
pad_sequence_tensor)
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# Track if SP communication has been warmed up
|
||||
_sp_warmup_done = False
|
||||
|
||||
|
||||
def tensor_model_parallel_all_reduce(input_: torch.Tensor) -> torch.Tensor:
|
||||
@@ -101,3 +108,86 @@ def sequence_model_parallel_shard(input_: torch.Tensor,
|
||||
input_ = input_.movedim(0, dim)
|
||||
|
||||
return input_, original_seq_len
|
||||
|
||||
|
||||
def warmup_sequence_parallel_communication(
|
||||
device: torch.device | None = None) -> None:
|
||||
"""Warmup NCCL communicators for sequence parallel all-to-all operations.
|
||||
|
||||
The first NCCL collective operation is slow due to lazy communicator
|
||||
initialization. This function runs dummy all-to-all operations to
|
||||
trigger the initialization upfront, before the first real forward pass.
|
||||
|
||||
Args:
|
||||
device: Device to use for warmup tensors. If None, uses CUDA device 0.
|
||||
"""
|
||||
global _sp_warmup_done
|
||||
|
||||
if _sp_warmup_done:
|
||||
return
|
||||
|
||||
if not model_parallel_is_initialized():
|
||||
return
|
||||
|
||||
sp_world_size = get_sp_world_size()
|
||||
if sp_world_size <= 1:
|
||||
_sp_warmup_done = True
|
||||
return
|
||||
|
||||
if device is None:
|
||||
device = torch.device("cuda")
|
||||
|
||||
logger.info("Warming up sequence parallel communication (SP=%d)...",
|
||||
sp_world_size)
|
||||
|
||||
# Use small but representative tensor shapes for warmup
|
||||
# Shape: [batch, seq_len, num_heads, head_dim]
|
||||
# The all-to-all patterns used in attention:
|
||||
# 1. scatter_dim=2 (heads), gather_dim=1 (seq) - before attention
|
||||
# 2. scatter_dim=1 (seq), gather_dim=2 (heads) - after attention
|
||||
batch_size = 1
|
||||
seq_len_per_rank = 16 # Small sequence per rank
|
||||
num_heads = sp_world_size * 4 # Must be divisible by sp_world_size
|
||||
head_dim = 64
|
||||
|
||||
# Create dummy tensor for warmup
|
||||
dummy = torch.zeros(batch_size,
|
||||
seq_len_per_rank,
|
||||
num_heads,
|
||||
head_dim,
|
||||
device=device,
|
||||
dtype=torch.bfloat16)
|
||||
|
||||
# Warmup pattern 1: scatter heads, gather sequence (before attention)
|
||||
_ = sequence_model_parallel_all_to_all_4D(dummy,
|
||||
scatter_dim=2,
|
||||
gather_dim=1)
|
||||
|
||||
# Warmup pattern 2: scatter sequence, gather heads (after attention)
|
||||
dummy2 = torch.zeros(batch_size,
|
||||
seq_len_per_rank * sp_world_size,
|
||||
num_heads // sp_world_size,
|
||||
head_dim,
|
||||
device=device,
|
||||
dtype=torch.bfloat16)
|
||||
_ = sequence_model_parallel_all_to_all_4D(dummy2,
|
||||
scatter_dim=1,
|
||||
gather_dim=2)
|
||||
|
||||
# Warmup all-gather (used for replicated tokens)
|
||||
dummy3 = torch.zeros(batch_size,
|
||||
8,
|
||||
num_heads // sp_world_size,
|
||||
head_dim,
|
||||
device=device,
|
||||
dtype=torch.bfloat16)
|
||||
_ = sequence_model_parallel_all_gather(dummy3, dim=2)
|
||||
|
||||
# Synchronize to ensure warmup completes
|
||||
torch.cuda.synchronize(device)
|
||||
|
||||
# Clean up
|
||||
del dummy, dummy2, dummy3
|
||||
|
||||
_sp_warmup_done = True
|
||||
logger.info("Sequence parallel communication warmup complete.")
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
@@ -931,6 +916,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 +1080,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")
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -33,7 +33,7 @@ from fastvideo.layers.visual_embedding import (PatchEmbed)
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.dits.base import BaseDiT
|
||||
from fastvideo.models.dits.wanvideo import WanT2VCrossAttention, WanTimeTextImageEmbedding
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
from fastvideo.platforms import AttentionBackendEnum, current_platform
|
||||
|
||||
logger = init_logger(__name__)
|
||||
class CausalWanSelfAttention(nn.Module):
|
||||
@@ -286,6 +286,8 @@ class CausalWanTransformerBlock(nn.Module):
|
||||
null_shift = null_scale = torch.tensor([0], device=hidden_states.device)
|
||||
norm_hidden_states, hidden_states = self.self_attn_residual_norm(
|
||||
hidden_states, attn_output, gate_msa, null_shift, null_scale)
|
||||
norm_hidden_states, hidden_states = norm_hidden_states.to(
|
||||
orig_dtype), hidden_states.to(orig_dtype)
|
||||
|
||||
# 2. Cross-attention
|
||||
attn_output = self.attn2(norm_hidden_states,
|
||||
@@ -452,7 +454,6 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
This function will be run for num_frame times.
|
||||
Process the latent frames one by one (1560 tokens each)
|
||||
"""
|
||||
from fastvideo.platforms import current_platform
|
||||
|
||||
orig_dtype = hidden_states.dtype
|
||||
if not isinstance(encoder_hidden_states, torch.Tensor):
|
||||
@@ -676,4 +677,7 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
# u = torch.einsum('fhwpqrc->cfphqwr', u.contiguous())
|
||||
u = u.reshape(c, *[i * j for i, j in zip(v, self.patch_size)])
|
||||
out.append(u)
|
||||
return out
|
||||
return out
|
||||
|
||||
# Entry point for model registry
|
||||
EntryClass = CausalWanTransformer3DModel
|
||||
|
||||
@@ -723,4 +723,7 @@ class CosmosTransformer3DModel(BaseDiT):
|
||||
hidden_states = hidden_states.permute(0, 7, 1, 6, 2, 4, 3, 5)
|
||||
hidden_states = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
|
||||
|
||||
return hidden_states
|
||||
return hidden_states
|
||||
|
||||
# Entry point for model registry
|
||||
EntryClass = CosmosTransformer3DModel
|
||||
|
||||
@@ -959,3 +959,5 @@ class Cosmos25Transformer3DModel(BaseDiT):
|
||||
|
||||
return hidden_states
|
||||
|
||||
# Entry point for model registry
|
||||
EntryClass = Cosmos25Transformer3DModel
|
||||
|
||||
@@ -940,3 +940,6 @@ class FinalLayer(nn.Module):
|
||||
x = self.norm_final(x) * (1.0 + scale.unsqueeze(1)) + shift.unsqueeze(1)
|
||||
x, _ = self.linear(x)
|
||||
return x
|
||||
|
||||
# Entry point for model registry
|
||||
EntryClass = HunyuanVideoTransformer3DModel
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -1134,3 +1134,5 @@ class LongCatTransformer3DModel(CachableDiT):
|
||||
)
|
||||
return x
|
||||
|
||||
# Entry point for model registry
|
||||
EntryClass = LongCatTransformer3DModel
|
||||
|
||||
+654
-46
@@ -18,12 +18,23 @@ import torch.nn as nn
|
||||
from einops import rearrange, repeat
|
||||
|
||||
from fastvideo.attention.backends.sdpa import SDPAMetadata
|
||||
from fastvideo.attention.layer import LocalAttention
|
||||
from fastvideo.attention.layer import DistributedAttention, LocalAttention
|
||||
from fastvideo.configs.models.dits import LTX2VideoConfig
|
||||
from fastvideo.forward_context import get_forward_context, set_forward_context
|
||||
from fastvideo.distributed.communication_op import (
|
||||
sequence_model_parallel_all_gather,
|
||||
sequence_model_parallel_all_gather_with_unpad,
|
||||
sequence_model_parallel_all_to_all_4D,
|
||||
sequence_model_parallel_shard,
|
||||
)
|
||||
from fastvideo.distributed.parallel_state import get_sp_parallel_rank, get_sp_world_size
|
||||
from fastvideo.distributed.utils import create_attention_mask_for_padding
|
||||
from fastvideo.forward_context import ForwardContext, get_forward_context, set_forward_context
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.dits.base import CachableDiT
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def get_timestep_embedding(
|
||||
timesteps: torch.Tensor,
|
||||
@@ -295,6 +306,12 @@ class VideoLatentPatchifier:
|
||||
output_shape: VideoLatentShape,
|
||||
device: Optional[torch.device] = None,
|
||||
) -> torch.Tensor:
|
||||
"""Get patch grid bounds for RoPE computation.
|
||||
|
||||
Args:
|
||||
output_shape: Shape of the video latent tensor
|
||||
device: Device to create tensors on
|
||||
"""
|
||||
frames = output_shape.frames
|
||||
height = output_shape.height
|
||||
width = output_shape.width
|
||||
@@ -366,6 +383,12 @@ class AudioLatentPatchifier:
|
||||
output_shape: AudioLatentShape,
|
||||
device: Optional[torch.device] = None,
|
||||
) -> torch.Tensor:
|
||||
"""Get patch grid bounds for audio RoPE computation.
|
||||
|
||||
Args:
|
||||
output_shape: Shape of the audio latent tensor
|
||||
device: Device to create tensors on
|
||||
"""
|
||||
start_timings = self._get_audio_latent_time_in_sec(
|
||||
self.shift,
|
||||
output_shape.frames + self.shift,
|
||||
@@ -515,6 +538,44 @@ def apply_ltx_rotary_emb(
|
||||
raise ValueError(f"Invalid rope type: {rope_type}")
|
||||
|
||||
|
||||
def apply_ltx_rotary_emb_4d(
|
||||
input_tensor: torch.Tensor,
|
||||
freqs_cis: tuple[torch.Tensor, torch.Tensor],
|
||||
rope_type: LTXRopeType = LTXRopeType.INTERLEAVED,
|
||||
) -> torch.Tensor:
|
||||
"""Apply LTX-2 rotary embeddings to a 4D tensor [B, H, T, D].
|
||||
|
||||
This is used for applying RoPE after all-to-all in distributed attention,
|
||||
where the tensor is already in [B, H, T, D] format.
|
||||
"""
|
||||
cos_freqs, sin_freqs = freqs_cis
|
||||
if rope_type == LTXRopeType.INTERLEAVED:
|
||||
# For interleaved, cos/sin have shape [B, T, inner_dim]
|
||||
# Need to reshape to [B, H, T, D] format
|
||||
# Actually interleaved doesn't have per-head rotations, so we broadcast
|
||||
t_dup = rearrange(input_tensor, "... (d r) -> ... d r", r=2)
|
||||
t1, t2 = t_dup.unbind(dim=-1)
|
||||
t_dup = torch.stack((-t2, t1), dim=-1)
|
||||
input_tensor_rot = rearrange(t_dup, "... d r -> ... (d r)")
|
||||
return input_tensor * cos_freqs + input_tensor_rot * sin_freqs
|
||||
if rope_type == LTXRopeType.SPLIT:
|
||||
# For split, cos/sin already have shape [B, H, T, D/2]
|
||||
# input_tensor is [B, H, T, D]
|
||||
split_input = rearrange(input_tensor, "... (d r) -> ... d r", d=2)
|
||||
first_half_input = split_input[..., :1, :]
|
||||
second_half_input = split_input[..., 1:, :]
|
||||
|
||||
output = split_input * cos_freqs.unsqueeze(-2)
|
||||
first_half_output = output[..., :1, :]
|
||||
second_half_output = output[..., 1:, :]
|
||||
|
||||
first_half_output.addcmul_(-sin_freqs.unsqueeze(-2), second_half_input)
|
||||
second_half_output.addcmul_(sin_freqs.unsqueeze(-2), first_half_input)
|
||||
|
||||
return rearrange(output, "... d r -> ... (d r)")
|
||||
raise ValueError(f"Invalid rope type: {rope_type}")
|
||||
|
||||
|
||||
def _apply_ltx_interleaved_rotary_emb(
|
||||
input_tensor: torch.Tensor, cos_freqs: torch.Tensor, sin_freqs: torch.Tensor
|
||||
) -> torch.Tensor:
|
||||
@@ -918,6 +979,205 @@ class TransformerConfig:
|
||||
context_dim: int
|
||||
|
||||
|
||||
class LTXDistributedAttention(DistributedAttention):
|
||||
"""LTX-2 specialized DistributedAttention that handles LTX-style RoPE internally."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
num_heads: int,
|
||||
head_size: int,
|
||||
rope_type: LTXRopeType,
|
||||
num_kv_heads: int | None = None,
|
||||
softmax_scale: float | None = None,
|
||||
causal: bool = False,
|
||||
supported_attention_backends: tuple[AttentionBackendEnum, ...] | None = None,
|
||||
prefix: str = "",
|
||||
**extra_impl_args,
|
||||
) -> None:
|
||||
super().__init__(
|
||||
num_heads=num_heads,
|
||||
head_size=head_size,
|
||||
num_kv_heads=num_kv_heads,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=causal,
|
||||
supported_attention_backends=supported_attention_backends,
|
||||
prefix=prefix,
|
||||
**extra_impl_args,
|
||||
)
|
||||
self.rope_type = rope_type
|
||||
|
||||
@torch.compiler.disable
|
||||
def forward(
|
||||
self,
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
replicated_q: torch.Tensor | None = None,
|
||||
replicated_k: torch.Tensor | None = None,
|
||||
replicated_v: torch.Tensor | None = None,
|
||||
attention_mask: torch.Tensor | None = None,
|
||||
ltx_freqs_cis: tuple[torch.Tensor, torch.Tensor] | None = None,
|
||||
) -> tuple[torch.Tensor, torch.Tensor | None]:
|
||||
"""Forward pass with LTX-2 style RoPE application.
|
||||
|
||||
Args:
|
||||
q: Query tensor [batch_size, seq_len, num_heads, head_dim]
|
||||
k: Key tensor [batch_size, seq_len, num_heads, head_dim]
|
||||
v: Value tensor [batch_size, seq_len, num_heads, head_dim]
|
||||
replicated_q: Replicated query tensor for text tokens
|
||||
replicated_k: Replicated key tensor
|
||||
replicated_v: Replicated value tensor
|
||||
attention_mask: Attention mask [batch_size, seq_len]
|
||||
ltx_freqs_cis: LTX-2 style RoPE (cos, sin) with shape [B, H, T, D]
|
||||
|
||||
Returns:
|
||||
Tuple of output tensor and optional replicated output
|
||||
"""
|
||||
assert q.dim() == 4 and k.dim() == 4 and v.dim() == 4, "Expected 4D tensors"
|
||||
batch_size, seq_len, num_heads, head_dim = q.shape
|
||||
local_rank = get_sp_parallel_rank()
|
||||
world_size = get_sp_world_size()
|
||||
|
||||
forward_context: ForwardContext = get_forward_context()
|
||||
ctx_attn_metadata = forward_context.attn_metadata
|
||||
|
||||
# Stack QKV
|
||||
qkv = torch.cat([q, k, v], dim=0) # [3*batch, seq_len, num_heads, head_dim]
|
||||
|
||||
# Redistribute heads across sequence dimension
|
||||
qkv = sequence_model_parallel_all_to_all_4D(qkv, scatter_dim=2, gather_dim=1)
|
||||
|
||||
# After all-to-all, each rank has the full sequence but only a subset of heads
|
||||
valid_seq_len = None
|
||||
if attention_mask is not None:
|
||||
valid_seq_len = (attention_mask[0] == 1).sum().item()
|
||||
qkv = qkv[:, :valid_seq_len, :, :]
|
||||
|
||||
# Apply LTX-2 style RoPE after all-to-all (when we have full sequence)
|
||||
if ltx_freqs_cis is not None:
|
||||
cos, sin = ltx_freqs_cis
|
||||
heads_per_rank = num_heads // world_size
|
||||
head_start = local_rank * heads_per_rank
|
||||
head_end = head_start + heads_per_rank
|
||||
# Slice to this rank's heads: [B, H/SP, T, D]
|
||||
cos_local = cos[:, head_start:head_end, :valid_seq_len, :]
|
||||
sin_local = sin[:, head_start:head_end, :valid_seq_len, :]
|
||||
|
||||
# Apply RoPE to Q and K together (first 2*batch_size in dim 0)
|
||||
qk_part = qkv[:batch_size * 2]
|
||||
# Transpose to [2*B, H/SP, T, D] for LTX RoPE application
|
||||
qk_part = qk_part.transpose(1, 2)
|
||||
qk_part = apply_ltx_rotary_emb_4d(qk_part, (cos_local, sin_local), self.rope_type)
|
||||
# Transpose back to [2*B, T, H/SP, D]
|
||||
qkv[:batch_size * 2] = qk_part.transpose(1, 2)
|
||||
|
||||
# Apply backend-specific preprocess_qkv
|
||||
qkv = self.attn_impl.preprocess_qkv(qkv, ctx_attn_metadata)
|
||||
|
||||
# Concatenate with replicated QKV if provided
|
||||
if replicated_q is not None:
|
||||
assert replicated_k is not None and replicated_v is not None
|
||||
replicated_qkv = torch.cat(
|
||||
[replicated_q, replicated_k, replicated_v],
|
||||
dim=0) # [3, seq_len, num_heads, head_dim]
|
||||
heads_per_rank = num_heads // world_size
|
||||
replicated_qkv = replicated_qkv[:, :, local_rank *
|
||||
heads_per_rank:(local_rank + 1) *
|
||||
heads_per_rank]
|
||||
qkv = torch.cat([qkv, replicated_qkv], dim=1)
|
||||
|
||||
q, k, v = qkv.chunk(3, dim=0)
|
||||
|
||||
output = self.attn_impl.forward(q, k, v, ctx_attn_metadata)
|
||||
|
||||
# Redistribute back if using sequence parallelism
|
||||
replicated_output = None
|
||||
if replicated_q is not None:
|
||||
split_idx = seq_len * world_size if valid_seq_len is None else valid_seq_len
|
||||
replicated_output = output[:, split_idx:]
|
||||
output = output[:, :split_idx]
|
||||
replicated_output = sequence_model_parallel_all_gather(
|
||||
replicated_output.contiguous(), dim=2)
|
||||
|
||||
# Apply backend-specific postprocess_output
|
||||
output = self.attn_impl.postprocess_output(output, ctx_attn_metadata)
|
||||
|
||||
if attention_mask is not None:
|
||||
pad_len = (attention_mask[0] == 0).sum().item()
|
||||
output = torch.nn.functional.pad(output, (0, 0, 0, 0, 0, pad_len))
|
||||
|
||||
output = sequence_model_parallel_all_to_all_4D(output, scatter_dim=1, gather_dim=2)
|
||||
|
||||
return output, replicated_output
|
||||
|
||||
|
||||
class LTXLocalAttention(LocalAttention):
|
||||
"""LTX-2 specialized LocalAttention that handles LTX-style RoPE internally."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
num_heads: int,
|
||||
head_size: int,
|
||||
rope_type: LTXRopeType,
|
||||
num_kv_heads: int | None = None,
|
||||
softmax_scale: float | None = None,
|
||||
causal: bool = False,
|
||||
supported_attention_backends: tuple[AttentionBackendEnum, ...] | None = None,
|
||||
**extra_impl_args,
|
||||
) -> None:
|
||||
super().__init__(
|
||||
num_heads=num_heads,
|
||||
head_size=head_size,
|
||||
num_kv_heads=num_kv_heads,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=causal,
|
||||
supported_attention_backends=supported_attention_backends,
|
||||
**extra_impl_args,
|
||||
)
|
||||
self.rope_type = rope_type
|
||||
|
||||
def forward(
|
||||
self,
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
ltx_freqs_cis: tuple[torch.Tensor, torch.Tensor] | None = None,
|
||||
ltx_k_freqs_cis: tuple[torch.Tensor, torch.Tensor] | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""Forward pass with LTX-2 style RoPE application.
|
||||
|
||||
Args:
|
||||
q: Query tensor [batch_size, seq_len, num_heads, head_dim]
|
||||
k: Key tensor [batch_size, seq_len, num_heads, head_dim]
|
||||
v: Value tensor [batch_size, seq_len, num_heads, head_dim]
|
||||
ltx_freqs_cis: LTX-2 style RoPE (cos, sin) for Q with shape [B, H, T, D]
|
||||
ltx_k_freqs_cis: LTX-2 style RoPE (cos, sin) for K (if different from Q)
|
||||
|
||||
Returns:
|
||||
Output tensor after local attention
|
||||
"""
|
||||
assert q.dim() == 4 and k.dim() == 4 and v.dim() == 4, "Expected 4D tensors"
|
||||
|
||||
forward_context: ForwardContext = get_forward_context()
|
||||
ctx_attn_metadata = forward_context.attn_metadata
|
||||
|
||||
# Apply LTX-2 style RoPE
|
||||
if ltx_freqs_cis is not None:
|
||||
# Use separate K RoPE if provided (for cross-attention), otherwise use same as Q
|
||||
k_freqs = ltx_k_freqs_cis if ltx_k_freqs_cis is not None else ltx_freqs_cis
|
||||
# Transpose to [B, H, T, D] for RoPE application
|
||||
q = q.transpose(1, 2)
|
||||
k = k.transpose(1, 2)
|
||||
q = apply_ltx_rotary_emb_4d(q, ltx_freqs_cis, self.rope_type)
|
||||
k = apply_ltx_rotary_emb_4d(k, k_freqs, self.rope_type)
|
||||
# Transpose back to [B, T, H, D]
|
||||
q = q.transpose(1, 2)
|
||||
k = k.transpose(1, 2)
|
||||
|
||||
output = self.attn_impl.forward(q, k, v, ctx_attn_metadata)
|
||||
return output
|
||||
|
||||
|
||||
class LTXSelfAttention(nn.Module):
|
||||
"""LTX-2 attention block with RMSNorm + FastVideo LocalAttention."""
|
||||
def __init__(
|
||||
@@ -945,17 +1205,19 @@ class LTXSelfAttention(nn.Module):
|
||||
self.to_v = nn.Linear(context_dim, inner_dim, bias=True)
|
||||
self.to_out = nn.Sequential(nn.Linear(inner_dim, query_dim, bias=True), nn.Identity())
|
||||
|
||||
self.attn = LocalAttention(
|
||||
self.attn = LTXLocalAttention(
|
||||
num_heads=heads,
|
||||
head_size=dim_head,
|
||||
rope_type=rope_type,
|
||||
dropout_rate=0,
|
||||
softmax_scale=None,
|
||||
causal=False,
|
||||
supported_attention_backends=supported_attention_backends,
|
||||
)
|
||||
self.attn_masked = LocalAttention(
|
||||
self.attn_masked = LTXLocalAttention(
|
||||
num_heads=heads,
|
||||
head_size=dim_head,
|
||||
rope_type=rope_type,
|
||||
dropout_rate=0,
|
||||
softmax_scale=None,
|
||||
causal=False,
|
||||
@@ -978,15 +1240,13 @@ class LTXSelfAttention(nn.Module):
|
||||
q = self.q_norm(q)
|
||||
k = self.k_norm(k)
|
||||
|
||||
if pe is not None:
|
||||
q = apply_ltx_rotary_emb(q, pe, self.rope_type)
|
||||
k = apply_ltx_rotary_emb(k, pe if k_pe is None else k_pe, self.rope_type)
|
||||
|
||||
# RoPE is applied inside LTXLocalAttention
|
||||
b, q_len, _ = q.shape
|
||||
k_len = k.shape[1]
|
||||
q = q.view(b, q_len, self.heads, self.dim_head)
|
||||
k = k.view(b, k_len, self.heads, self.dim_head)
|
||||
v = v.view(b, k_len, self.heads, self.dim_head)
|
||||
|
||||
if mask is not None:
|
||||
if mask.ndim == 2:
|
||||
mask = mask.unsqueeze(0)
|
||||
@@ -1008,9 +1268,93 @@ class LTXSelfAttention(nn.Module):
|
||||
attn_metadata=attn_metadata,
|
||||
forward_batch=forward_batch,
|
||||
):
|
||||
out = self.attn_masked(q, k, v)
|
||||
out = self.attn_masked(q, k, v, ltx_freqs_cis=pe, ltx_k_freqs_cis=k_pe)
|
||||
else:
|
||||
out = self.attn(q, k, v)
|
||||
out = self.attn(q, k, v, ltx_freqs_cis=pe, ltx_k_freqs_cis=k_pe)
|
||||
out = out.reshape(b, q_len, -1)
|
||||
return self.to_out(out)
|
||||
|
||||
|
||||
class LTXDistributedSelfAttention(nn.Module):
|
||||
"""LTX-2 attention block with RMSNorm + LTXDistributedAttention for SP."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
query_dim: int,
|
||||
context_dim: int | None,
|
||||
heads: int,
|
||||
dim_head: int,
|
||||
norm_eps: float,
|
||||
rope_type: LTXRopeType,
|
||||
supported_attention_backends: tuple[AttentionBackendEnum, ...],
|
||||
prefix: str = "",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
inner_dim = dim_head * heads
|
||||
context_dim = query_dim if context_dim is None else context_dim
|
||||
|
||||
self.heads = heads
|
||||
self.dim_head = dim_head
|
||||
self.rope_type = rope_type
|
||||
|
||||
self.q_norm = torch.nn.RMSNorm(inner_dim, eps=norm_eps)
|
||||
self.k_norm = torch.nn.RMSNorm(inner_dim, eps=norm_eps)
|
||||
self.to_q = nn.Linear(query_dim, inner_dim, bias=True)
|
||||
self.to_k = nn.Linear(context_dim, inner_dim, bias=True)
|
||||
self.to_v = nn.Linear(context_dim, inner_dim, bias=True)
|
||||
self.to_out = nn.Sequential(nn.Linear(inner_dim, query_dim, bias=True), nn.Identity())
|
||||
|
||||
self.attn = LTXDistributedAttention(
|
||||
num_heads=heads,
|
||||
head_size=dim_head,
|
||||
rope_type=rope_type,
|
||||
causal=False,
|
||||
supported_attention_backends=supported_attention_backends,
|
||||
prefix=f"{prefix}.attn",
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
context: torch.Tensor | None = None,
|
||||
pe: tuple[torch.Tensor, torch.Tensor] | None = None,
|
||||
k_pe: tuple[torch.Tensor, torch.Tensor] | None = None,
|
||||
attention_mask: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""Forward pass for distributed self-attention.
|
||||
|
||||
Args:
|
||||
x: Query tensor [B, L, C]
|
||||
context: Key/Value tensor [B, L, C], or None for self-attention
|
||||
pe: Rotary position embeddings for Q (cos, sin) with shape [B, H, T_full, D]
|
||||
NOTE: For SP, this must be the FULL sequence RoPE, not sharded!
|
||||
k_pe: Rotary position embeddings for K (cos, sin), or None to use pe
|
||||
attention_mask: Attention mask for padding [B, padded_seq_len]
|
||||
"""
|
||||
q = self.to_q(x)
|
||||
context = x if context is None else context
|
||||
k = self.to_k(context)
|
||||
v = self.to_v(context)
|
||||
|
||||
q = self.q_norm(q)
|
||||
k = self.k_norm(k)
|
||||
|
||||
# RoPE is applied inside LTXDistributedAttention AFTER the all-to-all,
|
||||
# when each rank has the full sequence (but subset of heads).
|
||||
b, q_len, _ = q.shape
|
||||
k_len = k.shape[1]
|
||||
q = q.view(b, q_len, self.heads, self.dim_head)
|
||||
k = k.view(b, k_len, self.heads, self.dim_head)
|
||||
v = v.view(b, k_len, self.heads, self.dim_head)
|
||||
|
||||
# Pass full RoPE to distributed attention - it will apply after all-to-all
|
||||
# and slice to this rank's heads
|
||||
out, _ = self.attn(
|
||||
q, k, v,
|
||||
attention_mask=attention_mask,
|
||||
ltx_freqs_cis=pe,
|
||||
)
|
||||
|
||||
out = out.reshape(b, q_len, -1)
|
||||
return self.to_out(out)
|
||||
|
||||
@@ -1025,12 +1369,31 @@ class BasicAVTransformerBlock(torch.nn.Module):
|
||||
audio: TransformerConfig | None = None,
|
||||
rope_type: LTXRopeType = LTXRopeType.INTERLEAVED,
|
||||
norm_eps: float = 1e-6,
|
||||
use_distributed_attention: bool = False,
|
||||
prefix: str = "",
|
||||
):
|
||||
super().__init__()
|
||||
self.idx = idx
|
||||
self.use_distributed_attention = use_distributed_attention
|
||||
|
||||
# Choose attention class based on SP mode
|
||||
# Self-attention and audio-video cross-attention use DistributedAttention when SP > 1
|
||||
# Text cross-attention uses LocalAttention (text embeddings are replicated)
|
||||
SelfAttnCls = LTXDistributedSelfAttention if use_distributed_attention else LTXSelfAttention
|
||||
CrossAttnCls = LTXSelfAttention # Text cross-attention is always local
|
||||
|
||||
if video is not None:
|
||||
self.attn1 = LTXSelfAttention(
|
||||
# Video self-attention - use distributed when SP > 1
|
||||
self.attn1 = SelfAttnCls(
|
||||
query_dim=video.dim,
|
||||
context_dim=None,
|
||||
heads=video.heads,
|
||||
dim_head=video.d_head,
|
||||
norm_eps=norm_eps,
|
||||
rope_type=rope_type,
|
||||
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA),
|
||||
prefix=f"{prefix}.blocks.{idx}.attn1" if use_distributed_attention else "",
|
||||
) if use_distributed_attention else LTXSelfAttention(
|
||||
query_dim=video.dim,
|
||||
context_dim=None,
|
||||
heads=video.heads,
|
||||
@@ -1039,7 +1402,8 @@ class BasicAVTransformerBlock(torch.nn.Module):
|
||||
rope_type=rope_type,
|
||||
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA),
|
||||
)
|
||||
self.attn2 = LTXSelfAttention(
|
||||
# Text cross-attention - always local (text is replicated)
|
||||
self.attn2 = CrossAttnCls(
|
||||
query_dim=video.dim,
|
||||
context_dim=video.context_dim,
|
||||
heads=video.heads,
|
||||
@@ -1052,7 +1416,17 @@ class BasicAVTransformerBlock(torch.nn.Module):
|
||||
self.scale_shift_table = torch.nn.Parameter(torch.empty(6, video.dim))
|
||||
|
||||
if audio is not None:
|
||||
self.audio_attn1 = LTXSelfAttention(
|
||||
# Audio self-attention - use distributed when SP > 1
|
||||
self.audio_attn1 = SelfAttnCls(
|
||||
query_dim=audio.dim,
|
||||
context_dim=None,
|
||||
heads=audio.heads,
|
||||
dim_head=audio.d_head,
|
||||
norm_eps=norm_eps,
|
||||
rope_type=rope_type,
|
||||
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA),
|
||||
prefix=f"{prefix}.blocks.{idx}.audio_attn1" if use_distributed_attention else "",
|
||||
) if use_distributed_attention else LTXSelfAttention(
|
||||
query_dim=audio.dim,
|
||||
context_dim=None,
|
||||
heads=audio.heads,
|
||||
@@ -1061,7 +1435,8 @@ class BasicAVTransformerBlock(torch.nn.Module):
|
||||
rope_type=rope_type,
|
||||
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA),
|
||||
)
|
||||
self.audio_attn2 = LTXSelfAttention(
|
||||
# Text cross-attention - always local (text is replicated)
|
||||
self.audio_attn2 = CrossAttnCls(
|
||||
query_dim=audio.dim,
|
||||
context_dim=audio.context_dim,
|
||||
heads=audio.heads,
|
||||
@@ -1074,6 +1449,8 @@ class BasicAVTransformerBlock(torch.nn.Module):
|
||||
self.audio_scale_shift_table = torch.nn.Parameter(torch.empty(6, audio.dim))
|
||||
|
||||
if audio is not None and video is not None:
|
||||
# Audio-to-video cross-attention
|
||||
# Uses local attention - context is gathered from all SP ranks in forward()
|
||||
self.audio_to_video_attn = LTXSelfAttention(
|
||||
query_dim=video.dim,
|
||||
context_dim=audio.dim,
|
||||
@@ -1083,6 +1460,8 @@ class BasicAVTransformerBlock(torch.nn.Module):
|
||||
rope_type=rope_type,
|
||||
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA),
|
||||
)
|
||||
# Video-to-audio cross-attention
|
||||
# Uses local attention - context is gathered from all SP ranks in forward()
|
||||
self.video_to_audio_attn = LTXSelfAttention(
|
||||
query_dim=audio.dim,
|
||||
context_dim=video.dim,
|
||||
@@ -1131,7 +1510,17 @@ class BasicAVTransformerBlock(torch.nn.Module):
|
||||
self,
|
||||
video: TransformerArgs | None,
|
||||
audio: TransformerArgs | None,
|
||||
video_attention_mask: torch.Tensor | None = None,
|
||||
audio_attention_mask: torch.Tensor | None = None,
|
||||
) -> tuple[TransformerArgs | None, TransformerArgs | None]:
|
||||
"""Forward pass for transformer block.
|
||||
|
||||
Args:
|
||||
video: Video transformer args
|
||||
audio: Audio transformer args
|
||||
video_attention_mask: SP padding attention mask for video [B, padded_seq_len]
|
||||
audio_attention_mask: SP padding attention mask for audio [B, padded_seq_len]
|
||||
"""
|
||||
vx = video.x if video is not None else None
|
||||
ax = audio.x if audio is not None else None
|
||||
|
||||
@@ -1145,7 +1534,12 @@ class BasicAVTransformerBlock(torch.nn.Module):
|
||||
self.scale_shift_table, vx.shape[0], video.timesteps, slice(0, 3)
|
||||
)
|
||||
norm_vx = torch.nn.functional.rms_norm(vx, (vx.shape[-1],), eps=self.norm_eps) * (1 + vscale_msa) + vshift_msa
|
||||
vx = vx + self.attn1(norm_vx, pe=video.positional_embeddings) * vgate_msa
|
||||
# Self-attention: pass SP attention mask for distributed attention
|
||||
if self.use_distributed_attention:
|
||||
vx = vx + self.attn1(norm_vx, pe=video.positional_embeddings, attention_mask=video_attention_mask) * vgate_msa
|
||||
else:
|
||||
vx = vx + self.attn1(norm_vx, pe=video.positional_embeddings) * vgate_msa
|
||||
# Text cross-attention: no SP mask needed (text is replicated, uses local attention)
|
||||
vx = vx + self.attn2(
|
||||
torch.nn.functional.rms_norm(vx, (vx.shape[-1],), eps=self.norm_eps),
|
||||
context=video.context,
|
||||
@@ -1157,7 +1551,12 @@ class BasicAVTransformerBlock(torch.nn.Module):
|
||||
self.audio_scale_shift_table, ax.shape[0], audio.timesteps, slice(0, 3)
|
||||
)
|
||||
norm_ax = torch.nn.functional.rms_norm(ax, (ax.shape[-1],), eps=self.norm_eps) * (1 + ascale_msa) + ashift_msa
|
||||
ax = ax + self.audio_attn1(norm_ax, pe=audio.positional_embeddings) * agate_msa
|
||||
# Self-attention: pass SP attention mask for distributed attention
|
||||
if self.use_distributed_attention:
|
||||
ax = ax + self.audio_attn1(norm_ax, pe=audio.positional_embeddings, attention_mask=audio_attention_mask) * agate_msa
|
||||
else:
|
||||
ax = ax + self.audio_attn1(norm_ax, pe=audio.positional_embeddings) * agate_msa
|
||||
# Text cross-attention: no SP mask needed (text is replicated, uses local attention)
|
||||
ax = ax + self.audio_attn2(
|
||||
torch.nn.functional.rms_norm(ax, (ax.shape[-1],), eps=self.norm_eps),
|
||||
context=audio.context,
|
||||
@@ -1197,28 +1596,91 @@ class BasicAVTransformerBlock(torch.nn.Module):
|
||||
if run_a2v:
|
||||
vx_scaled = vx_norm3 * (1 + scale_ca_video_hidden_states_a2v) + shift_ca_video_hidden_states_a2v
|
||||
ax_scaled = ax_norm3 * (1 + scale_ca_audio_hidden_states_a2v) + shift_ca_audio_hidden_states_a2v
|
||||
vx = vx + (
|
||||
self.audio_to_video_attn(
|
||||
vx_scaled,
|
||||
context=ax_scaled,
|
||||
pe=video.cross_positional_embeddings,
|
||||
k_pe=audio.cross_positional_embeddings,
|
||||
# Audio-to-video cross-attention: Q from video (sharded), K/V from audio
|
||||
# For SP, gather audio context from all ranks before cross-attention
|
||||
if self.use_distributed_attention:
|
||||
# Gather audio context from all SP ranks
|
||||
ax_context = sequence_model_parallel_all_gather(ax_scaled, dim=1)
|
||||
# Video Q: video tokens divide evenly (adjusted upstream), so simple slicing works
|
||||
video_q_pe = None
|
||||
if video.cross_positional_embeddings is not None:
|
||||
sp_rank = get_sp_parallel_rank()
|
||||
local_seq_len = vx_scaled.shape[1]
|
||||
pe_tuple = video.cross_positional_embeddings
|
||||
start_idx = sp_rank * local_seq_len
|
||||
end_idx = start_idx + local_seq_len
|
||||
video_q_pe = tuple(pe[:, :, start_idx:end_idx, :] for pe in pe_tuple)
|
||||
# Audio K: gathered context may have padding, trim to original length
|
||||
audio_k_pe = audio.cross_positional_embeddings
|
||||
if audio_k_pe is not None:
|
||||
original_audio_len = audio_k_pe[0].shape[2]
|
||||
ax_context = ax_context[:, :original_audio_len, :]
|
||||
vx = vx + (
|
||||
self.audio_to_video_attn(
|
||||
vx_scaled,
|
||||
context=ax_context,
|
||||
pe=video_q_pe,
|
||||
k_pe=audio_k_pe,
|
||||
)
|
||||
* gate_out_a2v
|
||||
)
|
||||
else:
|
||||
ax_context = ax_scaled
|
||||
video_q_pe = video.cross_positional_embeddings
|
||||
audio_k_pe = audio.cross_positional_embeddings
|
||||
vx = vx + (
|
||||
self.audio_to_video_attn(
|
||||
vx_scaled,
|
||||
context=ax_context,
|
||||
pe=video_q_pe,
|
||||
k_pe=audio_k_pe,
|
||||
)
|
||||
* gate_out_a2v
|
||||
)
|
||||
* gate_out_a2v
|
||||
)
|
||||
|
||||
if run_v2a:
|
||||
ax_scaled = ax_norm3 * (1 + scale_ca_audio_hidden_states_v2a) + shift_ca_audio_hidden_states_v2a
|
||||
vx_scaled = vx_norm3 * (1 + scale_ca_video_hidden_states_v2a) + shift_ca_video_hidden_states_v2a
|
||||
ax = ax + (
|
||||
self.video_to_audio_attn(
|
||||
ax_scaled,
|
||||
context=vx_scaled,
|
||||
pe=audio.cross_positional_embeddings,
|
||||
k_pe=video.cross_positional_embeddings,
|
||||
# Video-to-audio cross-attention: Q from audio, K/V from video
|
||||
# For SP, gather both modalities and run cross-attention on full sequences
|
||||
# (audio is short, so no need to keep it sharded for cross-attention)
|
||||
if self.use_distributed_attention:
|
||||
# Gather both modalities to full sequences
|
||||
ax_full = sequence_model_parallel_all_gather(ax_scaled, dim=1)
|
||||
vx_context = sequence_model_parallel_all_gather(vx_scaled, dim=1)
|
||||
# Use full positional embeddings directly
|
||||
audio_q_pe = audio.cross_positional_embeddings
|
||||
video_k_pe = video.cross_positional_embeddings
|
||||
# Trim gathered tensors to original lengths (may have padding from sharding)
|
||||
if audio_q_pe is not None:
|
||||
original_audio_len = audio_q_pe[0].shape[2]
|
||||
ax_full = ax_full[:, :original_audio_len, :]
|
||||
if video_k_pe is not None:
|
||||
original_video_len = video_k_pe[0].shape[2]
|
||||
vx_context = vx_context[:, :original_video_len, :]
|
||||
# Run cross-attention on full audio
|
||||
v2a_out_full = self.video_to_audio_attn(
|
||||
ax_full,
|
||||
context=vx_context,
|
||||
pe=audio_q_pe,
|
||||
k_pe=video_k_pe,
|
||||
)
|
||||
# Shard the output back to match ax's sharding
|
||||
v2a_out, _ = sequence_model_parallel_shard(v2a_out_full, dim=1)
|
||||
ax = ax + v2a_out * gate_out_v2a
|
||||
else:
|
||||
vx_context = vx_scaled
|
||||
audio_q_pe = audio.cross_positional_embeddings
|
||||
video_k_pe = video.cross_positional_embeddings
|
||||
ax = ax + (
|
||||
self.video_to_audio_attn(
|
||||
ax_scaled,
|
||||
context=vx_context,
|
||||
pe=audio_q_pe,
|
||||
k_pe=video_k_pe,
|
||||
)
|
||||
* gate_out_v2a
|
||||
)
|
||||
* gate_out_v2a
|
||||
)
|
||||
|
||||
if run_vx:
|
||||
vshift_mlp, vscale_mlp, vgate_mlp = self.get_ada_values(
|
||||
@@ -1289,6 +1751,8 @@ class LTXModel(torch.nn.Module):
|
||||
av_ca_timestep_scale_multiplier: int = 1,
|
||||
rope_type: LTXRopeType = LTXRopeType.INTERLEAVED,
|
||||
double_precision_rope: bool = False,
|
||||
use_distributed_attention: bool = False,
|
||||
prefix: str = "",
|
||||
):
|
||||
super().__init__()
|
||||
self._enable_gradient_checkpointing = False
|
||||
@@ -1340,6 +1804,8 @@ class LTXModel(torch.nn.Module):
|
||||
audio_attention_head_dim=audio_attention_head_dim if model_type.is_audio_enabled() else 0,
|
||||
audio_cross_attention_dim=audio_cross_attention_dim,
|
||||
norm_eps=norm_eps,
|
||||
use_distributed_attention=use_distributed_attention,
|
||||
prefix=prefix,
|
||||
)
|
||||
|
||||
def _init_video(
|
||||
@@ -1469,6 +1935,8 @@ class LTXModel(torch.nn.Module):
|
||||
audio_attention_head_dim: int,
|
||||
audio_cross_attention_dim: int,
|
||||
norm_eps: float,
|
||||
use_distributed_attention: bool = False,
|
||||
prefix: str = "",
|
||||
) -> None:
|
||||
video_config = (
|
||||
TransformerConfig(
|
||||
@@ -1490,6 +1958,7 @@ class LTXModel(torch.nn.Module):
|
||||
if self.model_type.is_audio_enabled()
|
||||
else None
|
||||
)
|
||||
self.use_distributed_attention = use_distributed_attention
|
||||
self.transformer_blocks = torch.nn.ModuleList(
|
||||
[
|
||||
BasicAVTransformerBlock(
|
||||
@@ -1498,6 +1967,8 @@ class LTXModel(torch.nn.Module):
|
||||
audio=audio_config,
|
||||
rope_type=self.rope_type,
|
||||
norm_eps=norm_eps,
|
||||
use_distributed_attention=use_distributed_attention,
|
||||
prefix=prefix,
|
||||
)
|
||||
for idx in range(num_layers)
|
||||
]
|
||||
@@ -1507,9 +1978,16 @@ class LTXModel(torch.nn.Module):
|
||||
self,
|
||||
video: TransformerArgs | None,
|
||||
audio: TransformerArgs | None,
|
||||
video_attention_mask: torch.Tensor | None = None,
|
||||
audio_attention_mask: torch.Tensor | None = None,
|
||||
) -> tuple[TransformerArgs | None, TransformerArgs | None]:
|
||||
for block in self.transformer_blocks:
|
||||
video, audio = block(video=video, audio=audio)
|
||||
video, audio = block(
|
||||
video=video,
|
||||
audio=audio,
|
||||
video_attention_mask=video_attention_mask,
|
||||
audio_attention_mask=audio_attention_mask,
|
||||
)
|
||||
return video, audio
|
||||
|
||||
def _process_output(
|
||||
@@ -1533,7 +2011,17 @@ class LTXModel(torch.nn.Module):
|
||||
self,
|
||||
video: Modality | None,
|
||||
audio: Modality | None,
|
||||
video_attention_mask: torch.Tensor | None = None,
|
||||
audio_attention_mask: torch.Tensor | None = None,
|
||||
) -> tuple[torch.Tensor | None, torch.Tensor | None]:
|
||||
"""Forward pass through the LTX model.
|
||||
|
||||
Args:
|
||||
video: Video modality input
|
||||
audio: Audio modality input
|
||||
video_attention_mask: SP padding attention mask for video [B, padded_seq_len]
|
||||
audio_attention_mask: SP padding attention mask for audio [B, padded_seq_len]
|
||||
"""
|
||||
if os.getenv("LTX2_PIPELINE_DEBUG_LOG", "0") == "1":
|
||||
_debug_block_log_line(
|
||||
"fastvideo:patchify_proj"
|
||||
@@ -1551,7 +2039,12 @@ class LTXModel(torch.nn.Module):
|
||||
audio_args = self.audio_args_preprocessor.prepare(audio) if audio is not None else None
|
||||
_debug_transformer_args("fastvideo:prep_video", video_args)
|
||||
_debug_transformer_args("fastvideo:prep_audio", audio_args)
|
||||
video_out, audio_out = self._process_transformer_blocks(video_args, audio_args)
|
||||
video_out, audio_out = self._process_transformer_blocks(
|
||||
video_args,
|
||||
audio_args,
|
||||
video_attention_mask=video_attention_mask,
|
||||
audio_attention_mask=audio_attention_mask,
|
||||
)
|
||||
|
||||
vx = (
|
||||
self._process_output(
|
||||
@@ -1588,6 +2081,24 @@ class LTX2Transformer3DModel(CachableDiT):
|
||||
|
||||
arch = config.arch_config
|
||||
|
||||
# Get SP world size for distributed attention
|
||||
sp_world_size = get_sp_world_size()
|
||||
use_distributed_attention = sp_world_size > 1
|
||||
|
||||
# Validate that attention heads are divisible by SP world size
|
||||
if sp_world_size > 1:
|
||||
assert arch.num_attention_heads % sp_world_size == 0, (
|
||||
f"The number of video attention heads ({arch.num_attention_heads}) "
|
||||
f"must be divisible by the sequence parallel size ({sp_world_size})"
|
||||
)
|
||||
assert arch.audio_num_attention_heads % sp_world_size == 0, (
|
||||
f"The number of audio attention heads ({arch.audio_num_attention_heads}) "
|
||||
f"must be divisible by the sequence parallel size ({sp_world_size})"
|
||||
)
|
||||
logger.info(
|
||||
f"LTX2 sequence parallelism enabled with SP world size {sp_world_size}"
|
||||
)
|
||||
|
||||
model_type = LTXModelType.AudioVideo
|
||||
self.model = LTXModel(
|
||||
model_type=model_type,
|
||||
@@ -1612,6 +2123,8 @@ class LTX2Transformer3DModel(CachableDiT):
|
||||
audio_cross_attention_dim=arch.audio_cross_attention_dim,
|
||||
audio_positional_embedding_max_pos=arch.audio_positional_embedding_max_pos,
|
||||
av_ca_timestep_scale_multiplier=arch.av_ca_timestep_scale_multiplier,
|
||||
use_distributed_attention=use_distributed_attention,
|
||||
prefix=config.prefix,
|
||||
)
|
||||
|
||||
self.patchifier = VideoLatentPatchifier(patch_size=arch.patch_size[1])
|
||||
@@ -1627,6 +2140,7 @@ class LTX2Transformer3DModel(CachableDiT):
|
||||
self.hidden_size = arch.num_attention_heads * arch.attention_head_dim
|
||||
self.num_attention_heads = arch.num_attention_heads
|
||||
self.num_channels_latents = arch.num_channels_latents
|
||||
self._logged_attention_mask = False
|
||||
|
||||
if os.getenv("LTX2_DEBUG_DETAIL", "0") == "1":
|
||||
detail_path = os.getenv("LTX2_PIPELINE_DEBUG_DETAIL_PATH", "")
|
||||
@@ -1702,10 +2216,12 @@ class LTX2Transformer3DModel(CachableDiT):
|
||||
if isinstance(encoder_hidden_states, list):
|
||||
encoder_hidden_states = encoder_hidden_states[0]
|
||||
|
||||
video_shape = VideoLatentShape.from_torch_shape(hidden_states.shape)
|
||||
positions = self.patchifier.get_patch_grid_bounds(
|
||||
video_shape, device=hidden_states.device
|
||||
)
|
||||
# Get SP parameters
|
||||
sp_world_size = get_sp_world_size()
|
||||
sp_rank = get_sp_parallel_rank()
|
||||
batch_size = hidden_states.shape[0]
|
||||
|
||||
# Get fps for position computation
|
||||
fps = None
|
||||
try:
|
||||
forward_ctx = get_forward_context()
|
||||
@@ -1717,17 +2233,60 @@ class LTX2Transformer3DModel(CachableDiT):
|
||||
fps_value = fps_value[0] if fps_value else None
|
||||
if fps_value is not None:
|
||||
fps = float(fps_value)
|
||||
|
||||
# Patchify video latents
|
||||
video_shape = VideoLatentShape.from_torch_shape(hidden_states.shape)
|
||||
latents = self.patchifier.patchify(hidden_states)
|
||||
|
||||
# Shard video latents and timestep across SP ranks
|
||||
video_original_seq_len = latents.shape[1]
|
||||
video_attention_mask = None
|
||||
video_timestep = timestep
|
||||
if sp_world_size > 1:
|
||||
latents, video_original_seq_len = sequence_model_parallel_shard(latents, dim=1)
|
||||
# Shard timestep along sequence dimension (timestep has shape [batch, seq_len])
|
||||
video_timestep, _ = sequence_model_parallel_shard(timestep, dim=1)
|
||||
# Create attention mask for padded tokens
|
||||
current_seq_len = latents.shape[1]
|
||||
padded_seq_len = current_seq_len * sp_world_size
|
||||
if padded_seq_len > video_original_seq_len:
|
||||
if not self._logged_attention_mask:
|
||||
logger.info(
|
||||
f"Video padding applied, original seq len: {video_original_seq_len}, "
|
||||
f"padded seq len: {padded_seq_len}"
|
||||
)
|
||||
self._logged_attention_mask = True
|
||||
video_attention_mask = create_attention_mask_for_padding(
|
||||
seq_len=video_original_seq_len,
|
||||
padded_seq_len=padded_seq_len,
|
||||
batch_size=batch_size,
|
||||
device=hidden_states.device,
|
||||
)
|
||||
|
||||
# Compute RoPE positions for the FULL sequence (before sharding)
|
||||
# This is necessary because sequence sharding may not align to frame boundaries.
|
||||
# The full RoPE will be passed to DistributedAttention which applies it after
|
||||
# the all-to-all (when each rank has the full sequence but subset of heads).
|
||||
positions = self.patchifier.get_patch_grid_bounds(
|
||||
video_shape, device=hidden_states.device
|
||||
)
|
||||
positions = _get_pixel_coords(
|
||||
positions,
|
||||
DEFAULT_LTX2_SCALE_FACTORS,
|
||||
fps=fps,
|
||||
causal_fix=True,
|
||||
).to(hidden_states.dtype)
|
||||
latents = self.patchifier.patchify(hidden_states)
|
||||
|
||||
# Pad positions to match padded sequence length for SP cross-attention
|
||||
if sp_world_size > 1 and padded_seq_len > video_original_seq_len:
|
||||
padding_needed = padded_seq_len - video_original_seq_len
|
||||
last_pos = positions[:, :, -1:, :].expand(-1, -1, padding_needed, -1)
|
||||
positions = torch.cat([positions, last_pos], dim=2)
|
||||
|
||||
video_modality = Modality(
|
||||
enabled=True,
|
||||
latent=latents,
|
||||
timesteps=timestep,
|
||||
timesteps=video_timestep,
|
||||
positions=positions,
|
||||
context=encoder_hidden_states,
|
||||
context_mask=encoder_attention_mask,
|
||||
@@ -1745,19 +2304,45 @@ class LTX2Transformer3DModel(CachableDiT):
|
||||
f"latent_head={video_head} "
|
||||
f"latent_checksum={video_checksum:.6f}"
|
||||
)
|
||||
|
||||
# Process audio modality if provided
|
||||
audio_modality = None
|
||||
audio_shape = None
|
||||
audio_original_seq_len = 0
|
||||
audio_attention_mask = None
|
||||
|
||||
if audio_hidden_states is not None and audio_encoder_hidden_states is not None and audio_timestep is not None:
|
||||
audio_shape = AudioLatentShape.from_torch_shape(
|
||||
audio_hidden_states.shape)
|
||||
audio_shape = AudioLatentShape.from_torch_shape(audio_hidden_states.shape)
|
||||
audio_latents = self.audio_patchifier.patchify(audio_hidden_states)
|
||||
|
||||
# Shard audio latents and timestep across SP ranks
|
||||
audio_original_seq_len = audio_latents.shape[1]
|
||||
sharded_audio_timestep = audio_timestep
|
||||
if sp_world_size > 1:
|
||||
audio_latents, audio_original_seq_len = sequence_model_parallel_shard(audio_latents, dim=1)
|
||||
# Shard audio timestep along sequence dimension
|
||||
sharded_audio_timestep, _ = sequence_model_parallel_shard(audio_timestep, dim=1)
|
||||
# Create attention mask for padded tokens
|
||||
audio_current_seq_len = audio_latents.shape[1]
|
||||
audio_padded_seq_len = audio_current_seq_len * sp_world_size
|
||||
if audio_padded_seq_len > audio_original_seq_len:
|
||||
audio_attention_mask = create_attention_mask_for_padding(
|
||||
seq_len=audio_original_seq_len,
|
||||
padded_seq_len=audio_padded_seq_len,
|
||||
batch_size=batch_size,
|
||||
device=audio_hidden_states.device,
|
||||
)
|
||||
|
||||
# Compute audio RoPE positions for the FULL sequence (before sharding)
|
||||
# Same as video: full RoPE is applied after all-to-all in DistributedAttention
|
||||
audio_positions = self.audio_patchifier.get_patch_grid_bounds(
|
||||
audio_shape, device=audio_hidden_states.device)
|
||||
audio_latents = self.audio_patchifier.patchify(
|
||||
audio_hidden_states)
|
||||
audio_shape, device=audio_hidden_states.device
|
||||
)
|
||||
|
||||
audio_modality = Modality(
|
||||
enabled=True,
|
||||
latent=audio_latents,
|
||||
timesteps=audio_timestep,
|
||||
timesteps=sharded_audio_timestep,
|
||||
positions=audio_positions,
|
||||
context=audio_encoder_hidden_states,
|
||||
context_mask=audio_encoder_attention_mask,
|
||||
@@ -1775,10 +2360,16 @@ class LTX2Transformer3DModel(CachableDiT):
|
||||
f"latent_head={audio_head} "
|
||||
f"latent_checksum={audio_checksum:.6f}"
|
||||
)
|
||||
|
||||
# Run transformer with attention masks
|
||||
video_out, audio_out = self.model(
|
||||
video=video_modality,
|
||||
audio=audio_modality,
|
||||
video_attention_mask=video_attention_mask,
|
||||
audio_attention_mask=audio_attention_mask,
|
||||
)
|
||||
|
||||
# Denoised prediction
|
||||
if video_out is not None and video_modality is not None:
|
||||
video_out = _to_denoised(
|
||||
video_modality.latent,
|
||||
@@ -1791,6 +2382,20 @@ class LTX2Transformer3DModel(CachableDiT):
|
||||
audio_out,
|
||||
audio_modality.timesteps,
|
||||
)
|
||||
|
||||
# Gather and unpad video output
|
||||
if sp_world_size > 1 and video_out is not None:
|
||||
video_out = sequence_model_parallel_all_gather_with_unpad(
|
||||
video_out, video_original_seq_len, dim=1
|
||||
)
|
||||
|
||||
# Gather and unpad audio output
|
||||
if sp_world_size > 1 and audio_out is not None:
|
||||
audio_out = sequence_model_parallel_all_gather_with_unpad(
|
||||
audio_out, audio_original_seq_len, dim=1
|
||||
)
|
||||
|
||||
# Unpatchify
|
||||
video_out = self.patchifier.unpatchify(
|
||||
video_out, output_shape=video_shape)
|
||||
if audio_out is None or audio_shape is None:
|
||||
@@ -1798,3 +2403,6 @@ class LTX2Transformer3DModel(CachableDiT):
|
||||
audio_out = self.audio_patchifier.unpatchify(
|
||||
audio_out, output_shape=audio_shape)
|
||||
return video_out, audio_out
|
||||
|
||||
# Entry point for model registry
|
||||
EntryClass = LTX2Transformer3DModel
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -561,3 +561,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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -351,3 +351,5 @@ class Reason1TextEncoder(TextEncoder):
|
||||
tensor.std(dim=-1, keepdim=True) + 1e-8
|
||||
)
|
||||
|
||||
# Entry point for model registry
|
||||
EntryClass = Reason1TextEncoder
|
||||
|
||||
@@ -417,3 +417,6 @@ class SiglipVisionModel(ImageEncoder):
|
||||
loaded_params.add(name)
|
||||
|
||||
return loaded_params
|
||||
|
||||
# Entry point for model registry
|
||||
EntryClass = SiglipVisionModel
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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."""
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -1319,3 +1319,5 @@ class AutoencoderKLWan(nn.Module, ParallelTiledVAE):
|
||||
dec = self.decode(z)
|
||||
return dec
|
||||
|
||||
# Entry point for model registry
|
||||
EntryClass = AutoencoderKLWan
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
@@ -16,7 +16,7 @@ from fastvideo.pipelines.stages import (DecodingStage, InputValidationStage,
|
||||
LTX2AudioDecodingStage,
|
||||
LTX2DenoisingStage,
|
||||
LTX2LatentPreparationStage,
|
||||
TextEncodingStage)
|
||||
LTX2TextEncodingStage)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -40,7 +40,7 @@ class LTX2Pipeline(ComposedPipelineBase):
|
||||
|
||||
self.add_stage(
|
||||
stage_name="prompt_encoding_stage",
|
||||
stage=TextEncodingStage(
|
||||
stage=LTX2TextEncodingStage(
|
||||
text_encoders=[self.get_module("text_encoder")],
|
||||
tokenizers=[self.get_module("tokenizer")],
|
||||
),
|
||||
|
||||
@@ -15,6 +15,8 @@ import torch
|
||||
from fastvideo.configs.pipelines import PipelineConfig
|
||||
from fastvideo.distributed import (
|
||||
maybe_init_distributed_environment_and_model_parallel, get_world_group)
|
||||
from fastvideo.distributed.communication_op import (
|
||||
warmup_sequence_parallel_communication)
|
||||
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.profiler import get_or_create_profiler
|
||||
@@ -99,6 +101,58 @@ class ComposedPipelineBase(ABC):
|
||||
module.requires_grad_(True)
|
||||
module.train()
|
||||
|
||||
@staticmethod
|
||||
def _compile_with_conditions(
|
||||
module: torch.nn.Module,
|
||||
compile_kwargs: dict[str, Any],
|
||||
) -> int:
|
||||
"""Compile submodules that match module._compile_conditions."""
|
||||
compile_conditions = getattr(module, "_compile_conditions", None)
|
||||
if not compile_conditions:
|
||||
return 0
|
||||
|
||||
compiled_count = 0
|
||||
for name, submodule in module.named_modules():
|
||||
if not name:
|
||||
continue
|
||||
if any(cond(name, submodule) for cond in compile_conditions):
|
||||
submodule.forward = torch.compile(submodule.forward,
|
||||
**compile_kwargs)
|
||||
compiled_count += 1
|
||||
return compiled_count
|
||||
|
||||
def _maybe_compile_pipeline_module(
|
||||
self,
|
||||
module_name: str,
|
||||
fsdp_module_cls: type | None,
|
||||
compile_kwargs: dict[str, Any],
|
||||
) -> None:
|
||||
if module_name not in self.modules:
|
||||
return
|
||||
|
||||
module = self.modules[module_name]
|
||||
if fsdp_module_cls is not None and isinstance(module, fsdp_module_cls):
|
||||
logger.info(
|
||||
"%s is already FSDP-wrapped; skipping torch.compile in pipeline",
|
||||
module_name.capitalize(),
|
||||
)
|
||||
return
|
||||
|
||||
compiled_count = self._compile_with_conditions(module, compile_kwargs)
|
||||
if compiled_count > 0:
|
||||
logger.info(
|
||||
"Enabled torch.compile for %d submodules in %s via _compile_conditions with kwargs=%s",
|
||||
compiled_count,
|
||||
module_name,
|
||||
compile_kwargs,
|
||||
)
|
||||
return
|
||||
|
||||
# Backward-compatible fallback: compile full module if no condition matched.
|
||||
logger.info("Enabling torch.compile for %s with kwargs=%s", module_name,
|
||||
compile_kwargs)
|
||||
self.modules[module_name] = torch.compile(module, **compile_kwargs)
|
||||
|
||||
def post_init(self) -> None:
|
||||
assert self.fastvideo_args is not None, "fastvideo_args must be set"
|
||||
if self.post_init_called:
|
||||
@@ -114,7 +168,6 @@ class ComposedPipelineBase(ABC):
|
||||
|
||||
self.initialize_pipeline(self.fastvideo_args)
|
||||
if self.fastvideo_args.enable_torch_compile:
|
||||
transformer_module = self.modules["transformer"]
|
||||
if self.fastvideo_args.training_mode:
|
||||
logger.info(
|
||||
"Torch Compile enabled via FSDP loader for training; skipping additional pipeline compile"
|
||||
@@ -128,35 +181,26 @@ class ComposedPipelineBase(ABC):
|
||||
fsdp_module_cls = None
|
||||
|
||||
compile_kwargs = self.fastvideo_args.torch_compile_kwargs or {}
|
||||
if fsdp_module_cls is not None and isinstance(
|
||||
transformer_module, fsdp_module_cls):
|
||||
logger.info(
|
||||
"Transformer is already FSDP-wrapped; skipping torch.compile in pipeline"
|
||||
)
|
||||
else:
|
||||
logger.info("Enabling torch.compile for DiT with kwargs=%s",
|
||||
compile_kwargs)
|
||||
self.modules["transformer"] = torch.compile(
|
||||
transformer_module, **compile_kwargs)
|
||||
if "transformer_2" in self.modules:
|
||||
transformer_module_2 = self.modules["transformer_2"]
|
||||
if fsdp_module_cls is not None and isinstance(
|
||||
transformer_module_2, fsdp_module_cls):
|
||||
logger.info(
|
||||
"Transformer_2 is already FSDP-wrapped; skipping torch.compile in pipeline"
|
||||
)
|
||||
else:
|
||||
logger.info(
|
||||
"Enabling torch.compile for Transformer_2 with kwargs=%s",
|
||||
compile_kwargs)
|
||||
self.modules["transformer_2"] = torch.compile(
|
||||
transformer_module_2, **compile_kwargs)
|
||||
self._maybe_compile_pipeline_module(
|
||||
module_name="transformer",
|
||||
fsdp_module_cls=fsdp_module_cls,
|
||||
compile_kwargs=compile_kwargs,
|
||||
)
|
||||
self._maybe_compile_pipeline_module(
|
||||
module_name="transformer_2",
|
||||
fsdp_module_cls=fsdp_module_cls,
|
||||
compile_kwargs=compile_kwargs,
|
||||
)
|
||||
logger.info("Torch Compile enabled for DiT")
|
||||
|
||||
if not self.fastvideo_args.training_mode:
|
||||
logger.info("Creating pipeline stages...")
|
||||
self.create_pipeline_stages(self.fastvideo_args)
|
||||
|
||||
# Warmup NCCL communicators for sequence parallelism to avoid
|
||||
# slow first forward pass due to lazy initialization
|
||||
warmup_sequence_parallel_communication()
|
||||
|
||||
def initialize_training_pipeline(self, training_args: TrainingArgs):
|
||||
raise NotImplementedError(
|
||||
"if training_mode is True, the pipeline must implement this method")
|
||||
@@ -287,6 +331,7 @@ class ComposedPipelineBase(ABC):
|
||||
# remove keys that are not pipeline modules
|
||||
model_index.pop("_class_name")
|
||||
model_index.pop("_diffusers_version")
|
||||
model_index.pop("_name_or_path", None)
|
||||
model_index.pop("workload_type", None)
|
||||
if "boundary_ratio" in model_index and model_index[
|
||||
"boundary_ratio"] is not None:
|
||||
|
||||
@@ -123,6 +123,7 @@ class ForwardBatch:
|
||||
|
||||
# Latent tensors
|
||||
latents: torch.Tensor | None = None
|
||||
lq_latents: torch.Tensor | None = None
|
||||
raw_latent_shape: tuple[int, ...] | None = None
|
||||
noise_pred: torch.Tensor | None = None
|
||||
image_latent: torch.Tensor | None = None
|
||||
@@ -144,6 +145,8 @@ class ForwardBatch:
|
||||
# Original dimensions (before VAE scaling)
|
||||
height: list[int] | int | None = None
|
||||
width: list[int] | int | None = None
|
||||
height_sr: list[int] | int | None = None
|
||||
width_sr: list[int] | int | None = None
|
||||
fps: list[int] | int | None = None
|
||||
|
||||
# Timesteps
|
||||
@@ -154,6 +157,7 @@ class ForwardBatch:
|
||||
|
||||
# Scheduler parameters
|
||||
num_inference_steps: int = 50
|
||||
num_inference_steps_sr: int = 50
|
||||
guidance_scale: float = 1.0
|
||||
guidance_scale_2: float | None = None
|
||||
guidance_rescale: float = 0.0
|
||||
|
||||
@@ -16,29 +16,6 @@ from fastvideo.pipelines.lora_pipeline import LoRAPipeline
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# map pipeline name to folder name
|
||||
_PIPELINE_NAME_TO_ARCHITECTURE_NAME: dict[str, str] = {
|
||||
"WanPipeline": "wan",
|
||||
"WanDMDPipeline": "wan",
|
||||
"WanImageToVideoPipeline": "wan",
|
||||
"WanVideoToVideoPipeline": "wan",
|
||||
"WanCausalDMDPipeline": "wan",
|
||||
"TurboDiffusionPipeline": "turbodiffusion",
|
||||
"TurboDiffusionI2VPipeline": "turbodiffusion",
|
||||
"StepVideoPipeline": "stepvideo",
|
||||
"HunyuanVideoPipeline": "hunyuan",
|
||||
"HunyuanVideo15Pipeline": "hunyuan15",
|
||||
"HYWorldPipeline": "hyworld",
|
||||
"Cosmos2VideoToWorldPipeline": "cosmos",
|
||||
"Cosmos2_5Pipeline": "cosmos",
|
||||
"MatrixGamePipeline": "matrixgame",
|
||||
"MatrixGameCausalDMDPipeline": "matrixgame",
|
||||
"LongCatPipeline": "longcat",
|
||||
"LongCatImageToVideoPipeline": "longcat",
|
||||
"LongCatVideoContinuationPipeline": "longcat",
|
||||
"LTX2Pipeline": "ltx2",
|
||||
}
|
||||
|
||||
_PREPROCESS_WORKLOAD_TYPE_TO_PIPELINE_NAME: dict[WorkloadType, str] = {
|
||||
WorkloadType.I2V: "PreprocessPipelineI2V",
|
||||
WorkloadType.T2V: "PreprocessPipelineT2V",
|
||||
@@ -73,44 +50,36 @@ class PipelineType(str, Enum):
|
||||
|
||||
@dataclass
|
||||
class _PipelineRegistry:
|
||||
# Keyed by pipeline_type -> architecture -> pipeline_name
|
||||
# pipelines[pipeline_type][architecture][pipeline_name] = pipeline_cls
|
||||
pipelines: dict[str, dict[str, dict[str, type[ComposedPipelineBase]
|
||||
| None]]] = field(default_factory=dict)
|
||||
# Keyed by pipeline_type -> pipeline_name
|
||||
# pipelines[pipeline_type][pipeline_name] = pipeline_cls
|
||||
pipelines: dict[str, dict[str, type[ComposedPipelineBase]
|
||||
| None]] = field(default_factory=dict)
|
||||
|
||||
def get_supported_archs(self, pipeline_name_in_config: str,
|
||||
pipeline_type: PipelineType) -> Set[str]:
|
||||
"""Get supported architectures, optionally filtered by pipeline type and workload type."""
|
||||
arch = _PIPELINE_NAME_TO_ARCHITECTURE_NAME[pipeline_name_in_config]
|
||||
return set(self.pipelines[pipeline_type.value][arch].keys())
|
||||
def get_supported_pipelines(self, pipeline_type: PipelineType) -> Set[str]:
|
||||
"""Get supported pipelines for the given pipeline type."""
|
||||
return set(self.pipelines.get(pipeline_type.value, {}).keys())
|
||||
|
||||
def _load_preprocess_pipeline_cls(
|
||||
self, workload_type: WorkloadType,
|
||||
arch: str) -> type[ComposedPipelineBase] | None:
|
||||
self,
|
||||
workload_type: WorkloadType) -> type[ComposedPipelineBase] | None:
|
||||
pipeline_name = _PREPROCESS_WORKLOAD_TYPE_TO_PIPELINE_NAME[
|
||||
workload_type]
|
||||
|
||||
return self.pipelines[
|
||||
PipelineType.PREPROCESS.value][arch][pipeline_name]
|
||||
return self.pipelines.get(PipelineType.PREPROCESS.value,
|
||||
{}).get(pipeline_name)
|
||||
|
||||
def _try_load_pipeline_cls(
|
||||
self, pipeline_name_in_config: str, pipeline_type: PipelineType,
|
||||
workload_type: WorkloadType
|
||||
) -> type[ComposedPipelineBase] | type[LoRAPipeline] | None:
|
||||
"""Try to load a pipeline class for the given architecture, pipeline type, and workload type."""
|
||||
arch = _PIPELINE_NAME_TO_ARCHITECTURE_NAME[pipeline_name_in_config]
|
||||
|
||||
if (pipeline_type.value not in self.pipelines
|
||||
or arch not in self.pipelines[pipeline_type.value]):
|
||||
if pipeline_type.value not in self.pipelines:
|
||||
return None
|
||||
|
||||
if pipeline_type == PipelineType.PREPROCESS:
|
||||
return self._load_preprocess_pipeline_cls(workload_type, arch)
|
||||
elif pipeline_type == PipelineType.BASIC:
|
||||
return self.pipelines[
|
||||
pipeline_type.value][arch][pipeline_name_in_config]
|
||||
elif pipeline_type == PipelineType.TRAINING:
|
||||
pass
|
||||
return self._load_preprocess_pipeline_cls(workload_type)
|
||||
elif pipeline_type == PipelineType.BASIC or pipeline_type == PipelineType.TRAINING:
|
||||
return self.pipelines[pipeline_type.value].get(
|
||||
pipeline_name_in_config)
|
||||
else:
|
||||
raise ValueError(f"Invalid pipeline type: {pipeline_type.value}")
|
||||
|
||||
@@ -130,18 +99,28 @@ class _PipelineRegistry:
|
||||
pipeline_type, workload_type)
|
||||
if pipeline_cls is not None:
|
||||
return pipeline_cls
|
||||
supported_archs = self.get_supported_archs(pipeline_name_in_config,
|
||||
pipeline_type)
|
||||
supported_pipelines = self.get_supported_pipelines(pipeline_type)
|
||||
raise ValueError(
|
||||
f"Pipeline architecture '{pipeline_name_in_config}' is not supported for pipeline type '{pipeline_type.value}' "
|
||||
f"Pipeline '{pipeline_name_in_config}' is not supported for pipeline type '{pipeline_type.value}' "
|
||||
f"and workload type '{workload_type.value}'. "
|
||||
f"Supported architectures: {supported_archs}")
|
||||
f"Supported pipelines: {supported_pipelines}")
|
||||
|
||||
|
||||
def import_pipeline_classes(
|
||||
pipeline_types: list[PipelineType] | PipelineType | None = None
|
||||
) -> dict[str, dict[str, type[ComposedPipelineBase] | None]]:
|
||||
pipeline_types_key: tuple[PipelineType, ...] | PipelineType | None
|
||||
if isinstance(pipeline_types, list):
|
||||
pipeline_types_key = tuple(pipeline_types)
|
||||
else:
|
||||
pipeline_types_key = pipeline_types
|
||||
return _import_pipeline_classes_cached(pipeline_types_key)
|
||||
|
||||
|
||||
@lru_cache
|
||||
def import_pipeline_classes(
|
||||
pipeline_types: list[PipelineType] | PipelineType | None = None
|
||||
) -> dict[str, dict[str, dict[str, type[ComposedPipelineBase] | None]]]:
|
||||
def _import_pipeline_classes_cached(
|
||||
pipeline_types: tuple[PipelineType, ...] | PipelineType | None = None
|
||||
) -> dict[str, dict[str, type[ComposedPipelineBase] | None]]:
|
||||
"""
|
||||
Import pipeline classes based on the pipeline type and workload type.
|
||||
|
||||
@@ -150,19 +129,16 @@ def import_pipeline_classes(
|
||||
If None, loads all types.
|
||||
|
||||
Returns:
|
||||
A three-level nested dictionary:
|
||||
{pipeline_type: {architecture_name: {pipeline_name: pipeline_cls}}}
|
||||
e.g., {"basic": {"wan": {"WanPipeline": WanPipeline}}}
|
||||
A two-level nested dictionary:
|
||||
{pipeline_type: {pipeline_name: pipeline_cls}}
|
||||
e.g., {"basic": {"WanPipeline": WanPipeline}}
|
||||
"""
|
||||
type_to_arch_to_pipeline_dict: dict[str,
|
||||
dict[str,
|
||||
dict[str,
|
||||
type[ComposedPipelineBase]
|
||||
| None]]] = {}
|
||||
type_to_pipeline_dict: dict[str, dict[str, type[ComposedPipelineBase]
|
||||
| None]] = {}
|
||||
package_name: str = "fastvideo.pipelines"
|
||||
|
||||
# Determine which pipeline types to scan
|
||||
if isinstance(pipeline_types, list):
|
||||
if isinstance(pipeline_types, tuple):
|
||||
pipeline_types_to_scan = [
|
||||
pipeline_type.value for pipeline_type in pipeline_types
|
||||
]
|
||||
@@ -174,8 +150,7 @@ def import_pipeline_classes(
|
||||
logger.info("Loading pipelines for types: %s", pipeline_types_to_scan)
|
||||
|
||||
for pipeline_type_str in pipeline_types_to_scan:
|
||||
arch_to_pipeline_dict: dict[str, dict[str, type[ComposedPipelineBase]
|
||||
| None]] = {}
|
||||
pipeline_dict: dict[str, type[ComposedPipelineBase] | None] = {}
|
||||
|
||||
# Try to load from pipeline-type-specific directory first
|
||||
pipeline_type_package_name = f"{package_name}.{pipeline_type_str}"
|
||||
@@ -187,50 +162,45 @@ def import_pipeline_classes(
|
||||
|
||||
for _, arch, ispkg in pkgutil.iter_modules(
|
||||
pipeline_type_package.__path__):
|
||||
pipeline_dict: dict[str, type[ComposedPipelineBase] | None] = {}
|
||||
|
||||
arch_package_name = f"{pipeline_type_package_name}.{arch}"
|
||||
if ispkg:
|
||||
arch_package = importlib.import_module(arch_package_name)
|
||||
for _, module_name, ispkg in pkgutil.walk_packages(
|
||||
arch_package.__path__, arch_package_name + "."):
|
||||
if not ispkg:
|
||||
pipeline_module = importlib.import_module(
|
||||
module_name)
|
||||
if hasattr(pipeline_module, "EntryClass"):
|
||||
if isinstance(pipeline_module.EntryClass, list):
|
||||
for pipeline in pipeline_module.EntryClass:
|
||||
pipeline_name = pipeline.__name__
|
||||
assert (
|
||||
pipeline_name not in pipeline_dict
|
||||
), f"Duplicated pipeline implementation for {pipeline_name} in {pipeline_type_str}.{arch_package_name}"
|
||||
pipeline_dict[pipeline_name] = pipeline
|
||||
else:
|
||||
pipeline_name = pipeline_module.EntryClass.__name__
|
||||
assert (
|
||||
pipeline_name not in pipeline_dict
|
||||
), f"Duplicated pipeline implementation for {pipeline_name} in {pipeline_type_str}.{arch_package_name}"
|
||||
pipeline_dict[
|
||||
pipeline_name] = pipeline_module.EntryClass
|
||||
if not ispkg:
|
||||
continue
|
||||
|
||||
arch_to_pipeline_dict[arch] = pipeline_dict
|
||||
arch_package = importlib.import_module(arch_package_name)
|
||||
for _, module_name, ispkg in pkgutil.walk_packages(
|
||||
arch_package.__path__, arch_package_name + "."):
|
||||
if ispkg:
|
||||
continue
|
||||
pipeline_module = importlib.import_module(module_name)
|
||||
if not hasattr(pipeline_module, "EntryClass"):
|
||||
continue
|
||||
entry_cls = pipeline_module.EntryClass
|
||||
entry_cls_list = ([
|
||||
entry_cls
|
||||
] if not isinstance(entry_cls, list) else entry_cls)
|
||||
|
||||
for pipeline in entry_cls_list:
|
||||
pipeline_name = pipeline.__name__
|
||||
if pipeline_name in pipeline_dict:
|
||||
logger.warning(
|
||||
"Duplicate pipeline name '%s' found in %s. Overwriting.",
|
||||
pipeline_name, pipeline_type_str)
|
||||
pipeline_dict[pipeline_name] = pipeline
|
||||
|
||||
except ImportError as e:
|
||||
raise ImportError(
|
||||
f"Could not import {pipeline_type_package_name} when importing pipeline classes: {e}"
|
||||
) from None
|
||||
|
||||
type_to_arch_to_pipeline_dict[pipeline_type_str] = arch_to_pipeline_dict
|
||||
type_to_pipeline_dict[pipeline_type_str] = pipeline_dict
|
||||
|
||||
# Log summary
|
||||
total_pipelines = sum(
|
||||
len(pipeline_dict)
|
||||
for arch_to_pipeline_dict in type_to_arch_to_pipeline_dict.values()
|
||||
for pipeline_dict in arch_to_pipeline_dict.values())
|
||||
len(pipeline_dict) for pipeline_dict in type_to_pipeline_dict.values())
|
||||
logger.info("Loaded %d pipeline classes across %d types", total_pipelines,
|
||||
len(pipeline_types_to_scan))
|
||||
|
||||
return type_to_arch_to_pipeline_dict
|
||||
return type_to_pipeline_dict
|
||||
|
||||
|
||||
def get_pipeline_registry(
|
||||
|
||||
@@ -10,10 +10,11 @@ from fastvideo.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.pipelines.stages.causal_denoising import CausalDMDDenosingStage
|
||||
from fastvideo.pipelines.stages.conditioning import ConditioningStage
|
||||
from fastvideo.pipelines.stages.decoding import DecodingStage
|
||||
from fastvideo.pipelines.stages.denoising import (Cosmos25DenoisingStage,
|
||||
CosmosDenoisingStage,
|
||||
DenoisingStage,
|
||||
DmdDenoisingStage)
|
||||
from fastvideo.pipelines.stages.denoising import (
|
||||
Cosmos25AutoDenoisingStage, Cosmos25DenoisingStage,
|
||||
Cosmos25V2WDenoisingStage, Cosmos25T2WDenoisingStage, CosmosDenoisingStage,
|
||||
DenoisingStage, DmdDenoisingStage)
|
||||
from fastvideo.pipelines.stages.sr_denoising import SRDenoisingStage
|
||||
from fastvideo.pipelines.stages.encoding import EncodingStage
|
||||
from fastvideo.pipelines.stages.image_encoding import (
|
||||
ImageEncodingStage, MatrixGameImageEncodingStage, RefImageEncodingStage,
|
||||
@@ -22,11 +23,13 @@ from fastvideo.pipelines.stages.image_encoding import (
|
||||
from fastvideo.pipelines.stages.input_validation import InputValidationStage
|
||||
from fastvideo.pipelines.stages.latent_preparation import (
|
||||
Cosmos25LatentPreparationStage, CosmosLatentPreparationStage,
|
||||
LatentPreparationStage)
|
||||
Cosmos25AutoLatentPreparationStage, Cosmos25T2WLatentPreparationStage,
|
||||
Cosmos25V2WLatentPreparationStage, LatentPreparationStage)
|
||||
from fastvideo.pipelines.stages.ltx2_audio_decoding import LTX2AudioDecodingStage
|
||||
from fastvideo.pipelines.stages.ltx2_denoising import LTX2DenoisingStage
|
||||
from fastvideo.pipelines.stages.ltx2_latent_preparation import (
|
||||
LTX2LatentPreparationStage)
|
||||
from fastvideo.pipelines.stages.ltx2_text_encoding import LTX2TextEncodingStage
|
||||
from fastvideo.pipelines.stages.matrixgame_denoising import (
|
||||
MatrixGameCausalDenoisingStage)
|
||||
from fastvideo.pipelines.stages.hyworld_denoising import HYWorldDenoisingStage
|
||||
@@ -50,6 +53,9 @@ __all__ = [
|
||||
"LatentPreparationStage",
|
||||
"CosmosLatentPreparationStage",
|
||||
"Cosmos25LatentPreparationStage",
|
||||
"Cosmos25T2WLatentPreparationStage",
|
||||
"Cosmos25V2WLatentPreparationStage",
|
||||
"Cosmos25AutoLatentPreparationStage",
|
||||
"LTX2LatentPreparationStage",
|
||||
"LTX2AudioDecodingStage",
|
||||
"ConditioningStage",
|
||||
@@ -60,7 +66,12 @@ __all__ = [
|
||||
"HYWorldDenoisingStage",
|
||||
"CosmosDenoisingStage",
|
||||
"Cosmos25DenoisingStage",
|
||||
"Cosmos25T2WDenoisingStage",
|
||||
"Cosmos25V2WDenoisingStage",
|
||||
"Cosmos25AutoDenoisingStage",
|
||||
"LTX2DenoisingStage",
|
||||
"LTX2TextEncodingStage",
|
||||
"SRDenoisingStage",
|
||||
"EncodingStage",
|
||||
"DecodingStage",
|
||||
"ImageEncodingStage",
|
||||
|
||||
@@ -242,8 +242,8 @@ class DecodingStage(PipelineStage):
|
||||
decoded_frames = self.decode(cur_latent, fastvideo_args)
|
||||
batch.trajectory_decoded.append(decoded_frames.cpu().float())
|
||||
|
||||
# Convert to CPU float32 for compatibility
|
||||
frames = frames.cpu().float()
|
||||
# Convert to float32 for compatibility
|
||||
frames = frames.to(torch.float32)
|
||||
|
||||
# Crop padding if this is a LongCat refinement
|
||||
if hasattr(batch, 'num_cond_frames_added') and hasattr(
|
||||
|
||||
@@ -314,6 +314,27 @@ class DenoisingStage(PipelineStage):
|
||||
t_expand = timestep.repeat(latent_model_input.shape[0], 1)
|
||||
else:
|
||||
t_expand = t.repeat(latent_model_input.shape[0])
|
||||
t_expand = t_expand.to(get_local_torch_device())
|
||||
|
||||
use_meanflow = getattr(self.transformer.config, "use_meanflow",
|
||||
False)
|
||||
if use_meanflow:
|
||||
if i == len(timesteps) - 1:
|
||||
timesteps_r = torch.tensor(
|
||||
[0.0], device=get_local_torch_device())
|
||||
else:
|
||||
timesteps_r = timesteps[i + 1]
|
||||
timesteps_r = timesteps_r.repeat(
|
||||
latent_model_input.shape[0])
|
||||
else:
|
||||
timesteps_r = None
|
||||
|
||||
timesteps_r_kwarg = self.prepare_extra_func_kwargs(
|
||||
self.transformer.forward,
|
||||
{
|
||||
"timestep_r": timesteps_r,
|
||||
},
|
||||
)
|
||||
|
||||
latent_model_input = self.scheduler.scale_model_input(
|
||||
latent_model_input, t)
|
||||
@@ -406,6 +427,7 @@ class DenoisingStage(PipelineStage):
|
||||
**image_kwargs,
|
||||
**pos_cond_kwargs,
|
||||
**action_kwargs,
|
||||
**timesteps_r_kwarg,
|
||||
)
|
||||
|
||||
if batch.do_classifier_free_guidance:
|
||||
@@ -423,6 +445,7 @@ class DenoisingStage(PipelineStage):
|
||||
**image_kwargs,
|
||||
**neg_cond_kwargs,
|
||||
**action_kwargs,
|
||||
**timesteps_r_kwarg,
|
||||
)
|
||||
|
||||
noise_pred_text = noise_pred
|
||||
@@ -487,7 +510,7 @@ class DenoisingStage(PipelineStage):
|
||||
mgr2.release_all()
|
||||
|
||||
# Save STA mask search results if needed
|
||||
if st_attn_available and self.attn_backend == SlidingTileAttentionBackend and fastvideo_args.STA_mode == STA_Mode.STA_SEARCHING:
|
||||
if st_attn_available and self.attn_backend == SlidingTileAttentionBackend and fastvideo_args.pipeline_config.STA_mode == STA_Mode.STA_SEARCHING:
|
||||
self.save_sta_search_results(batch)
|
||||
|
||||
# deallocate transformer if on mps
|
||||
@@ -580,8 +603,8 @@ class DenoisingStage(PipelineStage):
|
||||
"""
|
||||
# TODO(kevin): STA mask search, currently only support Wan2.1 with 69x768x1280
|
||||
from fastvideo.attention.backends.STA_configuration import configure_sta
|
||||
STA_mode = fastvideo_args.STA_mode
|
||||
skip_time_steps = fastvideo_args.skip_time_steps
|
||||
STA_mode = fastvideo_args.pipeline_config.STA_mode
|
||||
skip_time_steps = fastvideo_args.pipeline_config.skip_time_steps
|
||||
if batch.timesteps is None:
|
||||
raise ValueError("Timesteps must be provided")
|
||||
timesteps_num = batch.timesteps.shape[0]
|
||||
@@ -1057,7 +1080,6 @@ class Cosmos25DenoisingStage(CosmosDenoisingStage):
|
||||
"latents must be provided for Cosmos25DenoisingStage")
|
||||
guidance_scale = batch.guidance_scale
|
||||
|
||||
# Use timesteps prepared by Cosmos25TimestepPreparationStage when available.
|
||||
if batch.timesteps is None:
|
||||
self.scheduler.set_timesteps(batch.num_inference_steps,
|
||||
device=latents.device)
|
||||
@@ -1065,20 +1087,45 @@ class Cosmos25DenoisingStage(CosmosDenoisingStage):
|
||||
else:
|
||||
timesteps = batch.timesteps.to(latents.device)
|
||||
|
||||
# Match official behavior: pass fps as a tensor.
|
||||
fps_val = batch.fps if isinstance(batch.fps, int | float) else 24
|
||||
fps_tensor = torch.tensor([fps_val],
|
||||
device=latents.device,
|
||||
dtype=target_dtype)
|
||||
cfg = fastvideo_args.pipeline_config
|
||||
|
||||
if batch.fps is None:
|
||||
gen = batch.generator
|
||||
if isinstance(gen, list) and len(gen) > 0:
|
||||
gen = gen[0]
|
||||
fps_tensor = torch.randint(
|
||||
16,
|
||||
32,
|
||||
(1, ),
|
||||
generator=gen if isinstance(gen, torch.Generator) else None,
|
||||
device=latents.device,
|
||||
).float().to(dtype=target_dtype)
|
||||
else:
|
||||
fps_val = batch.fps
|
||||
fps_tensor = torch.tensor(
|
||||
[fps_val],
|
||||
device=latents.device,
|
||||
dtype=target_dtype,
|
||||
)
|
||||
|
||||
# Cosmos2.5 denoises a 4D latent (C,T,H,W) and the scheduler.step expects (B,C,T,H,W).
|
||||
latents_4d = latents[0]
|
||||
|
||||
# Masks from latent prep stage
|
||||
condition_mask = batch.cond_mask.to(target_dtype) if hasattr(
|
||||
batch, 'cond_mask') else None
|
||||
padding_mask = batch.padding_mask.to(target_dtype) if hasattr(
|
||||
batch, 'padding_mask') else None
|
||||
# Masks are optional for T2W.
|
||||
cond_mask = getattr(batch, "cond_mask", None)
|
||||
condition_mask = cond_mask.to(target_dtype) if isinstance(
|
||||
cond_mask, torch.Tensor) else None
|
||||
pad_mask = getattr(batch, "padding_mask", None)
|
||||
padding_mask = pad_mask.to(target_dtype) if isinstance(
|
||||
pad_mask, torch.Tensor) else None
|
||||
|
||||
# Conditioning fields are attached by latent preparation stage.
|
||||
conditioning_latents = getattr(batch, "conditioning_latents", None)
|
||||
cond_indicator = getattr(batch, "cond_indicator", None)
|
||||
# Infer whether this is a conditioned run (V2W/I2W) purely from the presence
|
||||
# of conditioning latents. Avoid carrying explicit mode flags on the batch.
|
||||
is_conditioned = (conditioning_latents is not None)
|
||||
|
||||
init_noise_4d = latents_4d.clone()
|
||||
if condition_mask is None:
|
||||
_, t, h, w = latents_4d.shape
|
||||
condition_mask = torch.zeros(1,
|
||||
@@ -1090,23 +1137,58 @@ class Cosmos25DenoisingStage(CosmosDenoisingStage):
|
||||
dtype=target_dtype)
|
||||
if padding_mask is None:
|
||||
_, _, h, w = latents_4d.shape
|
||||
padding_mask = torch.ones(1,
|
||||
1,
|
||||
h,
|
||||
w,
|
||||
device=latents.device,
|
||||
dtype=target_dtype)
|
||||
padding_default = 0.0 if is_conditioned else 1.0
|
||||
padding_mask = torch.full(
|
||||
(1, 1, h, w),
|
||||
float(padding_default),
|
||||
device=latents.device,
|
||||
dtype=target_dtype,
|
||||
)
|
||||
|
||||
# Cosmos2.5 timestep scaling (see compare_pipelines.py): t * 0.001
|
||||
timestep_scale = 0.001
|
||||
|
||||
state_dtype = torch.float32
|
||||
|
||||
conditional_frame_timestep = 0.1
|
||||
latents_4d = latents_4d.to(state_dtype)
|
||||
init_noise_4d = init_noise_4d.to(state_dtype)
|
||||
|
||||
clamp_every_step = bool(getattr(cfg, "cosmos25_clamp_every_step",
|
||||
True)) if is_conditioned else False
|
||||
|
||||
with self.progress_bar(total=len(timesteps)) as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
t_val = float(t)
|
||||
timestep_val = t_val * timestep_scale
|
||||
timestep = torch.tensor([[timestep_val]],
|
||||
device=latents.device,
|
||||
dtype=target_dtype)
|
||||
if is_conditioned:
|
||||
t_frames = int(latents_4d.shape[1])
|
||||
timestep = torch.full(
|
||||
(1, t_frames),
|
||||
float(t_val * timestep_scale),
|
||||
device=latents.device,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
if cond_indicator is not None and t_frames > 0:
|
||||
cond_t = cond_indicator[0, 0, :t_frames, 0, 0]
|
||||
cond_mask_t = (cond_t > 0.5)
|
||||
if bool(cond_mask_t.any().item()):
|
||||
timestep[0, cond_mask_t] = float(
|
||||
conditional_frame_timestep)
|
||||
else:
|
||||
timestep_val = t_val * timestep_scale
|
||||
timestep = torch.tensor(
|
||||
[[float(timestep_val)]],
|
||||
device=latents.device,
|
||||
dtype=target_dtype,
|
||||
)
|
||||
|
||||
# Conditioned runs: replace x_t with GT x0 on the conditioned frames.
|
||||
if (is_conditioned and cond_indicator is not None
|
||||
and conditioning_latents is not None
|
||||
and (clamp_every_step or i == 0)):
|
||||
cond_ind_4d = cond_indicator[0].to(state_dtype)
|
||||
gt_x0 = conditioning_latents[0].to(state_dtype)
|
||||
latents_4d = gt_x0 * cond_ind_4d + latents_4d * (
|
||||
1 - cond_ind_4d)
|
||||
|
||||
model_hidden_states = latents_4d.unsqueeze(0)
|
||||
|
||||
@@ -1140,10 +1222,22 @@ class Cosmos25DenoisingStage(CosmosDenoisingStage):
|
||||
padding_mask=padding_mask,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
v = uncond_v + guidance_scale * (cond_v - uncond_v)
|
||||
if is_conditioned:
|
||||
v = cond_v + guidance_scale * (cond_v - uncond_v)
|
||||
else:
|
||||
v = uncond_v + guidance_scale * (cond_v - uncond_v)
|
||||
else:
|
||||
v = cond_v
|
||||
|
||||
# Conditioned runs: replace velocity on conditioned frames with GT velocity.
|
||||
if (is_conditioned and cond_indicator is not None
|
||||
and conditioning_latents is not None):
|
||||
cond_ind_4d = cond_indicator[0].to(state_dtype)
|
||||
gt_x0 = conditioning_latents[0].to(state_dtype)
|
||||
gt_v = init_noise_4d.to(state_dtype) - gt_x0
|
||||
v = cond_ind_4d * gt_v + (1 -
|
||||
cond_ind_4d) * v.to(state_dtype)
|
||||
|
||||
prev = self.scheduler.step(v.unsqueeze(0),
|
||||
t,
|
||||
latents_4d.unsqueeze(0),
|
||||
@@ -1153,10 +1247,79 @@ class Cosmos25DenoisingStage(CosmosDenoisingStage):
|
||||
|
||||
progress_bar.update()
|
||||
|
||||
batch.latents = latents_4d.unsqueeze(0)
|
||||
batch.latents = latents_4d.to(target_dtype).unsqueeze(0)
|
||||
return batch
|
||||
|
||||
|
||||
class Cosmos25T2WDenoisingStage(Cosmos25DenoisingStage):
|
||||
"""Cosmos 2.5 Text2World denoising stage."""
|
||||
|
||||
_CONDITIONING_FIELDS = (
|
||||
"conditioning_latents",
|
||||
"cond_indicator",
|
||||
"uncond_indicator",
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
for name in self._CONDITIONING_FIELDS:
|
||||
if hasattr(batch, name):
|
||||
setattr(batch, name, None)
|
||||
return super().forward(batch, fastvideo_args)
|
||||
|
||||
|
||||
class Cosmos25V2WDenoisingStage(Cosmos25DenoisingStage):
|
||||
"""Cosmos 2.5 Video2World denoising stage."""
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
return super().forward(batch, fastvideo_args)
|
||||
|
||||
|
||||
class Cosmos25AutoDenoisingStage(PipelineStage):
|
||||
"""Route Cosmos 2.5 denoising to T2W vs V2W/I2W."""
|
||||
|
||||
def __init__(self, transformer, scheduler) -> None:
|
||||
super().__init__()
|
||||
self._t2w = Cosmos25T2WDenoisingStage(transformer=transformer,
|
||||
scheduler=scheduler)
|
||||
self._v2w = Cosmos25V2WDenoisingStage(transformer=transformer,
|
||||
scheduler=scheduler)
|
||||
|
||||
def pipeline(self):
|
||||
return self._v2w.pipeline() if self._v2w.pipeline else None
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
conditioning_latents = getattr(batch, "conditioning_latents", None)
|
||||
if conditioning_latents is not None:
|
||||
return self._v2w.forward(batch, fastvideo_args)
|
||||
return self._t2w.forward(batch, fastvideo_args)
|
||||
|
||||
def verify_input(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
conditioning_latents = getattr(batch, "conditioning_latents", None)
|
||||
if conditioning_latents is not None:
|
||||
return self._v2w.verify_input(batch, fastvideo_args)
|
||||
return self._t2w.verify_input(batch, fastvideo_args)
|
||||
|
||||
def verify_output(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
conditioning_latents = getattr(batch, "conditioning_latents", None)
|
||||
if conditioning_latents is not None:
|
||||
return self._v2w.verify_output(batch, fastvideo_args)
|
||||
return self._t2w.verify_output(batch, fastvideo_args)
|
||||
|
||||
|
||||
class DmdDenoisingStage(DenoisingStage):
|
||||
"""
|
||||
Denoising stage for DMD.
|
||||
|
||||
@@ -20,6 +20,7 @@ from fastvideo.pipelines.stages.validators import VerificationResult
|
||||
from fastvideo.utils import dict_to_3d_list
|
||||
from fastvideo.models.dits.hyworld.retrieval_context import (
|
||||
generate_points_in_sphere, select_aligned_memory_frames)
|
||||
from fastvideo.models.dits.hyworld.pose import pose_to_input, compute_latent_num
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -58,7 +59,7 @@ class HYWorldDenoisingStage(DenoisingStage):
|
||||
- viewmats: torch.Tensor | None - Camera view matrices (B, T, 4, 4)
|
||||
- Ks: torch.Tensor | None - Camera intrinsics (B, T, 3, 3)
|
||||
- action: torch.Tensor | None - Action conditioning (B, T)
|
||||
- chunk_latent_frames: int - Number of frames per chunk (default: 4)
|
||||
- chunk_latent_frames: int - Number of frames per chunk (default: 16 for bidirectional model)
|
||||
These can be passed via batch.extra dict or as direct attributes.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
@@ -82,16 +83,38 @@ class HYWorldDenoisingStage(DenoisingStage):
|
||||
action = getattr(batch, "action", None) or batch.extra.get(
|
||||
"action", None)
|
||||
chunk_latent_frames = (getattr(batch, "chunk_latent_frames", None)
|
||||
or batch.extra.get("chunk_latent_frames", 4))
|
||||
or batch.extra.get("chunk_latent_frames", 16)
|
||||
) # 16 for bidirectional model
|
||||
stabilization_level = 15
|
||||
points_local = (getattr(batch, "points_local", None)
|
||||
or batch.extra.get("points_local", None))
|
||||
|
||||
# If viewmats/Ks/action not provided, convert from pose string
|
||||
if viewmats is None or Ks is None or action is None:
|
||||
raise ValueError(
|
||||
"viewmats, Ks, and action are required for HYWorld denoising. "
|
||||
"Please provide them in batch.extra['viewmats'], batch.extra['Ks'], "
|
||||
"and batch.extra['action']")
|
||||
pose = getattr(batch, "pose", None) or batch.extra.get("pose", None)
|
||||
if pose is None:
|
||||
raise ValueError(
|
||||
"Either pose string or (viewmats, Ks, action) must be provided. "
|
||||
"Provide pose in batch.pose or batch.extra['pose'], or "
|
||||
"viewmats/Ks/action in batch.extra.")
|
||||
|
||||
# Get num_frames from batch
|
||||
num_frames = batch.num_frames
|
||||
if isinstance(num_frames, list):
|
||||
num_frames = num_frames[0]
|
||||
|
||||
# Calculate number of latents and convert pose to tensors
|
||||
latent_num = compute_latent_num(num_frames)
|
||||
viewmats, Ks, action = pose_to_input(pose, latent_num)
|
||||
|
||||
# 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)
|
||||
|
||||
logger.info(
|
||||
"Converted pose '%s' to viewmats, Ks, and action for %d latent frames",
|
||||
pose, latent_num)
|
||||
|
||||
# Prepare extra step kwargs for scheduler
|
||||
extra_step_kwargs = self.prepare_extra_func_kwargs(
|
||||
@@ -266,12 +289,9 @@ class HYWorldDenoisingStage(DenoisingStage):
|
||||
latents_concat = torch.concat(
|
||||
[latent_model_input, cond_latents_input], dim=1)
|
||||
|
||||
# Note: Unlike some other pipelines, HYWorld runs CFG sequentially (two passes)
|
||||
# rather than batching pos/neg together, following the original implementation
|
||||
latents_concat = self.scheduler.scale_model_input(
|
||||
latents_concat, t)
|
||||
|
||||
# Keep batch size 1 for sequential CFG
|
||||
t_expand_txt = t.unsqueeze(0)
|
||||
t_expand = timestep_input
|
||||
viewmats_input = viewmats_input.to(device)
|
||||
@@ -287,7 +307,6 @@ class HYWorldDenoisingStage(DenoisingStage):
|
||||
batch.is_cfg_negative = False
|
||||
|
||||
# Prepare transformer kwargs with HYWorld-specific inputs
|
||||
# Note: batch size 1 for sequential CFG (matching original HY-WorldPlay)
|
||||
transformer_kwargs = {
|
||||
**image_kwargs,
|
||||
"timestep": t_expand,
|
||||
@@ -397,34 +416,26 @@ class HYWorldDenoisingStage(DenoisingStage):
|
||||
[V.is_tensor, V.with_dims(5)])
|
||||
result.add_check("prompt_embeds", batch.prompt_embeds, V.list_not_empty)
|
||||
|
||||
# Check for HYWorld-specific inputs
|
||||
# Check for HYWorld-specific inputs - either pose string OR
|
||||
# (viewmats, Ks, action) must be provided
|
||||
pose = getattr(batch, "pose", None) or batch.extra.get("pose", None)
|
||||
viewmats = getattr(batch, "viewmats", None) or batch.extra.get(
|
||||
"viewmats", None)
|
||||
Ks = getattr(batch, "Ks", None) or batch.extra.get("Ks", None)
|
||||
action = getattr(batch, "action", None) or batch.extra.get(
|
||||
"action", None)
|
||||
|
||||
if viewmats is None:
|
||||
result.add_failure(
|
||||
"viewmats",
|
||||
"viewmats must be provided in batch.extra['viewmats'] or as batch.viewmats",
|
||||
)
|
||||
else:
|
||||
result.add_check("viewmats", viewmats, V.is_tensor)
|
||||
has_pose = pose is not None
|
||||
has_direct_inputs = (viewmats is not None and Ks is not None
|
||||
and action is not None)
|
||||
|
||||
if Ks is None:
|
||||
if not has_pose and not has_direct_inputs:
|
||||
result.add_failure(
|
||||
"Ks", "Ks must be provided in batch.extra['Ks'] or as batch.Ks")
|
||||
else:
|
||||
result.add_check("Ks", Ks, V.is_tensor)
|
||||
|
||||
if action is None:
|
||||
result.add_failure(
|
||||
"action",
|
||||
"action must be provided in batch.extra['action'] or as batch.action",
|
||||
"pose_or_camera_inputs",
|
||||
"Either pose string OR (viewmats, Ks, action) must be provided. "
|
||||
"Provide pose in batch.pose or batch.extra['pose'], or "
|
||||
"viewmats/Ks/action in batch.extra.",
|
||||
)
|
||||
else:
|
||||
result.add_check("action", action, V.is_tensor)
|
||||
|
||||
result.add_check("num_inference_steps", batch.num_inference_steps,
|
||||
V.positive_int)
|
||||
|
||||
@@ -238,7 +238,8 @@ class InputValidationStage(PipelineStage):
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""Verify input validation stage inputs."""
|
||||
result = VerificationResult()
|
||||
result.add_check("seed", batch.seed, [V.not_none, V.positive_int])
|
||||
# Cosmos-Predict2.5 default seed is 0; allow non-negative seeds here.
|
||||
result.add_check("seed", batch.seed, [V.not_none, V.non_negative_int])
|
||||
result.add_check("num_videos_per_prompt", batch.num_videos_per_prompt,
|
||||
V.positive_int)
|
||||
result.add_check(
|
||||
|
||||
@@ -3,6 +3,7 @@
|
||||
Latent preparation stage for diffusion pipelines.
|
||||
"""
|
||||
|
||||
import os
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
@@ -450,10 +451,6 @@ class Cosmos25LatentPreparationStage(CosmosLatentPreparationStage):
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
# Differences vs `CosmosLatentPreparationStage`: channel convention, seed usage,
|
||||
# and `padding_mask` for concat_padding_mask=True.
|
||||
|
||||
# Determine batch size
|
||||
if isinstance(batch.prompt, list):
|
||||
batch_size = len(batch.prompt)
|
||||
elif batch.prompt is not None:
|
||||
@@ -463,8 +460,6 @@ class Cosmos25LatentPreparationStage(CosmosLatentPreparationStage):
|
||||
|
||||
batch_size *= batch.num_videos_per_prompt
|
||||
|
||||
# Match `compare_pipelines.py`: initialize noise in fp32, then run the
|
||||
# denoising computation in bf16.
|
||||
dtype = torch.float32
|
||||
device = get_local_torch_device()
|
||||
generator = batch.generator
|
||||
@@ -483,7 +478,6 @@ class Cosmos25LatentPreparationStage(CosmosLatentPreparationStage):
|
||||
latent_width = width // vae_scale_factor_spatial
|
||||
num_latent_frames = (num_frames - 1) // vae_scale_factor_temporal + 1
|
||||
|
||||
# Cosmos 2.5 convention: transformer in_channels == latent channels
|
||||
num_channels_latents = self.transformer.config.in_channels
|
||||
|
||||
shape = (batch_size, num_channels_latents, num_latent_frames,
|
||||
@@ -492,16 +486,16 @@ class Cosmos25LatentPreparationStage(CosmosLatentPreparationStage):
|
||||
init_latents = None
|
||||
conditioning_latents = None
|
||||
video = None
|
||||
image_conditioning = False
|
||||
|
||||
if hasattr(batch, 'video') and batch.video is not None:
|
||||
video = batch.video
|
||||
elif hasattr(batch, 'pil_image') and batch.pil_image is not None:
|
||||
if hasattr(batch, 'pil_image') and batch.pil_image is not None:
|
||||
vae_scale_factor_spatial = 8
|
||||
image_processor = ImageProcessor(
|
||||
vae_scale_factor=vae_scale_factor_spatial)
|
||||
processed_image = image_processor.preprocess(
|
||||
processed_image = image_processor._preprocess_cosmos25(
|
||||
batch.pil_image, height, width)
|
||||
video = processed_image.unsqueeze(2)
|
||||
image_conditioning = True
|
||||
video = video.to(device=device, dtype=torch.bfloat16)
|
||||
elif hasattr(
|
||||
batch,
|
||||
@@ -509,25 +503,77 @@ class Cosmos25LatentPreparationStage(CosmosLatentPreparationStage):
|
||||
if isinstance(batch.preprocessed_image, torch.Tensor):
|
||||
if batch.preprocessed_image.dim() == 4:
|
||||
video = batch.preprocessed_image.unsqueeze(2)
|
||||
image_conditioning = True
|
||||
elif batch.preprocessed_image.dim() == 5:
|
||||
video = batch.preprocessed_image
|
||||
elif hasattr(batch, "video_latent") and isinstance(
|
||||
batch.video_latent, torch.Tensor):
|
||||
if batch.video_latent.dim() == 5:
|
||||
video = batch.video_latent
|
||||
else:
|
||||
logger.info(
|
||||
"CosmosLatentPreparationStage - No video input sources found")
|
||||
|
||||
if video is not None:
|
||||
num_cond_frames = video.size(2)
|
||||
if num_cond_frames >= num_frames:
|
||||
# Avoid missing CUDA fallback kernels in some builds by using fp16 contiguous input.
|
||||
if isinstance(video, torch.Tensor):
|
||||
video = video.to(device=device,
|
||||
dtype=torch.float16).contiguous()
|
||||
num_video_frames = int(video.size(2))
|
||||
|
||||
num_cond_frames_eff: int | None = None
|
||||
try:
|
||||
ncf = getattr(batch, "num_cond_frames", None)
|
||||
if isinstance(ncf, int) and ncf > 0:
|
||||
num_cond_frames_eff = min(int(ncf), num_video_frames)
|
||||
except Exception:
|
||||
num_cond_frames_eff = None
|
||||
if num_cond_frames_eff is None:
|
||||
num_cond_frames_eff = 1 if (
|
||||
not image_conditioning) else num_video_frames
|
||||
|
||||
if num_video_frames >= num_cond_frames_eff:
|
||||
cond_video = video[:, :, -num_cond_frames_eff:]
|
||||
else:
|
||||
pad = video[:, :,
|
||||
-1:].repeat(1, 1,
|
||||
num_cond_frames_eff - num_video_frames,
|
||||
1, 1)
|
||||
cond_video = torch.cat([video, pad], dim=2)
|
||||
|
||||
if num_cond_frames_eff >= num_frames:
|
||||
cond_video = cond_video[:, :, -num_frames:]
|
||||
num_cond_latent_frames = (num_frames -
|
||||
1) // vae_scale_factor_temporal + 1
|
||||
video = video[:, :, -num_frames:]
|
||||
video = cond_video
|
||||
else:
|
||||
num_cond_latent_frames = (num_cond_frames -
|
||||
num_cond_latent_frames = (num_cond_frames_eff -
|
||||
1) // vae_scale_factor_temporal + 1
|
||||
num_padding_frames = num_frames - num_cond_frames
|
||||
last_frame = video[:, :, -1:]
|
||||
num_padding_frames = num_frames - num_cond_frames_eff
|
||||
last_frame = cond_video[:, :, -1:]
|
||||
padding = last_frame.repeat(1, 1, num_padding_frames, 1, 1)
|
||||
video = torch.cat([video, padding], dim=2)
|
||||
if image_conditioning:
|
||||
padding = cond_video.new_full(
|
||||
(cond_video.size(0), cond_video.size(1),
|
||||
num_padding_frames, cond_video.size(3),
|
||||
cond_video.size(4)),
|
||||
-1.0,
|
||||
)
|
||||
video = torch.cat([cond_video, padding], dim=2)
|
||||
|
||||
if os.environ.get("FASTVIDEO_COSMOS25_LOG_KNOBS",
|
||||
"0") in ("1", "true", "True"):
|
||||
logger.info(
|
||||
"[Cosmos2.5 latent_prep] using conditioning input: image=%s video_latent=%s video_path=%s | "
|
||||
"num_video_frames=%s num_cond_frames=%s num_frames=%s",
|
||||
bool(image_conditioning),
|
||||
isinstance(getattr(batch, "video_latent", None),
|
||||
torch.Tensor),
|
||||
getattr(batch, "video_path", None),
|
||||
num_video_frames,
|
||||
num_cond_frames_eff,
|
||||
num_frames,
|
||||
)
|
||||
|
||||
if self.vae is not None:
|
||||
self.vae = self.vae.to(device)
|
||||
@@ -588,18 +634,14 @@ class Cosmos25LatentPreparationStage(CosmosLatentPreparationStage):
|
||||
|
||||
if latents is None:
|
||||
seed = int(batch.seed if batch.seed is not None else 0)
|
||||
# Use arch-invariant RNG to match Cosmos2.5 reference sampling.
|
||||
latents_fp32 = self._arch_invariant_randn(shape,
|
||||
seed=seed,
|
||||
device=device,
|
||||
dtype=torch.float32)
|
||||
latents = latents_fp32.to(torch.bfloat16)
|
||||
else:
|
||||
# If latents are supplied, keep compute dtype consistent with Cosmos sampling.
|
||||
latents = latents.to(device=device, dtype=torch.bfloat16)
|
||||
|
||||
# Cosmos2.5 starts from unit Gaussian noise (no extra sigma_max scaling).
|
||||
|
||||
padding_shape = (batch_size, 1, num_latent_frames, latent_height,
|
||||
latent_width)
|
||||
ones_padding = latents.new_ones(padding_shape)
|
||||
@@ -618,9 +660,7 @@ class Cosmos25LatentPreparationStage(CosmosLatentPreparationStage):
|
||||
uncond_mask = uncond_indicator * ones_padding + (
|
||||
1 - uncond_indicator) * zeros_padding
|
||||
|
||||
# Cosmos 2.5 requires a spatial padding mask when concat_padding_mask=True
|
||||
padding_mask = latents.new_ones(batch_size, 1, latent_height,
|
||||
latent_width)
|
||||
padding_mask = latents.new_zeros(batch_size, 1, height, width)
|
||||
|
||||
batch.latents = latents
|
||||
batch.raw_latent_shape = latents.shape
|
||||
@@ -681,3 +721,127 @@ class Cosmos25LatentPreparationStage(CosmosLatentPreparationStage):
|
||||
[V.is_tensor, V.with_dims(5)])
|
||||
result.add_check("raw_latent_shape", batch.raw_latent_shape, V.is_tuple)
|
||||
return result
|
||||
|
||||
|
||||
class Cosmos25T2WLatentPreparationStage(PipelineStage):
|
||||
"""Cosmos 2.5 Text2World latent preparation."""
|
||||
|
||||
def __init__(self, scheduler, transformer) -> None:
|
||||
super().__init__()
|
||||
self.scheduler = scheduler
|
||||
self.transformer = transformer
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
if isinstance(batch.prompt, list):
|
||||
batch_size = len(batch.prompt)
|
||||
elif batch.prompt is not None:
|
||||
batch_size = 1
|
||||
else:
|
||||
batch_size = batch.prompt_embeds[0].shape[0]
|
||||
batch_size *= batch.num_videos_per_prompt
|
||||
|
||||
device = get_local_torch_device()
|
||||
latents = batch.latents
|
||||
num_frames = batch.num_frames
|
||||
height = batch.height
|
||||
width = batch.width
|
||||
if height is None or width is None:
|
||||
raise ValueError("Height and width must be provided")
|
||||
|
||||
vae_scale_factor_spatial = 8
|
||||
vae_scale_factor_temporal = 4
|
||||
latent_height = height // vae_scale_factor_spatial
|
||||
latent_width = width // vae_scale_factor_spatial
|
||||
num_latent_frames = (num_frames - 1) // vae_scale_factor_temporal + 1
|
||||
|
||||
# Cosmos 2.5 convention: transformer in_channels == latent channels
|
||||
num_channels_latents = self.transformer.config.in_channels
|
||||
shape = (batch_size, num_channels_latents, num_latent_frames,
|
||||
latent_height, latent_width)
|
||||
|
||||
if latents is None:
|
||||
seed = int(batch.seed if batch.seed is not None else 0)
|
||||
latents_fp32 = Cosmos25LatentPreparationStage._arch_invariant_randn(
|
||||
shape,
|
||||
seed=seed,
|
||||
device=device,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
latents = latents_fp32.to(torch.bfloat16)
|
||||
else:
|
||||
latents = latents.to(device=device, dtype=torch.bfloat16)
|
||||
|
||||
batch.latents = latents
|
||||
batch.raw_latent_shape = latents.shape
|
||||
|
||||
batch.conditioning_latents = None
|
||||
batch.cond_indicator = None
|
||||
batch.uncond_indicator = None
|
||||
batch.cond_mask = None
|
||||
batch.uncond_mask = None
|
||||
if hasattr(batch, "padding_mask"):
|
||||
batch.padding_mask = None
|
||||
|
||||
return batch
|
||||
|
||||
def verify_input(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
return Cosmos25LatentPreparationStage.verify_input(
|
||||
self, batch, fastvideo_args) # type: ignore[misc]
|
||||
|
||||
def verify_output(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
return Cosmos25LatentPreparationStage.verify_output(
|
||||
self, batch, fastvideo_args) # type: ignore[misc]
|
||||
|
||||
|
||||
class Cosmos25V2WLatentPreparationStage(Cosmos25LatentPreparationStage):
|
||||
"""Cosmos 2.5 V2W/I2W latent preparation stage (conditioning-aware)."""
|
||||
|
||||
|
||||
class Cosmos25AutoLatentPreparationStage(PipelineStage):
|
||||
"""Route Cosmos 2.5 latent prep to T2W vs V2W/I2W."""
|
||||
|
||||
def __init__(self, scheduler, transformer, vae=None) -> None:
|
||||
super().__init__()
|
||||
self._t2w = Cosmos25T2WLatentPreparationStage(
|
||||
scheduler=scheduler,
|
||||
transformer=transformer,
|
||||
)
|
||||
self._v2w = Cosmos25V2WLatentPreparationStage(
|
||||
scheduler=scheduler,
|
||||
transformer=transformer,
|
||||
vae=vae,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _has_conditioning_input(batch: ForwardBatch) -> bool:
|
||||
return (getattr(batch, "pil_image", None) is not None
|
||||
or getattr(batch, "preprocessed_image", None) is not None
|
||||
or bool(getattr(batch, "video_path", None)) or isinstance(
|
||||
getattr(batch, "video_latent", None), torch.Tensor))
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
if self._has_conditioning_input(batch):
|
||||
return self._v2w.forward(batch, fastvideo_args)
|
||||
return self._t2w.forward(batch, fastvideo_args)
|
||||
|
||||
def verify_input(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
if self._has_conditioning_input(batch):
|
||||
return self._v2w.verify_input(batch, fastvideo_args)
|
||||
return self._t2w.verify_input(batch, fastvideo_args)
|
||||
|
||||
def verify_output(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
if self._has_conditioning_input(batch):
|
||||
return self._v2w.verify_output(batch, fastvideo_args)
|
||||
return self._t2w.verify_output(batch, fastvideo_args)
|
||||
|
||||
@@ -0,0 +1,186 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
LTX2-specific text encoding stage with sequence parallelism broadcast support.
|
||||
|
||||
When running with sequence parallelism (SP), the Gemma text encoder is only
|
||||
executed on rank 0, and the embeddings are broadcast to all other ranks.
|
||||
This avoids I/O contention from all ranks loading the Gemma model simultaneously.
|
||||
"""
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.distributed.parallel_state import get_sp_group
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.text_encoding import TextEncodingStage
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class LTX2TextEncodingStage(TextEncodingStage):
|
||||
"""
|
||||
LTX2 text encoding stage with sequence parallelism support.
|
||||
|
||||
When SP is enabled (sp_world_size > 1), only rank 0 runs the text encoder
|
||||
and broadcasts embeddings to other ranks. This avoids I/O contention from
|
||||
all ranks loading the Gemma model simultaneously, which can cause text
|
||||
encoding to take 100+ seconds instead of ~5 seconds.
|
||||
"""
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
sp_group = get_sp_group()
|
||||
sp_world_size = sp_group.world_size
|
||||
sp_rank = sp_group.rank_in_group
|
||||
|
||||
# Single GPU or no SP: use parent implementation
|
||||
if sp_world_size <= 1:
|
||||
return super().forward(batch, fastvideo_args)
|
||||
|
||||
# SP enabled: only rank 0 encodes, then broadcasts
|
||||
if sp_rank == 0:
|
||||
logger.info(
|
||||
"[LTX2TextEncodingStage] SP rank 0: running text encoding")
|
||||
# Run encoding on rank 0
|
||||
result_batch = super().forward(batch, fastvideo_args)
|
||||
|
||||
# Build broadcast dict from batch
|
||||
broadcast_dict = self._build_broadcast_dict(result_batch)
|
||||
|
||||
# Broadcast to other ranks
|
||||
logger.info(
|
||||
"[LTX2TextEncodingStage] SP rank 0: broadcasting %d tensors",
|
||||
len(broadcast_dict))
|
||||
sp_group.broadcast_tensor_dict(broadcast_dict, src=0)
|
||||
|
||||
return result_batch
|
||||
else:
|
||||
logger.info(
|
||||
"[LTX2TextEncodingStage] SP rank %d: receiving broadcast",
|
||||
sp_rank)
|
||||
# Other ranks: receive broadcast and populate batch
|
||||
broadcast_dict = sp_group.broadcast_tensor_dict(None, src=0)
|
||||
|
||||
# Unpack into batch
|
||||
self._unpack_broadcast_to_batch(batch, broadcast_dict)
|
||||
|
||||
logger.info(
|
||||
"[LTX2TextEncodingStage] SP rank %d: received %d prompt embeds",
|
||||
sp_rank, len(batch.prompt_embeds))
|
||||
|
||||
return batch
|
||||
|
||||
def _build_broadcast_dict(self,
|
||||
batch: ForwardBatch) -> dict[str, torch.Tensor]:
|
||||
"""Build dict of tensors to broadcast from rank 0."""
|
||||
d: dict[str, torch.Tensor] = {}
|
||||
|
||||
# Use a dummy tensor for metadata since broadcast_tensor_dict expects tensors
|
||||
device = batch.prompt_embeds[0].device if batch.prompt_embeds else "cuda"
|
||||
|
||||
# Prompt embeddings
|
||||
num_prompt_embeds = len(batch.prompt_embeds)
|
||||
d["_num_prompt_embeds"] = torch.tensor([num_prompt_embeds],
|
||||
device=device)
|
||||
for i, pe in enumerate(batch.prompt_embeds):
|
||||
d[f"prompt_embed_{i}"] = pe
|
||||
|
||||
# Prompt attention masks
|
||||
has_prompt_masks = (batch.prompt_attention_mask is not None
|
||||
and len(batch.prompt_attention_mask) > 0)
|
||||
d["_has_prompt_masks"] = torch.tensor([1 if has_prompt_masks else 0],
|
||||
device=device)
|
||||
if has_prompt_masks:
|
||||
for i, pm in enumerate(batch.prompt_attention_mask):
|
||||
d[f"prompt_mask_{i}"] = pm
|
||||
|
||||
# Negative embeddings (CFG)
|
||||
has_neg_embeds = (batch.negative_prompt_embeds is not None
|
||||
and len(batch.negative_prompt_embeds) > 0)
|
||||
d["_has_neg_embeds"] = torch.tensor([1 if has_neg_embeds else 0],
|
||||
device=device)
|
||||
if has_neg_embeds:
|
||||
d["_num_neg_embeds"] = torch.tensor(
|
||||
[len(batch.negative_prompt_embeds)], device=device)
|
||||
for i, ne in enumerate(batch.negative_prompt_embeds):
|
||||
d[f"neg_embed_{i}"] = ne
|
||||
|
||||
# Negative attention masks
|
||||
has_neg_masks = (batch.negative_attention_mask is not None
|
||||
and len(batch.negative_attention_mask) > 0)
|
||||
d["_has_neg_masks"] = torch.tensor([1 if has_neg_masks else 0],
|
||||
device=device)
|
||||
if has_neg_masks:
|
||||
for i, nm in enumerate(batch.negative_attention_mask):
|
||||
d[f"neg_mask_{i}"] = nm
|
||||
|
||||
# LTX2 audio embeddings
|
||||
has_audio_embeds = "ltx2_audio_prompt_embeds" in batch.extra
|
||||
d["_has_audio_embeds"] = torch.tensor([1 if has_audio_embeds else 0],
|
||||
device=device)
|
||||
if has_audio_embeds:
|
||||
audio_embeds = batch.extra["ltx2_audio_prompt_embeds"]
|
||||
d["_num_audio_embeds"] = torch.tensor([len(audio_embeds)],
|
||||
device=device)
|
||||
for i, ae in enumerate(audio_embeds):
|
||||
d[f"audio_embed_{i}"] = ae
|
||||
|
||||
# LTX2 audio negative embeddings
|
||||
has_audio_neg = "ltx2_audio_negative_embeds" in batch.extra
|
||||
d["_has_audio_neg"] = torch.tensor([1 if has_audio_neg else 0],
|
||||
device=device)
|
||||
if has_audio_neg:
|
||||
audio_neg = batch.extra["ltx2_audio_negative_embeds"]
|
||||
for i, audio_neg_embed in enumerate(audio_neg):
|
||||
d[f"audio_neg_embed_{i}"] = audio_neg_embed
|
||||
|
||||
return d
|
||||
|
||||
def _unpack_broadcast_to_batch(self, batch: ForwardBatch,
|
||||
d: dict[str, torch.Tensor]) -> None:
|
||||
"""Unpack broadcast dict into batch on non-rank-0 processes."""
|
||||
# Prompt embeddings
|
||||
num_embeds = int(d["_num_prompt_embeds"].item())
|
||||
for i in range(num_embeds):
|
||||
batch.prompt_embeds.append(d[f"prompt_embed_{i}"])
|
||||
|
||||
# Prompt attention masks
|
||||
has_prompt_masks = int(d["_has_prompt_masks"].item()) == 1
|
||||
if has_prompt_masks and batch.prompt_attention_mask is not None:
|
||||
for i in range(num_embeds):
|
||||
batch.prompt_attention_mask.append(d[f"prompt_mask_{i}"])
|
||||
|
||||
# Negative embeddings (CFG)
|
||||
has_neg_embeds = int(d["_has_neg_embeds"].item()) == 1
|
||||
if has_neg_embeds:
|
||||
num_neg = int(d["_num_neg_embeds"].item())
|
||||
if batch.negative_prompt_embeds is not None:
|
||||
for i in range(num_neg):
|
||||
batch.negative_prompt_embeds.append(d[f"neg_embed_{i}"])
|
||||
|
||||
# Negative attention masks
|
||||
has_neg_masks = int(
|
||||
d.get("_has_neg_masks", torch.tensor([0])).item()) == 1
|
||||
if has_neg_masks and batch.negative_attention_mask is not None:
|
||||
for i in range(num_neg):
|
||||
batch.negative_attention_mask.append(d[f"neg_mask_{i}"])
|
||||
|
||||
# LTX2 audio embeddings
|
||||
has_audio_embeds = int(d["_has_audio_embeds"].item()) == 1
|
||||
if has_audio_embeds:
|
||||
num_audio = int(d["_num_audio_embeds"].item())
|
||||
audio_embeds = [d[f"audio_embed_{i}"] for i in range(num_audio)]
|
||||
batch.extra["ltx2_audio_prompt_embeds"] = audio_embeds
|
||||
|
||||
# LTX2 audio negative embeddings
|
||||
has_audio_neg = int(d["_has_audio_neg"].item()) == 1
|
||||
if has_audio_neg:
|
||||
# Use same count as audio embeds
|
||||
num_audio = int(d["_num_audio_embeds"].item())
|
||||
audio_neg = [d[f"audio_neg_embed_{i}"] for i in range(num_audio)]
|
||||
batch.extra["ltx2_audio_negative_embeds"] = audio_neg
|
||||
@@ -0,0 +1,431 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Denoising stage for diffusion pipelines.
|
||||
"""
|
||||
|
||||
import inspect
|
||||
import weakref
|
||||
from collections.abc import Iterable
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from einops import rearrange
|
||||
import numpy as np
|
||||
from tqdm.auto import tqdm
|
||||
|
||||
from fastvideo.attention import get_attn_backend
|
||||
from fastvideo.configs.pipelines.base import STA_Mode
|
||||
from fastvideo.distributed import (get_local_torch_device, get_world_group)
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.loader.component_loader import TransformerLoader, UpsamplerLoader
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.pipelines.stages.validators import StageValidators as V
|
||||
from fastvideo.pipelines.stages.validators import VerificationResult
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
|
||||
try:
|
||||
from fastvideo.attention.backends.sliding_tile_attn import (
|
||||
SlidingTileAttentionBackend)
|
||||
st_attn_available = True
|
||||
except ImportError:
|
||||
st_attn_available = False
|
||||
|
||||
try:
|
||||
from fastvideo.attention.backends.vmoba import VMOBAAttentionBackend
|
||||
from fastvideo.utils import is_vmoba_available
|
||||
vmoba_attn_available = is_vmoba_available()
|
||||
except ImportError:
|
||||
vmoba_attn_available = False
|
||||
|
||||
try:
|
||||
from fastvideo.attention.backends.video_sparse_attn import (
|
||||
VideoSparseAttentionBackend)
|
||||
vsa_available = True
|
||||
except ImportError:
|
||||
vsa_available = False
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class SRDenoisingStage(PipelineStage):
|
||||
"""
|
||||
Stage for running the denoising loop in SR diffusion pipelines. Used by Hunyuan15 SR pipeline.
|
||||
|
||||
This stage handles the iterative denoising process that transforms
|
||||
the initial noise into the final output.
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
transformer,
|
||||
scheduler,
|
||||
upsampler,
|
||||
pipeline=None,
|
||||
vae=None) -> None:
|
||||
super().__init__()
|
||||
self.transformer = transformer
|
||||
self.scheduler = scheduler
|
||||
self.vae = vae
|
||||
self.upsampler = upsampler
|
||||
self.pipeline = weakref.ref(pipeline) if pipeline else None
|
||||
attn_head_size = self.transformer.hidden_size // self.transformer.num_attention_heads
|
||||
self.attn_backend = get_attn_backend(
|
||||
head_size=attn_head_size,
|
||||
dtype=torch.float16, # TODO(will): hack
|
||||
supported_attention_backends=(
|
||||
AttentionBackendEnum.SLIDING_TILE_ATTN,
|
||||
AttentionBackendEnum.VIDEO_SPARSE_ATTN,
|
||||
AttentionBackendEnum.VMOBA_ATTN,
|
||||
AttentionBackendEnum.FLASH_ATTN,
|
||||
AttentionBackendEnum.TORCH_SDPA,
|
||||
AttentionBackendEnum.SAGE_ATTN_THREE) # hack
|
||||
)
|
||||
|
||||
def add_noise_to_lq(self,
|
||||
lq_latents: torch.Tensor,
|
||||
strength: float = 0.7) -> torch.Tensor:
|
||||
|
||||
def expand_dims(tensor: torch.Tensor, ndim: int):
|
||||
shape = tensor.shape + (1, ) * (ndim - tensor.ndim)
|
||||
return tensor.reshape(shape)
|
||||
|
||||
noise = torch.randn_like(lq_latents)
|
||||
timestep = torch.tensor([1000.0],
|
||||
device=get_local_torch_device()) * strength
|
||||
t = expand_dims(timestep, lq_latents.ndim)
|
||||
return (1 - t / 1000.0) * lq_latents + (t / 1000.0) * noise
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Run the denoising loop.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The batch with denoised latents.
|
||||
"""
|
||||
pipeline = self.pipeline() if self.pipeline else None
|
||||
if not fastvideo_args.model_loaded["transformer"]:
|
||||
loader = TransformerLoader()
|
||||
self.transformer = loader.load(
|
||||
fastvideo_args.model_paths["transformer"], fastvideo_args)
|
||||
if pipeline:
|
||||
pipeline.add_module("transformer", self.transformer)
|
||||
fastvideo_args.model_loaded["transformer"] = True
|
||||
|
||||
if not fastvideo_args.model_loaded["upsampler"]:
|
||||
loader = UpsamplerLoader()
|
||||
self.upsampler = loader.load(
|
||||
fastvideo_args.model_paths["upsampler"], fastvideo_args)
|
||||
if pipeline:
|
||||
pipeline.add_module("upsampler", self.upsampler)
|
||||
fastvideo_args.model_loaded["upsampler"] = True
|
||||
|
||||
# Setup precision and autocast settings
|
||||
# TODO(will): make the precision configurable for inference
|
||||
# target_dtype = PRECISION_TO_TYPE[fastvideo_args.precision]
|
||||
target_dtype = torch.bfloat16
|
||||
autocast_enabled = (target_dtype != torch.float32
|
||||
) and not fastvideo_args.disable_autocast
|
||||
|
||||
self.scheduler.set_shift(fastvideo_args.pipeline_config.flow_shift_sr)
|
||||
sigmas = np.linspace(1.0, 0.0, batch.num_inference_steps_sr + 1)[:-1]
|
||||
self.scheduler.set_timesteps(sigmas=sigmas,
|
||||
device=get_local_torch_device())
|
||||
timesteps = self.scheduler.timesteps
|
||||
logger.info("timesteps: %s", timesteps)
|
||||
num_inference_steps = len(timesteps)
|
||||
num_warmup_steps = len(
|
||||
timesteps) - num_inference_steps * self.scheduler.order
|
||||
assert num_inference_steps == batch.num_inference_steps_sr, "num_inference_steps_sr must match the number of timesteps"
|
||||
|
||||
pos_cond_kwargs = self.prepare_extra_func_kwargs(
|
||||
self.transformer.forward,
|
||||
{
|
||||
"encoder_hidden_states_2": batch.clip_embedding_pos,
|
||||
"encoder_attention_mask": batch.prompt_attention_mask,
|
||||
},
|
||||
)
|
||||
|
||||
image_embeds = batch.image_embeds
|
||||
if len(image_embeds) > 0:
|
||||
assert not torch.isnan(
|
||||
image_embeds[0]).any(), "image_embeds contains nan"
|
||||
image_embeds = [
|
||||
image_embed.to(target_dtype) for image_embed in image_embeds
|
||||
]
|
||||
|
||||
image_kwargs = self.prepare_extra_func_kwargs(
|
||||
self.transformer.forward,
|
||||
{
|
||||
"encoder_hidden_states_image": image_embeds,
|
||||
},
|
||||
)
|
||||
|
||||
# Prepare STA parameters
|
||||
if st_attn_available and self.attn_backend == SlidingTileAttentionBackend:
|
||||
self.prepare_sta_param(batch, fastvideo_args)
|
||||
|
||||
# Get latents and embeddings
|
||||
prompt_embeds = batch.prompt_embeds
|
||||
assert not torch.isnan(
|
||||
prompt_embeds[0]).any(), "prompt_embeds contains nan"
|
||||
|
||||
latents = batch.latents
|
||||
lq_latents = batch.lq_latents
|
||||
logger.info("lq_latents: %s", lq_latents.shape)
|
||||
logger.info("latents: %s", latents.shape)
|
||||
tgt_shape = latents.shape[-2:] # (h w)
|
||||
bsz = lq_latents.shape[0]
|
||||
lq_latents = rearrange(lq_latents, "b c f h w -> (b f) c h w")
|
||||
lq_latents = F.interpolate(lq_latents,
|
||||
size=tgt_shape,
|
||||
mode="bilinear",
|
||||
align_corners=False)
|
||||
lq_latents = rearrange(lq_latents, "(b f) c h w -> b c f h w", b=bsz)
|
||||
lq_latents = self.upsampler(
|
||||
lq_latents.to(dtype=torch.float32, device=get_local_torch_device()))
|
||||
lq_latents = lq_latents.to(dtype=latents.dtype)
|
||||
lq_latents = self.add_noise_to_lq(lq_latents, 0.7)
|
||||
b, c, f, h, w = lq_latents.shape
|
||||
mask_ones = torch.ones(b, 1, f, h, w).to(lq_latents.device)
|
||||
lq_cond_latents = torch.concat([lq_latents, mask_ones],
|
||||
dim=1).to(target_dtype)
|
||||
cond_latents = torch.cat(
|
||||
[batch.video_latent, torch.zeros_like(latents)],
|
||||
dim=1).to(target_dtype)
|
||||
condition = torch.concat([cond_latents, lq_cond_latents], dim=1)
|
||||
zero_lq_condition = condition.clone()
|
||||
zero_lq_condition[:, c + 1:2 * c + 1] = torch.zeros_like(lq_latents)
|
||||
zero_lq_condition[:, 2 * c + 1] = 0
|
||||
|
||||
latent_model_input = latents.to(target_dtype)
|
||||
assert latent_model_input.shape[0] == 1, "only support batch size 1"
|
||||
|
||||
# Run denoising loop
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
# Skip if interrupted
|
||||
if hasattr(self, 'interrupt') and self.interrupt:
|
||||
continue
|
||||
|
||||
latent_model_input = latents.to(target_dtype)
|
||||
if t < 1000 * 0.7:
|
||||
condition = zero_lq_condition
|
||||
|
||||
latent_model_input = torch.concat([latents, condition], dim=1)
|
||||
|
||||
assert not torch.isnan(
|
||||
latent_model_input).any(), "latent_model_input contains nan"
|
||||
t_expand = t.repeat(latent_model_input.shape[0])
|
||||
|
||||
if i == len(timesteps) - 1:
|
||||
timesteps_r = torch.tensor([0.0],
|
||||
device=get_local_torch_device())
|
||||
else:
|
||||
timesteps_r = timesteps[i + 1]
|
||||
timesteps_r = timesteps_r.repeat(latent_model_input.shape[0])
|
||||
|
||||
latent_model_input = self.scheduler.scale_model_input(
|
||||
latent_model_input, t)
|
||||
|
||||
# Prepare inputs for transformer
|
||||
guidance_expand = (
|
||||
torch.tensor(
|
||||
[fastvideo_args.pipeline_config.embedded_cfg_scale] *
|
||||
latent_model_input.shape[0],
|
||||
dtype=torch.float32,
|
||||
device=get_local_torch_device(),
|
||||
).to(target_dtype) *
|
||||
1000.0 if fastvideo_args.pipeline_config.embedded_cfg_scale
|
||||
is not None else None)
|
||||
|
||||
# Predict noise residual
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=target_dtype,
|
||||
enabled=autocast_enabled):
|
||||
if (st_attn_available
|
||||
and self.attn_backend == SlidingTileAttentionBackend
|
||||
) or (vsa_available and self.attn_backend
|
||||
== VideoSparseAttentionBackend):
|
||||
self.attn_metadata_builder_cls = self.attn_backend.get_builder_cls(
|
||||
)
|
||||
|
||||
if self.attn_metadata_builder_cls is not None:
|
||||
self.attn_metadata_builder = self.attn_metadata_builder_cls(
|
||||
)
|
||||
# TODO(will): clean this up
|
||||
attn_metadata = self.attn_metadata_builder.build( # type: ignore
|
||||
current_timestep=i, # type: ignore
|
||||
raw_latent_shape=batch.
|
||||
raw_latent_shape[2:5], # type: ignore
|
||||
patch_size=fastvideo_args.
|
||||
pipeline_config. # type: ignore
|
||||
dit_config.patch_size, # type: ignore
|
||||
STA_param=batch.STA_param, # type: ignore
|
||||
VSA_sparsity=fastvideo_args.
|
||||
VSA_sparsity, # type: ignore
|
||||
device=get_local_torch_device(),
|
||||
)
|
||||
assert attn_metadata is not None, "attn_metadata cannot be None"
|
||||
else:
|
||||
attn_metadata = None
|
||||
elif (vmoba_attn_available
|
||||
and self.attn_backend == VMOBAAttentionBackend):
|
||||
self.attn_metadata_builder_cls = self.attn_backend.get_builder_cls(
|
||||
)
|
||||
if self.attn_metadata_builder_cls is not None:
|
||||
self.attn_metadata_builder = self.attn_metadata_builder_cls(
|
||||
)
|
||||
# Prepare V-MoBA parameters from config
|
||||
moba_params = fastvideo_args.moba_config.copy()
|
||||
moba_params.update({
|
||||
"current_timestep":
|
||||
i,
|
||||
"raw_latent_shape":
|
||||
batch.raw_latent_shape[2:5],
|
||||
"patch_size":
|
||||
fastvideo_args.pipeline_config.dit_config.
|
||||
patch_size,
|
||||
"device":
|
||||
get_local_torch_device(),
|
||||
})
|
||||
attn_metadata = self.attn_metadata_builder.build(
|
||||
**moba_params)
|
||||
assert attn_metadata is not None, "attn_metadata cannot be None"
|
||||
else:
|
||||
attn_metadata = None
|
||||
else:
|
||||
attn_metadata = None
|
||||
# TODO(will): finalize the interface. vLLM uses this to
|
||||
# support torch dynamo compilation. They pass in
|
||||
# attn_metadata, vllm_config, and num_tokens. We can pass in
|
||||
# fastvideo_args or training_args, and attn_metadata.
|
||||
batch.is_cfg_negative = False
|
||||
with set_forward_context(
|
||||
current_timestep=i,
|
||||
attn_metadata=attn_metadata,
|
||||
forward_batch=batch,
|
||||
# fastvideo_args=fastvideo_args
|
||||
):
|
||||
# Run transformer
|
||||
noise_pred = self.transformer(latent_model_input,
|
||||
prompt_embeds,
|
||||
t_expand,
|
||||
guidance=guidance_expand,
|
||||
timestep_r=timesteps_r,
|
||||
**pos_cond_kwargs,
|
||||
**image_kwargs)
|
||||
|
||||
# Compute the previous noisy sample
|
||||
latents = self.scheduler.step(noise_pred,
|
||||
t,
|
||||
latents,
|
||||
return_dict=False)[0]
|
||||
|
||||
# Update progress bar
|
||||
if i == len(timesteps) - 1 or (
|
||||
(i + 1) > num_warmup_steps and
|
||||
(i + 1) % self.scheduler.order == 0
|
||||
and progress_bar is not None):
|
||||
progress_bar.update()
|
||||
|
||||
# Update batch with final latents
|
||||
batch.latents = latents
|
||||
|
||||
# Save STA mask search results if needed
|
||||
if st_attn_available and self.attn_backend == SlidingTileAttentionBackend and fastvideo_args.pipeline_config.STA_mode == STA_Mode.STA_SEARCHING:
|
||||
self.save_sta_search_results(batch)
|
||||
|
||||
# deallocate transformer if on mps
|
||||
if torch.backends.mps.is_available():
|
||||
logger.info("Memory before deallocating transformer: %s",
|
||||
torch.mps.current_allocated_memory())
|
||||
del self.transformer
|
||||
if pipeline is not None and "transformer" in pipeline.modules:
|
||||
del pipeline.modules["transformer"]
|
||||
fastvideo_args.model_loaded["transformer"] = False
|
||||
logger.info("Memory after deallocating transformer: %s",
|
||||
torch.mps.current_allocated_memory())
|
||||
|
||||
return batch
|
||||
|
||||
def prepare_extra_func_kwargs(self, func, kwargs) -> dict[str, Any]:
|
||||
"""
|
||||
Prepare extra kwargs for the scheduler step / denoise step.
|
||||
|
||||
Args:
|
||||
func: The function to prepare kwargs for.
|
||||
kwargs: The kwargs to prepare.
|
||||
|
||||
Returns:
|
||||
The prepared kwargs.
|
||||
"""
|
||||
extra_step_kwargs = {}
|
||||
for k, v in kwargs.items():
|
||||
accepts = k in set(inspect.signature(func).parameters.keys())
|
||||
if accepts:
|
||||
extra_step_kwargs[k] = v
|
||||
return extra_step_kwargs
|
||||
|
||||
def progress_bar(self,
|
||||
iterable: Iterable | None = None,
|
||||
total: int | None = None) -> tqdm:
|
||||
"""
|
||||
Create a progress bar for the denoising process.
|
||||
|
||||
Args:
|
||||
iterable: The iterable to iterate over.
|
||||
total: The total number of items.
|
||||
|
||||
Returns:
|
||||
A tqdm progress bar.
|
||||
"""
|
||||
local_rank = get_world_group().local_rank
|
||||
if local_rank == 0:
|
||||
return tqdm(iterable=iterable, total=total)
|
||||
else:
|
||||
return tqdm(iterable=iterable, total=total, disable=True)
|
||||
|
||||
def verify_input(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""Verify denoising stage inputs."""
|
||||
result = VerificationResult()
|
||||
result.add_check("timesteps", batch.timesteps,
|
||||
[V.is_tensor, V.min_dims(1)])
|
||||
result.add_check("latents", batch.latents,
|
||||
[V.is_tensor, V.with_dims(5)])
|
||||
result.add_check("prompt_embeds", batch.prompt_embeds, V.list_not_empty)
|
||||
result.add_check("image_embeds", batch.image_embeds, V.is_list)
|
||||
result.add_check("image_latent", batch.image_latent,
|
||||
V.none_or_tensor_with_dims(5))
|
||||
result.add_check("num_inference_steps", batch.num_inference_steps,
|
||||
V.positive_int)
|
||||
result.add_check("guidance_scale", batch.guidance_scale,
|
||||
V.positive_float)
|
||||
result.add_check("eta", batch.eta, V.non_negative_float)
|
||||
result.add_check("generator", batch.generator,
|
||||
V.generator_or_list_generators)
|
||||
result.add_check("do_classifier_free_guidance",
|
||||
batch.do_classifier_free_guidance, V.bool_value)
|
||||
result.add_check(
|
||||
"negative_prompt_embeds", batch.negative_prompt_embeds, lambda x:
|
||||
not batch.do_classifier_free_guidance or V.list_not_empty(x))
|
||||
return result
|
||||
|
||||
def verify_output(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""Verify denoising stage outputs."""
|
||||
result = VerificationResult()
|
||||
result.add_check("latents", batch.latents,
|
||||
[V.is_tensor, V.with_dims(5)])
|
||||
return result
|
||||
@@ -25,6 +25,11 @@ class StageValidators:
|
||||
"""Check if value is a positive integer."""
|
||||
return isinstance(value, int) and value > 0
|
||||
|
||||
@staticmethod
|
||||
def non_negative_int(value: Any) -> bool:
|
||||
"""Check if value is a non-negative integer (allows 0)."""
|
||||
return isinstance(value, int) and value >= 0
|
||||
|
||||
@staticmethod
|
||||
def positive_float(value: Any) -> bool:
|
||||
"""Check if value is a positive float."""
|
||||
@@ -405,6 +410,8 @@ class VerificationResult:
|
||||
error_msg = f"tensor contains {torch.isnan(value).sum().item()} NaN values"
|
||||
elif validator_name == 'positive_int':
|
||||
expected = "positive integer"
|
||||
elif validator_name == 'non_negative_int':
|
||||
expected = "non-negative integer"
|
||||
elif validator_name == 'not_none':
|
||||
expected = "non-None value"
|
||||
elif validator_name == 'list_not_empty':
|
||||
|
||||
@@ -155,7 +155,7 @@ class CudaPlatformBase(Platform):
|
||||
)
|
||||
elif selected_backend == AttentionBackendEnum.SAGE_ATTN_THREE:
|
||||
try:
|
||||
from fastvideo.attention.backends.sageattn.api import sageattn_blackwell # noqa: F401
|
||||
from sageattn3 import sageattn3_blackwell # noqa: F401
|
||||
|
||||
from fastvideo.attention.backends.sage_attn3 import ( # noqa: F401
|
||||
SageAttention3Backend)
|
||||
|
||||
@@ -0,0 +1,619 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Central registry for FastVideo pipelines and model configuration discovery.
|
||||
|
||||
This module mirrors the organization of sglang's registry while keeping
|
||||
FastVideo's legacy behavior and mappings intact.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import dataclasses
|
||||
import os
|
||||
from collections.abc import Callable
|
||||
from functools import lru_cache
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
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, Hunyuan15I2V480PStepDistilledConfig,
|
||||
Hunyuan15T2V720PConfig, Hunyuan15I2V720PConfig, Hunyuan15SR1080PConfig)
|
||||
from fastvideo.configs.pipelines.hyworld import HYWorldConfig
|
||||
from fastvideo.configs.pipelines.longcat import LongCatT2V480PConfig
|
||||
from fastvideo.configs.pipelines.ltx2 import LTX2T2VConfig
|
||||
from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
|
||||
from fastvideo.configs.pipelines.turbodiffusion import (
|
||||
TurboDiffusionI2V_A14B_Config,
|
||||
TurboDiffusionT2V_14B_Config,
|
||||
TurboDiffusionT2V_1_3B_Config,
|
||||
)
|
||||
from fastvideo.configs.pipelines.wan import (
|
||||
FastWan2_1_T2V_480P_Config,
|
||||
FastWan2_2_TI2V_5B_Config,
|
||||
MatrixGameI2V480PConfig,
|
||||
SelfForcingWan2_2_T2V480PConfig,
|
||||
SelfForcingWanT2V480PConfig,
|
||||
WANV2VConfig,
|
||||
Wan2_2_I2V_A14B_Config,
|
||||
Wan2_2_T2V_A14B_Config,
|
||||
Wan2_2_TI2V_5B_Config,
|
||||
WanI2V480PConfig,
|
||||
WanI2V720PConfig,
|
||||
WanT2V480PConfig,
|
||||
WanT2V720PConfig,
|
||||
)
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
from fastvideo.configs.sample.cosmos import (
|
||||
Cosmos_Predict2_2B_Video2World_SamplingParam, )
|
||||
from fastvideo.configs.sample.cosmos2_5 import Cosmos25SamplingParamBase
|
||||
from fastvideo.configs.sample.hunyuan import (FastHunyuanSamplingParam,
|
||||
HunyuanSamplingParam)
|
||||
from fastvideo.configs.sample.hunyuan15 import (
|
||||
Hunyuan15_480P_SamplingParam,
|
||||
Hunyuan15_480P_StepDistilled_I2V_SamplingParam,
|
||||
Hunyuan15_720P_SamplingParam, Hunyuan15_720P_Distilled_I2V_SamplingParam,
|
||||
Hunyuan15_SR_1080P_SamplingParam)
|
||||
from fastvideo.configs.sample.hyworld import HYWorld_SamplingParam
|
||||
from fastvideo.configs.sample.ltx2 import LTX2SamplingParam
|
||||
from fastvideo.configs.sample.stepvideo import StepVideoT2VSamplingParam
|
||||
from fastvideo.configs.sample.turbodiffusion import (
|
||||
TurboDiffusionI2V_A14B_SamplingParam,
|
||||
TurboDiffusionT2V_14B_SamplingParam,
|
||||
TurboDiffusionT2V_1_3B_SamplingParam,
|
||||
)
|
||||
from fastvideo.configs.sample.wan import (
|
||||
FastWanT2V480P_SamplingParam,
|
||||
MatrixGame2_SamplingParam,
|
||||
SelfForcingWan2_1_T2V_1_3B_480P_SamplingParam,
|
||||
SelfForcingWan2_2_T2V_A14B_480P_SamplingParam,
|
||||
Wan2_1_Fun_1_3B_Control_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_14B_SamplingParam,
|
||||
WanT2V_1_3B_SamplingParam,
|
||||
)
|
||||
from fastvideo.fastvideo_args import WorkloadType
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.utils import (maybe_download_model_index,
|
||||
verify_model_config_and_directory)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
from fastvideo.pipelines.pipeline_registry import PipelineType
|
||||
|
||||
# --- Part 1: Pipeline Discovery ---
|
||||
|
||||
_PIPELINE_REGISTRY: dict[str, dict[str, type[ComposedPipelineBase]]] = {}
|
||||
|
||||
# Registry for pipeline configuration classes (for single-file weights without
|
||||
# model_index.json). Maps pipeline_class_name -> (PipelineConfig, SamplingParam)
|
||||
_PIPELINE_CONFIG_REGISTRY: dict[str, tuple[type[PipelineConfig],
|
||||
type[SamplingParam]]] = {}
|
||||
|
||||
|
||||
def _discover_and_register_pipelines() -> None:
|
||||
if _PIPELINE_REGISTRY:
|
||||
return
|
||||
|
||||
from fastvideo.pipelines.pipeline_registry import import_pipeline_classes
|
||||
|
||||
pipeline_classes = import_pipeline_classes()
|
||||
for pipeline_type, pipeline_dict in pipeline_classes.items():
|
||||
_PIPELINE_REGISTRY[pipeline_type] = pipeline_dict
|
||||
for pipeline_cls in pipeline_dict.values():
|
||||
if pipeline_cls is None:
|
||||
continue
|
||||
if hasattr(pipeline_cls, "pipeline_config_cls") and hasattr(
|
||||
pipeline_cls, "sampling_params_cls"):
|
||||
_PIPELINE_CONFIG_REGISTRY[pipeline_cls.__name__] = (
|
||||
pipeline_cls.pipeline_config_cls,
|
||||
pipeline_cls.sampling_params_cls,
|
||||
)
|
||||
|
||||
|
||||
def get_pipeline_config_classes(
|
||||
pipeline_class_name: str
|
||||
) -> tuple[type[PipelineConfig], type[SamplingParam]] | None:
|
||||
_discover_and_register_pipelines()
|
||||
return _PIPELINE_CONFIG_REGISTRY.get(pipeline_class_name)
|
||||
|
||||
|
||||
# --- Part 2: Config Registration ---
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
class ConfigInfo:
|
||||
"""Encapsulates sampling + pipeline config classes for a model family."""
|
||||
|
||||
sampling_param_cls: type[SamplingParam] | None
|
||||
pipeline_config_cls: type[PipelineConfig]
|
||||
|
||||
|
||||
# The central registry mapping a model name to its configuration information
|
||||
_CONFIG_REGISTRY: dict[str, ConfigInfo] = {}
|
||||
|
||||
# Mappings from Hugging Face model paths to our internal model names
|
||||
_MODEL_HF_PATH_TO_NAME: dict[str, str] = {}
|
||||
|
||||
# Detectors to identify model families from paths or class names
|
||||
_MODEL_NAME_DETECTORS: list[tuple[str, Callable[[str], bool]]] = []
|
||||
|
||||
|
||||
def register_configs(
|
||||
sampling_param_cls: type[SamplingParam] | None,
|
||||
pipeline_config_cls: type[PipelineConfig],
|
||||
hf_model_paths: list[str] | None = None,
|
||||
model_detectors: list[Callable[[str], bool]] | None = None,
|
||||
) -> None:
|
||||
"""Register config classes for a model family."""
|
||||
model_id = str(len(_CONFIG_REGISTRY))
|
||||
|
||||
_CONFIG_REGISTRY[model_id] = ConfigInfo(
|
||||
sampling_param_cls=sampling_param_cls,
|
||||
pipeline_config_cls=pipeline_config_cls,
|
||||
)
|
||||
|
||||
if hf_model_paths:
|
||||
for path in hf_model_paths:
|
||||
if path in _MODEL_HF_PATH_TO_NAME:
|
||||
logger.warning(
|
||||
"Model path '%s' is already mapped to '%s' and will be overwritten by '%s'.",
|
||||
path, _MODEL_HF_PATH_TO_NAME[path], model_id)
|
||||
_MODEL_HF_PATH_TO_NAME[path] = model_id
|
||||
|
||||
if model_detectors:
|
||||
for detector in model_detectors:
|
||||
_MODEL_NAME_DETECTORS.append((model_id, detector))
|
||||
|
||||
|
||||
def get_model_short_name(model_id: str) -> str:
|
||||
if "/" in model_id:
|
||||
return model_id.split("/")[-1]
|
||||
return model_id
|
||||
|
||||
|
||||
def _get_config_info(
|
||||
model_path: str,
|
||||
*,
|
||||
raise_on_missing: bool = True,
|
||||
) -> ConfigInfo | None:
|
||||
# 1. Exact match
|
||||
if model_path in _MODEL_HF_PATH_TO_NAME:
|
||||
model_id = _MODEL_HF_PATH_TO_NAME[model_path]
|
||||
logger.debug("Resolved model path '%s' from exact path match.",
|
||||
model_path)
|
||||
return _CONFIG_REGISTRY.get(model_id)
|
||||
|
||||
# 2. Partial match: use short model name.
|
||||
model_name = get_model_short_name(model_path.lower())
|
||||
all_model_hf_paths = sorted(_MODEL_HF_PATH_TO_NAME.keys(),
|
||||
key=len,
|
||||
reverse=True)
|
||||
for registered_model_hf_id in all_model_hf_paths:
|
||||
registered_model_name = get_model_short_name(
|
||||
registered_model_hf_id.lower())
|
||||
if registered_model_name == model_name:
|
||||
logger.debug("Resolved model name '%s' from partial path match.",
|
||||
registered_model_hf_id)
|
||||
model_id = _MODEL_HF_PATH_TO_NAME[registered_model_hf_id]
|
||||
return _CONFIG_REGISTRY.get(model_id)
|
||||
|
||||
# 3. Use detectors (path or model_index pipeline name).
|
||||
if os.path.exists(model_path):
|
||||
config = verify_model_config_and_directory(model_path)
|
||||
else:
|
||||
config = maybe_download_model_index(model_path)
|
||||
|
||||
pipeline_name = config.get("_class_name", "").lower()
|
||||
|
||||
matched_model_names: list[str] = []
|
||||
for model_id, detector in _MODEL_NAME_DETECTORS:
|
||||
if detector(model_path.lower()) or detector(pipeline_name):
|
||||
logger.debug("Matched model name '%s' using a registered detector.",
|
||||
model_id)
|
||||
matched_model_names.append(model_id)
|
||||
|
||||
if matched_model_names:
|
||||
if len(matched_model_names) > 1:
|
||||
logger.warning(
|
||||
"Multiple models matched for path '%s': %s. Using the first matched: '%s'.",
|
||||
model_path,
|
||||
matched_model_names,
|
||||
matched_model_names[0],
|
||||
)
|
||||
model_id = matched_model_names[0]
|
||||
return _CONFIG_REGISTRY.get(model_id)
|
||||
|
||||
if raise_on_missing:
|
||||
raise RuntimeError(f"No model info found for model path: {model_path}")
|
||||
return None
|
||||
|
||||
|
||||
def _register_configs() -> None:
|
||||
# LTX-2
|
||||
register_configs(
|
||||
sampling_param_cls=LTX2SamplingParam,
|
||||
pipeline_config_cls=LTX2T2VConfig,
|
||||
hf_model_paths=[
|
||||
"Lightricks/LTX-2",
|
||||
"converted/ltx2_diffusers",
|
||||
"FastVideo/LTX2-Distilled-Diffusers",
|
||||
],
|
||||
model_detectors=[
|
||||
lambda path: "ltx2" in path.lower() or "ltx-2" in path.lower(),
|
||||
],
|
||||
)
|
||||
|
||||
# Hunyuan 1.5 (specific)
|
||||
register_configs(
|
||||
sampling_param_cls=Hunyuan15_480P_SamplingParam,
|
||||
pipeline_config_cls=Hunyuan15T2V480PConfig,
|
||||
hf_model_paths=[
|
||||
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v",
|
||||
],
|
||||
model_detectors=[
|
||||
lambda path: any(token in path.lower() for token in (
|
||||
"hunyuan15",
|
||||
"hunyuanvideo15",
|
||||
"hunyuanvideo-1.5",
|
||||
"hunyuanvideo_1.5",
|
||||
)),
|
||||
],
|
||||
)
|
||||
register_configs(
|
||||
sampling_param_cls=Hunyuan15_480P_StepDistilled_I2V_SamplingParam,
|
||||
pipeline_config_cls=Hunyuan15I2V480PStepDistilledConfig,
|
||||
hf_model_paths=[
|
||||
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_i2v_step_distilled",
|
||||
],
|
||||
)
|
||||
register_configs(
|
||||
sampling_param_cls=Hunyuan15_720P_SamplingParam,
|
||||
pipeline_config_cls=Hunyuan15T2V720PConfig,
|
||||
hf_model_paths=[
|
||||
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-720p_t2v",
|
||||
],
|
||||
)
|
||||
register_configs(
|
||||
sampling_param_cls=Hunyuan15_720P_Distilled_I2V_SamplingParam,
|
||||
pipeline_config_cls=Hunyuan15I2V720PConfig,
|
||||
hf_model_paths=[
|
||||
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-720p_i2v_distilled",
|
||||
],
|
||||
)
|
||||
register_configs(
|
||||
sampling_param_cls=Hunyuan15_SR_1080P_SamplingParam,
|
||||
pipeline_config_cls=Hunyuan15SR1080PConfig,
|
||||
hf_model_paths=[
|
||||
"weizhou03/HunyuanVideo-1.5-Diffusers-1080p",
|
||||
"weizhou03/HunyuanVideo-1.5-Diffusers-1080p-2SR"
|
||||
],
|
||||
)
|
||||
|
||||
# Hunyuan
|
||||
register_configs(
|
||||
sampling_param_cls=HunyuanSamplingParam,
|
||||
pipeline_config_cls=HunyuanConfig,
|
||||
hf_model_paths=[
|
||||
"hunyuanvideo-community/HunyuanVideo",
|
||||
],
|
||||
model_detectors=[lambda path: "hunyuan" in path.lower()],
|
||||
)
|
||||
register_configs(
|
||||
sampling_param_cls=FastHunyuanSamplingParam,
|
||||
pipeline_config_cls=FastHunyuanConfig,
|
||||
hf_model_paths=[
|
||||
"FastVideo/FastHunyuan-diffusers",
|
||||
],
|
||||
)
|
||||
|
||||
# HYWorld
|
||||
register_configs(
|
||||
sampling_param_cls=HYWorld_SamplingParam,
|
||||
pipeline_config_cls=HYWorldConfig,
|
||||
hf_model_paths=[
|
||||
"FastVideo/HY-WorldPlay-Bidirectional-Diffusers",
|
||||
],
|
||||
model_detectors=[lambda path: "hyworld" in path.lower()],
|
||||
)
|
||||
|
||||
# LongCat
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
pipeline_config_cls=LongCatT2V480PConfig,
|
||||
hf_model_paths=[
|
||||
"FastVideo/LongCat-Video-T2V-Diffusers",
|
||||
"FastVideo/LongCat-Video-I2V-Diffusers",
|
||||
"FastVideo/LongCat-Video-VC-Diffusers",
|
||||
],
|
||||
model_detectors=[
|
||||
lambda path: "longcatimagetovideo" in path.lower(),
|
||||
lambda path: "longcatvideocontinuation" in path.lower(),
|
||||
lambda path: "longcat" in path.lower(),
|
||||
],
|
||||
)
|
||||
|
||||
# StepVideo
|
||||
register_configs(
|
||||
sampling_param_cls=StepVideoT2VSamplingParam,
|
||||
pipeline_config_cls=StepVideoT2VConfig,
|
||||
hf_model_paths=[
|
||||
"FastVideo/stepvideo-t2v-diffusers",
|
||||
],
|
||||
model_detectors=[lambda path: "stepvideo" in path.lower()],
|
||||
)
|
||||
|
||||
# MatrixGame
|
||||
register_configs(
|
||||
sampling_param_cls=MatrixGame2_SamplingParam,
|
||||
pipeline_config_cls=MatrixGameI2V480PConfig,
|
||||
hf_model_paths=[
|
||||
"FastVideo/Matrix-Game-2.0-Base-Diffusers",
|
||||
"FastVideo/Matrix-Game-2.0-GTA-Diffusers",
|
||||
"FastVideo/Matrix-Game-2.0-TempleRun-Diffusers",
|
||||
],
|
||||
model_detectors=[
|
||||
lambda path: "matrix-game" in path.lower() or "matrixgame" in path.
|
||||
lower(),
|
||||
],
|
||||
)
|
||||
|
||||
# Cosmos 2.5
|
||||
register_configs(
|
||||
sampling_param_cls=Cosmos25SamplingParamBase,
|
||||
pipeline_config_cls=Cosmos25Config,
|
||||
hf_model_paths=[
|
||||
"KyleShao/Cosmos-Predict2.5-2B-Diffusers",
|
||||
],
|
||||
model_detectors=[
|
||||
lambda path: any(token in path.lower() for token in (
|
||||
"cosmos25",
|
||||
"cosmos2_5",
|
||||
"cosmos2.5",
|
||||
)),
|
||||
],
|
||||
)
|
||||
|
||||
# Cosmos 2
|
||||
register_configs(
|
||||
sampling_param_cls=Cosmos_Predict2_2B_Video2World_SamplingParam,
|
||||
pipeline_config_cls=CosmosConfig,
|
||||
hf_model_paths=[
|
||||
"nvidia/Cosmos-Predict2-2B-Video2World",
|
||||
],
|
||||
model_detectors=[
|
||||
lambda path: "cosmos" in path.lower() and ("2.5" not in path.lower(
|
||||
) and "2_5" not in path.lower() and "25" not in path.lower()),
|
||||
],
|
||||
)
|
||||
|
||||
# TurboDiffusion
|
||||
register_configs(
|
||||
sampling_param_cls=TurboDiffusionT2V_1_3B_SamplingParam,
|
||||
pipeline_config_cls=TurboDiffusionT2V_1_3B_Config,
|
||||
hf_model_paths=[
|
||||
"loayrashid/TurboWan2.1-T2V-1.3B-Diffusers",
|
||||
],
|
||||
model_detectors=[
|
||||
lambda path: "turbodiffusion" in path.lower() or "turbowan" in path.
|
||||
lower()
|
||||
],
|
||||
)
|
||||
register_configs(
|
||||
sampling_param_cls=TurboDiffusionT2V_14B_SamplingParam,
|
||||
pipeline_config_cls=TurboDiffusionT2V_14B_Config,
|
||||
hf_model_paths=[
|
||||
"loayrashid/TurboWan2.1-T2V-14B-Diffusers",
|
||||
],
|
||||
)
|
||||
register_configs(
|
||||
sampling_param_cls=TurboDiffusionI2V_A14B_SamplingParam,
|
||||
pipeline_config_cls=TurboDiffusionI2V_A14B_Config,
|
||||
hf_model_paths=[
|
||||
"loayrashid/TurboWan2.2-I2V-A14B-Diffusers",
|
||||
],
|
||||
)
|
||||
|
||||
# Wan
|
||||
register_configs(
|
||||
sampling_param_cls=WanT2V_1_3B_SamplingParam,
|
||||
pipeline_config_cls=WanT2V480PConfig,
|
||||
hf_model_paths=[
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
],
|
||||
model_detectors=[lambda path: "wanpipeline" in path.lower()],
|
||||
)
|
||||
register_configs(
|
||||
sampling_param_cls=WanT2V_14B_SamplingParam,
|
||||
pipeline_config_cls=WanT2V720PConfig,
|
||||
hf_model_paths=[
|
||||
"Wan-AI/Wan2.1-T2V-14B-Diffusers",
|
||||
"FastVideo/Wan2.1-VSA-T2V-14B-720P-Diffusers",
|
||||
],
|
||||
)
|
||||
register_configs(
|
||||
sampling_param_cls=WanI2V_14B_480P_SamplingParam,
|
||||
pipeline_config_cls=WanI2V480PConfig,
|
||||
hf_model_paths=[
|
||||
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers",
|
||||
],
|
||||
model_detectors=[lambda path: "wanimagetovideo" in path.lower()],
|
||||
)
|
||||
register_configs(
|
||||
sampling_param_cls=WanI2V_14B_720P_SamplingParam,
|
||||
pipeline_config_cls=WanI2V720PConfig,
|
||||
hf_model_paths=[
|
||||
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers",
|
||||
],
|
||||
)
|
||||
register_configs(
|
||||
sampling_param_cls=Wan2_1_Fun_1_3B_InP_SamplingParam,
|
||||
pipeline_config_cls=WanI2V480PConfig,
|
||||
hf_model_paths=[
|
||||
"weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers",
|
||||
],
|
||||
)
|
||||
register_configs(
|
||||
sampling_param_cls=Wan2_1_Fun_1_3B_Control_SamplingParam,
|
||||
pipeline_config_cls=WANV2VConfig,
|
||||
hf_model_paths=[
|
||||
"IRMChen/Wan2.1-Fun-1.3B-Control-Diffusers",
|
||||
],
|
||||
)
|
||||
register_configs(
|
||||
sampling_param_cls=FastWanT2V480P_SamplingParam,
|
||||
pipeline_config_cls=FastWan2_1_T2V_480P_Config,
|
||||
hf_model_paths=[
|
||||
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers",
|
||||
"FastVideo/FastWan2.1-T2V-14B-480P-Diffusers",
|
||||
],
|
||||
model_detectors=[lambda path: "wandmdpipeline" in path.lower()],
|
||||
)
|
||||
register_configs(
|
||||
sampling_param_cls=Wan2_2_TI2V_5B_SamplingParam,
|
||||
pipeline_config_cls=Wan2_2_TI2V_5B_Config,
|
||||
hf_model_paths=[
|
||||
"Wan-AI/Wan2.2-TI2V-5B-Diffusers",
|
||||
],
|
||||
)
|
||||
register_configs(
|
||||
sampling_param_cls=Wan2_2_TI2V_5B_SamplingParam,
|
||||
pipeline_config_cls=FastWan2_2_TI2V_5B_Config,
|
||||
hf_model_paths=[
|
||||
"FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers",
|
||||
"FastVideo/FastWan2.2-TI2V-5B-Diffusers",
|
||||
],
|
||||
)
|
||||
register_configs(
|
||||
sampling_param_cls=Wan2_2_T2V_A14B_SamplingParam,
|
||||
pipeline_config_cls=Wan2_2_T2V_A14B_Config,
|
||||
hf_model_paths=[
|
||||
"Wan-AI/Wan2.2-T2V-A14B-Diffusers",
|
||||
],
|
||||
)
|
||||
register_configs(
|
||||
sampling_param_cls=Wan2_2_I2V_A14B_SamplingParam,
|
||||
pipeline_config_cls=Wan2_2_I2V_A14B_Config,
|
||||
hf_model_paths=[
|
||||
"Wan-AI/Wan2.2-I2V-A14B-Diffusers",
|
||||
],
|
||||
)
|
||||
register_configs(
|
||||
sampling_param_cls=SelfForcingWan2_1_T2V_1_3B_480P_SamplingParam,
|
||||
pipeline_config_cls=SelfForcingWanT2V480PConfig,
|
||||
hf_model_paths=[
|
||||
"wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers",
|
||||
],
|
||||
model_detectors=[lambda path: "wancausaldmdpipeline" in path.lower()],
|
||||
)
|
||||
register_configs(
|
||||
sampling_param_cls=SelfForcingWan2_2_T2V_A14B_480P_SamplingParam,
|
||||
pipeline_config_cls=SelfForcingWan2_2_T2V480PConfig,
|
||||
hf_model_paths=[
|
||||
"rand0nmr/SFWan2.2-T2V-A14B-Diffusers",
|
||||
"FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers",
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
# --- Part 3: Main Resolver ---
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
class ModelInfo:
|
||||
pipeline_cls: type[ComposedPipelineBase]
|
||||
sampling_param_cls: type[SamplingParam]
|
||||
pipeline_config_cls: type[PipelineConfig]
|
||||
|
||||
|
||||
@lru_cache(maxsize=32)
|
||||
def get_model_info(
|
||||
model_path: str,
|
||||
pipeline_type: PipelineType | str | None = None,
|
||||
workload_type: WorkloadType | None = None,
|
||||
override_pipeline_cls_name: str | None = None,
|
||||
) -> ModelInfo:
|
||||
from fastvideo.pipelines.pipeline_registry import (PipelineType,
|
||||
get_pipeline_registry)
|
||||
|
||||
if pipeline_type is None:
|
||||
pipeline_type = PipelineType.BASIC
|
||||
elif isinstance(pipeline_type, str):
|
||||
pipeline_type = PipelineType.from_string(pipeline_type)
|
||||
|
||||
if workload_type is None:
|
||||
workload_type = WorkloadType.T2V
|
||||
|
||||
if os.path.exists(model_path):
|
||||
config = verify_model_config_and_directory(model_path)
|
||||
else:
|
||||
config = maybe_download_model_index(model_path)
|
||||
|
||||
pipeline_name = config.get("_class_name")
|
||||
if override_pipeline_cls_name:
|
||||
logger.info("Overriding pipeline class name from %s to %s",
|
||||
pipeline_name, override_pipeline_cls_name)
|
||||
pipeline_name = 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.")
|
||||
|
||||
pipeline_registry = get_pipeline_registry(pipeline_type)
|
||||
pipeline_cls = pipeline_registry.resolve_pipeline_cls(
|
||||
pipeline_name, pipeline_type, workload_type)
|
||||
|
||||
config_info = _get_config_info(model_path, raise_on_missing=True)
|
||||
assert config_info is not None, "config_info must be resolved"
|
||||
|
||||
sampling_param_cls = config_info.sampling_param_cls or SamplingParam
|
||||
|
||||
return ModelInfo(
|
||||
pipeline_cls=pipeline_cls,
|
||||
sampling_param_cls=sampling_param_cls,
|
||||
pipeline_config_cls=config_info.pipeline_config_cls,
|
||||
)
|
||||
|
||||
|
||||
def get_pipeline_config_cls_from_name(
|
||||
pipeline_name_or_path: str) -> type[PipelineConfig]:
|
||||
config_info = _get_config_info(pipeline_name_or_path,
|
||||
raise_on_missing=False)
|
||||
if config_info is None:
|
||||
raise ValueError(
|
||||
f"No match found for pipeline {pipeline_name_or_path}, please check the pipeline name or path."
|
||||
)
|
||||
return config_info.pipeline_config_cls
|
||||
|
||||
|
||||
def get_sampling_param_cls_for_name(pipeline_name_or_path: str) -> Any | None:
|
||||
config_info = _get_config_info(pipeline_name_or_path,
|
||||
raise_on_missing=False)
|
||||
if config_info is None:
|
||||
logger.warning(
|
||||
"No match found for pipeline %s, using default sampling param.",
|
||||
pipeline_name_or_path)
|
||||
return None
|
||||
return config_info.sampling_param_cls
|
||||
|
||||
|
||||
_register_configs()
|
||||
|
||||
__all__ = [
|
||||
"ConfigInfo",
|
||||
"ModelInfo",
|
||||
"get_model_info",
|
||||
"get_pipeline_config_cls_from_name",
|
||||
"get_sampling_param_cls_for_name",
|
||||
"get_pipeline_config_classes",
|
||||
]
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user