Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
c4a0f789da | ||
|
|
5ac14938d2 | ||
|
|
9d239e9f8b |
@@ -31,9 +31,9 @@ log "Setting up Modal authentication from Buildkite secrets..."
|
||||
MODAL_TOKEN_ID=$(buildkite-agent secret get modal_token_id)
|
||||
MODAL_TOKEN_SECRET=$(buildkite-agent secret get modal_token_secret)
|
||||
|
||||
# Retrieve other secrets
|
||||
WANDB_API_KEY=$(buildkite-agent secret get wandb_api_key)
|
||||
HF_API_KEY=$(buildkite-agent secret get hf_api_key)
|
||||
|
||||
WANDB_API_KEY=$(buildkite-agent secret get wandb_api_key)
|
||||
|
||||
if [ -n "$MODAL_TOKEN_ID" ] && [ -n "$MODAL_TOKEN_SECRET" ]; then
|
||||
log "Retrieved Modal credentials from Buildkite secrets"
|
||||
@@ -63,19 +63,19 @@ MODAL_ENV="BUILDKITE_REPO=$BUILDKITE_REPO BUILDKITE_COMMIT=$BUILDKITE_COMMIT BUI
|
||||
case "$TEST_TYPE" in
|
||||
"encoder")
|
||||
log "Running encoder tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV HF_API_KEY=$HF_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_encoder_tests"
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_encoder_tests"
|
||||
;;
|
||||
"vae")
|
||||
log "Running VAE tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV HF_API_KEY=$HF_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_vae_tests"
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_vae_tests"
|
||||
;;
|
||||
"transformer")
|
||||
log "Running transformer tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV HF_API_KEY=$HF_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_transformer_tests"
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_transformer_tests"
|
||||
;;
|
||||
"ssim")
|
||||
log "Running SSIM tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV HF_API_KEY=$HF_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_ssim_tests"
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_ssim_tests"
|
||||
;;
|
||||
"training")
|
||||
log "Running training tests..."
|
||||
|
||||
@@ -1,65 +1,82 @@
|
||||
name: Deploy Documentation
|
||||
# Sample workflow for building and deploying a Hugo site to GitHub Pages
|
||||
name: Deploy FastVideo Docs to Pages
|
||||
|
||||
on:
|
||||
# Runs on pushes targeting the default branch
|
||||
push:
|
||||
branches: [ main ]
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- 'docs/**'
|
||||
- 'mkdocs.yml'
|
||||
- 'requirements-mkdocs.txt'
|
||||
- '.github/workflows/docs.yml'
|
||||
- "docs/**/*.md"
|
||||
- "fastvideo/examples/**/*.py"
|
||||
pull_request:
|
||||
branches: [ main ]
|
||||
branches:
|
||||
- main
|
||||
types: [opened, ready_for_review, synchronize, reopened]
|
||||
paths:
|
||||
- 'docs/**'
|
||||
- 'mkdocs.yml'
|
||||
- 'requirements-mkdocs.txt'
|
||||
- '.github/workflows/docs.yml'
|
||||
- "docs/**/*.md"
|
||||
- "fastvideo/examples/**/*.py"
|
||||
|
||||
# Allows you to run this workflow manually from the Actions tab
|
||||
workflow_dispatch:
|
||||
|
||||
# Sets permissions of the GITHUB_TOKEN to allow deployment to GitHub Pages
|
||||
permissions:
|
||||
contents: read
|
||||
pages: write
|
||||
id-token: write
|
||||
|
||||
# Allow only one concurrent deployment, skipping runs queued between the run in-progress and latest queued.
|
||||
# However, do NOT cancel in-progress runs as we want to allow these production deployments to complete.
|
||||
concurrency:
|
||||
group: "pages"
|
||||
cancel-in-progress: false
|
||||
|
||||
# Default to bash
|
||||
defaults:
|
||||
run:
|
||||
shell: bash
|
||||
|
||||
jobs:
|
||||
pre-commit:
|
||||
uses: ./.github/workflows/pre-commit.yml
|
||||
|
||||
# Build job
|
||||
build:
|
||||
runs-on: ubuntu-latest
|
||||
needs: pre-commit
|
||||
steps:
|
||||
- name: Checkout
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Setup Python
|
||||
- name: Setup Pages
|
||||
id: pages
|
||||
uses: actions/configure-pages@v5
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: '3.12'
|
||||
|
||||
python-version: "3.10"
|
||||
- name: Install dependencies
|
||||
run: |
|
||||
python -m pip install --upgrade pip
|
||||
pip install -r requirements-mkdocs.txt
|
||||
|
||||
- name: Setup Pages
|
||||
uses: actions/configure-pages@v4
|
||||
|
||||
- name: Build documentation
|
||||
run: mkdocs build
|
||||
|
||||
cd docs
|
||||
pip install -r requirements-docs.txt
|
||||
- name: Build docs
|
||||
run: |
|
||||
cd docs
|
||||
make clean
|
||||
make html
|
||||
- name: Upload artifact
|
||||
uses: actions/upload-pages-artifact@v3
|
||||
with:
|
||||
path: ./site
|
||||
path: ./docs/build/html
|
||||
|
||||
# Deployment job
|
||||
deploy:
|
||||
environment:
|
||||
name: github-pages
|
||||
url: ${{ steps.deployment.outputs.page_url }}
|
||||
if: ${{ github.event_name == 'push' }}
|
||||
runs-on: ubuntu-latest
|
||||
needs: build
|
||||
if: github.ref == 'refs/heads/main'
|
||||
steps:
|
||||
- name: Deploy to GitHub Pages
|
||||
id: deployment
|
||||
|
||||
@@ -14,8 +14,6 @@ wandb/
|
||||
*.pt
|
||||
cache_dir/
|
||||
wandb/
|
||||
venv/
|
||||
.venv/
|
||||
runs/
|
||||
samples/
|
||||
*validation/
|
||||
@@ -39,13 +37,12 @@ dist/
|
||||
eggs/
|
||||
.eggs/
|
||||
|
||||
# MkDocs documentation
|
||||
site/
|
||||
docs/getting_started/examples/
|
||||
docs/inference/examples/
|
||||
docs/training/examples/
|
||||
docs/distillation/examples/
|
||||
!requirements-mkdocs.txt
|
||||
# Sphinx documentation
|
||||
docs/_build/
|
||||
docs/source/getting_started/examples/
|
||||
docs/source/inference/examples/
|
||||
docs/source/training/examples/
|
||||
docs/source/distillation/examples/
|
||||
|
||||
# VSCode
|
||||
.vscode/
|
||||
@@ -64,7 +61,7 @@ docs/distillation/examples/
|
||||
!fastvideo/tests/ssim/reference_videos/**/*.mp4
|
||||
|
||||
# Static images
|
||||
!docs/assets/images/**/*.png
|
||||
!docs/source/_static/images/**/*.png
|
||||
!comfyui/assets/**/*.png
|
||||
!comfyui/assets/**/*.gif
|
||||
|
||||
|
||||
@@ -10,7 +10,6 @@ exclude: |
|
||||
demo/.*|
|
||||
predict\.py|
|
||||
scripts/.*|
|
||||
prompts/.*|
|
||||
fastvideo/data_preprocess/.*|
|
||||
fastvideo/dataset/.*|
|
||||
fastvideo/models/.*|
|
||||
|
||||
@@ -1,12 +1,13 @@
|
||||
<div align="center">
|
||||
<img src=assets/logos/logo.svg width="30%"/>
|
||||
</div>
|
||||
|
||||
**FastVideo is a unified post-training and inference framework for accelerated video generation.**
|
||||
|
||||
FastVideo features an end-to-end unified pipeline for accelerating diffusion models, starting from data preprocessing to model training, finetuning, distillation, and inference. FastVideo is designed to be modular and extensible, allowing users to easily add new optimizations and techniques. Whether it is training-free optimizations or post-training optimizations, FastVideo has you covered.
|
||||
|
||||
<p align="center">
|
||||
| 🕹️ <a href="https://fastwan.fastvideo.org/"<b>Online Demo</b></a> | <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://huggingface.co/collections/FastVideo/fastwan-6886a305d9799c8cd1496408" target="_blank"><b>FastWan</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/TM8JyJCd" target="_blank"> <b> WeChat </b> </a> |
|
||||
| 🕹️ <a href="https://fastwan.fastvideo.org/"<b>Online Demo</b></a> | <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.html"><b> Quick Start</b></a> | 🤗 <a href="https://huggingface.co/collections/FastVideo/fastwan-6886a305d9799c8cd1496408" target="_blank"><b>FastWan</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-3csdw1isz-Euq8_Q8~baewG8hxjXs2gQ" target="_blank"> <b>Slack</b> </a> | 🟣💬 <a href="https://ibb.co/tMwknPLY" target="_blank"> <b> WeChat </b> </a> |
|
||||
</p>
|
||||
|
||||
<div align="center">
|
||||
@@ -14,7 +15,6 @@ FastVideo features an end-to-end unified pipeline for accelerating diffusion mod
|
||||
</div>
|
||||
|
||||
## 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.html) 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!
|
||||
@@ -49,10 +49,10 @@ conda activate fastvideo
|
||||
pip install fastvideo
|
||||
```
|
||||
|
||||
Please see our [docs](https://hao-ai-lab.github.io/FastVideo/getting_started/installation/) for more detailed installation instructions.
|
||||
Please see our [docs](https://hao-ai-lab.github.io/FastVideo/getting_started/installation.html) 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/).
|
||||
For our sparse distillation techniques, please see our [distillation docs](https://hao-ai-lab.github.io/FastVideo/distillation/dmd.html) and check out our [blog](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/).
|
||||
|
||||
See below for recipes and datasets:
|
||||
|
||||
@@ -64,7 +64,7 @@ See below for recipes and datasets:
|
||||
|
||||
## 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:
|
||||
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.html). Create a file called `example.py` with the following code:
|
||||
|
||||
```python
|
||||
import os
|
||||
@@ -100,32 +100,35 @@ Run the script with:
|
||||
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/).
|
||||
For a more detailed guide, please see our [inference quick start](https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start.html).
|
||||
|
||||
### Other docs:
|
||||
|
||||
- [Design Overview](https://hao-ai-lab.github.io/FastVideo/design/overview/)
|
||||
- [Contribution Guide](https://hao-ai-lab.github.io/FastVideo/getting_started/installation/)
|
||||
- [Design Overview](https://hao-ai-lab.github.io/FastVideo/design/overview.html)
|
||||
- [Contribution Guide](https://hao-ai-lab.github.io/FastVideo/getting_started/installation.html)
|
||||
|
||||
## Distillation and Finetuning
|
||||
- [Distillation Guide](https://hao-ai-lab.github.io/FastVideo/distillation/dmd/)
|
||||
- [Distillation Guide](https://hao-ai-lab.github.io/FastVideo/distillation/dmd.html)
|
||||
<!-- - [Finetuning Guide](https://hao-ai-lab.github.io/FastVideo/training/finetune.html) -->
|
||||
|
||||
## Awesome work using FastVideo or our research projects
|
||||
## 📑 Development Plan
|
||||
<!-- - More distillation methods -->
|
||||
<!-- - [ ] Add Distribution Matching Distillation -->
|
||||
More FastWan Models Coming Soon!
|
||||
- [ ] Add FastWan2.1-T2V-14B
|
||||
- [ ] Add FastWan2.2-T2V-14B
|
||||
- [ ] Add FastWan2.2-I2V-14B
|
||||
<!-- - Optimization features
|
||||
- Code updates -->
|
||||
<!-- - [ ] fp8 support -->
|
||||
<!-- - [ ] faster load model and save model support -->
|
||||
|
||||
- [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)
|
||||
See details in [development roadmap](https://github.com/hao-ai-lab/FastVideo/issues/468).
|
||||
|
||||
## 🤝 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/468).
|
||||
We welcome all contributions. Please check out our guide [here](https://hao-ai-lab.github.io/FastVideo/contributing/overview.html)
|
||||
|
||||
## Acknowledgement
|
||||
We learned and reused code from the following projects:
|
||||
- [Wan-Video](https://github.com/Wan-Video)
|
||||
|
||||
@@ -77,7 +77,7 @@ https://github.com/user-attachments/assets/f3b6dd79-7b43-4b60-a0fa-3d6495ec5747
|
||||
Here is a diagram of how the window is configured and passed through the FastVideo pipeline:
|
||||
|
||||
<div align="center">
|
||||
<img src="../../../docs/assets/images/STA_configuration.png" width="80%"/>
|
||||
<img src="../../../docs/source/_static/images/STA_configuration.png" width="80%"/>
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
@@ -247,11 +247,11 @@ def _attn_bwd_dq(dq, q, K, V, #
|
||||
|
||||
kv_blocks = tl.load(q2k_num + meta_base) # int32
|
||||
kv_ptr = q2k_index + meta_base * max_kv_blks # ptr to list
|
||||
block_size = tl.load(variable_block_sizes + q_blk)
|
||||
|
||||
|
||||
for blk_idx in range(kv_blocks*2):
|
||||
block_sparse_offset = (tl.load(kv_ptr + blk_idx//2).to(tl.int32)*2 + blk_idx%2) *step_n * stride_tok
|
||||
block_size = tl.load(variable_block_sizes + blk_idx//2) - (blk_idx%2) * step_n
|
||||
kT = tl.load(kT_ptrs + block_sparse_offset)
|
||||
vT = tl.load(vT_ptrs + block_sparse_offset)
|
||||
qk = tl.dot(q, kT)
|
||||
|
||||
@@ -0,0 +1,26 @@
|
||||
# Minimal makefile for Sphinx documentation
|
||||
#
|
||||
|
||||
# You can set these variables from the command line, and also
|
||||
# from the environment for the first two.
|
||||
SPHINXOPTS ?=
|
||||
SPHINXBUILD ?= sphinx-build
|
||||
SOURCEDIR = source
|
||||
BUILDDIR = build
|
||||
|
||||
# Put it first so that "make" without argument is like "make help".
|
||||
help:
|
||||
@$(SPHINXBUILD) -M help "$(SOURCEDIR)" "$(BUILDDIR)" $(SPHINXOPTS) $(O)
|
||||
|
||||
.PHONY: help Makefile
|
||||
|
||||
# Catch-all target: route all unknown targets to Sphinx using the new
|
||||
# "make mode" option. $(O) is meant as a shortcut for $(SPHINXOPTS).
|
||||
%: Makefile
|
||||
@$(SPHINXBUILD) -M $@ "$(SOURCEDIR)" "$(BUILDDIR)" $(SPHINXOPTS) $(O)
|
||||
|
||||
clean:
|
||||
@$(SPHINXBUILD) -M clean "$(SOURCEDIR)" "$(BUILDDIR)" $(SPHINXOPTS) $(O)
|
||||
rm -rf "$(SOURCEDIR)/getting_started/examples"
|
||||
rm -rf "$(SOURCEDIR)/inference/examples"
|
||||
rm -rf "$(SOURCEDIR)/training/examples"
|
||||
@@ -1,39 +1,20 @@
|
||||
# FastVideo Documentation
|
||||
# FastVideo documents
|
||||
|
||||
This directory contains the FastVideo documentation built with MkDocs.
|
||||
|
||||
## Build the docs locally
|
||||
## Build the docs
|
||||
|
||||
```bash
|
||||
# Install dependencies
|
||||
pip install -r docs/requirements-mkdocs.txt
|
||||
# Install dependencies.
|
||||
pip install -r requirements-docs.txt
|
||||
|
||||
# Serve docs with live reload (recommended for development)
|
||||
mkdocs serve
|
||||
|
||||
# Or build static site
|
||||
mkdocs build
|
||||
# Build the docs.
|
||||
make clean
|
||||
make html
|
||||
```
|
||||
|
||||
## View the docs
|
||||
|
||||
### Development server (with live reload)
|
||||
## Open the docs with your browser
|
||||
|
||||
```bash
|
||||
mkdocs serve
|
||||
python -m http.server -d build/html/
|
||||
```
|
||||
|
||||
Then open your browser to: http://127.0.0.1:8000
|
||||
|
||||
### Static build
|
||||
|
||||
```bash
|
||||
mkdocs build
|
||||
python -m http.server -d site/
|
||||
```
|
||||
|
||||
Then open your browser to: http://localhost:8000
|
||||
|
||||
## Automatic Deployment
|
||||
|
||||
Documentation is automatically built and deployed to GitHub Pages when changes are pushed to the `main` branch via the `.github/workflows/docs.yml` workflow.
|
||||
Launch your browser and open localhost:8000.
|
||||
|
||||
@@ -1,248 +0,0 @@
|
||||
# FastVideo API Reference
|
||||
|
||||
This page contains the complete API reference for the FastVideo library.
|
||||
|
||||
## fastvideo
|
||||
|
||||
### Modules
|
||||
|
||||
| Name | Description |
|
||||
|------|-------------|
|
||||
| [attention](#fastvideoattention) | Attention mechanisms and backends for video generation |
|
||||
| [configs](#fastvideoconfigs) | Configuration classes for pipelines, models, and sampling |
|
||||
| [distributed](#fastvideodistributed) | Distributed execution and communication utilities |
|
||||
| [entrypoints](#fastvideoentrypoints) | Main API entry points for video generation |
|
||||
| [models](#fastvideomodels) | Model implementations (transformers, VAEs, schedulers) |
|
||||
| [pipelines](#fastvideopipelines) | Core pipeline classes for video diffusion |
|
||||
| [training](#fastvideotraining) | Training utilities and helpers |
|
||||
| [workflow](#fastvideoworkflow) | Workflow management and orchestration |
|
||||
| [dataset](#fastvideodataset) | Dataset handling and preprocessing |
|
||||
| [layers](#fastvideolayers) | Custom neural network layers |
|
||||
| [platforms](#fastvideoplatforms) | Platform-specific implementations |
|
||||
| [utils](#fastvideoutils) | Utility functions and helpers |
|
||||
| [worker](#fastvideoworker) | Execution workers for video generation |
|
||||
|
||||
## fastvideo.attention
|
||||
|
||||
::: fastvideo.attention
|
||||
options:
|
||||
show_source: true
|
||||
show_root_heading: true
|
||||
show_root_toc_entry: true
|
||||
show_submodules: true
|
||||
heading_level: 3
|
||||
|
||||
## fastvideo.configs
|
||||
|
||||
::: fastvideo.configs
|
||||
options:
|
||||
show_source: true
|
||||
show_root_heading: true
|
||||
show_root_toc_entry: true
|
||||
show_submodules: true
|
||||
heading_level: 3
|
||||
|
||||
### Submodules
|
||||
|
||||
#### fastvideo.configs.pipelines
|
||||
|
||||
::: fastvideo.configs.pipelines
|
||||
options:
|
||||
show_source: true
|
||||
show_root_heading: true
|
||||
show_root_toc_entry: true
|
||||
heading_level: 4
|
||||
|
||||
#### fastvideo.configs.models
|
||||
|
||||
::: fastvideo.configs.models
|
||||
options:
|
||||
show_source: true
|
||||
show_root_heading: true
|
||||
show_root_toc_entry: true
|
||||
heading_level: 4
|
||||
|
||||
#### fastvideo.configs.sample
|
||||
|
||||
::: fastvideo.configs.sample
|
||||
options:
|
||||
show_source: true
|
||||
show_root_heading: true
|
||||
show_root_toc_entry: true
|
||||
heading_level: 4
|
||||
|
||||
## fastvideo.distributed
|
||||
|
||||
::: fastvideo.distributed
|
||||
options:
|
||||
show_source: true
|
||||
show_root_heading: true
|
||||
show_root_toc_entry: true
|
||||
show_submodules: true
|
||||
heading_level: 3
|
||||
|
||||
## fastvideo.entrypoints
|
||||
|
||||
::: fastvideo.entrypoints
|
||||
options:
|
||||
show_source: true
|
||||
show_root_heading: true
|
||||
show_root_toc_entry: true
|
||||
show_submodules: true
|
||||
heading_level: 3
|
||||
|
||||
## fastvideo.models
|
||||
|
||||
::: fastvideo.models
|
||||
options:
|
||||
show_source: true
|
||||
show_root_heading: true
|
||||
show_root_toc_entry: true
|
||||
show_submodules: true
|
||||
heading_level: 3
|
||||
|
||||
### Submodules
|
||||
|
||||
#### fastvideo.models.registry
|
||||
|
||||
::: fastvideo.models.registry
|
||||
options:
|
||||
show_source: true
|
||||
show_root_heading: true
|
||||
show_root_toc_entry: true
|
||||
heading_level: 4
|
||||
|
||||
#### fastvideo.models.loader
|
||||
|
||||
::: fastvideo.models.loader
|
||||
options:
|
||||
show_source: true
|
||||
show_root_heading: true
|
||||
show_root_toc_entry: true
|
||||
heading_level: 4
|
||||
|
||||
## fastvideo.pipelines
|
||||
|
||||
::: fastvideo.pipelines
|
||||
options:
|
||||
show_source: true
|
||||
show_root_heading: true
|
||||
show_root_toc_entry: true
|
||||
show_submodules: true
|
||||
heading_level: 3
|
||||
|
||||
### Submodules
|
||||
|
||||
#### fastvideo.pipelines.composed_pipeline_base
|
||||
|
||||
::: fastvideo.pipelines.composed_pipeline_base
|
||||
options:
|
||||
show_source: true
|
||||
show_root_heading: true
|
||||
show_root_toc_entry: true
|
||||
heading_level: 4
|
||||
|
||||
#### fastvideo.pipelines.lora_pipeline
|
||||
|
||||
::: fastvideo.pipelines.lora_pipeline
|
||||
options:
|
||||
show_source: true
|
||||
show_root_heading: true
|
||||
show_root_toc_entry: true
|
||||
heading_level: 4
|
||||
|
||||
#### fastvideo.pipelines.pipeline_batch_info
|
||||
|
||||
::: fastvideo.pipelines.pipeline_batch_info
|
||||
options:
|
||||
show_source: true
|
||||
show_root_heading: true
|
||||
show_root_toc_entry: true
|
||||
heading_level: 4
|
||||
|
||||
#### fastvideo.pipelines.pipeline_registry
|
||||
|
||||
::: fastvideo.pipelines.pipeline_registry
|
||||
options:
|
||||
show_source: true
|
||||
show_root_heading: true
|
||||
show_root_toc_entry: true
|
||||
heading_level: 4
|
||||
|
||||
#### fastvideo.pipelines.stages
|
||||
|
||||
::: fastvideo.pipelines.stages
|
||||
options:
|
||||
show_source: true
|
||||
show_root_heading: true
|
||||
show_root_toc_entry: true
|
||||
heading_level: 4
|
||||
|
||||
## fastvideo.training
|
||||
|
||||
::: fastvideo.training
|
||||
options:
|
||||
show_source: true
|
||||
show_root_heading: true
|
||||
show_root_toc_entry: true
|
||||
show_submodules: true
|
||||
heading_level: 3
|
||||
|
||||
## fastvideo.workflow
|
||||
|
||||
::: fastvideo.workflow
|
||||
options:
|
||||
show_source: true
|
||||
show_root_heading: true
|
||||
show_root_toc_entry: true
|
||||
show_submodules: true
|
||||
heading_level: 3
|
||||
|
||||
## fastvideo.dataset
|
||||
|
||||
::: fastvideo.dataset
|
||||
options:
|
||||
show_source: true
|
||||
show_root_heading: true
|
||||
show_root_toc_entry: true
|
||||
show_submodules: true
|
||||
heading_level: 3
|
||||
|
||||
## fastvideo.layers
|
||||
|
||||
::: fastvideo.layers
|
||||
options:
|
||||
show_source: true
|
||||
show_root_heading: true
|
||||
show_root_toc_entry: true
|
||||
show_submodules: true
|
||||
heading_level: 3
|
||||
|
||||
## fastvideo.platforms
|
||||
|
||||
::: fastvideo.platforms
|
||||
options:
|
||||
show_source: true
|
||||
show_root_heading: true
|
||||
show_root_toc_entry: true
|
||||
show_submodules: true
|
||||
heading_level: 3
|
||||
|
||||
## fastvideo.utils
|
||||
|
||||
::: fastvideo.utils
|
||||
options:
|
||||
show_source: true
|
||||
show_root_heading: true
|
||||
show_root_toc_entry: true
|
||||
heading_level: 3
|
||||
|
||||
## fastvideo.worker
|
||||
|
||||
::: fastvideo.worker
|
||||
options:
|
||||
show_source: true
|
||||
show_root_heading: true
|
||||
show_root_toc_entry: true
|
||||
show_submodules: true
|
||||
heading_level: 3
|
||||
@@ -1,27 +0,0 @@
|
||||
# API Summary
|
||||
|
||||
This page provides a quick overview of the main FastVideo API components.
|
||||
|
||||
## Video Generator
|
||||
|
||||
::: fastvideo.VideoGenerator
|
||||
options:
|
||||
show_root_heading: false
|
||||
show_source: false
|
||||
heading_level: 3
|
||||
|
||||
## Initialization Configuration
|
||||
|
||||
::: fastvideo.PipelineConfig
|
||||
options:
|
||||
show_root_heading: false
|
||||
show_source: false
|
||||
heading_level: 3
|
||||
|
||||
## Sampling Configuration
|
||||
|
||||
::: fastvideo.SamplingParam
|
||||
options:
|
||||
show_root_heading: false
|
||||
show_source: false
|
||||
heading_level: 3
|
||||
@@ -1,41 +0,0 @@
|
||||
.vertical-table-header th.head:not(.stub) {
|
||||
writing-mode: sideways-lr;
|
||||
white-space: nowrap;
|
||||
max-width: 0;
|
||||
p {
|
||||
margin: 0;
|
||||
}
|
||||
}
|
||||
|
||||
/* Image sizing classes */
|
||||
.image-small {
|
||||
max-width: 200px;
|
||||
height: auto;
|
||||
}
|
||||
|
||||
.image-medium {
|
||||
max-width: 400px;
|
||||
height: auto;
|
||||
}
|
||||
|
||||
.image-large {
|
||||
max-width: 600px;
|
||||
height: auto;
|
||||
}
|
||||
|
||||
.image-full {
|
||||
max-width: 100%;
|
||||
height: auto;
|
||||
}
|
||||
|
||||
/* Responsive images */
|
||||
img {
|
||||
max-width: 100%;
|
||||
height: auto;
|
||||
}
|
||||
|
||||
/* Center images */
|
||||
.image-center {
|
||||
display: block;
|
||||
margin: 0 auto;
|
||||
}
|
||||
|
Before Width: | Height: | Size: 122 KiB |
|
Before Width: | Height: | Size: 378 KiB |
|
Before Width: | Height: | Size: 575 KiB |
@@ -1,6 +0,0 @@
|
||||
<svg width="160" height="93" viewBox="0 0 160 93" fill="none" xmlns="http://www.w3.org/2000/svg">
|
||||
<path d="M28.8511 91.66L57.6319 1.86368H64.5394L35.7585 91.66H28.8511Z" fill="#356CFF" stroke="#356CFF" stroke-width="2.30244"/>
|
||||
<path d="M15.0376 91.66L43.8185 1.86368H46.1209L17.3401 91.66H15.0376Z" fill="#356CFF" stroke="#356CFF" stroke-width="2.30244"/>
|
||||
<path d="M1.22217 91.66L30.003 1.86366H31.1543L2.3734 91.66H1.22217Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.15122"/>
|
||||
<path d="M71.4465 1.86483L42.666 91.6599H69.144L78.3538 58.2746H123.251L129.007 39.855H84.1099L89.866 22.5868H152.032L157.788 1.86483H71.4465Z" fill="#356CFF" stroke="#356CFF" stroke-width="2.30244"/>
|
||||
</svg>
|
||||
|
Before Width: | Height: | Size: 691 B |
@@ -1,18 +0,0 @@
|
||||
<svg width="252" height="105" viewBox="0 0 252 105" fill="none" xmlns="http://www.w3.org/2000/svg">
|
||||
<path d="M89.4843 55.5457H101.361L87.7028 101H74.638L89.4843 55.5457Z" fill="#356CFF"/>
|
||||
<path fill-rule="evenodd" clip-rule="evenodd" d="M96.0167 1.00057H112.645L118.583 48.273H104.924L103.737 39.7882H79.9827L67.5117 55.5457H85.3273L43.1638 101H28.3174L22.3789 55.5457H33.6621L38.4129 91.3031L58.604 68.2729H44.3515L96.0167 1.00057ZM100.768 13.1217L87.7028 29.4852H103.143L100.768 13.1217Z" fill="#356CFF"/>
|
||||
<path d="M37.2252 1.00057L22.3789 48.273H36.0375L40.7884 30.6974L62.6727 30.6974L69.6727 21.0005L43.7576 21.0004L46.7269 11.9096L77.6727 11.9096L86 1.00057L37.2252 1.00057Z" fill="#356CFF"/>
|
||||
<path fill-rule="evenodd" clip-rule="evenodd" d="M108.488 55.5457L94.2351 101C94.2351 101 105.518 101 120.959 101C136.399 101 144.078 93.0133 148.276 79.788C152.432 68.0157 153.027 55.5457 136.399 55.5457C119.771 55.5457 108.488 55.5457 108.488 55.5457ZM109.081 90.697L116.802 65.8487C116.802 65.8487 120.959 65.8487 132.242 65.8487C143.525 65.8487 137.586 78.5759 135.211 84.0304C133.307 88.4021 127.491 90.697 122.74 90.697C117.989 90.697 109.081 90.697 109.081 90.697Z" fill="#356CFF"/>
|
||||
<path d="M173.188 1.00056L168.625 11.9096C168.625 11.9096 149.386 11.9092 142.525 11.9095C135.664 11.9098 136.586 20.3944 141.337 20.3944H159.747C168.654 20.3944 166.961 33.6899 163.904 38.5761C160.467 44.0675 157.371 48.273 148.463 48.273L125.188 48.273L124 37.97L147.87 37.97C153.808 37.97 156.184 29.4852 151.433 29.4852H131.836C120.142 29.4852 125.897 1.00043 141.337 1.00043L173.188 1.00056Z" fill="#356CFF"/>
|
||||
<path d="M179.938 1.00056L175.688 11.9096L191.221 11.9096L179.938 48.273H192.409L203.692 11.9096L219.132 11.9095L223.289 1.00043L179.938 1.00056Z" fill="#356CFF"/>
|
||||
<path d="M161.341 55.5457H202.845L198.5 65.8487H169.654L167.279 73.7268H188.5L184.749 82.8177H164.31L161.934 90.697H190.251L186.624 101H146.494L161.341 55.5457Z" fill="#356CFF"/>
|
||||
<path fill-rule="evenodd" clip-rule="evenodd" d="M230.821 54.9391C255.169 54.9391 251.776 67.0602 249.231 77.9692C246.686 88.8783 240.917 101 217.757 101C194.596 101 195.606 88.8783 199.347 77.9692C203.089 67.0602 206.473 54.9391 230.821 54.9391ZM237.948 77.9692C239.984 70.6965 240.917 65.242 228.446 65.242C215.975 65.242 211.818 71.9087 210.037 77.9692C208.255 84.0298 208.255 91.3025 219.538 91.3025C230.821 91.3025 235.911 85.2419 237.948 77.9692Z" fill="#356CFF"/>
|
||||
<path d="M173.188 1.00056L168.625 11.9096C168.625 11.9096 149.386 11.9092 142.525 11.9095C135.664 11.9098 136.586 20.3944 141.337 20.3944M173.188 1.00056C173.188 1.00056 156.777 1.00043 141.337 1.00043M173.188 1.00056L141.337 1.00043M141.337 20.3944C146.088 20.3944 150.839 20.3944 159.747 20.3944M141.337 20.3944H159.747M159.747 20.3944C168.654 20.3944 166.961 33.6899 163.904 38.5761C160.467 44.0675 157.371 48.273 148.463 48.273M148.463 48.273C139.556 48.273 125.188 48.273 125.188 48.273M148.463 48.273L125.188 48.273M125.188 48.273L124 37.97M124 37.97C124 37.97 141.931 37.97 147.87 37.97M124 37.97L147.87 37.97M147.87 37.97C153.808 37.97 156.184 29.4852 151.433 29.4852M151.433 29.4852C146.682 29.4852 138.962 29.4852 131.836 29.4852M151.433 29.4852H131.836M131.836 29.4852C120.142 29.4852 125.897 1.00043 141.337 1.00043M37.2252 1.00057L22.3789 48.273H36.0375L40.7884 30.6974L62.6727 30.6974L69.6727 21.0005L43.7576 21.0004L46.7269 11.9096L77.6727 11.9096L86 1.00057L37.2252 1.00057ZM96.0167 1.00057H112.645L118.583 48.273H104.924L103.737 39.7882H79.9827L67.5117 55.5457H85.3273L43.1638 101H28.3174L22.3789 55.5457H33.6621L38.4129 91.3031L58.604 68.2729H44.3515L96.0167 1.00057ZM87.7028 29.4852L100.768 13.1217L103.143 29.4852H87.7028ZM89.4843 55.5457H101.361L87.7028 101H74.638L89.4843 55.5457ZM108.488 55.5457L94.2351 101C94.2351 101 105.518 101 120.959 101C136.399 101 144.078 93.0133 148.276 79.788C152.432 68.0157 153.027 55.5457 136.399 55.5457C119.771 55.5457 108.488 55.5457 108.488 55.5457ZM116.802 65.8487L109.081 90.697C109.081 90.697 117.989 90.697 122.74 90.697C127.491 90.697 133.307 88.4021 135.211 84.0304C137.586 78.5759 143.525 65.8487 132.242 65.8487C120.959 65.8487 116.802 65.8487 116.802 65.8487ZM179.938 1.00056L175.688 11.9096L191.221 11.9096L179.938 48.273H192.409L203.692 11.9096L219.132 11.9095L223.289 1.00043L179.938 1.00056ZM161.341 55.5457H202.845L198.5 65.8487H169.654L167.279 73.7268H188.5L184.749 82.8177H164.31L161.934 90.697H190.251L186.624 101H146.494L161.341 55.5457ZM230.821 54.9391C255.169 54.9391 251.776 67.0602 249.231 77.9692C246.686 88.8783 240.917 101 217.757 101C194.596 101 195.606 88.8783 199.347 77.9692C203.089 67.0602 206.473 54.9391 230.821 54.9391ZM228.446 65.242C240.917 65.242 239.984 70.6965 237.948 77.9692C235.911 85.2419 230.821 91.3025 219.538 91.3025C208.255 91.3025 208.255 84.0298 210.037 77.9692C211.818 71.9087 215.975 65.242 228.446 65.242Z" stroke="#356CFF" stroke-width="1.18771"/>
|
||||
<path d="M15.2524 55.5451L21.191 100.999L24.7541 100.999L18.8156 55.5451L15.2524 55.5451Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.18771"/>
|
||||
<path d="M8.12646 55.5451L14.065 100.999L15.2527 100.999L9.31417 55.5451L8.12646 55.5451Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.18771"/>
|
||||
<path d="M1 55.5451L6.93853 100.999L7.53239 100.999L1.59385 55.5451L1 55.5451Z" fill="#356CFF" stroke="#356CFF" stroke-width="0.593853"/>
|
||||
<path d="M15.2524 48.2724L30.0988 1H33.6619L18.8156 48.2724H15.2524Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.18771"/>
|
||||
<path d="M8.12646 48.2724L22.9728 1H24.1605L9.31417 48.2724H8.12646Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.18771"/>
|
||||
<path d="M1 48.2724L15.8463 1H16.4402L1.59385 48.2724H1Z" fill="#356CFF" stroke="#356CFF" stroke-width="0.593853"/>
|
||||
<path d="M85.3271 55.5457H67.5116L87 12.7363L44.3513 68.2729H58.6038L43.1636 101L85.3271 55.5457Z" fill="#FDC717" stroke="#FDC717" stroke-width="1.18771" stroke-miterlimit="16"/>
|
||||
</svg>
|
||||
|
Before Width: | Height: | Size: 5.7 KiB |
@@ -1,129 +0,0 @@
|
||||
# Testing in FastVideo
|
||||
|
||||
This guide explains how to add and run tests in FastVideo. The testing suite is divided into several categories to ensure correctness across components, training workflows, and inference quality.
|
||||
|
||||
## Test Types
|
||||
|
||||
* **Unit Tests**: Located in `fastvideo/tests/dataset`, `fastvideo/tests/entrypoints`, and `fastvideo/tests/workflow`. These test individual functions and classes.
|
||||
* **Component Tests**: Located in `fastvideo/tests/encoders`, `fastvideo/tests/transformers`, and `fastvideo/tests/vaes`. These verify the loading and basic functionality of model components.
|
||||
* **SSIM Tests**: Located in `fastvideo/tests/ssim`. These are regression tests that compare generated videos against reference videos using the Structural Similarity Index Measure (SSIM) to detect quality degradation.
|
||||
* **Training Tests**: Located in `fastvideo/tests/training`. These validate training loops, loss calculations, and specific training techniques like LoRA, Distillation, and VSA.
|
||||
* **Inference Tests**: Located in `fastvideo/tests/inference`. These test specialized inference pipelines and optimizations (e.g., STA, V-MoBA).
|
||||
|
||||
For now, we will focus on **SSIM Tests**.
|
||||
|
||||
## SSIM Tests
|
||||
|
||||
SSIM tests are located in `fastvideo/tests/ssim`. These tests generate videos using specific models and parameters, and compare them against reference videos to ensure that changes in the codebase do not degrade generation quality or alter the output unexpectedly.
|
||||
|
||||
!!! note
|
||||
If you are adding an SSIM test, this serves as a safeguard. Any future code changes that break or cause errors with the specific arguments and configurations you defined will trigger a failure. Therefore, it is important to include multiple settings and arguments that cover the core features of your new pipeline to ensure robust regression testing.
|
||||
|
||||
### Directory Structure
|
||||
|
||||
```
|
||||
fastvideo/tests/ssim/
|
||||
├── <GPU>_reference_videos/ # Reference videos organized by GPU type (e.g., L40S_reference_videos)
|
||||
│ ├── <Model_Name>/
|
||||
│ │ ├── <Backend>/ # e.g., FLASH_ATTN, TORCH_SDPA
|
||||
│ │ │ └── <Video_File>
|
||||
├── test_causal_similarity.py
|
||||
├── test_inference_similarity.py
|
||||
├── update_reference_videos.sh
|
||||
└── ...
|
||||
```
|
||||
|
||||
### Adding a New SSIM Test
|
||||
|
||||
To add a new SSIM test, follow these steps:
|
||||
|
||||
1. **Create or Update a Test File**: You can add a new test function to an existing file (like `test_inference_similarity.py`) or create a new one if testing a distinct category of models.
|
||||
|
||||
2. **Define Model Parameters**: Define the configuration for the model you want to test. This includes model path, dimensions, inference steps, and other generation parameters. **Note:** Consider using lower `num_inference_steps` or reduced resolution (e.g., 480p instead of 720p) to keep test execution time reasonable, provided it doesn't compromise the test's ability to detect regression.
|
||||
|
||||
```python
|
||||
MY_MODEL_PARAMS = {
|
||||
"num_gpus": 1,
|
||||
"model_path": "organization/model-name",
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 45,
|
||||
"num_inference_steps": 20,
|
||||
# ... other parameters
|
||||
}
|
||||
```
|
||||
|
||||
3. **Implement the Test Function**:
|
||||
* Use `pytest.mark.parametrize` to run the test with different prompts, backends, and models.
|
||||
* Set the attention backend environment variable.
|
||||
* Initialize the `VideoGenerator`.
|
||||
* Generate the video.
|
||||
* Compare the generated video with the reference video using `compute_video_ssim_torchvision`.
|
||||
|
||||
Example structure:
|
||||
|
||||
```python
|
||||
@pytest.mark.parametrize("prompt", TEST_PROMPTS)
|
||||
@pytest.mark.parametrize("ATTENTION_BACKEND", ["FLASH_ATTN"])
|
||||
def test_my_model_similarity(prompt, ATTENTION_BACKEND):
|
||||
# Setup output directories
|
||||
# ...
|
||||
|
||||
# Initialize Generator
|
||||
generator = VideoGenerator.from_pretrained(...)
|
||||
generator.generate_video(prompt, ...)
|
||||
|
||||
# Compare with Reference
|
||||
ssim_values = compute_video_ssim_torchvision(reference_path, generated_path, use_ms_ssim=True)
|
||||
assert ssim_values[0] >= 0.98 # Threshold
|
||||
```
|
||||
|
||||
4. **Reference Videos**:
|
||||
* When running the test for the first time (or when updating the reference), the test will fail because the reference video is missing. The generated video will be saved in `fastvideo/tests/ssim/generated_videos`.
|
||||
* Inspect the generated video to ensure it meets quality expectations.
|
||||
* Move the generated video to the appropriate reference folder: `fastvideo/tests/ssim/<GPU>_reference_videos/<Model>/<Backend>/`.
|
||||
* You can use the helper script `update_reference_videos.sh` to automate copying videos from `generated_videos` to `L40S_reference_videos`. Note: Check the script to ensure paths match your environment (it defaults to `L40S_reference_videos`).
|
||||
|
||||
### Running Tests Locally
|
||||
|
||||
To run the SSIM tests locally:
|
||||
|
||||
```bash
|
||||
pytest fastvideo/tests/ssim/ -vs
|
||||
```
|
||||
|
||||
Ensure you have the necessary GPUs available as defined in your test parameters.
|
||||
|
||||
## Modal Workflow
|
||||
|
||||
FastVideo uses [Modal](https://modal.com/) for running tests in a CI environment. The workflow scripts are located in `fastvideo/tests/modal/`.
|
||||
|
||||
### `pr_test.py`
|
||||
|
||||
The main entry point for CI tests is `fastvideo/tests/modal/pr_test.py`. This script defines Modal functions that execute the pytest suites on specific hardware (e.g., L40S, H100).
|
||||
|
||||
### Updating Modal Configuration
|
||||
|
||||
If you add a new test that requires:
|
||||
* **Different GPU Hardware**: You may need to change the `@app.function(gpu=...)` decorator.
|
||||
* **Longer Execution Time**: Increase the `timeout` parameter.
|
||||
* **New Environment Variables/Secrets**: Add them to `secrets=[...]` or the image environment. For example, if your model is gated on Hugging Face, ensure `HF_API_KEY` is passed.
|
||||
|
||||
For SSIM tests, the `run_ssim_tests` function in `pr_test.py` currently runs:
|
||||
|
||||
```python
|
||||
@app.function(gpu="L40S:2", image=image, timeout=2700, secrets=[modal.Secret.from_dict({"HF_API_KEY": os.environ.get("HF_API_KEY", "")})])
|
||||
def run_ssim_tests():
|
||||
run_test("hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/ssim -vs")
|
||||
```
|
||||
|
||||
If your new test file is inside `fastvideo/tests/ssim`, it will automatically be picked up by this command. However, ensure that the `gpu="L40S:2"` configuration is sufficient for your model. If your model requires more GPUs (e.g., 4 or 8), you might need to create a separate Modal function or update the existing one.
|
||||
|
||||
### Workflow Scripts
|
||||
|
||||
The shell script that triggers these tests in the CI pipeline is located at `.buildkite/scripts/pr_test.sh`. If you add a new test category (e.g., a new folder outside of `ssim`), you will need to:
|
||||
1. Add a new function in `fastvideo/tests/modal/pr_test.py`.
|
||||
2. Add a new case in `.buildkite/scripts/pr_test.sh` to handle the new test type.
|
||||
|
||||
!!! note
|
||||
If you are a maintainer, you'll need to finally manually update the workflow script in Buildkite. Otherwise, a maintainer will help you update.
|
||||
@@ -1,12 +0,0 @@
|
||||
# 💡 Examples
|
||||
|
||||
A collection of examples demonstrating usage of FastVideo.
|
||||
|
||||
All documented examples are autogenerated using [generate_examples.py](https://github.com/hao-ai-lab/FastVideo/blob/main/docs/generate_examples.py) from examples found in the [examples](https://github.com/hao-ai-lab/FastVideo/tree/main/examples) directory.
|
||||
|
||||
## Examples
|
||||
|
||||
- [Examples Distillation Index](distillation/examples/examples_distillation_index.md)
|
||||
- [Examples Training Index](training/examples/examples_training_index.md)
|
||||
- [Examples Inference Index](inference/examples/examples_inference_index.md)
|
||||
|
||||
@@ -1,41 +0,0 @@
|
||||
|
||||
# 🔧 Installation
|
||||
|
||||
FastVideo supports the following hardware platforms:
|
||||
|
||||
- [NVIDIA CUDA](installation/gpu.md)
|
||||
- [Apple silicon](installation/mps.md)
|
||||
|
||||
## Quick Installation
|
||||
|
||||
### Using pip
|
||||
|
||||
```bash
|
||||
pip install fastvideo
|
||||
```
|
||||
|
||||
### Using conda
|
||||
|
||||
```bash
|
||||
conda install -c conda-forge fastvideo
|
||||
```
|
||||
|
||||
### From source
|
||||
|
||||
```bash
|
||||
git clone https://github.com/hao-ai-lab/FastVideo.git
|
||||
cd FastVideo
|
||||
pip install -e .
|
||||
```
|
||||
|
||||
## Hardware Requirements
|
||||
|
||||
- **NVIDIA GPUs**: CUDA 11.8+ with compute capability 7.0+
|
||||
- **Apple Silicon**: macOS 12.0+ with M1/M2/M3 chips
|
||||
- **CPU**: x86_64 architecture (for CPU-only inference)
|
||||
|
||||
## Next Steps
|
||||
|
||||
- [Quick Start Guide](quick_start.md) - Get started with your first video generation
|
||||
- [Configuration](../inference/configuration.md) - Learn about configuration options
|
||||
- [Examples](../inference/examples/) - Explore example scripts and notebooks
|
||||
@@ -1,55 +0,0 @@
|
||||
# 🚀 Quick Start
|
||||
|
||||
Get up and running with FastVideo in minutes!
|
||||
|
||||
## Installation
|
||||
|
||||
First, install FastVideo:
|
||||
|
||||
```bash
|
||||
pip install fastvideo
|
||||
```
|
||||
|
||||
## Basic Usage
|
||||
|
||||
### Text-to-Video Generation
|
||||
|
||||
```python
|
||||
from fastvideo import FastVideoPipeline
|
||||
|
||||
# Initialize the pipeline
|
||||
pipe = FastVideoPipeline.from_pretrained("wan2.1-t2v-1.3B")
|
||||
|
||||
# Generate a video
|
||||
prompt = "A cat playing with a ball of yarn"
|
||||
video = pipe(prompt, num_frames=16, height=512, width=512)
|
||||
|
||||
# Save the video
|
||||
video.save("output.mp4")
|
||||
```
|
||||
|
||||
### Image-to-Video Generation
|
||||
|
||||
```python
|
||||
from fastvideo import FastVideoPipeline
|
||||
from PIL import Image
|
||||
|
||||
# Load an image
|
||||
image = Image.open("input.jpg")
|
||||
|
||||
# Initialize the pipeline
|
||||
pipe = FastVideoPipeline.from_pretrained("wan2.1-i2v-14B-480p")
|
||||
|
||||
# Generate a video from the image
|
||||
video = pipe(image, num_frames=16, height=480, width=480)
|
||||
|
||||
# Save the video
|
||||
video.save("output.mp4")
|
||||
```
|
||||
|
||||
## Next Steps
|
||||
|
||||
- [Installation Guide](installation.md) - Detailed installation instructions
|
||||
- [Configuration](../inference/configuration.md) - Learn about configuration options
|
||||
- [Examples](../inference/examples/) - Explore more examples
|
||||
- [Optimizations](../inference/optimizations.md) - Performance optimization tips
|
||||
@@ -1,42 +0,0 @@
|
||||
# V1 API
|
||||
|
||||
FastVideo's V1 API provides a streamlined interface for video generation tasks with powerful customization options. This page documents the primary components of the API.
|
||||
|
||||
## Video Generator
|
||||
|
||||
This class will be the primary Python API for generating videos and images.
|
||||
|
||||
::: fastvideo.entrypoints.video_generator.VideoGenerator
|
||||
options:
|
||||
show_root_heading: true
|
||||
show_source: false
|
||||
members:
|
||||
- from_pretrained
|
||||
heading_level: 3
|
||||
|
||||
`VideoGenerator.from_pretrained()` should be the primary way of creating a new video generator.
|
||||
|
||||
## Configuring FastVideo
|
||||
|
||||
The following two classes `PipelineConfig` and `SamplingParam` are used to configure initialization and sampling parameters, respectively.
|
||||
|
||||
### PipelineConfig
|
||||
|
||||
::: fastvideo.configs.pipelines.base.PipelineConfig
|
||||
options:
|
||||
show_root_heading: true
|
||||
show_source: false
|
||||
members:
|
||||
- from_pretrained
|
||||
- dump_to_json
|
||||
heading_level: 4
|
||||
|
||||
### SamplingParam
|
||||
|
||||
::: fastvideo.configs.sample.base.SamplingParam
|
||||
options:
|
||||
show_root_heading: true
|
||||
show_source: false
|
||||
members:
|
||||
- from_pretrained
|
||||
heading_level: 4
|
||||
@@ -1,66 +0,0 @@
|
||||
# Compatibility Matrix
|
||||
|
||||
The table below shows every supported model and optimizations supported for them.
|
||||
|
||||
The symbols used have the following meanings:
|
||||
|
||||
- ✅ = Full compatibility
|
||||
- ❌ = No compatibility
|
||||
- ⭕ = Does not apply to this model
|
||||
|
||||
## Models x Optimization
|
||||
|
||||
The `HuggingFace Model ID` can be directly pass to `from_pretrained()` methods and FastVideo will use the optimal default parameters when initializing and generating videos.
|
||||
|
||||
<style>
|
||||
/* Target tables in this section */
|
||||
#models-x-optimization + p + table {
|
||||
display: block;
|
||||
overflow-x: auto;
|
||||
width: 100%;
|
||||
font-size: 0.85rem;
|
||||
}
|
||||
|
||||
#models-x-optimization + p + table td,
|
||||
#models-x-optimization + p + table th {
|
||||
text-align: center;
|
||||
white-space: nowrap;
|
||||
padding: 0.5em;
|
||||
}
|
||||
|
||||
/* First two columns can wrap */
|
||||
#models-x-optimization + p + table td:nth-child(1),
|
||||
#models-x-optimization + p + table td:nth-child(2) {
|
||||
white-space: normal;
|
||||
min-width: 120px;
|
||||
}
|
||||
|
||||
#models-x-optimization + p + table td:nth-child(2) code {
|
||||
font-size: 0.75rem;
|
||||
}
|
||||
</style>
|
||||
|
||||
| Model Name | HuggingFace Model ID | Resolutions | TeaCache | Sliding Tile Attn | Sage Attn | VSA |
|
||||
|------------|---------------------|-------------|----------|-------------------|-----------|-----|
|
||||
| FastWan2.1 T2V 1.3B | `FastVideo/FastWan2.1-T2V-1.3B-Diffusers` | 480P | ⭕ | ⭕ | ⭕ | ✅ |
|
||||
| FastWan2.2 TI2V 5B Full Attn* | `FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers` | 720P | ⭕ | ⭕ | ⭕ | ✅ |
|
||||
| Wan2.2 TI2V 5B | `Wan-AI/Wan2.2-TI2V-5B-Diffusers` | 720P | ⭕ | ⭕ | ✅ | ⭕ |
|
||||
| Wan2.2 T2V A14B | `Wan-AI/Wan2.2-T2V-A14B-Diffusers` | 480P<br>720P | ❌ | ❌ | ✅ | ⭕ |
|
||||
| Wan2.2 I2V A14B | `Wan-AI/Wan2.2-I2V-A14B-Diffusers` | 480P<br>720P | ❌ | ❌ | ✅ | ⭕ |
|
||||
| HunyuanVideo | `hunyuanvideo-community/HunyuanVideo` | 720px1280p<br>544px960p | ❌ | ✅ | ✅ | ⭕ |
|
||||
| FastHunyuan | `FastVideo/FastHunyuan-diffusers` | 720px1280p<br>544px960p | ❌ | ✅ | ✅ | ⭕ |
|
||||
| Wan2.1 T2V 1.3B | `Wan-AI/Wan2.1-T2V-1.3B-Diffusers` | 480P | ✅ | ✅* | ✅ | ⭕ |
|
||||
| Wan2.1 T2V 14B | `Wan-AI/Wan2.1-T2V-14B-Diffusers` | 480P, 720P | ✅ | ✅* | ✅ | ⭕ |
|
||||
| Wan2.1 I2V 480P | `Wan-AI/Wan2.1-I2V-14B-480P-Diffusers` | 480P | ✅ | ✅* | ✅ | ⭕ |
|
||||
| Wan2.1 I2V 720P | `Wan-AI/Wan2.1-I2V-14B-720P-Diffusers` | 720P | ✅ | ✅ | ✅ | ⭕ |
|
||||
| StepVideo T2V | `FastVideo/stepvideo-t2v-diffusers` | 768px768px204f<br>544px992px204f<br>544px992px136f | ❌ | ❌ | ✅ | ⭕ |
|
||||
|
||||
**Note**: Wan2.2 TI2V 5B has some quality issues when performing I2V generation. We are working on fixing this issue.
|
||||
|
||||
## Special requirements
|
||||
|
||||
### StepVideo T2V
|
||||
- The self-attention in text-encoder (step_llm) only supports CUDA capabilities sm_80 sm_86 and sm_90
|
||||
|
||||
### Sliding Tile Attention
|
||||
- Currently only Hopper GPUs (H100s) are supported.
|
||||
@@ -0,0 +1,35 @@
|
||||
@ECHO OFF
|
||||
|
||||
pushd %~dp0
|
||||
|
||||
REM Command file for Sphinx documentation
|
||||
|
||||
if "%SPHINXBUILD%" == "" (
|
||||
set SPHINXBUILD=sphinx-build
|
||||
)
|
||||
set SOURCEDIR=source
|
||||
set BUILDDIR=build
|
||||
|
||||
%SPHINXBUILD% >NUL 2>NUL
|
||||
if errorlevel 9009 (
|
||||
echo.
|
||||
echo.The 'sphinx-build' command was not found. Make sure you have Sphinx
|
||||
echo.installed, then set the SPHINXBUILD environment variable to point
|
||||
echo.to the full path of the 'sphinx-build' executable. Alternatively you
|
||||
echo.may add the Sphinx directory to PATH.
|
||||
echo.
|
||||
echo.If you don't have Sphinx installed, grab it from
|
||||
echo.https://www.sphinx-doc.org/
|
||||
exit /b 1
|
||||
)
|
||||
|
||||
if "%1" == "" goto help
|
||||
|
||||
%SPHINXBUILD% -M %1 %SOURCEDIR% %BUILDDIR% %SPHINXOPTS% %O%
|
||||
goto end
|
||||
|
||||
:help
|
||||
%SPHINXBUILD% -M help %SOURCEDIR% %BUILDDIR% %SPHINXOPTS% %O%
|
||||
|
||||
:end
|
||||
popd
|
||||
@@ -0,0 +1,15 @@
|
||||
sphinx==7.4.7
|
||||
sphinx-argparse==0.5.2
|
||||
sphinx-autodoc2==0.5.0
|
||||
sphinx-book-theme==1.1.4
|
||||
sphinx-copybutton==0.5.2
|
||||
sphinx-design==0.6.1
|
||||
sphinx-togglebutton==0.3.2
|
||||
myst-parser==3.0.1
|
||||
msgspec
|
||||
commonmark # Required by sphinx-argparse when using :markdownhelp:
|
||||
|
||||
# packages to install to build the documentation
|
||||
cachetools
|
||||
# -f https://download.pytorch.org/whl/cpu
|
||||
torch
|
||||
@@ -0,0 +1,51 @@
|
||||
# Seed Parameter Behavior in vLLM
|
||||
|
||||
## Overview
|
||||
|
||||
The `seed` parameter in vLLM is used to control the random states for various random number generators. This parameter can affect the behavior of random operations in user code, especially when working with models in vLLM.
|
||||
|
||||
## Default Behavior
|
||||
|
||||
By default, the `seed` parameter is set to `None`. When the `seed` parameter is `None`, the global random states for `random`, `np.random`, and `torch.manual_seed` are not set. This means that the random operations will behave as expected, without any fixed random states.
|
||||
|
||||
## Specifying a Seed
|
||||
|
||||
If a specific seed value is provided, the global random states for `random`, `np.random`, and `torch.manual_seed` will be set accordingly. This can be useful for reproducibility, as it ensures that the random operations produce the same results across multiple runs.
|
||||
|
||||
## Example Usage
|
||||
|
||||
### Without Specifying a Seed
|
||||
|
||||
```python
|
||||
import random
|
||||
from vllm import LLM
|
||||
|
||||
# Initialize a vLLM model without specifying a seed
|
||||
model = LLM(model="Qwen/Qwen2.5-0.5B-Instruct")
|
||||
|
||||
# Try generating random numbers
|
||||
print(random.randint(0, 100)) # Outputs different numbers across runs
|
||||
```
|
||||
|
||||
### Specifying a Seed
|
||||
|
||||
```python
|
||||
import random
|
||||
from vllm import LLM
|
||||
|
||||
# Initialize a vLLM model with a specific seed
|
||||
model = LLM(model="Qwen/Qwen2.5-0.5B-Instruct", seed=42)
|
||||
|
||||
# Try generating random numbers
|
||||
print(random.randint(0, 100)) # Outputs the same number across runs
|
||||
```
|
||||
|
||||
## Important Notes
|
||||
|
||||
- If the `seed` parameter is not specified, the behavior of global random states remains unaffected.
|
||||
- If a specific seed value is provided, the global random states for `random`, `np.random`, and `torch.manual_seed` will be set to that value.
|
||||
- This behavior can be useful for reproducibility but may lead to non-intuitive behavior if the user is not explicitly aware of it.
|
||||
|
||||
## Conclusion
|
||||
|
||||
Understanding the behavior of the `seed` parameter in vLLM is crucial for ensuring the expected behavior of random operations in your code. By default, the `seed` parameter is set to `None`, which means that the global random states are not affected. However, specifying a seed value can help achieve reproducibility in your experiments.
|
||||
@@ -0,0 +1,8 @@
|
||||
.vertical-table-header th.head:not(.stub) {
|
||||
writing-mode: sideways-lr;
|
||||
white-space: nowrap;
|
||||
max-width: 0;
|
||||
p {
|
||||
margin: 0;
|
||||
}
|
||||
}
|
||||
|
Before Width: | Height: | Size: 98 KiB After Width: | Height: | Size: 98 KiB |
|
Before Width: | Height: | Size: 194 KiB After Width: | Height: | Size: 194 KiB |
|
Before Width: | Height: | Size: 303 KiB After Width: | Height: | Size: 303 KiB |
|
Before Width: | Height: | Size: 18 KiB After Width: | Height: | Size: 18 KiB |
|
Before Width: | Height: | Size: 27 KiB After Width: | Height: | Size: 27 KiB |
|
Before Width: | Height: | Size: 40 KiB After Width: | Height: | Size: 40 KiB |
@@ -0,0 +1,39 @@
|
||||
<style>
|
||||
.notification-bar {
|
||||
width: 100vw;
|
||||
display: flex;
|
||||
justify-content: center;
|
||||
align-items: center;
|
||||
font-size: 16px;
|
||||
padding: 0 6px 0 6px;
|
||||
}
|
||||
.notification-bar p {
|
||||
margin: 0;
|
||||
}
|
||||
.notification-bar a {
|
||||
font-weight: bold;
|
||||
text-decoration: none;
|
||||
}
|
||||
|
||||
/* Light mode styles (default) */
|
||||
.notification-bar {
|
||||
background-color: #fff3cd;
|
||||
color: #856404;
|
||||
}
|
||||
.notification-bar a {
|
||||
color: #d97706;
|
||||
}
|
||||
|
||||
/* Dark mode styles */
|
||||
html[data-theme=dark] .notification-bar {
|
||||
background-color: #333;
|
||||
color: #ddd;
|
||||
}
|
||||
html[data-theme=dark] .notification-bar a {
|
||||
color: #ffa500; /* Brighter color for visibility */
|
||||
}
|
||||
</style>
|
||||
|
||||
<!-- <div class="notification-bar">
|
||||
<p>You are viewing the latest developer preview docs. <a href="https://docs.vllm.ai/en/stable/">Click here</a> to view docs for the latest stable release.</p>
|
||||
</div> -->
|
||||
@@ -0,0 +1,19 @@
|
||||
# Summary
|
||||
|
||||
## Video Generator
|
||||
|
||||
```{autodoc2-summary}
|
||||
fastvideo.VideoGenerator
|
||||
```
|
||||
|
||||
## Initialization Configuration
|
||||
|
||||
```{autodoc2-summary}
|
||||
fastvideo.configs.pipelines.PipelineConfig
|
||||
```
|
||||
|
||||
## Sampling Configuration
|
||||
|
||||
```{autodoc2-summary}
|
||||
fastvideo.configs.sample.SamplingParam
|
||||
```
|
||||
@@ -0,0 +1,22 @@
|
||||
# type: ignore
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from docutils import nodes
|
||||
from myst_parser.parsers.sphinx_ import MystParser
|
||||
from sphinx.ext.napoleon import docstring
|
||||
|
||||
|
||||
class NapoleonParser(MystParser):
|
||||
|
||||
def parse(self, input_string: str, document: nodes.document) -> None:
|
||||
# Get the Sphinx configuration
|
||||
config = document.settings.env.config
|
||||
|
||||
parsed_content = str(
|
||||
docstring.GoogleDocstring(
|
||||
str(docstring.NumpyDocstring(input_string, config)),
|
||||
config,
|
||||
))
|
||||
return super().parse(parsed_content, document)
|
||||
|
||||
|
||||
Parser = NapoleonParser
|
||||
@@ -0,0 +1,275 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
# Configuration file for the Sphinx documentation builder.
|
||||
#
|
||||
# This file only contains a selection of the most common options. For a full
|
||||
# list see the documentation:
|
||||
# https://www.sphinx-doc.org/en/master/usage/configuration.html
|
||||
|
||||
# -- Path setup --------------------------------------------------------------
|
||||
|
||||
# If extensions (or modules to document with autodoc) are in another directory,
|
||||
# add these directories to sys.path here. If the directory is relative to the
|
||||
# documentation root, use os.path.abspath to make it absolute, like shown here.
|
||||
|
||||
import datetime
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import requests
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
REPO_ROOT = Path(__file__).resolve().parent.parent.parent
|
||||
print(os.path.abspath(REPO_ROOT))
|
||||
sys.path.append(os.path.abspath(REPO_ROOT))
|
||||
|
||||
# -- Project information -----------------------------------------------------
|
||||
|
||||
project = 'FastVideo'
|
||||
copyright = f'{datetime.datetime.now().year}, FastVideo Team'
|
||||
author = 'the FastVideo Team'
|
||||
|
||||
# -- General configuration ---------------------------------------------------
|
||||
|
||||
# Add any Sphinx extension module names here, as strings. They can be
|
||||
# extensions coming with Sphinx (named 'sphinx.ext.*') or your custom
|
||||
# ones.
|
||||
extensions = [
|
||||
"sphinx.ext.napoleon",
|
||||
"sphinx.ext.linkcode",
|
||||
"sphinx.ext.intersphinx",
|
||||
"sphinx_copybutton",
|
||||
"autodoc2",
|
||||
"myst_parser",
|
||||
"sphinxarg.ext",
|
||||
"sphinx_design",
|
||||
"sphinx_togglebutton",
|
||||
]
|
||||
myst_enable_extensions = [
|
||||
"colon_fence",
|
||||
"fieldlist",
|
||||
]
|
||||
autodoc2_packages = [
|
||||
{
|
||||
"path": "../../fastvideo",
|
||||
"exclude_dirs": ["__pycache__", "third_party"],
|
||||
},
|
||||
]
|
||||
autodoc2_output_dir = "api"
|
||||
autodoc2_render_plugin = "myst"
|
||||
autodoc2_hidden_objects = ["dunder", "private", "inherited"]
|
||||
autodoc2_docstring_parser_regexes = [
|
||||
(".*", "docs.source.autodoc2_docstring_parser"),
|
||||
]
|
||||
autodoc2_sort_names = True
|
||||
autodoc2_index_template = None
|
||||
autodoc2_skip_module_regexes = [
|
||||
"fastvideo.dataset",
|
||||
"fastvideo.distill",
|
||||
"fastvideo.data_preprocess",
|
||||
"fastvideo.models",
|
||||
"fastvideo.sample",
|
||||
"fastvideo.utils",
|
||||
"fastvideo.distill_adv",
|
||||
"fastvideo.train",
|
||||
]
|
||||
|
||||
# Add any paths that contain templates here, relative to this directory.
|
||||
templates_path = ['_templates']
|
||||
|
||||
# List of patterns, relative to source directory, that match files and
|
||||
# directories to ignore when looking for source files.
|
||||
# This pattern also affects html_static_path and html_extra_path.
|
||||
exclude_patterns: list[str] = ["**/*.template.md", "**/*.inc.md"]
|
||||
|
||||
# Exclude the prompt "$" when copying code
|
||||
copybutton_prompt_text = r"\$ "
|
||||
copybutton_prompt_is_regexp = True
|
||||
|
||||
# -- Options for HTML output -------------------------------------------------
|
||||
|
||||
# The theme to use for HTML and HTML Help pages. See the documentation for
|
||||
# a list of builtin themes.
|
||||
#
|
||||
html_title = project
|
||||
html_theme = 'sphinx_book_theme'
|
||||
html_logo = '../../assets/logos/icon_simple.svg'
|
||||
html_theme_options = {
|
||||
'path_to_docs': 'docs/source',
|
||||
'repository_url': 'https://github.com/hao-ai-lab/FastVideo/',
|
||||
'use_repository_button': True,
|
||||
'use_edit_page_button': True,
|
||||
# Prevents the full API being added to the left sidebar of every page.
|
||||
# Reduces build time by 2.5x and reduces build size from ~225MB to ~95MB.
|
||||
'collapse_navbar': True,
|
||||
# Makes API visible in the right sidebar on API reference pages.
|
||||
'show_toc_level': 3,
|
||||
}
|
||||
# Add any paths that contain custom static files (such as style sheets) here,
|
||||
# relative to this directory. They are copied after the builtin static files,
|
||||
# so a file named "default.css" will overwrite the builtin "default.css".
|
||||
html_static_path = ["_static"]
|
||||
html_js_files = ["custom.js"]
|
||||
html_css_files = ["custom.css"]
|
||||
|
||||
myst_url_schemes = {
|
||||
'http': None,
|
||||
'https': None,
|
||||
'mailto': None,
|
||||
'ftp': None,
|
||||
"gh-issue": {
|
||||
"url":
|
||||
"https://github.com/hao-ai-lab/FastVideo/issues/{{path}}#{{fragment}}",
|
||||
"title": "Issue #{{path}}",
|
||||
"classes": ["github"],
|
||||
},
|
||||
"gh-pr": {
|
||||
"url":
|
||||
"https://github.com/hao-ai-lab/FastVideo/pull/{{path}}#{{fragment}}",
|
||||
"title": "Pull Request #{{path}}",
|
||||
"classes": ["github"],
|
||||
},
|
||||
"gh-dir": {
|
||||
"url": "https://github.com/hao-ai-lab/FastVideo/tree/main/{{path}}",
|
||||
"title": "{{path}}",
|
||||
"classes": ["github"],
|
||||
},
|
||||
"gh-file": {
|
||||
"url": "https://github.com/hao-ai-lab/FastVideo/blob/main/{{path}}",
|
||||
"title": "{{path}}",
|
||||
"classes": ["github"],
|
||||
},
|
||||
}
|
||||
|
||||
# see https://docs.readthedocs.io/en/stable/reference/environment-variables.html # noqa
|
||||
READTHEDOCS_VERSION_TYPE = os.environ.get('READTHEDOCS_VERSION_TYPE')
|
||||
if READTHEDOCS_VERSION_TYPE == "tag":
|
||||
# remove the warning banner if the version is a tagged release
|
||||
header_file = os.path.join(os.path.dirname(__file__),
|
||||
"_templates/sections/header.html")
|
||||
# The file might be removed already if the build is triggered multiple times
|
||||
# (readthedocs build both HTML and PDF versions separately)
|
||||
if os.path.exists(header_file):
|
||||
os.remove(header_file)
|
||||
|
||||
|
||||
# Generate additional rst documentation here.
|
||||
def setup(app):
|
||||
from docs.source.generate_examples import generate_examples
|
||||
generate_examples()
|
||||
|
||||
|
||||
_cached_base: str = ""
|
||||
_cached_branch: str = ""
|
||||
|
||||
|
||||
def get_repo_base_and_branch(pr_number: str) -> tuple[str | None, str | None]:
|
||||
global _cached_base, _cached_branch
|
||||
if _cached_base and _cached_branch:
|
||||
return _cached_base, _cached_branch
|
||||
|
||||
url = f"https://api.github.com/repos/hao-ai-lab/FastVideo/pulls/{pr_number}"
|
||||
response = requests.get(url)
|
||||
if response.status_code == 200:
|
||||
data = response.json()
|
||||
_cached_base = data['head']['repo']['full_name']
|
||||
_cached_branch = data['head']['ref']
|
||||
return _cached_base, _cached_branch
|
||||
else:
|
||||
logger.error("Failed to fetch PR details: %s", response)
|
||||
return None, None
|
||||
|
||||
|
||||
def linkcode_resolve(domain, info):
|
||||
if domain != 'py':
|
||||
return None
|
||||
if not info['module']:
|
||||
return None
|
||||
|
||||
# Get path from module name
|
||||
file = Path(f"{info['module'].replace('.', '/')}.py")
|
||||
path = REPO_ROOT / file
|
||||
if not path.exists():
|
||||
path = REPO_ROOT / file.with_suffix("") / "__init__.py"
|
||||
if not path.exists():
|
||||
return None
|
||||
|
||||
# Get the line number of the object
|
||||
with open(path) as f:
|
||||
lines = f.readlines()
|
||||
name = info['fullname'].split(".")[-1]
|
||||
pattern = fr"^( {{4}})*((def|class) )?{name}\b.*"
|
||||
for lineno, line in enumerate(lines, 1):
|
||||
if not line or line.startswith("#"):
|
||||
continue
|
||||
if re.match(pattern, line):
|
||||
break
|
||||
|
||||
# If the line number is not found, return None
|
||||
if lineno == len(lines):
|
||||
return None
|
||||
|
||||
# If the line number is found, create the URL
|
||||
filename = path.relative_to(REPO_ROOT)
|
||||
if "checkouts" in path.parts:
|
||||
# a PR build on readthedocs
|
||||
pr_number = REPO_ROOT.name
|
||||
base, branch = get_repo_base_and_branch(pr_number)
|
||||
if base and branch:
|
||||
return f"https://github.com/{base}/blob/{branch}/{filename}#L{lineno}"
|
||||
# Otherwise, link to the source file on the main branch
|
||||
return f"https://github.com/hao-ai-lab/FastVideo/blob/main/{filename}#L{lineno}"
|
||||
|
||||
|
||||
# Mock out external dependencies here, otherwise the autodoc pages may be blank.
|
||||
autodoc_mock_imports = [
|
||||
"blake3",
|
||||
"compressed_tensors",
|
||||
"cpuinfo",
|
||||
"cv2",
|
||||
"torch",
|
||||
"huggingface_hub",
|
||||
"torchvision",
|
||||
"transformers",
|
||||
"psutil",
|
||||
"prometheus_client",
|
||||
"sentencepiece",
|
||||
"vllm._C",
|
||||
"PIL",
|
||||
"numpy",
|
||||
'triton',
|
||||
"tqdm",
|
||||
"tensorizer",
|
||||
"pynvml",
|
||||
"outlines",
|
||||
"xgrammar",
|
||||
"librosa",
|
||||
"soundfile",
|
||||
"gguf",
|
||||
"lark",
|
||||
"decord",
|
||||
]
|
||||
|
||||
for mock_target in autodoc_mock_imports:
|
||||
if mock_target in sys.modules:
|
||||
logger.info(
|
||||
"Potentially problematic mock target (%s) found; "
|
||||
"autodoc_mock_imports cannot mock modules that have already "
|
||||
"been loaded into sys.modules when the sphinx build starts.",
|
||||
mock_target)
|
||||
|
||||
intersphinx_mapping = {
|
||||
"python": ("https://docs.python.org/3", None),
|
||||
"typing_extensions":
|
||||
("https://typing-extensions.readthedocs.io/en/latest", None),
|
||||
"aiohttp": ("https://docs.aiohttp.org/en/stable", None),
|
||||
"pillow": ("https://pillow.readthedocs.io/en/stable", None),
|
||||
"numpy": ("https://numpy.org/doc/stable", None),
|
||||
"torch": ("https://pytorch.org/docs/stable", None),
|
||||
"psutil": ("https://psutil.readthedocs.io/en/stable", None),
|
||||
}
|
||||
|
||||
navigation_with_keys = False
|
||||
@@ -1,4 +1,4 @@
|
||||
|
||||
(docker)=
|
||||
# 🐳 Using the FastVideo Docker Image
|
||||
|
||||
If you prefer a containerized development environment or want to avoid managing dependencies manually, you can use our prebuilt Docker image:
|
||||
@@ -3,3 +3,11 @@
|
||||
# 🧰 Developer Environment
|
||||
|
||||
Accelerate your FastVideo development workflow by leveraging Docker images and cloud GPUs for efficient experimentation and reproducible environments.
|
||||
|
||||
:::{toctree}
|
||||
:caption: Contents
|
||||
:maxdepth: 1
|
||||
|
||||
docker
|
||||
runpod
|
||||
:::
|
||||
@@ -1,3 +1,4 @@
|
||||
(runpod)=
|
||||
|
||||
# 📦 Developing FastVideo on RunPod
|
||||
|
||||
@@ -9,7 +10,7 @@ Choose a GPU that supports CUDA 12.8
|
||||
|
||||
Pick 1 or 2 L40S GPU(s)
|
||||
|
||||

|
||||

|
||||
|
||||
When creating your pod template, use this image:
|
||||
|
||||
@@ -23,11 +24,11 @@ Paste Container Start Command to support SSH ([RunPod Docs](https://docs.runpod.
|
||||
bash -c "apt update;DEBIAN_FRONTEND=noninteractive apt-get install openssh-server -y;mkdir -p ~/.ssh;cd $_;chmod 700 ~/.ssh;echo \"$PUBLIC_KEY\" >> authorized_keys;chmod 700 authorized_keys;service ssh start;sleep infinity"
|
||||
```
|
||||
|
||||

|
||||

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

|
||||

|
||||
|
||||
## Working with the pod
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
(developer-overview)=
|
||||
|
||||
# 🛠️ Contributing to FastVideo
|
||||
|
||||
@@ -70,7 +71,3 @@ uv pip install ninja
|
||||
|
||||
python setup.py install
|
||||
```
|
||||
|
||||
## Testing
|
||||
|
||||
Please refer to the [Testing Guide](testing.md) for more information on how to add and run tests in FastVideo.
|
||||
@@ -29,6 +29,7 @@ FastVideo separates model components from execution logic with these principles:
|
||||
- **Custom Attention Backends**: Components can support and use different Attention implementations
|
||||
- **Pipeline Abstraction**: Consistent interface across diffusion models
|
||||
|
||||
(design-fastvideo-args)=
|
||||
## FastVideoArgs
|
||||
|
||||
The `FastVideoArgs` class in `fastvideo/fastvideo_args.py` serves as the central configuration system for FastVideo. It contains all parameters needed to control model loading, inference configuration, performance optimization settings, and more.
|
||||
@@ -60,6 +61,7 @@ with set_current_fastvideo_args(fastvideo_args):
|
||||
result = generate_video()
|
||||
```
|
||||
|
||||
(design-pipeline-system)=
|
||||
## Pipeline System
|
||||
|
||||
### `ComposedPipelineBase`
|
||||
@@ -106,8 +108,7 @@ def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> Forward
|
||||
return batch
|
||||
```
|
||||
|
||||

|
||||
|
||||
(design-forwardbatch)=
|
||||
### ForwardBatch
|
||||
|
||||
Defined in `fastvideo/pipelines/pipeline_batch_info.py`, `ForwardBatch` encapsulates the data payload passed between pipeline stages. It typically holds:
|
||||
@@ -119,10 +120,12 @@ Defined in `fastvideo/pipelines/pipeline_batch_info.py`, `ForwardBatch` encapsul
|
||||
|
||||
This structure facilitates clear state transitions between stages.
|
||||
|
||||
(design-model-components)=
|
||||
## Model Components
|
||||
|
||||
The `fastvideo/models/` directory contains implementations of the core neural network models used in video diffusion:
|
||||
|
||||
(design-transformer-models)=
|
||||
### Transformer Models
|
||||
|
||||
Transformer networks perform the actual denoising during diffusion:
|
||||
@@ -149,6 +152,7 @@ def forward(
|
||||
return noise_pred # Predicted noise residual
|
||||
```
|
||||
|
||||
(design-vae-variational-auto-encoder)=
|
||||
### VAE (Variational Auto-Encoder)
|
||||
|
||||
VAEs handle conversion between pixel space and latent space:
|
||||
@@ -166,6 +170,7 @@ FastVideo's VAE implementations include:
|
||||
- Optional tiling for large frames
|
||||
- Distributed weight support
|
||||
|
||||
(design-text-and-image-encoders)=
|
||||
### Text and Image Encoders
|
||||
|
||||
Encoders process conditioning inputs into embeddings:
|
||||
@@ -183,6 +188,7 @@ FastVideo implements optimizations such as:
|
||||
- Caching for common prompts
|
||||
- Precision-tuned computation
|
||||
|
||||
(design-schedulers)=
|
||||
### Schedulers
|
||||
|
||||
Schedulers manage the diffusion sampling process:
|
||||
@@ -210,10 +216,7 @@ def step(
|
||||
return prev_sample
|
||||
```
|
||||
|
||||
This diagram shows how models are discovered, validated, and loaded across entrypoints, executors, pipelines, and model loaders.
|
||||
|
||||

|
||||
|
||||
(design-optimized-attention)=
|
||||
## Optimized Attention
|
||||
|
||||
The `fastvideo/attention/` directory contains optimized attention implementations crucial for efficient video diffusion:
|
||||
@@ -237,17 +240,17 @@ self.attn = LocalAttention(
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
|
||||
```
|
||||
|
||||

|
||||
|
||||
### Attention Patterns
|
||||
Supports various patterns with memory optimization techniques:
|
||||
- **Cross/Self/Temporal/Global-Local Attention**
|
||||
- Chunking, progressive computation, optimized masking
|
||||
|
||||
(design-distributed-processing)=
|
||||
## Distributed Processing
|
||||
|
||||
The `fastvideo/distributed/` directory contains implementations for distributed model execution:
|
||||
|
||||
(design-tensor-parallelism)=
|
||||
### Tensor Parallelism
|
||||
|
||||
Tensor parallelism splits model weights across devices:
|
||||
@@ -304,6 +307,7 @@ Efficient communication primitives minimize distributed overhead:
|
||||
- **Tensor-Parallel AllReduce**: Combines partial results
|
||||
- **Distributed Synchronization**: Coordinates execution
|
||||
|
||||
(design-forwardcontext)=
|
||||
## Forward Context Management
|
||||
|
||||
### ForwardContext
|
||||
@@ -326,6 +330,7 @@ with set_forward_context(current_timestep, attn_metadata, fastvideo_args):
|
||||
output = model(inputs)
|
||||
```
|
||||
|
||||
(design-executor-and-worker-abstractions)=
|
||||
## Executor and Worker System
|
||||
|
||||
The `fastvideo/worker/` directory contains the distributed execution framework:
|
||||
@@ -352,6 +357,7 @@ Each GPU worker:
|
||||
|
||||
This design allows FastVideo to efficiently utilize multiple GPUs while providing a simple, unified interface for model execution.
|
||||
|
||||
(design-platforms)=
|
||||
## Platforms
|
||||
|
||||
The `fastvideo/platforms/` directory provides hardware platform abstractions that enable FastVideo to run efficiently on different hardware configurations:
|
||||
@@ -382,6 +388,7 @@ else:
|
||||
|
||||
The platform system is designed to be extensible for future hardware targets.
|
||||
|
||||
(design-logger)=
|
||||
## Logger
|
||||
See [PR](https://github.com/hao-ai-lab/FastVideo/pull/356)
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
(v0-data-preprocess)=
|
||||
|
||||
# 🧱 Data Preprocess for Distillation
|
||||
|
||||
@@ -6,11 +6,10 @@ import re
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
|
||||
ROOT_DIR = Path(__file__).parent.parent.resolve()
|
||||
ROOT_DIR_RELATIVE = '../..'
|
||||
ROOT_DIR = Path(__file__).parent.parent.parent.resolve()
|
||||
ROOT_DIR_RELATIVE = '../../../..'
|
||||
EXAMPLE_DIR = ROOT_DIR / "examples"
|
||||
EXAMPLE_DOC_DIR = ROOT_DIR / "docs/getting_started/examples"
|
||||
GITHUB_REPO = "hao-ai-lab/FastVideo" # Update this to your repo
|
||||
EXAMPLE_DOC_DIR = ROOT_DIR / "docs/source/getting_started/examples"
|
||||
|
||||
|
||||
def fix_case(text: str) -> str:
|
||||
@@ -72,16 +71,9 @@ class Index:
|
||||
|
||||
def generate(self) -> str:
|
||||
content = f"# {self.title}\n\n{self.description}\n\n"
|
||||
if self.caption:
|
||||
content += f"## {self.caption}\n\n"
|
||||
# Generate a simple list of links for MkDocs
|
||||
for doc in self.documents:
|
||||
# Convert document path to proper link
|
||||
doc_link = doc.replace("\\", "/")
|
||||
# Get just the filename for the link text
|
||||
doc_title = fix_case(Path(doc).stem.replace("_", " ").title())
|
||||
content += f"- [{doc_title}]({doc_link}.md)\n"
|
||||
content += "\n"
|
||||
content += ":::{toctree}\n"
|
||||
content += f":caption: {self.caption}\n:maxdepth: {self.maxdepth}\n"
|
||||
content += "\n".join(self.documents) + "\n:::\n"
|
||||
return content
|
||||
|
||||
|
||||
@@ -150,66 +142,30 @@ class Example:
|
||||
return fix_case(self.path.stem.replace("_", " ").title())
|
||||
|
||||
def generate(self) -> str:
|
||||
# Create GitHub link to source
|
||||
github_path = str(self.path.relative_to(ROOT_DIR)).replace("\\", "/")
|
||||
github_url = f"https://github.com/{GITHUB_REPO}/blob/main/{github_path}"
|
||||
content = f"**Source:** [{github_path}]({github_url})\n\n"
|
||||
# Convert the path to a relative path from __file__
|
||||
make_relative = lambda path: ROOT_DIR_RELATIVE / path.relative_to(
|
||||
ROOT_DIR)
|
||||
|
||||
# Add title for code files
|
||||
if self.main_file.suffix != ".md":
|
||||
content = f"Source <gh-file:{self.path.relative_to(ROOT_DIR)}>.\n\n"
|
||||
include = "include" if self.main_file.suffix == ".md" else \
|
||||
"literalinclude"
|
||||
if include == "literalinclude":
|
||||
content += f"# {self.title}\n\n"
|
||||
|
||||
# Include main file content
|
||||
if self.main_file.suffix == ".md":
|
||||
# For markdown files, include the content directly
|
||||
with open(self.main_file, encoding='utf-8') as f:
|
||||
content += f.read() + "\n\n"
|
||||
else:
|
||||
# For code files, use code blocks
|
||||
language = self.main_file.suffix[1:] if self.main_file.suffix else ""
|
||||
with open(self.main_file, encoding='utf-8') as f:
|
||||
file_content = f.read()
|
||||
content += f"```{language}\n{file_content}\n```\n\n"
|
||||
content += f":::{{{include}}} {make_relative(self.main_file)}\n" # type: ignore[no-untyped-call]
|
||||
if include == "literalinclude":
|
||||
content += f":language: {self.main_file.suffix[1:]}\n"
|
||||
content += ":::\n\n"
|
||||
|
||||
if not self.other_files:
|
||||
return content
|
||||
|
||||
content += "## Additional Files\n\n"
|
||||
# Define binary/non-text file extensions to skip
|
||||
binary_extensions = {
|
||||
'.mp4', '.avi', '.mov', '.mkv', '.gif', '.jpg', '.jpeg', '.png',
|
||||
'.webp', '.bmp', '.pdf', '.zip', '.tar', '.gz', '.mp3', '.wav'
|
||||
}
|
||||
|
||||
content += "## Example materials\n\n"
|
||||
for file in sorted(self.other_files):
|
||||
# Skip binary files
|
||||
if file.suffix.lower() in binary_extensions:
|
||||
continue
|
||||
|
||||
file_rel_path = file.relative_to(self.path)
|
||||
# Use collapsible admonition syntax for MkDocs
|
||||
content += f"??? note \"{file_rel_path}\"\n\n"
|
||||
|
||||
try:
|
||||
if file.suffix == ".md":
|
||||
# Include markdown content with indentation
|
||||
with open(file, encoding='utf-8') as f:
|
||||
for line in f:
|
||||
content += f" {line}"
|
||||
else:
|
||||
# Include code with proper formatting
|
||||
language = file.suffix[1:] if file.suffix else ""
|
||||
with open(file, encoding='utf-8') as f:
|
||||
file_content = f.read()
|
||||
# Indent the code block for the admonition
|
||||
content += f" ```{language}\n"
|
||||
for line in file_content.split('\n'):
|
||||
content += f" {line}\n"
|
||||
content += " ```\n"
|
||||
content += "\n"
|
||||
except UnicodeDecodeError:
|
||||
# Skip files that can't be decoded as UTF-8
|
||||
continue
|
||||
include = "include" if file.suffix == ".md" else "literalinclude"
|
||||
content += f":::{{admonition}} {file.relative_to(self.path)}\n"
|
||||
content += ":class: dropdown\n\n"
|
||||
content += f":::{{{include}}} {make_relative(file)}\n:::\n" # type: ignore[no-untyped-call]
|
||||
content += ":::\n\n"
|
||||
|
||||
return content
|
||||
|
||||
@@ -239,7 +195,7 @@ class NestedStructure:
|
||||
|
||||
def create_category_indices() -> dict[str, Index]:
|
||||
"""Create category indices with their respective configurations."""
|
||||
main_index_dir = ROOT_DIR / "docs/examples"
|
||||
main_index_dir = ROOT_DIR / "docs/source/examples"
|
||||
if not main_index_dir.exists():
|
||||
main_index_dir.mkdir(parents=True)
|
||||
|
||||
@@ -247,16 +203,17 @@ def create_category_indices() -> dict[str, Index]:
|
||||
"inference":
|
||||
Index(
|
||||
path=ROOT_DIR /
|
||||
"docs/inference/examples/examples_inference_index.md",
|
||||
"docs/source/inference/examples/examples_inference_index.md",
|
||||
title="🚀 Examples",
|
||||
description=
|
||||
"Inference examples demonstrate how to use FastVideo inference. We recommend starting with [basic.md](basic.md).",
|
||||
"Inference examples demonstrate how to use FastVideo inference. We recommend starting with <project:basic.md>.",
|
||||
caption="Examples",
|
||||
maxdepth=1,
|
||||
),
|
||||
"training":
|
||||
Index(
|
||||
path=ROOT_DIR / "docs/training/examples/examples_training_index.md",
|
||||
path=ROOT_DIR /
|
||||
"docs/source/training/examples/examples_training_index.md",
|
||||
title="🚀 Examples",
|
||||
description=
|
||||
"Training examples demonstrate how to use FastVideo training.",
|
||||
@@ -266,7 +223,7 @@ def create_category_indices() -> dict[str, Index]:
|
||||
"distillation":
|
||||
Index(
|
||||
path=ROOT_DIR /
|
||||
"docs/distillation/examples/examples_distillation_index.md",
|
||||
"docs/source/distillation/examples/examples_distillation_index.md",
|
||||
title="🚀 Examples",
|
||||
description=
|
||||
"Distillation examples demonstrate how to use FastVideo distillation.",
|
||||
@@ -289,21 +246,9 @@ def find_examples(category_indices: dict[str, Index],
|
||||
examples = []
|
||||
glob_patterns = ["*.py", "*.md", "*.sh"]
|
||||
|
||||
# Map category names to actual directory names
|
||||
category_dir_mapping = {
|
||||
"distillation": "distill", # examples/distill/ -> distillation category
|
||||
}
|
||||
|
||||
# Find categorised examples
|
||||
for category in category_indices:
|
||||
# Use mapped directory name if available, otherwise use category name
|
||||
dir_name = category_dir_mapping.get(category, category)
|
||||
category_dir = EXAMPLE_DIR / dir_name
|
||||
|
||||
# Skip if directory doesn't exist
|
||||
if not category_dir.exists():
|
||||
continue
|
||||
|
||||
category_dir = EXAMPLE_DIR / category
|
||||
globs = [category_dir.glob(pattern) for pattern in glob_patterns]
|
||||
for path in itertools.chain(*globs):
|
||||
examples.append(Example(path, category))
|
||||
@@ -334,18 +279,11 @@ def create_nested_structures(
|
||||
dict[str,
|
||||
NestedStructure]]]] = {}
|
||||
|
||||
# Map category names to actual directory names
|
||||
category_dir_mapping = {
|
||||
"distillation": "distill",
|
||||
}
|
||||
|
||||
for example in examples:
|
||||
if example.category not in ["training", "distillation"]:
|
||||
continue
|
||||
|
||||
# Use mapped directory name if available
|
||||
dir_name = category_dir_mapping.get(example.category, example.category)
|
||||
category_dir = EXAMPLE_DIR / dir_name
|
||||
category_dir = EXAMPLE_DIR / example.category
|
||||
relative_path = example.path.relative_to(category_dir)
|
||||
path_parts = relative_path.parts
|
||||
|
||||
@@ -477,7 +415,7 @@ def generate_nested_examples(nested_structures: dict[str, dict[str, dict[
|
||||
category_index.documents.append(method)
|
||||
|
||||
|
||||
def generate_examples(generate_main_index: bool = False) -> None:
|
||||
def generate_examples(generate_main_index=False):
|
||||
"""
|
||||
Generate example documentation.
|
||||
|
||||
@@ -491,14 +429,12 @@ def generate_examples(generate_main_index: bool = False) -> None:
|
||||
# Create the main examples index only if requested
|
||||
examples_index = None
|
||||
if generate_main_index:
|
||||
main_index_dir = ROOT_DIR / "docs/examples"
|
||||
main_index_dir = ROOT_DIR / "docs/source/examples"
|
||||
examples_index = Index(
|
||||
path=main_index_dir / "examples_index.md",
|
||||
title="💡 Examples",
|
||||
description=
|
||||
"A collection of examples demonstrating usage of FastVideo.\n\n"
|
||||
f"All documented examples are autogenerated using [generate_examples.py](https://github.com/{GITHUB_REPO}/blob/main/docs/generate_examples.py) "
|
||||
f"from examples found in the [examples](https://github.com/{GITHUB_REPO}/tree/main/examples) directory.",
|
||||
"A collection of examples demonstrating usage of FastVideo.\nAll documented examples are autogenerated using <gh-file:docs/source/generate_examples.py> from examples found in <gh-file:examples>.",
|
||||
caption="Examples",
|
||||
maxdepth=2)
|
||||
|
||||
@@ -535,19 +471,3 @@ def generate_examples(generate_main_index: bool = False) -> None:
|
||||
if generate_main_index and examples_index:
|
||||
with open(examples_index.path, "w+") as f:
|
||||
f.write(examples_index.generate())
|
||||
|
||||
|
||||
def on_pre_build_hook(config, **kwargs):
|
||||
"""
|
||||
MkDocs hook to generate examples before building the documentation.
|
||||
This function is called automatically by the mkdocs-simple-hooks plugin.
|
||||
"""
|
||||
print("Generating example documentation...")
|
||||
generate_examples(generate_main_index=True)
|
||||
print("Example documentation generated successfully!")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
print("Generating example documentation...")
|
||||
generate_examples(generate_main_index=True)
|
||||
print("Example documentation generated successfully!")
|
||||
@@ -0,0 +1,18 @@
|
||||
(installation-index)=
|
||||
|
||||
# 🔧 Installation
|
||||
|
||||
FastVideo supports the following hardware platforms:
|
||||
|
||||
:::{toctree}
|
||||
:maxdepth: 1
|
||||
:hidden:
|
||||
|
||||
installation/gpu
|
||||
installation/mps
|
||||
:::
|
||||
|
||||
- <project:installation/gpu.md>
|
||||
- NVIDIA CUDA
|
||||
- <project:installation/mps.md>
|
||||
- Apple silicon
|
||||
@@ -30,8 +30,17 @@ conda create -n fastvideo python=3.12 -y
|
||||
conda activate fastvideo
|
||||
```
|
||||
|
||||
:::{note}
|
||||
[PyTorch has deprecated the conda release channel](https://github.com/pytorch/pytorch/issues/138506). If you use `conda`, please only use it to create Python environment rather than installing packages.
|
||||
:::
|
||||
|
||||
#### uv
|
||||
|
||||
:::{tip}
|
||||
We highly recommend using `uv` to install FastVideo. In our experience, `uv` speeds up installation by at least 3x.
|
||||
Note that you can also use `uv` to install FastVideo in a Conda environment.
|
||||
:::
|
||||
|
||||
Or you can create a new Python environment using [uv](https://docs.astral.sh/uv/), a very fast Python environment manager. Please follow the [documentation](https://docs.astral.sh/uv/#getting-started) to install `uv`. After installing `uv`, you can create a new Python environment using the following command:
|
||||
|
||||
```console
|
||||
@@ -31,8 +31,17 @@ conda create -n fastvideo python=3.12.4 -y
|
||||
conda activate fastvideo
|
||||
```
|
||||
|
||||
:::{note}
|
||||
[PyTorch has deprecated the conda release channel](https://github.com/pytorch/pytorch/issues/138506). If you use `conda`, please only use it to create Python environment rather than installing packages.
|
||||
:::
|
||||
|
||||
#### uv
|
||||
|
||||
:::{tip}
|
||||
We highly recommend using `uv` to install FastVideo. In our experience, `uv` speeds up installation by at least 3x.
|
||||
Note that you can also use `uv` to install FastVideo in a Conda environment.
|
||||
:::
|
||||
|
||||
Or you can create a new Python environment using [uv](https://docs.astral.sh/uv/), a very fast Python environment manager. Please follow the [documentation](https://docs.astral.sh/uv/#getting-started) to install `uv`. After installing `uv`, you can create a new Python environment using the following command:
|
||||
|
||||
```console
|
||||
@@ -0,0 +1,83 @@
|
||||
# V1 API
|
||||
|
||||
FastVideo's V1 API provides a streamlined interface for video generation tasks with powerful customization options. This page documents the primary components of the API.
|
||||
|
||||
## Video Generator
|
||||
|
||||
This class will be the primary Python API for generating videos and images.
|
||||
|
||||
```{autodoc2-summary}
|
||||
fastvideo.VideoGenerator
|
||||
```
|
||||
|
||||
`````{py:class} VideoGenerator(fastvideo_args: fastvideo.fastvideo_args.FastVideoArgs, executor_class: type[fastvideo.worker.executor.Executor], log_stats: bool)
|
||||
:canonical: fastvideo.entrypoints.video_generator.VideoGenerator
|
||||
|
||||
```{autodoc2-docstring} fastvideo.entrypoints.video_generator.VideoGenerator
|
||||
:parser: docs.source.autodoc2_docstring_parser
|
||||
```
|
||||
|
||||
`VideoGenerator.from_pretrained()` should be the primary way of creating a new video generator.
|
||||
|
||||
````{py:method} from_pretrained(model_path: str, device: typing.Optional[str] = None, torch_dtype: typing.Optional[torch.dtype] = None, pipeline_config: typing.Optional[typing.Union[str | fastvideo.configs.pipelines.PipelineConfig]] = None, **kwargs) -> fastvideo.entrypoints.video_generator.VideoGenerator
|
||||
:canonical: fastvideo.entrypoints.video_generator.VideoGenerator.from_pretrained
|
||||
:classmethod:
|
||||
|
||||
```{autodoc2-docstring} fastvideo.entrypoints.video_generator.VideoGenerator.from_pretrained
|
||||
:parser: docs.source.autodoc2_docstring_parser
|
||||
```
|
||||
|
||||
|
||||
## Configuring FastVideo
|
||||
|
||||
The follow two classes `PipelineConfig` and `SamplingParam` are used to configure initialization and sampling parameters, respectively.
|
||||
|
||||
### PipelineConfig
|
||||
```{autodoc2-summary}
|
||||
fastvideo.PipelineConfig
|
||||
```
|
||||
|
||||
`````{py:class} PipelineConfig
|
||||
:canonical: fastvideo.configs.pipelines.base.PipelineConfig
|
||||
|
||||
```{autodoc2-docstring} fastvideo.configs.pipelines.base.PipelineConfig
|
||||
:parser: docs.source.autodoc2_docstring_parser
|
||||
```
|
||||
|
||||
````{py:method} from_pretrained(model_path: str) -> fastvideo.configs.pipelines.base.PipelineConfig
|
||||
:canonical: fastvideo.configs.pipelines.base.PipelineConfig.from_pretrained
|
||||
:classmethod:
|
||||
|
||||
```{autodoc2-docstring} fastvideo.configs.pipelines.base.PipelineConfig.from_pretrained
|
||||
:parser: docs.source.autodoc2_docstring_parser
|
||||
```
|
||||
|
||||
|
||||
````{py:method} dump_to_json(file_path: str)
|
||||
:canonical: fastvideo.configs.pipelines.base.PipelineConfig.dump_to_json
|
||||
|
||||
```{autodoc2-docstring} fastvideo.configs.pipelines.base.PipelineConfig.dump_to_json
|
||||
:parser: docs.source.autodoc2_docstring_parser
|
||||
```
|
||||
|
||||
|
||||
### SamplingParam
|
||||
|
||||
```{autodoc2-summary}
|
||||
fastvideo.SamplingParam
|
||||
```
|
||||
|
||||
`````{py:class} SamplingParam
|
||||
:canonical: fastvideo.configs.sample.base.SamplingParam
|
||||
|
||||
```{autodoc2-docstring} fastvideo.configs.sample.base.SamplingParam
|
||||
:parser: docs.source.autodoc2_docstring_parser
|
||||
```
|
||||
|
||||
````{py:method} from_pretrained(model_path: str) -> fastvideo.configs.sample.base.SamplingParam
|
||||
:canonical: fastvideo.configs.sample.base.SamplingParam.from_pretrained
|
||||
:classmethod:
|
||||
|
||||
```{autodoc2-docstring} fastvideo.configs.sample.base.SamplingParam.from_pretrained
|
||||
:parser: docs.source.autodoc2_docstring_parser
|
||||
```
|
||||
@@ -1,52 +1,132 @@
|
||||
# Welcome to FastVideo
|
||||
|
||||
<div style="text-align: center;">
|
||||
<img src="assets/logos/logo.svg" alt="FastVideo" style="width: 60%;" />
|
||||
</div>
|
||||
:::{figure} ../../assets/logos/logo.svg
|
||||
:align: center
|
||||
:alt: FastVideo
|
||||
:class: no-scaled-link
|
||||
:width: 60%
|
||||
:::
|
||||
|
||||
<div style="text-align: center;">
|
||||
<strong>FastVideo is a unified inference and post-training framework for accelerated video generation.</strong>
|
||||
</div>
|
||||
:::{raw} html
|
||||
<p style="text-align:center">
|
||||
<strong>FastVideo is a unified inference and post-training framework for accelerated video generation.
|
||||
</strong>
|
||||
</p>
|
||||
|
||||
<div style="text-align: center;">
|
||||
<p style="text-align:center">
|
||||
<script async defer src="https://buttons.github.io/buttons.js"></script>
|
||||
<a class="github-button" href="https://github.com/hao-ai-lab/FastVideo/" data-show-count="true" data-size="large" aria-label="Star">Star</a>
|
||||
<a class="github-button" href="https://github.com/hao-ai-lab/FastVideo/subscription" data-icon="octicon-eye" data-size="large" aria-label="Watch">Watch</a>
|
||||
<a class="github-button" href="https://github.com/hao-ai-lab/FastVideo/fork" data-icon="octicon-repo-forked" data-size="large" aria-label="Fork">Fork</a>
|
||||
</div>
|
||||
</p>
|
||||
:::
|
||||
|
||||
FastVideo is an inference and post-training framework for diffusion models. It features an end-to-end unified pipeline for accelerating diffusion models, starting from data preprocessing to model training, finetuning, distillation, and inference. FastVideo is designed to be modular and extensible, allowing users to easily add new optimizations and techniques. Whether it is training-free optimizations or post-training optimizations, FastVideo has you covered.
|
||||
|
||||
<div style="text-align: center;">
|
||||
<img src="assets/images/fastwan.png" style="width: 100%;"/>
|
||||
<img src=_static/images/fastwan.png width="100%"/>
|
||||
</div>
|
||||
|
||||
## Key Features
|
||||
|
||||
FastVideo has the following features:
|
||||
|
||||
- State-of-the-art performance optimizations for inference
|
||||
- [Sliding Tile Attention](https://arxiv.org/pdf/2502.04507)
|
||||
- [TeaCache](https://arxiv.org/pdf/2411.19108)
|
||||
- [Sage Attention](https://arxiv.org/abs/2410.02367)
|
||||
- E2E post-training support
|
||||
- Data preprocessing pipeline for video data
|
||||
- Data preprocessing pipeline for video data.
|
||||
- [Sparse distillation](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/) for Wan2.1 and Wan2.2 using [Video Sparse Attention](https://arxiv.org/pdf/2505.13389) and [Distribution Matching Distillation](https://tianweiy.github.io/dmd2/)
|
||||
- Support full finetuning and LoRA finetuning for state-of-the-art open video DiTs.
|
||||
- Scalable training with FSDP2, sequence parallelism, and selective activation checkpointing, with near linear scaling to 64 GPUs.
|
||||
|
||||
## Documentation
|
||||
|
||||
Welcome to FastVideo! This documentation will help you get started with our unified inference and post-training framework for accelerated video generation.
|
||||
% How to start using FastVideo?
|
||||
|
||||
Use the navigation menu on the left to explore different sections:
|
||||
:::{toctree}
|
||||
:caption: Getting Started
|
||||
:maxdepth: 1
|
||||
|
||||
- **Getting Started**: Installation and quick start guides
|
||||
- **Inference**: Learn how to use FastVideo for video generation
|
||||
- **Training**: Data preprocessing and fine-tuning workflows
|
||||
- **Distillation**: Post-training optimization techniques
|
||||
- **Sliding Tile Attention**: Advanced attention mechanisms
|
||||
- **Video Sparse Attention**: Efficient attention for video models
|
||||
- **Design**: Framework architecture and design principles
|
||||
- **Developer Guide**: Contributing and development setup
|
||||
- **API Reference**: Complete API documentation
|
||||
getting_started/installation
|
||||
<!-- getting_started/v1_api -->
|
||||
:::
|
||||
|
||||
:::{toctree}
|
||||
:caption: Inference
|
||||
:maxdepth: 1
|
||||
|
||||
inference/inference_quick_start
|
||||
inference/examples/examples_inference_index
|
||||
inference/configuration
|
||||
inference/optimizations
|
||||
inference/comfyui
|
||||
inference/support_matrix
|
||||
inference/cli
|
||||
inference/add_pipeline
|
||||
:::
|
||||
|
||||
:::{toctree}
|
||||
:caption: Training
|
||||
:maxdepth: 1
|
||||
|
||||
training/examples/examples_training_index
|
||||
training/data_preprocess
|
||||
<!-- training/finetune -->
|
||||
:::
|
||||
|
||||
:::{toctree}
|
||||
:caption: Distillation
|
||||
:maxdepth: 1
|
||||
|
||||
distillation/examples/examples_distillation_index
|
||||
distillation/data_preprocess
|
||||
distillation/dmd
|
||||
:::
|
||||
|
||||
% What is STA Kernel?
|
||||
|
||||
:::{toctree}
|
||||
:caption: Sliding Tile Attention
|
||||
:maxdepth: 1
|
||||
|
||||
sliding_tile_attention/installation
|
||||
sliding_tile_attention/demo
|
||||
:::
|
||||
|
||||
% What is VSA Kernel?
|
||||
|
||||
:::{toctree}
|
||||
:caption: Video Sparse Attention
|
||||
:maxdepth: 1
|
||||
|
||||
video_sparse_attention/installation
|
||||
:::
|
||||
|
||||
:::{toctree}
|
||||
:caption: Design
|
||||
:maxdepth: 1
|
||||
design/overview
|
||||
:::
|
||||
|
||||
:::{toctree}
|
||||
:caption: Developer Guide
|
||||
:maxdepth: 2
|
||||
|
||||
contributing/overview
|
||||
contributing/developer_env/index
|
||||
contributing/profiling
|
||||
:::
|
||||
|
||||
:::{toctree}
|
||||
:caption: API Reference
|
||||
:maxdepth: 2
|
||||
|
||||
<!-- api/summary -->
|
||||
api/fastvideo/fastvideo
|
||||
:::
|
||||
|
||||
## Indices and tables
|
||||
|
||||
- {ref}`genindex`
|
||||
- {ref}`modindex`
|
||||
@@ -1,3 +1,4 @@
|
||||
(add-pipeline)=
|
||||
|
||||
# 🏗️ Adding a New Pipeline
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
|
||||
(inference-configuration)=
|
||||
# Configuration
|
||||
|
||||
## Multi-GPU Setup
|
||||
@@ -1,3 +1,4 @@
|
||||
(inference-optimizations)=
|
||||
|
||||
# Optimizations
|
||||
|
||||
@@ -15,6 +16,8 @@ This page describes the various options for speeding up generation times in Fast
|
||||
- Caching Techniques
|
||||
- [TeaCache](#optimizations-teacache)
|
||||
|
||||
(optimizations-backends)=
|
||||
|
||||
## Attention Backends
|
||||
|
||||
### Available Backends
|
||||
@@ -46,6 +49,8 @@ You can also set the environment variable on the command line:
|
||||
FASTVIDEO_ATTENTION_BACKEND=SAGE_ATTN python example.py
|
||||
```
|
||||
|
||||
(optimizations-flash)=
|
||||
|
||||
### Flash Attention
|
||||
|
||||
**`FLASH_ATTN`**
|
||||
@@ -66,6 +71,12 @@ pip install ninja
|
||||
python setup.py install
|
||||
```
|
||||
|
||||
:::{note}
|
||||
FastVideo will automatically detect and use `FA3` if it is installed when using `FLASH_ATTN` backend.
|
||||
:::
|
||||
|
||||
(optimizations-sta)=
|
||||
|
||||
### Sliding Tile Attention
|
||||
|
||||
**`SLIDING_TILE_ATTN`**
|
||||
@@ -76,6 +87,8 @@ pip install st_attn==0.0.4
|
||||
|
||||
Please see [this page](#sta-installation) for more installation instructions.
|
||||
|
||||
(optimizations-vsa)=
|
||||
|
||||
### Video Sparse Attention
|
||||
|
||||
**`VIDEO_SPARSE_ATTN`**
|
||||
@@ -87,6 +100,8 @@ python setup_vsa.py install
|
||||
|
||||
Please see [this page](#vsa-installation) for more installation instructions.
|
||||
|
||||
(optimizations-sage)=
|
||||
|
||||
### Sage Attention
|
||||
|
||||
**`SAGE_ATTN`**
|
||||
@@ -99,6 +114,8 @@ cd sageattention
|
||||
python setup.py install # or pip install -e .
|
||||
```
|
||||
|
||||
(optimizations-sage3)=
|
||||
|
||||
### Sage Attention 3
|
||||
|
||||
**`SAGE_ATTN_THREE`**
|
||||
@@ -119,6 +136,8 @@ To use Sage Attention 3 in FastVideo, first get access to the SageAttention3 cod
|
||||
python setup.py install
|
||||
```
|
||||
|
||||
(optimizations-teacache)=
|
||||
|
||||
## Teacache
|
||||
|
||||
TeaCache is an optimization technique supported in FastVideo that can significantly speed up video generation by skipping redundant calculations across diffusion steps. This guide explains how to enable and configure TeaCache for optimal performance in FastVideo.
|
||||
@@ -0,0 +1,136 @@
|
||||
(support-matrix)=
|
||||
# Compatibility Matrix
|
||||
The table below shows every supported model and optimizations supported for them.
|
||||
|
||||
The symbols used have the following meanings:
|
||||
|
||||
- ✅ = Full compatibility
|
||||
- ❌ = No compatibility
|
||||
- ⭕ = Does not apply to this model
|
||||
|
||||
## Models x Optimization
|
||||
The `HuggingFace Model ID` can be directly pass to `from_pretrained()` methods and FastVideo will use the optimal default parameters when initializing and generating videos.
|
||||
|
||||
:::{raw} html
|
||||
<style>
|
||||
/* Make smaller to try to improve readability */
|
||||
td {
|
||||
font-size: 0.9rem;
|
||||
text-align: center;
|
||||
}
|
||||
|
||||
th {
|
||||
text-align: center;
|
||||
font-size: 0.9rem;
|
||||
}
|
||||
</style>
|
||||
:::
|
||||
|
||||
:::{list-table}
|
||||
:header-rows: 1
|
||||
:stub-columns: 3
|
||||
:widths: auto
|
||||
:class: vertical-table-header
|
||||
|
||||
- * Model Name
|
||||
* HuggingFace Model ID
|
||||
* Resolutions
|
||||
* TeaCache
|
||||
* Sliding Tile Attn
|
||||
* Sage Attn
|
||||
* Video Sparse Attention (VSA)
|
||||
- * FastWan2.1 T2V 1.3B
|
||||
* `FastVideo/FastWan2.1-T2V-1.3B-Diffusers`
|
||||
* 480P
|
||||
* ⭕
|
||||
* ⭕
|
||||
* ⭕
|
||||
* ✅
|
||||
- * FastWan2.2 TI2V 5B Full Attn*
|
||||
* `FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers`
|
||||
* 720P
|
||||
* ⭕
|
||||
* ⭕
|
||||
* ⭕
|
||||
* ✅
|
||||
- * Wan2.2 TI2V 5B
|
||||
* `Wan-AI/Wan2.2-TI2V-5B-Diffusers`
|
||||
* 720P
|
||||
* ⭕
|
||||
* ⭕
|
||||
* ✅
|
||||
* ⭕
|
||||
- * Wan2.2 T2V A14B
|
||||
* `Wan-AI/Wan2.2-T2V-A14B-Diffusers`
|
||||
* 480P<br>720P
|
||||
* ❌
|
||||
* ❌
|
||||
* ✅
|
||||
* ⭕
|
||||
- * Wan2.2 I2V A14B
|
||||
* `Wan-AI/Wan2.2-I2V-A14B-Diffusers`
|
||||
* 480P<br>720P
|
||||
* ❌
|
||||
* ❌
|
||||
* ✅
|
||||
* ⭕
|
||||
- * HunyuanVideo
|
||||
* `hunyuanvideo-community/HunyuanVideo`
|
||||
* 720px1280p<br>544px960p
|
||||
* ❌
|
||||
* ✅
|
||||
* ✅
|
||||
* ⭕
|
||||
- * FastHunyuan
|
||||
* `FastVideo/FastHunyuan-diffusers`
|
||||
* 720px1280p<br>544px960p
|
||||
* ❌
|
||||
* ✅
|
||||
* ✅
|
||||
* ⭕
|
||||
- * Wan2.1 T2V 1.3B
|
||||
* `Wan-AI/Wan2.1-T2V-1.3B-Diffusers`
|
||||
* 480P
|
||||
* ✅
|
||||
* ✅*
|
||||
* ✅
|
||||
* ⭕
|
||||
- * Wan2.1 T2V 14B
|
||||
* `Wan-AI/Wan2.1-T2V-14B-Diffusers`
|
||||
* 480P, 720P
|
||||
* ✅
|
||||
* ✅*
|
||||
* ✅
|
||||
* ⭕
|
||||
- * Wan2.1 I2V 480P
|
||||
* `Wan-AI/Wan2.1-I2V-14B-480P-Diffusers`
|
||||
* 480P
|
||||
* ✅
|
||||
* ✅*
|
||||
* ✅
|
||||
* ⭕
|
||||
- * Wan2.1 I2V 720P
|
||||
* `Wan-AI/Wan2.1-I2V-14B-720P-Diffusers`
|
||||
* 720P
|
||||
* ✅
|
||||
* ✅
|
||||
* ✅
|
||||
* ⭕
|
||||
- * StepVideo T2V
|
||||
* `FastVideo/stepvideo-t2v-diffusers`
|
||||
* 768px768px204f<br>544px992px204f<br>544px992px136f
|
||||
* ❌
|
||||
* ❌
|
||||
* ✅
|
||||
* ⭕
|
||||
:::
|
||||
|
||||
**Note**: Wan2.2 TI2V 5B has some quality issues when performing I2V generation. We are working on fixing this issue.
|
||||
|
||||
## Special requirements
|
||||
|
||||
### StepVideo T2V
|
||||
- The self-attention in text-encoder (step_llm) only supports CUDA capabilities sm_80 sm_86 and sm_90
|
||||
|
||||
### Sliding Tile Attention
|
||||
- Currently only Hopper GPUs (H100s) are supported.
|
||||
@@ -1,3 +1,4 @@
|
||||
(sta-demo)=
|
||||
|
||||
# 🔍 Demo
|
||||
This is is a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
|
||||
@@ -1,3 +1,4 @@
|
||||
(sta-installation)=
|
||||
|
||||
# 🔧 Installation
|
||||
You can install the Sliding Tile Attention package using
|
||||
@@ -1,3 +1,4 @@
|
||||
(v0-data-preprocess)=
|
||||
|
||||
# 🧱 Data Preprocess
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
|
||||
(v0-finetune)=
|
||||
# 🧠 Finetune
|
||||
## ⚡ Full Finetune
|
||||
Ensure your data is prepared and preprocessed in the format specified in [data_preprocess.md](#v0-data-preprocess). For convenience, we also provide a mochi preprocessed Black Myth Wukong data that can be downloaded directly:
|
||||
@@ -1,3 +1,4 @@
|
||||
(vsa-installation)=
|
||||
|
||||
# 🔧 Installation
|
||||
You can install the Video Sparse Attention package using
|
||||
@@ -1,44 +0,0 @@
|
||||
# NOTE: This is still a work in progress, and the checkpoints are not released yet.
|
||||
|
||||
from fastvideo import VideoGenerator, SamplingParam
|
||||
import json
|
||||
# from fastvideo.configs.sample import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples_self_forcing_causal_wan2_2_14B_i2v"
|
||||
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(
|
||||
"FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
dit_cpu_offload=True, # DiT need to be offloaded for MoE
|
||||
dit_precision="fp32",
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
dmd_denoising_steps=[1000, 850, 700, 550, 350, 275, 200, 125],
|
||||
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
|
||||
pin_cpu_memory=True,
|
||||
# image_encoder_cpu_offload=False,
|
||||
)
|
||||
|
||||
sampling_param = SamplingParam.from_pretrained("FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers")
|
||||
sampling_param.num_frames = 81
|
||||
sampling_param.width = 832
|
||||
sampling_param.height = 480
|
||||
sampling_param.seed = 1000
|
||||
|
||||
with open("prompts/mixkit_i2v.jsonl", "r") as f:
|
||||
prompt_image_pairs = json.load(f)
|
||||
|
||||
for prompt_image_pair in prompt_image_pairs:
|
||||
prompt = prompt_image_pair["prompt"]
|
||||
image_path = prompt_image_pair["image_path"]
|
||||
_ = generator.generate_video(prompt, image_path=image_path, output_path=OUTPUT_PATH, save_video=True, sampling_param=sampling_param)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -3,9 +3,6 @@ import os
|
||||
import requests
|
||||
import base64
|
||||
import time
|
||||
import json
|
||||
from pathlib import Path
|
||||
import tempfile
|
||||
|
||||
import gradio as gr
|
||||
|
||||
@@ -15,7 +12,6 @@ from fastvideo.configs.sample.base import SamplingParam
|
||||
MODEL_PATH_MAPPING = {
|
||||
"FastWan2.1-T2V-1.3B": "FastVideo/FastWan2.1-T2V-1.3B-Diffusers",
|
||||
"FastWan2.2-TI2V-5B-FullAttn": "FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers",
|
||||
"CausalWan2.2-I2V-A14B-Preview": "FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers",
|
||||
}
|
||||
|
||||
|
||||
@@ -41,7 +37,7 @@ class RayServeClient:
|
||||
f"{self.backend_url}/generate_video",
|
||||
json=request_data,
|
||||
headers=headers,
|
||||
timeout=900 # 15 minutes timeout for longer video generation
|
||||
timeout=300
|
||||
)
|
||||
|
||||
round_trip_time = time.time() - start_time
|
||||
@@ -85,78 +81,49 @@ def save_video_from_base64(video_data: str, output_dir: str, prompt: str) -> str
|
||||
return None
|
||||
|
||||
|
||||
def encode_image_to_base64(image_path: str) -> str:
|
||||
"""Encode an image file to base64 string."""
|
||||
if not image_path or not os.path.exists(image_path):
|
||||
return None
|
||||
def create_timing_display(inference_time, encoding_time, network_time, total_time, stage_execution_times, num_frames):
|
||||
dit_denoising_time = f"{stage_execution_times[5]:.2f}s" if len(stage_execution_times) > 5 else "N/A"
|
||||
|
||||
try:
|
||||
with open(image_path, 'rb') as f:
|
||||
image_bytes = f.read()
|
||||
|
||||
image_base64 = base64.b64encode(image_bytes).decode('utf-8')
|
||||
|
||||
# Determine image type from extension
|
||||
ext = os.path.splitext(image_path)[1].lower()
|
||||
mime_types = {
|
||||
'.jpg': 'image/jpeg',
|
||||
'.jpeg': 'image/jpeg',
|
||||
'.png': 'image/png',
|
||||
'.gif': 'image/gif',
|
||||
'.webp': 'image/webp',
|
||||
}
|
||||
mime_type = mime_types.get(ext, 'image/jpeg')
|
||||
|
||||
return f"data:{mime_type};base64,{image_base64}"
|
||||
|
||||
except Exception as e:
|
||||
print(f"Failed to encode image: {e}")
|
||||
return None
|
||||
|
||||
|
||||
# def create_timing_display(inference_time, encoding_time, network_time, total_time, stage_execution_times, num_frames):
|
||||
# dit_denoising_time = f"{stage_execution_times[5]:.2f}s" if len(stage_execution_times) > 5 else "N/A"
|
||||
#
|
||||
# timing_html = f"""
|
||||
# <div style="margin: 10px 0;">
|
||||
# <h3 style="text-align: center; margin-bottom: 10px;">⏱️ Timing Breakdown</h3>
|
||||
# <div style="display: grid; grid-template-columns: repeat(5, 1fr); gap: 10px; margin-bottom: 10px;">
|
||||
# <div class="timing-card timing-card-highlight">
|
||||
# <div style="font-size: 20px;">🚀</div>
|
||||
# <div style="font-weight: bold; margin: 3px 0; font-size: 14px;">DiT Denoising</div>
|
||||
# <div style="font-size: 16px; color: #ffa200; font-weight: bold;">{dit_denoising_time}</div>
|
||||
# </div>
|
||||
# <div class="timing-card">
|
||||
# <div style="font-size: 20px;">🧠</div>
|
||||
# <div style="font-weight: bold; margin: 3px 0; font-size: 14px;">E2E (w. vae/text encoder)</div>
|
||||
# <div style="font-size: 16px; color: #2563eb;">{inference_time:.2f}s</div>
|
||||
# </div>
|
||||
# <div class="timing-card">
|
||||
# <div style="font-size: 20px;">🎬</div>
|
||||
# <div style="font-weight: bold; margin: 3px 0; font-size: 14px;">Video Encoding</div>
|
||||
# <div style="font-size: 16px; color: #dc2626;">{encoding_time:.2f}s</div>
|
||||
# </div>
|
||||
# <div class="timing-card">
|
||||
# <div style="font-size: 20px;">🌐</div>
|
||||
# <div style="font-weight: bold; margin: 3px 0; font-size: 14px;">Network Transfer</div>
|
||||
# <div style="font-size: 16px; color: #059669;">{network_time:.2f}s</div>
|
||||
# </div>
|
||||
# <div class="timing-card">
|
||||
# <div style="font-size: 20px;">📊</div>
|
||||
# <div style="font-weight: bold; margin: 3px 0; font-size: 14px;">Total Processing</div>
|
||||
# <div style="font-size: 18px; color: #0277bd;">{total_time:.2f}s</div>
|
||||
# </div>
|
||||
# </div>"""
|
||||
#
|
||||
# if inference_time > 0:
|
||||
# fps = num_frames / inference_time
|
||||
# timing_html += f"""
|
||||
# <div class="performance-card" style="margin-top: 15px;">
|
||||
# <span style="font-weight: bold;">Generation Speed: </span>
|
||||
# <span style="font-size: 18px; color: #6366f1; font-weight: bold;">{fps:.1f} frames/second</span>
|
||||
# </div>"""
|
||||
#
|
||||
# return timing_html + "</div>"
|
||||
timing_html = f"""
|
||||
<div style="margin: 10px 0;">
|
||||
<h3 style="text-align: center; margin-bottom: 10px;">⏱️ Timing Breakdown</h3>
|
||||
<div style="display: grid; grid-template-columns: repeat(5, 1fr); gap: 10px; margin-bottom: 10px;">
|
||||
<div class="timing-card timing-card-highlight">
|
||||
<div style="font-size: 20px;">🚀</div>
|
||||
<div style="font-weight: bold; margin: 3px 0; font-size: 14px;">DiT Denoising</div>
|
||||
<div style="font-size: 16px; color: #ffa200; font-weight: bold;">{dit_denoising_time}</div>
|
||||
</div>
|
||||
<div class="timing-card">
|
||||
<div style="font-size: 20px;">🧠</div>
|
||||
<div style="font-weight: bold; margin: 3px 0; font-size: 14px;">E2E (w. vae/text encoder)</div>
|
||||
<div style="font-size: 16px; color: #2563eb;">{inference_time:.2f}s</div>
|
||||
</div>
|
||||
<div class="timing-card">
|
||||
<div style="font-size: 20px;">🎬</div>
|
||||
<div style="font-weight: bold; margin: 3px 0; font-size: 14px;">Video Encoding</div>
|
||||
<div style="font-size: 16px; color: #dc2626;">{encoding_time:.2f}s</div>
|
||||
</div>
|
||||
<div class="timing-card">
|
||||
<div style="font-size: 20px;">🌐</div>
|
||||
<div style="font-weight: bold; margin: 3px 0; font-size: 14px;">Network Transfer</div>
|
||||
<div style="font-size: 16px; color: #059669;">{network_time:.2f}s</div>
|
||||
</div>
|
||||
<div class="timing-card">
|
||||
<div style="font-size: 20px;">📊</div>
|
||||
<div style="font-weight: bold; margin: 3px 0; font-size: 14px;">Total Processing</div>
|
||||
<div style="font-size: 18px; color: #0277bd;">{total_time:.2f}s</div>
|
||||
</div>
|
||||
</div>"""
|
||||
|
||||
if inference_time > 0:
|
||||
fps = num_frames / inference_time
|
||||
timing_html += f"""
|
||||
<div class="performance-card" style="margin-top: 15px;">
|
||||
<span style="font-weight: bold;">Generation Speed: </span>
|
||||
<span style="font-size: 18px; color: #6366f1; font-weight: bold;">{fps:.1f} frames/second</span>
|
||||
</div>"""
|
||||
|
||||
return timing_html + "</div>"
|
||||
|
||||
|
||||
def load_example_prompts():
|
||||
@@ -177,83 +144,26 @@ def load_example_prompts():
|
||||
print(f"Warning: Could not read {filepath}: {e}")
|
||||
return prompts, labels
|
||||
|
||||
# Load prompts from prompts.txt
|
||||
examples, example_labels = load_from_file("examples/inference/gradio/serving/prompts.txt")
|
||||
examples, example_labels = load_from_file("prompts/prompts_final.txt")
|
||||
|
||||
if not examples:
|
||||
examples = ["A crowded rooftop bar buzzes with energy, the city skyline twinkling like a field of stars in the background."]
|
||||
example_labels = ["Crowded rooftop bar at night"]
|
||||
|
||||
# Load image mappings from JSON file
|
||||
prompt_to_image = {}
|
||||
# Try to find the JSON file relative to project root
|
||||
possible_json_paths = [
|
||||
Path("prompts/mixkit_i2v.jsonl"),
|
||||
Path(__file__).parent.parent.parent.parent / "prompts" / "mixkit_i2v.jsonl",
|
||||
]
|
||||
json_path = None
|
||||
for path in possible_json_paths:
|
||||
if path.exists():
|
||||
json_path = path
|
||||
break
|
||||
|
||||
if json_path and json_path.exists():
|
||||
try:
|
||||
with open(json_path, "r", encoding='utf-8') as f:
|
||||
data = json.load(f)
|
||||
# Get the project root directory (parent of prompts directory)
|
||||
project_root = json_path.parent.parent
|
||||
for item in data:
|
||||
prompt_text = item.get("prompt", "").strip()
|
||||
image_path = item.get("image_path", "")
|
||||
if prompt_text and image_path:
|
||||
# Resolve image path relative to project root
|
||||
full_image_path = project_root / image_path
|
||||
if full_image_path.exists():
|
||||
prompt_to_image[prompt_text] = str(full_image_path.absolute())
|
||||
except Exception as e:
|
||||
print(f"Warning: Could not load image mappings from {json_path}: {e}")
|
||||
|
||||
# Create image paths list matching the prompts
|
||||
example_images = []
|
||||
for prompt in examples:
|
||||
# Try exact match first
|
||||
image_path = prompt_to_image.get(prompt)
|
||||
if not image_path:
|
||||
# Try fuzzy match (case-insensitive, whitespace normalized)
|
||||
normalized_prompt = " ".join(prompt.split())
|
||||
for json_prompt, img_path in prompt_to_image.items():
|
||||
normalized_json = " ".join(json_prompt.split())
|
||||
if normalized_prompt.lower() == normalized_json.lower():
|
||||
image_path = img_path
|
||||
break
|
||||
example_images.append(image_path if image_path and os.path.exists(image_path) else None)
|
||||
|
||||
return examples, example_labels, example_images
|
||||
return examples, example_labels
|
||||
|
||||
|
||||
def create_gradio_interface(backend_url: str, default_params: dict[str, SamplingParam]):
|
||||
|
||||
client = RayServeClient(backend_url)
|
||||
|
||||
def is_i2v_model(model_name: str) -> bool:
|
||||
"""Check if the model is an I2V model."""
|
||||
return "I2V" in model_name
|
||||
|
||||
def generate_video(
|
||||
prompt, negative_prompt, use_negative_prompt, guidance_scale,
|
||||
num_frames, height, width, model_selection, input_image, progress
|
||||
prompt, negative_prompt, use_negative_prompt, seed, guidance_scale,
|
||||
num_frames, height, width, randomize_seed, model_selection, progress
|
||||
):
|
||||
# Use default seed value (randomize_seed disabled)
|
||||
seed = 1000
|
||||
randomize_seed = False
|
||||
if not client.check_health():
|
||||
return None, f"Backend is not available. Please check if Ray Serve is running at {backend_url}", ""
|
||||
|
||||
# Check if I2V model requires an image
|
||||
if is_i2v_model(model_selection) and not input_image:
|
||||
return None, "I2V models require an input image. Please upload an image.", ""
|
||||
|
||||
# Validate dimensions
|
||||
max_pixels = 720 * 1280
|
||||
if height * width > max_pixels:
|
||||
@@ -262,15 +172,6 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
|
||||
if progress:
|
||||
progress(0.1, desc="Checking backend health...")
|
||||
|
||||
# Encode image if provided
|
||||
image_data = None
|
||||
if input_image:
|
||||
if progress:
|
||||
progress(0.2, desc="Encoding input image...")
|
||||
image_data = encode_image_to_base64(input_image)
|
||||
if not image_data:
|
||||
return None, "Failed to encode input image", ""
|
||||
|
||||
request_data = {
|
||||
"prompt": prompt,
|
||||
"negative_prompt": negative_prompt,
|
||||
@@ -282,7 +183,7 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
|
||||
"width": width,
|
||||
"randomize_seed": randomize_seed,
|
||||
"return_frames": False,
|
||||
"image_data": image_data,
|
||||
"image_path": None,
|
||||
"model_path": MODEL_PATH_MAPPING.get(model_selection, "FastVideo/FastWan2.1-T2V-1.3B-Diffusers")
|
||||
}
|
||||
|
||||
@@ -297,16 +198,16 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
|
||||
if response.get("success", False):
|
||||
video_data = response.get("video_data", "")
|
||||
used_seed = response.get("seed", seed)
|
||||
# inference_time = response.get("inference_time", 0.0)
|
||||
# encoding_time = response.get("encoding_time", 0.0)
|
||||
# total_time = response.get("total_time", 0.0)
|
||||
# network_time = response.get("network_time", 0.0)
|
||||
# stage_execution_times = response.get("stage_execution_times", [])
|
||||
inference_time = response.get("inference_time", 0.0)
|
||||
encoding_time = response.get("encoding_time", 0.0)
|
||||
total_time = response.get("total_time", 0.0)
|
||||
network_time = response.get("network_time", 0.0)
|
||||
stage_execution_times = response.get("stage_execution_times", [])
|
||||
|
||||
# timing_details = create_timing_display(
|
||||
# inference_time, encoding_time, network_time, total_time,
|
||||
# stage_execution_times, num_frames
|
||||
# )
|
||||
timing_details = create_timing_display(
|
||||
inference_time, encoding_time, network_time, total_time,
|
||||
stage_execution_times, num_frames
|
||||
)
|
||||
|
||||
if video_data:
|
||||
if progress:
|
||||
@@ -318,7 +219,7 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
|
||||
progress(1.0, desc="Generation complete!")
|
||||
|
||||
if video_path and os.path.exists(video_path):
|
||||
return video_path, used_seed, ""
|
||||
return video_path, used_seed, timing_details
|
||||
else:
|
||||
return None, "Failed to save video", ""
|
||||
else:
|
||||
@@ -327,7 +228,7 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
|
||||
error_msg = response.get("error_message", "Unknown error occurred")
|
||||
return None, f"Generation failed: {error_msg}", ""
|
||||
|
||||
examples, example_labels, example_images = load_example_prompts()
|
||||
examples, example_labels = load_example_prompts()
|
||||
|
||||
theme = gr.themes.Base().set(
|
||||
button_primary_background_fill="#2563eb",
|
||||
@@ -338,39 +239,33 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
|
||||
)
|
||||
|
||||
def get_default_values(model_name):
|
||||
# model_path = MODEL_PATH_MAPPING.get(model_name)
|
||||
# if model_path and model_path in default_params:
|
||||
# params = default_params[model_path]
|
||||
# return {
|
||||
# 'height': params.height,
|
||||
# 'width': params.width,
|
||||
# 'num_frames': params.num_frames,
|
||||
# 'guidance_scale': params.guidance_scale,
|
||||
# }
|
||||
model_path = MODEL_PATH_MAPPING.get(model_name)
|
||||
if model_path and model_path in default_params:
|
||||
params = default_params[model_path]
|
||||
return {
|
||||
'height': params.height,
|
||||
'width': params.width,
|
||||
'num_frames': params.num_frames,
|
||||
'guidance_scale': params.guidance_scale,
|
||||
'seed': params.seed,
|
||||
}
|
||||
|
||||
return {
|
||||
'height': 480,
|
||||
'height': 448,
|
||||
'width': 832,
|
||||
'num_frames': 73,
|
||||
'num_frames': 61,
|
||||
'guidance_scale': 3.0,
|
||||
'seed': 1024,
|
||||
}
|
||||
|
||||
# Get available models based on what's loaded
|
||||
available_models = []
|
||||
for model_name, model_path in MODEL_PATH_MAPPING.items():
|
||||
if model_path in default_params:
|
||||
available_models.append(model_name)
|
||||
initial_values = get_default_values("FastWan2.1-T2V-1.3B")
|
||||
|
||||
# Select first available model as default
|
||||
default_model = available_models[0] if available_models else "FastWan2.1-T2V-1.3B"
|
||||
initial_values = get_default_values(default_model)
|
||||
initial_show_image = is_i2v_model(default_model)
|
||||
|
||||
with gr.Blocks(title="CausalWan", theme=theme) as demo:
|
||||
with gr.Blocks(title="FastWan", theme=theme) as demo:
|
||||
gr.Image("assets/logos/logo.svg", show_label=False, container=False, height=80)
|
||||
gr.HTML("""
|
||||
<div style="text-align: center; margin-bottom: 10px;">
|
||||
<p style="font-size: 18px;"> Make Video Generation Go Blurrrrrrr </p>
|
||||
<p style="font-size: 18px;"> <a href="https://github.com/hao-ai-lab/FastVideo/tree/main" target="_blank">Code</a> | <a href="https://hao-ai-lab.github.io/blogs/fastvideo_causalwan_preview/" target="_blank">Blog</a> | <a href="https://hao-ai-lab.github.io/FastVideo/" target="_blank">Docs</a> </p>
|
||||
<p style="font-size: 18px;"> <a href="https://github.com/hao-ai-lab/FastVideo/tree/main" target="_blank">Code</a> | <a href="https://hao-ai-lab.github.io/blogs/fastvideo_post_training/" target="_blank">Blog</a> | <a href="https://hao-ai-lab.github.io/FastVideo/" target="_blank">Docs</a> </p>
|
||||
</div>
|
||||
""")
|
||||
|
||||
@@ -385,8 +280,8 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
|
||||
|
||||
with gr.Row():
|
||||
model_selection = gr.Dropdown(
|
||||
choices=available_models,
|
||||
value=default_model,
|
||||
choices=list(MODEL_PATH_MAPPING.keys()),
|
||||
value="FastWan2.1-T2V-1.3B",
|
||||
label="Select Model",
|
||||
interactive=True
|
||||
)
|
||||
@@ -417,70 +312,69 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
|
||||
with gr.Row():
|
||||
with gr.Column():
|
||||
error_output = gr.Text(label="Error", visible=False)
|
||||
# timing_display = gr.Markdown(label="Timing Breakdown", visible=False)
|
||||
timing_display = gr.Markdown(label="Timing Breakdown", visible=False)
|
||||
|
||||
with gr.Row(equal_height=False):
|
||||
with gr.Column(scale=1):
|
||||
with gr.Tabs():
|
||||
with gr.Tab("Input Image", visible=initial_show_image) as image_tab:
|
||||
gr.Markdown("**Please make sure you upload a 480x832 image**")
|
||||
input_image = gr.Image(
|
||||
label="",
|
||||
type="filepath",
|
||||
height=400,
|
||||
with gr.Row(equal_height=True, elem_classes="main-content-row"):
|
||||
with gr.Column(scale=1, elem_classes="advanced-options-column"):
|
||||
with gr.Group():
|
||||
gr.HTML("<div style='margin: 0 0 15px 0; text-align: center; font-size: 16px;'>Advanced Options</div>")
|
||||
with gr.Row():
|
||||
height = gr.Number(
|
||||
label="Height",
|
||||
value=initial_values['height'],
|
||||
interactive=False,
|
||||
container=True
|
||||
)
|
||||
width = gr.Number(
|
||||
label="Width",
|
||||
value=initial_values['width'],
|
||||
interactive=False,
|
||||
container=True
|
||||
)
|
||||
|
||||
with gr.Tab("Advanced Options"):
|
||||
with gr.Group():
|
||||
with gr.Row():
|
||||
height = gr.Number(
|
||||
label="Height",
|
||||
value=initial_values['height'],
|
||||
interactive=False,
|
||||
container=True
|
||||
)
|
||||
width = gr.Number(
|
||||
label="Width",
|
||||
value=initial_values['width'],
|
||||
interactive=False,
|
||||
container=True
|
||||
)
|
||||
|
||||
with gr.Row():
|
||||
num_frames = gr.Number(
|
||||
label="Number of Frames",
|
||||
value=initial_values['num_frames'],
|
||||
interactive=False,
|
||||
container=True
|
||||
)
|
||||
guidance_scale = gr.Number(
|
||||
label="Guidance Scale",
|
||||
value=1.0,
|
||||
interactive=False,
|
||||
container=True
|
||||
)
|
||||
|
||||
with gr.Row():
|
||||
use_negative_prompt = gr.Checkbox(
|
||||
label="Use negative prompt", value=False)
|
||||
negative_prompt = gr.Text(
|
||||
label="Negative prompt",
|
||||
max_lines=3,
|
||||
lines=3,
|
||||
placeholder="Enter a negative prompt",
|
||||
visible=False,
|
||||
)
|
||||
with gr.Row():
|
||||
num_frames = gr.Number(
|
||||
label="Number of Frames",
|
||||
value=initial_values['num_frames'],
|
||||
interactive=False,
|
||||
container=True
|
||||
)
|
||||
guidance_scale = gr.Slider(
|
||||
label="Guidance Scale",
|
||||
minimum=1,
|
||||
maximum=12,
|
||||
value=initial_values['guidance_scale'],
|
||||
)
|
||||
|
||||
with gr.Row():
|
||||
use_negative_prompt = gr.Checkbox(
|
||||
label="Use negative prompt", value=False)
|
||||
negative_prompt = gr.Text(
|
||||
label="Negative prompt",
|
||||
max_lines=3,
|
||||
lines=3,
|
||||
placeholder="Enter a negative prompt",
|
||||
visible=False,
|
||||
)
|
||||
|
||||
# randomize_seed = gr.Checkbox(label="Randomize seed", value=False)
|
||||
seed_output = gr.Number(label="Used Seed", value=1000)
|
||||
seed = gr.Slider(
|
||||
label="Seed",
|
||||
minimum=0,
|
||||
maximum=1000000,
|
||||
step=1,
|
||||
value=initial_values['seed'],
|
||||
)
|
||||
randomize_seed = gr.Checkbox(label="Randomize seed", value=False)
|
||||
seed_output = gr.Number(label="Used Seed")
|
||||
|
||||
with gr.Column(scale=1):
|
||||
with gr.Column(scale=1, elem_classes="video-column"):
|
||||
result = gr.Video(
|
||||
label="Generated Video",
|
||||
show_label=True,
|
||||
height=500,
|
||||
height=466,
|
||||
width=600,
|
||||
container=True,
|
||||
autoplay=True,
|
||||
elem_classes="video-component"
|
||||
)
|
||||
|
||||
gr.HTML("""
|
||||
@@ -493,10 +387,116 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
|
||||
}
|
||||
|
||||
.gradio-container {
|
||||
max-width: 1400px !important;
|
||||
max-width: 1200px !important;
|
||||
margin: 0 auto !important;
|
||||
}
|
||||
|
||||
.main {
|
||||
max-width: 1200px !important;
|
||||
margin: 0 auto !important;
|
||||
}
|
||||
|
||||
.gr-form, .gr-box, .gr-group {
|
||||
max-width: 1200px !important;
|
||||
}
|
||||
|
||||
.gr-video {
|
||||
max-width: 500px !important;
|
||||
margin: 0 auto !important;
|
||||
}
|
||||
|
||||
.main-content-row {
|
||||
display: flex !important;
|
||||
align-items: flex-start !important;
|
||||
min-height: 500px !important;
|
||||
gap: 20px !important;
|
||||
}
|
||||
|
||||
.advanced-options-column,
|
||||
.video-column {
|
||||
display: flex !important;
|
||||
flex-direction: column !important;
|
||||
flex: 1 !important;
|
||||
min-height: 400px !important;
|
||||
align-items: stretch !important;
|
||||
}
|
||||
|
||||
.video-column > * {
|
||||
margin-top: 0 !important;
|
||||
}
|
||||
|
||||
.video-column .gr-video,
|
||||
.video-component {
|
||||
margin-top: 0 !important;
|
||||
padding-top: 0 !important;
|
||||
}
|
||||
|
||||
.video-column .gr-video .gr-form {
|
||||
margin-top: 0 !important;
|
||||
}
|
||||
|
||||
.advanced-options-column .gr-group,
|
||||
.video-column .gr-video {
|
||||
margin-top: 0 !important;
|
||||
vertical-align: top !important;
|
||||
}
|
||||
|
||||
.advanced-options-column > *:last-child,
|
||||
.video-column > *:last-child {
|
||||
flex-grow: 0 !important;
|
||||
}
|
||||
|
||||
@media (max-width: 1400px) {
|
||||
.main-content-row {
|
||||
min-height: 600px !important;
|
||||
}
|
||||
|
||||
.advanced-options-column,
|
||||
.video-column {
|
||||
min-height: 600px !important;
|
||||
}
|
||||
}
|
||||
|
||||
@media (max-width: 1200px) {
|
||||
.main-content-row {
|
||||
flex-direction: column !important;
|
||||
align-items: stretch !important;
|
||||
}
|
||||
|
||||
.advanced-options-column,
|
||||
.video-column {
|
||||
min-height: auto !important;
|
||||
width: 100% !important;
|
||||
}
|
||||
}
|
||||
|
||||
.timing-card {
|
||||
background: var(--background-fill-secondary) !important;
|
||||
border: 1px solid var(--border-color-primary) !important;
|
||||
color: var(--body-text-color) !important;
|
||||
padding: 10px;
|
||||
border-radius: 8px;
|
||||
text-align: center;
|
||||
min-height: 80px;
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
justify-content: center;
|
||||
}
|
||||
|
||||
.timing-card-highlight {
|
||||
background: var(--background-fill-primary) !important;
|
||||
border: 2px solid var(--color-accent) !important;
|
||||
}
|
||||
|
||||
.performance-card {
|
||||
background: var(--background-fill-secondary) !important;
|
||||
border: 1px solid var(--border-color-primary) !important;
|
||||
color: var(--body-text-color) !important;
|
||||
padding: 10px;
|
||||
border-radius: 6px;
|
||||
text-align: center;
|
||||
}
|
||||
|
||||
.gr-number input[readonly] {
|
||||
background-color: var(--background-fill-secondary) !important;
|
||||
border: 1px solid var(--border-color-primary) !important;
|
||||
@@ -511,20 +511,18 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
|
||||
def on_example_select(example_label):
|
||||
if example_label and example_label in example_labels:
|
||||
index = example_labels.index(example_label)
|
||||
selected_prompt = examples[index]
|
||||
selected_image = example_images[index] if index < len(example_images) else None
|
||||
return selected_prompt, selected_image
|
||||
return "", None
|
||||
return examples[index]
|
||||
return ""
|
||||
|
||||
example_dropdown.change(
|
||||
fn=on_example_select,
|
||||
inputs=example_dropdown,
|
||||
outputs=[prompt, input_image],
|
||||
outputs=prompt,
|
||||
)
|
||||
|
||||
gr.HTML("""
|
||||
<div style="text-align: center; margin-top: 10px; margin-bottom: 15px;">
|
||||
<p style="font-size: 16px; margin: 0;">The compute for this demo is generously provided by <a href="https://www.gmicloud.ai/" target="_blank">GMI Cloud</a>. Note that this demo is meant as a preview of our distilled I2V model. Outside of few-step distillation, we have not yet fully optimized it for speed. Stay tuned for updates!</p>
|
||||
<p style="font-size: 16px; margin: 0;">The compute for this demo is generously provided by <a href="https://www.gmicloud.ai/" target="_blank">GMI Cloud</a>. Note that this demo is meant to showcase FastWan's quality and that under a large number of requests, generation speed may be affected. We are also rate-limiting users to 3 requests per minute.</p>
|
||||
</div>
|
||||
""")
|
||||
|
||||
@@ -539,7 +537,6 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
|
||||
selected_model = "FastWan2.1-T2V-1.3B"
|
||||
|
||||
model_path = MODEL_PATH_MAPPING.get(selected_model)
|
||||
show_image_input = is_i2v_model(selected_model)
|
||||
|
||||
if model_path and model_path in default_params:
|
||||
params = default_params[model_path]
|
||||
@@ -548,29 +545,29 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
|
||||
gr.update(value=params.width),
|
||||
gr.update(value=params.num_frames),
|
||||
gr.update(value=params.guidance_scale),
|
||||
gr.update(visible=show_image_input),
|
||||
gr.update(value=params.seed),
|
||||
)
|
||||
|
||||
return (
|
||||
gr.update(value=448),
|
||||
gr.update(value=832),
|
||||
gr.update(value=20),
|
||||
gr.update(value=61),
|
||||
gr.update(value=3.0),
|
||||
gr.update(visible=show_image_input),
|
||||
gr.update(value=1024),
|
||||
)
|
||||
|
||||
model_selection.change(
|
||||
fn=on_model_selection_change,
|
||||
inputs=model_selection,
|
||||
outputs=[height, width, num_frames, guidance_scale, image_tab],
|
||||
outputs=[height, width, num_frames, guidance_scale, seed],
|
||||
)
|
||||
|
||||
def handle_generation(*args, progress=None, request: gr.Request = None):
|
||||
model_selection, prompt, negative_prompt, use_negative_prompt, guidance_scale, num_frames, height, width, input_image = args
|
||||
model_selection, prompt, negative_prompt, use_negative_prompt, seed, guidance_scale, num_frames, height, width, randomize_seed = args
|
||||
|
||||
result_path, seed_or_error, _ = generate_video(
|
||||
prompt, negative_prompt, use_negative_prompt, guidance_scale,
|
||||
num_frames, height, width, model_selection, input_image, progress
|
||||
result_path, seed_or_error, timing_details = generate_video(
|
||||
prompt, negative_prompt, use_negative_prompt, seed, guidance_scale,
|
||||
num_frames, height, width, randomize_seed, model_selection, progress
|
||||
)
|
||||
|
||||
if result_path and os.path.exists(result_path):
|
||||
@@ -578,12 +575,14 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
|
||||
result_path,
|
||||
seed_or_error,
|
||||
gr.update(visible=False),
|
||||
gr.update(visible=True, value=timing_details),
|
||||
)
|
||||
else:
|
||||
return (
|
||||
None,
|
||||
seed_or_error,
|
||||
gr.update(visible=True, value=seed_or_error),
|
||||
gr.update(visible=False),
|
||||
)
|
||||
|
||||
run_button.click(
|
||||
@@ -593,14 +592,14 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
|
||||
prompt,
|
||||
negative_prompt,
|
||||
use_negative_prompt,
|
||||
seed,
|
||||
guidance_scale,
|
||||
num_frames,
|
||||
height,
|
||||
width,
|
||||
# randomize_seed,
|
||||
input_image,
|
||||
randomize_seed,
|
||||
],
|
||||
outputs=[result, seed_output, error_output], # timing_display removed
|
||||
outputs=[result, seed_output, error_output, timing_display],
|
||||
concurrency_limit=20,
|
||||
)
|
||||
|
||||
@@ -612,11 +611,8 @@ def main():
|
||||
parser.add_argument("--backend_url", type=str, default="http://localhost:8000",
|
||||
help="URL of the Ray Serve backend")
|
||||
parser.add_argument("--t2v_model_paths", type=str,
|
||||
default="",
|
||||
default="FastVideo/FastWan2.1-T2V-1.3B-Diffusers,FastVideo/FastWan2.1-T2V-14B-Diffusers",
|
||||
help="Comma separated list of paths to the T2V model(s)")
|
||||
parser.add_argument("--i2v_model_paths", type=str,
|
||||
default="FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers",
|
||||
help="Comma separated list of paths to the I2V model(s)")
|
||||
parser.add_argument("--host", type=str, default="0.0.0.0",
|
||||
help="Host to bind to")
|
||||
parser.add_argument("--port", type=int, default=7860,
|
||||
@@ -625,15 +621,8 @@ def main():
|
||||
args = parser.parse_args()
|
||||
|
||||
default_params = {}
|
||||
|
||||
# Load T2V model params
|
||||
t2v_paths = [p.strip() for p in args.t2v_model_paths.split(",") if p.strip()]
|
||||
for model_path in t2v_paths:
|
||||
default_params[model_path] = SamplingParam.from_pretrained(model_path)
|
||||
|
||||
# Load I2V model params
|
||||
i2v_paths = [p.strip() for p in args.i2v_model_paths.split(",") if p.strip()]
|
||||
for model_path in i2v_paths:
|
||||
model_paths = args.t2v_model_paths.split(",")
|
||||
for model_path in model_paths:
|
||||
default_params[model_path] = SamplingParam.from_pretrained(model_path)
|
||||
|
||||
demo = create_gradio_interface(args.backend_url, default_params)
|
||||
@@ -641,8 +630,6 @@ def main():
|
||||
print(f"Starting Gradio frontend at http://{args.host}:{args.port}")
|
||||
print(f"Backend URL: {args.backend_url}")
|
||||
print(f"T2V Models: {args.t2v_model_paths}")
|
||||
if args.i2v_model_paths:
|
||||
print(f"I2V Models: {args.i2v_model_paths}")
|
||||
|
||||
from fastapi import FastAPI, Request, HTTPException
|
||||
from fastapi.responses import HTMLResponse, FileResponse
|
||||
@@ -687,23 +674,23 @@ def main():
|
||||
<meta charset="UTF-8" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
|
||||
|
||||
<title>CausalWan</title>
|
||||
<meta name="title" content="CausalWan">
|
||||
<title>FastWan</title>
|
||||
<meta name="title" content="FastWan">
|
||||
<meta name="description" content="Make video generation go blurrrrrrr">
|
||||
<meta name="keywords" content="FastVideo, video generation, AI, machine learning, CausalWan">
|
||||
<meta name="keywords" content="FastVideo, video generation, AI, machine learning, FastWan">
|
||||
|
||||
<meta property="og:type" content="website">
|
||||
<meta property="og:url" content="{base_url}/">
|
||||
<meta property="og:title" content="CausalWan">
|
||||
<meta property="og:title" content="FastWan">
|
||||
<meta property="og:description" content="Make video generation go blurrrrrrr">
|
||||
<meta property="og:image" content="{base_url}/logo.svg">
|
||||
<meta property="og:image:width" content="1200">
|
||||
<meta property="og:image:height" content="630">
|
||||
<meta property="og:site_name" content="CausalWan">
|
||||
<meta property="og:site_name" content="FastWan">
|
||||
|
||||
<meta property="twitter:card" content="summary_large_image">
|
||||
<meta property="twitter:url" content="{base_url}/">
|
||||
<meta property="twitter:title" content="CausalWan">
|
||||
<meta property="twitter:title" content="FastWan">
|
||||
<meta property="twitter:description" content="Make video generation go blurrrrrrr">
|
||||
<meta property="twitter:image" content="{base_url}/logo.svg">
|
||||
<link rel="icon" type="image/png" sizes="32x32" href="/favicon.ico">
|
||||
@@ -733,14 +720,7 @@ def main():
|
||||
app,
|
||||
demo,
|
||||
path="/gradio",
|
||||
allowed_paths=[
|
||||
os.path.abspath("outputs"),
|
||||
os.path.abspath("fastvideo-logos"),
|
||||
os.path.abspath("prompts"),
|
||||
os.path.abspath("images"),
|
||||
os.path.abspath(tempfile.gettempdir()),
|
||||
os.path.abspath(os.path.join(tempfile.gettempdir(), "gradio")),
|
||||
]
|
||||
allowed_paths=[os.path.abspath("outputs"), os.path.abspath("fastvideo-logos")]
|
||||
)
|
||||
|
||||
uvicorn.run(app, host=args.host, port=args.port)
|
||||
|
||||
@@ -1,15 +0,0 @@
|
||||
Some friends dancing and having fun together in circles, at a party surrounded by colored lights at a party, in a fancy old place, in a view from below them.
|
||||
A man wearing grey shorts jumps rope in a gym, weights and gym equipment in the background.
|
||||
Flying over a peninsula covered in bushy trees, while discovering the sea around it, painted a beautiful turquoise blue, on a sunny day.
|
||||
Skillful cyclist doing a wheelie on a bike while riding through a forest, on a dirt road, surrounded by many trees, in the morning.
|
||||
In the video, a lone rider guides a majestic horse across an expansive, open field as the sun sets in the background. The rider, dressed in a classic blue shirt and wide-brimmed hat, sits confidently in the saddle, silhouetted against the warm glow of the evening sky. The horse moves gracefully, its mane and tail flowing with each step, creating a sense of harmony between horse and rider. Surrounding the pair, towering trees form a natural border, their leaves gently rustling in the breeze. The shadows lengthen on the ground, accentuating the serene and timeless feel of the scene. The distant hills and wooden fences frame the horizon, adding depth to the tranquil landscape. A few horses graze peacefully in the background, blending into the pastoral setting. The overall ambiance evokes a sense of calmness and quietude, capturing a perfect moment in the golden light of dusk.
|
||||
A silver SUV drives along a winding, snow-covered mountain road, with dense pine trees blanketed in snow lining both sides. The scene is serene, with the vehicle moving smoothly, possibly on a winter journey or vacation. As the SUV disappears around the bend, another, darker SUV follows, creating a sense of motion and perspective on the snow-dusted asphalt. The towering, snow-laden rock formation to the right contrasts with the dark green of the pines, highlighting the peacefulness of the wintry landscape.
|
||||
A man and a woman playing in a field with grass, during a bright afternoon, while cars pass by in the distance.
|
||||
A saxophonist wearing a blazer dances while playing a song in a park.
|
||||
Romantic couple embracing and looking at each other in the middle of a forest, during a break on a road trip through nature.
|
||||
Young woman cleaning her house decorated with plants and decorations, while dancing happily to music in her headphones.
|
||||
Man dressed in 80's style dances very happily in his kitchen while listening to music on his radio and drinking wine.
|
||||
Pair of jazz musicians performing a song with their saxophone and trombone on an abandoned train.
|
||||
A young woman with short hair wearing pink sunglasses chews gum and makes a bubble gum with the city in the background.
|
||||
Natural aerial landscape with a relief covered with abundant trees and vegetation and a thick layer of mist.
|
||||
Loving couple sitting on a log on the shore of a lake outside, sharing an affectionate hug.
|
||||
@@ -26,7 +26,6 @@ SEED_RANGE_MAX = 1_000_000
|
||||
SUPPORTED_MODELS = [
|
||||
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers",
|
||||
"FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers",
|
||||
"FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers",
|
||||
]
|
||||
|
||||
MODEL_CONFIGS = {
|
||||
@@ -43,13 +42,6 @@ MODEL_CONFIGS = {
|
||||
"dit_cpu_offload": True,
|
||||
"vae_cpu_offload": False,
|
||||
"VSA_sparsity": 0.9,
|
||||
},
|
||||
"I2V-A14B": {
|
||||
"num_cpus": 15,
|
||||
"text_encoder_cpu_offload": True,
|
||||
"dit_cpu_offload": True,
|
||||
"vae_cpu_offload": False,
|
||||
"VSA_sparsity": 0.0,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -66,7 +58,6 @@ class VideoGenerationRequest(BaseModel):
|
||||
randomize_seed: bool = False
|
||||
return_frames: bool = False
|
||||
model_path: Optional[str] = None
|
||||
image_data: Optional[str] = None # Base64 encoded image for I2V
|
||||
|
||||
|
||||
class VideoGenerationResponse(BaseModel):
|
||||
@@ -100,38 +91,11 @@ def encode_video_to_base64(frames: List[np.ndarray], fps: int = DEFAULT_FPS) ->
|
||||
return ""
|
||||
|
||||
|
||||
def save_image_from_base64(image_data: str, output_dir: str) -> Optional[str]:
|
||||
"""Save base64 image data to a temporary file and return the path."""
|
||||
if not image_data:
|
||||
return None
|
||||
|
||||
try:
|
||||
# Remove data URL prefix if present
|
||||
if image_data.startswith('data:image/'):
|
||||
image_data = image_data.split(',')[1]
|
||||
|
||||
image_bytes = base64.b64decode(image_data)
|
||||
|
||||
# Save to temporary file
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
temp_image_path = os.path.join(output_dir, f"temp_input_{int(time.time() * 1000)}.png")
|
||||
|
||||
with open(temp_image_path, 'wb') as f:
|
||||
f.write(image_bytes)
|
||||
|
||||
return temp_image_path
|
||||
|
||||
except Exception as e:
|
||||
print(f"Warning: Failed to save image: {e}")
|
||||
return None
|
||||
|
||||
|
||||
def setup_model_environment(model_path: str) -> None:
|
||||
# if "fullattn" in model_path.lower():
|
||||
# os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "FLASH_ATTN"
|
||||
# else:
|
||||
# os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "VIDEO_SPARSE_ATTN"
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "FLASH_ATTN"
|
||||
if "fullattn" in model_path.lower():
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "FLASH_ATTN"
|
||||
else:
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "VIDEO_SPARSE_ATTN"
|
||||
os.environ["FASTVIDEO_STAGE_LOGGING"] = "1"
|
||||
|
||||
|
||||
@@ -193,41 +157,22 @@ class BaseModelDeployment:
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
text_encoder_cpu_offload=config["text_encoder_cpu_offload"],
|
||||
dmd_denoising_steps=[1000, 850, 700, 550, 350, 275, 200, 125], # TODO: hardocde for I2V
|
||||
dit_precision="fp32", # TODO: hardocde for I2V
|
||||
dit_cpu_offload=config["dit_cpu_offload"],
|
||||
vae_cpu_offload=config["vae_cpu_offload"],
|
||||
VSA_sparsity=config["VSA_sparsity"],
|
||||
enable_stage_verification=False,
|
||||
)
|
||||
self.default_params = SamplingParam.from_pretrained(self.model_path)
|
||||
self.default_params.seed = 1000
|
||||
self.default_params.num_frames = 73
|
||||
self.default_params.width = 832
|
||||
self.default_params.height = 480
|
||||
|
||||
def generate_video(self, video_request: VideoGenerationRequest) -> VideoGenerationResponse:
|
||||
total_start_time = time.time()
|
||||
|
||||
params = prepare_sampling_params(video_request, self.default_params)
|
||||
|
||||
# Save image if provided (for I2V)
|
||||
image_path = None
|
||||
if video_request.image_data:
|
||||
image_path = save_image_from_base64(video_request.image_data, self.output_path)
|
||||
if image_path is None:
|
||||
return VideoGenerationResponse(
|
||||
video_data=None,
|
||||
seed=params.seed,
|
||||
success=False,
|
||||
error_message="Failed to save input image",
|
||||
)
|
||||
|
||||
inference_start_time = time.time()
|
||||
result = self.generator.generate_video(
|
||||
prompt=video_request.prompt,
|
||||
sampling_param=params,
|
||||
image_path=image_path,
|
||||
save_video=False,
|
||||
return_frames=False,
|
||||
)
|
||||
@@ -240,13 +185,6 @@ class BaseModelDeployment:
|
||||
encoding_time = time.time() - encoding_start_time
|
||||
|
||||
total_time = time.time() - total_start_time
|
||||
|
||||
# Clean up temporary image file
|
||||
if image_path and os.path.exists(image_path):
|
||||
try:
|
||||
os.remove(image_path)
|
||||
except Exception as e:
|
||||
print(f"Warning: Failed to remove temporary image file {image_path}: {e}")
|
||||
|
||||
return VideoGenerationResponse(
|
||||
video_data=video_data,
|
||||
@@ -262,7 +200,7 @@ class BaseModelDeployment:
|
||||
|
||||
|
||||
@serve.deployment(
|
||||
ray_actor_options={"num_cpus": 2, "num_gpus": 1, "runtime_env": {"conda": "demo-fv"}},
|
||||
ray_actor_options={"num_cpus": 2, "num_gpus": 1, "runtime_env": {"conda": "fv"}},
|
||||
)
|
||||
class T2VModelDeployment(BaseModelDeployment):
|
||||
def __init__(self, t2v_model_path: str, output_path: str = "outputs"):
|
||||
@@ -272,7 +210,7 @@ class T2VModelDeployment(BaseModelDeployment):
|
||||
|
||||
|
||||
@serve.deployment(
|
||||
ray_actor_options={"num_cpus": 16, "num_gpus": 1, "runtime_env": {"conda": "demo-fv"}},
|
||||
ray_actor_options={"num_cpus": 16, "num_gpus": 1, "runtime_env": {"conda": "fv"}},
|
||||
)
|
||||
class T2V14BModelDeployment(BaseModelDeployment):
|
||||
def __init__(self, t2v_14b_model_path: str, output_path: str = "outputs"):
|
||||
@@ -283,32 +221,18 @@ class T2V14BModelDeployment(BaseModelDeployment):
|
||||
print("✅ T2V 14B model initialized successfully")
|
||||
|
||||
|
||||
@serve.deployment(
|
||||
ray_actor_options={"num_cpus": 15, "num_gpus": 1, "runtime_env": {"conda": "demo-fv"}},
|
||||
)
|
||||
class I2VModelDeployment(BaseModelDeployment):
|
||||
def __init__(self, i2v_model_path: str, output_path: str = "outputs"):
|
||||
super().__init__(i2v_model_path, output_path)
|
||||
# Override environment for I2V model
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "FLASH_ATTN"
|
||||
self._initialize_generator(MODEL_CONFIGS["I2V-A14B"])
|
||||
print("✅ I2V model initialized successfully")
|
||||
|
||||
|
||||
app = FastAPI()
|
||||
limiter = Limiter(key_func=get_remote_address)
|
||||
app.state.limiter = limiter
|
||||
app.add_exception_handler(RateLimitExceeded, _rate_limit_exceeded_handler)
|
||||
|
||||
|
||||
@serve.deployment(num_replicas=1, ray_actor_options={"num_cpus": 1})
|
||||
@serve.deployment(num_replicas=50, ray_actor_options={"num_cpus": 2})
|
||||
@serve.ingress(app)
|
||||
class FastVideoAPI:
|
||||
|
||||
def __init__(self, t2v_deployments: Dict[str, DeploymentHandle], i2v_deployments: Dict[str, DeploymentHandle] = None):
|
||||
def __init__(self, t2v_deployments: Dict[str, DeploymentHandle]):
|
||||
self.t2v_deployments = t2v_deployments
|
||||
self.i2v_deployments = i2v_deployments or {}
|
||||
self.all_deployments = {**self.t2v_deployments, **self.i2v_deployments}
|
||||
|
||||
# Initialize Prometheus metrics
|
||||
self.request_count = Counter('fastvideo_requests_total', 'Total FastVideo requests', ['model_type', 'status'])
|
||||
@@ -333,10 +257,10 @@ class FastVideoAPI:
|
||||
model_name = self._get_model_name(video_request.model_path)
|
||||
|
||||
try:
|
||||
if video_request.model_path not in self.all_deployments:
|
||||
if video_request.model_path not in self.t2v_deployments:
|
||||
raise ValueError(f"Model {video_request.model_path} not found")
|
||||
|
||||
response_ref = self.all_deployments[video_request.model_path].generate_video.remote(video_request)
|
||||
response_ref = self.t2v_deployments[video_request.model_path].generate_video.remote(video_request)
|
||||
response = await response_ref
|
||||
|
||||
self._record_metrics(model_name, "success", time.time() - start_time, response)
|
||||
@@ -367,21 +291,18 @@ class FastVideoAPI:
|
||||
|
||||
|
||||
def validate_configuration(model_paths: List[str], replicas: List[int]) -> None:
|
||||
assert len(model_paths) > 0, "At least one model must be specified"
|
||||
assert len(model_paths) == len(replicas), "Number of models and replicas must match"
|
||||
assert sum(replicas) <= NUM_GPUS, f"Total replicas ({sum(replicas)}) must be <= {NUM_GPUS}"
|
||||
|
||||
for model, replica_count in zip(model_paths, replicas):
|
||||
assert model in SUPPORTED_MODELS, f"Model {model} not supported. Supported models: {SUPPORTED_MODELS}"
|
||||
assert model in SUPPORTED_MODELS, f"Model {model} not supported"
|
||||
assert replica_count > 0, f"Replicas must be greater than 0"
|
||||
|
||||
|
||||
def start_ray_serve(
|
||||
*,
|
||||
t2v_model_paths: str = "",
|
||||
t2v_model_replicas: str = "",
|
||||
i2v_model_paths: str = "",
|
||||
i2v_model_replicas: str = "",
|
||||
t2v_model_paths: str,
|
||||
t2v_model_replicas: str,
|
||||
output_path: str = "outputs",
|
||||
host: str = "0.0.0.0",
|
||||
port: int = 8000,
|
||||
@@ -389,39 +310,21 @@ def start_ray_serve(
|
||||
if not ray.is_initialized():
|
||||
ray.init()
|
||||
|
||||
# Parse T2V models
|
||||
t2v_paths = [p.strip() for p in t2v_model_paths.split(",") if p.strip()]
|
||||
t2v_reps = [int(r.strip()) for r in t2v_model_replicas.split(",") if r.strip()] if t2v_model_replicas else []
|
||||
|
||||
# Parse I2V models
|
||||
i2v_paths = [p.strip() for p in i2v_model_paths.split(",") if p.strip()]
|
||||
i2v_reps = [int(r.strip()) for r in i2v_model_replicas.split(",") if r.strip()] if i2v_model_replicas else []
|
||||
|
||||
# Validate configurations
|
||||
all_paths = t2v_paths + i2v_paths
|
||||
all_replicas = t2v_reps + i2v_reps
|
||||
validate_configuration(all_paths, all_replicas)
|
||||
model_paths = t2v_model_paths.split(",")
|
||||
replicas = [int(r) for r in t2v_model_replicas.split(",")]
|
||||
validate_configuration(model_paths, replicas)
|
||||
|
||||
# Create T2V deployments
|
||||
t2v_deps = {}
|
||||
for model_path, replica_count in zip(t2v_paths, t2v_reps):
|
||||
for model_path, replica_count in zip(model_paths, replicas):
|
||||
t2v_dep = T2VModelDeployment.options(num_replicas=replica_count).bind(model_path, output_path)
|
||||
t2v_deps[model_path] = t2v_dep
|
||||
|
||||
# Create I2V deployments
|
||||
i2v_deps = {}
|
||||
for model_path, replica_count in zip(i2v_paths, i2v_reps):
|
||||
i2v_dep = I2VModelDeployment.options(num_replicas=replica_count).bind(model_path, output_path)
|
||||
i2v_deps[model_path] = i2v_dep
|
||||
|
||||
api = FastVideoAPI.bind(t2v_deps, i2v_deps)
|
||||
api = FastVideoAPI.bind(t2v_deps)
|
||||
serve.run(api, route_prefix="/", name="fast_video")
|
||||
|
||||
print(f"Ray Serve backend started at http://{host}:{port}")
|
||||
for model_path, replica_count in zip(t2v_paths, t2v_reps):
|
||||
for model_path, replica_count in zip(model_paths, replicas):
|
||||
print(f"T2V Model: {model_path} | Replicas: {replica_count}")
|
||||
for model_path, replica_count in zip(i2v_paths, i2v_reps):
|
||||
print(f"I2V Model: {model_path} | Replicas: {replica_count}")
|
||||
print(f"Health check: http://{host}:{port}/health")
|
||||
print(f"Video generation endpoint: http://{host}:{port}/generate_video")
|
||||
|
||||
@@ -437,20 +340,12 @@ if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser(description="FastVideo Ray Serve Backend")
|
||||
parser.add_argument("--t2v_model_paths",
|
||||
type=str,
|
||||
default="",
|
||||
default="FastVideo/FastWan2.1-T2V-1.3B-Diffusers,FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers",
|
||||
help="Comma separated list of paths to the T2V model(s)")
|
||||
parser.add_argument("--t2v_model_replicas",
|
||||
type=str,
|
||||
default="",
|
||||
default="4,4",
|
||||
help="Comma separated list of number of replicas for the T2V model(s)")
|
||||
parser.add_argument("--i2v_model_paths",
|
||||
type=str,
|
||||
default="FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers",
|
||||
help="Comma separated list of paths to the I2V model(s)")
|
||||
parser.add_argument("--i2v_model_replicas",
|
||||
type=str,
|
||||
default="1",
|
||||
help="Comma separated list of number of replicas for the I2V model(s)")
|
||||
parser.add_argument("--output_path",
|
||||
type=str,
|
||||
default="outputs",
|
||||
@@ -466,21 +361,13 @@ if __name__ == "__main__":
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
# Parse and validate all models
|
||||
t2v_paths = [p.strip() for p in args.t2v_model_paths.split(",") if p.strip()]
|
||||
t2v_reps = [int(r.strip()) for r in args.t2v_model_replicas.split(",") if r.strip()] if args.t2v_model_replicas else []
|
||||
i2v_paths = [p.strip() for p in args.i2v_model_paths.split(",") if p.strip()]
|
||||
i2v_reps = [int(r.strip()) for r in args.i2v_model_replicas.split(",") if r.strip()] if args.i2v_model_replicas else []
|
||||
|
||||
all_paths = t2v_paths + i2v_paths
|
||||
all_replicas = t2v_reps + i2v_reps
|
||||
validate_configuration(all_paths, all_replicas)
|
||||
model_paths = args.t2v_model_paths.split(",")
|
||||
replicas = [int(r) for r in args.t2v_model_replicas.split(",")]
|
||||
validate_configuration(model_paths, replicas)
|
||||
|
||||
start_ray_serve(
|
||||
t2v_model_paths=args.t2v_model_paths,
|
||||
t2v_model_replicas=args.t2v_model_replicas,
|
||||
i2v_model_paths=args.i2v_model_paths,
|
||||
i2v_model_replicas=args.i2v_model_replicas,
|
||||
output_path=args.output_path,
|
||||
host=args.host,
|
||||
port=args.port,
|
||||
@@ -489,4 +376,4 @@ if __name__ == "__main__":
|
||||
setup_signal_handlers()
|
||||
print("✅ FastVideo backend is running. Press Ctrl-C to stop.")
|
||||
while True:
|
||||
time.sleep(3600)
|
||||
time.sleep(3600)
|
||||
@@ -1,5 +1,3 @@
|
||||
python examples/inference/gradio/serving/start_ray_serve_app.py \
|
||||
--t2v_model_paths "" \
|
||||
--t2v_model_replicas "" \
|
||||
--i2v_model_paths "FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers" \
|
||||
--i2v_model_replicas "1"
|
||||
python examples/inference/gradio/start_ray_serve_app.py \
|
||||
--t2v_model_paths "FastVideo/FastWan2.1-T2V-1.3B-Diffusers,FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers" \
|
||||
--t2v_model_replicas "4,4"
|
||||
@@ -20,10 +20,8 @@ DEFAULT_BACKEND_PORT = 8000
|
||||
DEFAULT_FRONTEND_HOST = "0.0.0.0"
|
||||
DEFAULT_FRONTEND_PORT = 7860
|
||||
DEFAULT_OUTPUT_PATH = "outputs"
|
||||
DEFAULT_T2V_MODELS = ""
|
||||
DEFAULT_T2V_REPLICAS = ""
|
||||
DEFAULT_I2V_MODELS = "FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers"
|
||||
DEFAULT_I2V_REPLICAS = "1"
|
||||
DEFAULT_T2V_MODELS = "FastVideo/FastWan2.1-T2V-1.3B-Diffusers,FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers"
|
||||
DEFAULT_T2V_REPLICAS = "4,4"
|
||||
|
||||
HEALTH_CHECK_TIMEOUT = 5
|
||||
HEALTH_CHECK_MAX_RETRIES = 100
|
||||
@@ -102,12 +100,6 @@ class ServiceManager:
|
||||
"port": self.args.backend_port
|
||||
}
|
||||
|
||||
# Add I2V parameters if provided
|
||||
if self.args.i2v_model_paths:
|
||||
backend_args["i2v_model_paths"] = self.args.i2v_model_paths
|
||||
if self.args.i2v_model_replicas:
|
||||
backend_args["i2v_model_replicas"] = self.args.i2v_model_replicas
|
||||
|
||||
self.backend_process = self._start_service("ray_serve_backend.py", backend_args, "backend")
|
||||
return self.backend_process
|
||||
|
||||
@@ -119,10 +111,6 @@ class ServiceManager:
|
||||
"port": self.args.frontend_port
|
||||
}
|
||||
|
||||
# Add I2V parameters if provided
|
||||
if self.args.i2v_model_paths:
|
||||
frontend_args["i2v_model_paths"] = self.args.i2v_model_paths
|
||||
|
||||
self.frontend_process = self._start_service("gradio_frontend.py", frontend_args, "frontend")
|
||||
return self.frontend_process
|
||||
|
||||
@@ -185,9 +173,6 @@ def print_startup_info(args: argparse.Namespace) -> None:
|
||||
print("=" * 50)
|
||||
print(f"T2V Models: {args.t2v_model_paths}")
|
||||
print(f"T2V Model Replicas: {args.t2v_model_replicas}")
|
||||
if args.i2v_model_paths:
|
||||
print(f"I2V Models: {args.i2v_model_paths}")
|
||||
print(f"I2V Model Replicas: {args.i2v_model_replicas}")
|
||||
print(f"Output: {args.output_path}")
|
||||
print(f"Backend: http://{args.backend_host}:{args.backend_port}")
|
||||
print(f"Frontend: http://{args.frontend_host}:{args.frontend_port}")
|
||||
@@ -205,14 +190,6 @@ def parse_arguments() -> argparse.Namespace:
|
||||
type=str,
|
||||
default=DEFAULT_T2V_REPLICAS,
|
||||
help="Comma separated list of number of replicas for the T2V model(s)")
|
||||
parser.add_argument("--i2v_model_paths",
|
||||
type=str,
|
||||
default=DEFAULT_I2V_MODELS,
|
||||
help="Comma separated list of paths to the I2V model(s)")
|
||||
parser.add_argument("--i2v_model_replicas",
|
||||
type=str,
|
||||
default=DEFAULT_I2V_REPLICAS,
|
||||
help="Comma separated list of number of replicas for the I2V model(s)")
|
||||
parser.add_argument("--output_path",
|
||||
type=str,
|
||||
default=DEFAULT_OUTPUT_PATH,
|
||||
|
||||
@@ -5,12 +5,12 @@ These are e2e example scripts for finetuning Wan2.1 T2V 1.3B on the crush-smol d
|
||||
|
||||
### Download crush-smol dataset:
|
||||
|
||||
`bash examples/training/finetune/wan_t2v_1.3B/crush_smol/download_dataset.sh`
|
||||
`bash examples/training/finetune/wan_t2v_1_3b/crush_smol/download_dataset.sh`
|
||||
|
||||
### Preprocess the videos and captions into latents:
|
||||
|
||||
`bash examples/training/finetune/wan_t2v_1.3B/crush_smol/preprocess_wan_data_t2v.sh`
|
||||
`bash examples/training/finetune/wan_t2v_1_3b/crush_smol/preprocess_wan_data_t2v.sh`
|
||||
|
||||
### Edit the following file and run finetuning:
|
||||
|
||||
`bash examples/training/finetune/wan_t2v_1.3B/crush_smol/finetune_t2v.sh`
|
||||
`bash examples/training/finetune/wan_t2v_1_3b/crush_smol/finetune_t2v.sh`
|
||||
|
||||
@@ -54,7 +54,7 @@ validation_args=(
|
||||
--validation_dataset_file $VALIDATION_DATASET_FILE
|
||||
--validation_steps 200
|
||||
--validation_sampling_steps "50"
|
||||
--validation_guidance_scale "3.0"
|
||||
--validation_guidance_scale "6.0"
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
|
||||
@@ -3,4 +3,4 @@ from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.entrypoints.video_generator import VideoGenerator
|
||||
from fastvideo.version import __version__
|
||||
|
||||
__all__ = ["VideoGenerator", "PipelineConfig", "SamplingParam", "__version__"]
|
||||
__all__ = ["VideoGenerator", "PipelineConfig", "SamplingParam", "__version__"]
|
||||
@@ -83,9 +83,6 @@ class PreprocessConfig:
|
||||
speed_factor: float = 1.0
|
||||
drop_short_ratio: float = 1.0
|
||||
do_temporal_sample: bool = False
|
||||
enable_smart_resize: bool = False
|
||||
smart_resize_max_area: int | None = None
|
||||
hw_aspect_threshold: float = 1.5
|
||||
|
||||
# Model configuration
|
||||
training_cfg_rate: float = 0.0
|
||||
@@ -187,23 +184,6 @@ class PreprocessConfig:
|
||||
action=StoreBoolean,
|
||||
default=PreprocessConfig.do_temporal_sample,
|
||||
help="Whether to do temporal sampling")
|
||||
preprocess_args.add_argument(
|
||||
f"--{prefix_with_dot}enable-smart-resize",
|
||||
action=StoreBoolean,
|
||||
default=PreprocessConfig.enable_smart_resize,
|
||||
help="Whether to enable smart resizing")
|
||||
preprocess_args.add_argument(
|
||||
f"--{prefix_with_dot}smart-resize-max-area",
|
||||
type=int,
|
||||
default=PreprocessConfig.smart_resize_max_area,
|
||||
help="Maximum area for smart resizing")
|
||||
preprocess_args.add_argument(
|
||||
f"--{prefix_with_dot}hw-aspect-threshold",
|
||||
type=float,
|
||||
default=PreprocessConfig.hw_aspect_threshold,
|
||||
help=
|
||||
"Height/Width aspect ratio threshold. Allowed range is [1/threshold * target_aspect, threshold * target_aspect]."
|
||||
)
|
||||
|
||||
# Model Training configuration
|
||||
preprocess_args.add_argument(f"--{prefix_with_dot}training-cfg-rate",
|
||||
|
||||
@@ -1,10 +1,5 @@
|
||||
from fastvideo.configs.models.dits.cosmos import CosmosVideoConfig
|
||||
from fastvideo.configs.models.dits.cosmos2_5 import Cosmos25VideoConfig
|
||||
from fastvideo.configs.models.dits.hunyuanvideo import HunyuanVideoConfig
|
||||
from fastvideo.configs.models.dits.stepvideo import StepVideoConfig
|
||||
from fastvideo.configs.models.dits.wanvideo import WanVideoConfig
|
||||
|
||||
__all__ = [
|
||||
"HunyuanVideoConfig", "WanVideoConfig", "StepVideoConfig",
|
||||
"CosmosVideoConfig", "Cosmos25VideoConfig"
|
||||
]
|
||||
__all__ = ["HunyuanVideoConfig", "WanVideoConfig", "StepVideoConfig"]
|
||||
|
||||
@@ -1,104 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
|
||||
|
||||
def is_transformer_blocks(n: str, m) -> bool:
|
||||
return "transformer_blocks" in n and str.isdigit(n.split(".")[-1])
|
||||
|
||||
|
||||
@dataclass
|
||||
class CosmosArchConfig(DiTArchConfig):
|
||||
_fsdp_shard_conditions: list = field(
|
||||
default_factory=lambda: [is_transformer_blocks])
|
||||
|
||||
param_names_mapping: dict = field(
|
||||
default_factory=lambda: {
|
||||
r"^patch_embed\.(.*)$": r"patch_embed.\1",
|
||||
r"^time_embed\.time_proj\.(.*)$": r"time_embed.time_proj.\1",
|
||||
r"^time_embed\.t_embedder\.(.*)$": r"time_embed.t_embedder.\1",
|
||||
r"^time_embed\.norm\.(.*)$": r"time_embed.norm.\1",
|
||||
r"^transformer_blocks\.(\d+)\.attn1\.to_q\.(.*)$":
|
||||
r"transformer_blocks.\1.attn1.to_q.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn1\.to_k\.(.*)$":
|
||||
r"transformer_blocks.\1.attn1.to_k.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn1\.to_v\.(.*)$":
|
||||
r"transformer_blocks.\1.attn1.to_v.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn1\.to_out\.0\.(.*)$":
|
||||
r"transformer_blocks.\1.attn1.to_out.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn1\.norm_q\.(.*)$":
|
||||
r"transformer_blocks.\1.attn1.norm_q.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn1\.norm_k\.(.*)$":
|
||||
r"transformer_blocks.\1.attn1.norm_k.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn2\.to_q\.(.*)$":
|
||||
r"transformer_blocks.\1.attn2.to_q.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn2\.to_k\.(.*)$":
|
||||
r"transformer_blocks.\1.attn2.to_k.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn2\.to_v\.(.*)$":
|
||||
r"transformer_blocks.\1.attn2.to_v.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn2\.to_out\.0\.(.*)$":
|
||||
r"transformer_blocks.\1.attn2.to_out.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn2\.norm_q\.(.*)$":
|
||||
r"transformer_blocks.\1.attn2.norm_q.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn2\.norm_k\.(.*)$":
|
||||
r"transformer_blocks.\1.attn2.norm_k.\2",
|
||||
r"^transformer_blocks\.(\d+)\.ff\.net\.0\.proj\.(.*)$":
|
||||
r"transformer_blocks.\1.ff.fc_in.\2",
|
||||
r"^transformer_blocks\.(\d+)\.ff\.net\.2\.(.*)$":
|
||||
r"transformer_blocks.\1.ff.fc_out.\2",
|
||||
r"^norm_out\.(.*)$": r"norm_out.\1",
|
||||
r"^proj_out\.(.*)$": r"proj_out.\1",
|
||||
})
|
||||
|
||||
lora_param_names_mapping: dict = field(
|
||||
default_factory=lambda: {
|
||||
r"^transformer_blocks\.(\d+)\.attn1\.to_q\.(.*)$":
|
||||
r"transformer_blocks.\1.attn1.to_q.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn1\.to_k\.(.*)$":
|
||||
r"transformer_blocks.\1.attn1.to_k.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn1\.to_v\.(.*)$":
|
||||
r"transformer_blocks.\1.attn1.to_v.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn1\.to_out\.(.*)$":
|
||||
r"transformer_blocks.\1.attn1.to_out.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn2\.to_q\.(.*)$":
|
||||
r"transformer_blocks.\1.attn2.to_q.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn2\.to_k\.(.*)$":
|
||||
r"transformer_blocks.\1.attn2.to_k.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn2\.to_v\.(.*)$":
|
||||
r"transformer_blocks.\1.attn2.to_v.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn2\.to_out\.(.*)$":
|
||||
r"transformer_blocks.\1.attn2.to_out.\2",
|
||||
r"^transformer_blocks\.(\d+)\.ff\.(.*)$":
|
||||
r"transformer_blocks.\1.ff.\2",
|
||||
})
|
||||
|
||||
# Cosmos-specific config parameters based on transformer_cosmos.py
|
||||
in_channels: int = 16
|
||||
out_channels: int = 16
|
||||
num_attention_heads: int = 16
|
||||
attention_head_dim: int = 128
|
||||
num_layers: int = 28
|
||||
mlp_ratio: float = 4.0
|
||||
text_embed_dim: int = 1024
|
||||
adaln_lora_dim: int = 256
|
||||
max_size: tuple[int, int, int] = (128, 240, 240)
|
||||
patch_size: tuple[int, int, int] = (1, 2, 2)
|
||||
rope_scale: tuple[float, float, float] = (1.0, 3.0, 3.0)
|
||||
concat_padding_mask: bool = True
|
||||
extra_pos_embed_type: str | None = None
|
||||
qk_norm: str = "rms_norm"
|
||||
eps: float = 1e-6
|
||||
exclude_lora_layers: list[str] = field(default_factory=lambda: ["embedder"])
|
||||
|
||||
def __post_init__(self):
|
||||
super().__post_init__()
|
||||
self.out_channels = self.out_channels or self.in_channels
|
||||
self.hidden_size = self.num_attention_heads * self.attention_head_dim
|
||||
self.num_channels_latents = self.in_channels
|
||||
|
||||
|
||||
@dataclass
|
||||
class CosmosVideoConfig(DiTConfig):
|
||||
arch_config: DiTArchConfig = field(default_factory=CosmosArchConfig)
|
||||
prefix: str = "Cosmos"
|
||||
@@ -1,181 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
|
||||
|
||||
def is_transformer_blocks(n: str, m) -> bool:
|
||||
return "transformer_blocks" in n and str.isdigit(n.split(".")[-1])
|
||||
|
||||
|
||||
@dataclass
|
||||
class Cosmos25ArchConfig(DiTArchConfig):
|
||||
"""Configuration for Cosmos 2.5 architecture (MiniTrainDIT)."""
|
||||
|
||||
_fsdp_shard_conditions: list = field(
|
||||
default_factory=lambda: [is_transformer_blocks])
|
||||
|
||||
param_names_mapping: dict = field(
|
||||
default_factory=lambda: {
|
||||
# Remove "net." prefix and map official structure to FastVideo
|
||||
# Patch embedding: net.x_embedder.proj.1.weight -> patch_embed.proj.weight
|
||||
r"^net\.x_embedder\.proj\.1\.(.*)$":
|
||||
r"patch_embed.proj.\1",
|
||||
|
||||
# Time embedding: net.t_embedder.1.linear_1.weight -> time_embed.t_embedder.linear_1.weight
|
||||
r"^net\.t_embedder\.1\.linear_1\.(.*)$":
|
||||
r"time_embed.t_embedder.linear_1.\1",
|
||||
r"^net\.t_embedder\.1\.linear_2\.(.*)$":
|
||||
r"time_embed.t_embedder.linear_2.\1",
|
||||
# Time embedding norm: net.t_embedding_norm.weight -> time_embed.norm.weight
|
||||
# Note: This also handles _extra_state if present
|
||||
r"^net\.t_embedding_norm\.(.*)$":
|
||||
r"time_embed.norm.\1",
|
||||
|
||||
# Cross-attention projection (optional): net.crossattn_proj.0.weight -> crossattn_proj.0.weight
|
||||
r"^net\.crossattn_proj\.0\.weight$":
|
||||
r"crossattn_proj.0.weight",
|
||||
r"^net\.crossattn_proj\.0\.bias$":
|
||||
r"crossattn_proj.0.bias",
|
||||
|
||||
# Transformer blocks: net.blocks.N -> transformer_blocks.N
|
||||
# Self-attention (self_attn -> attn1)
|
||||
r"^net\.blocks\.(\d+)\.self_attn\.q_proj\.(.*)$":
|
||||
r"transformer_blocks.\1.attn1.to_q.\2",
|
||||
r"^net\.blocks\.(\d+)\.self_attn\.k_proj\.(.*)$":
|
||||
r"transformer_blocks.\1.attn1.to_k.\2",
|
||||
r"^net\.blocks\.(\d+)\.self_attn\.v_proj\.(.*)$":
|
||||
r"transformer_blocks.\1.attn1.to_v.\2",
|
||||
r"^net\.blocks\.(\d+)\.self_attn\.output_proj\.(.*)$":
|
||||
r"transformer_blocks.\1.attn1.to_out.\2",
|
||||
r"^net\.blocks\.(\d+)\.self_attn\.q_norm\.weight$":
|
||||
r"transformer_blocks.\1.attn1.norm_q.weight",
|
||||
r"^net\.blocks\.(\d+)\.self_attn\.k_norm\.weight$":
|
||||
r"transformer_blocks.\1.attn1.norm_k.weight",
|
||||
# RMSNorm _extra_state keys (internal PyTorch state, will be recomputed automatically)
|
||||
r"^net\.blocks\.(\d+)\.self_attn\.q_norm\._extra_state$":
|
||||
r"transformer_blocks.\1.attn1.norm_q._extra_state",
|
||||
r"^net\.blocks\.(\d+)\.self_attn\.k_norm\._extra_state$":
|
||||
r"transformer_blocks.\1.attn1.norm_k._extra_state",
|
||||
|
||||
# Cross-attention (cross_attn -> attn2)
|
||||
r"^net\.blocks\.(\d+)\.cross_attn\.q_proj\.(.*)$":
|
||||
r"transformer_blocks.\1.attn2.to_q.\2",
|
||||
r"^net\.blocks\.(\d+)\.cross_attn\.k_proj\.(.*)$":
|
||||
r"transformer_blocks.\1.attn2.to_k.\2",
|
||||
r"^net\.blocks\.(\d+)\.cross_attn\.v_proj\.(.*)$":
|
||||
r"transformer_blocks.\1.attn2.to_v.\2",
|
||||
r"^net\.blocks\.(\d+)\.cross_attn\.output_proj\.(.*)$":
|
||||
r"transformer_blocks.\1.attn2.to_out.\2",
|
||||
r"^net\.blocks\.(\d+)\.cross_attn\.q_norm\.weight$":
|
||||
r"transformer_blocks.\1.attn2.norm_q.weight",
|
||||
r"^net\.blocks\.(\d+)\.cross_attn\.k_norm\.weight$":
|
||||
r"transformer_blocks.\1.attn2.norm_k.weight",
|
||||
# RMSNorm _extra_state keys for cross-attention
|
||||
r"^net\.blocks\.(\d+)\.cross_attn\.q_norm\._extra_state$":
|
||||
r"transformer_blocks.\1.attn2.norm_q._extra_state",
|
||||
r"^net\.blocks\.(\d+)\.cross_attn\.k_norm\._extra_state$":
|
||||
r"transformer_blocks.\1.attn2.norm_k._extra_state",
|
||||
|
||||
# MLP: net.blocks.N.mlp.layer1 -> transformer_blocks.N.mlp.fc_in
|
||||
r"^net\.blocks\.(\d+)\.mlp\.layer1\.(.*)$":
|
||||
r"transformer_blocks.\1.mlp.fc_in.\2",
|
||||
r"^net\.blocks\.(\d+)\.mlp\.layer2\.(.*)$":
|
||||
r"transformer_blocks.\1.mlp.fc_out.\2",
|
||||
|
||||
# AdaLN-LoRA modulations: net.blocks.N.adaln_modulation_* -> transformer_blocks.N.adaln_modulation_*
|
||||
# These are now at the block level, not inside norm layers
|
||||
r"^net\.blocks\.(\d+)\.adaln_modulation_self_attn\.1\.(.*)$":
|
||||
r"transformer_blocks.\1.adaln_modulation_self_attn.1.\2",
|
||||
r"^net\.blocks\.(\d+)\.adaln_modulation_self_attn\.2\.(.*)$":
|
||||
r"transformer_blocks.\1.adaln_modulation_self_attn.2.\2",
|
||||
r"^net\.blocks\.(\d+)\.adaln_modulation_cross_attn\.1\.(.*)$":
|
||||
r"transformer_blocks.\1.adaln_modulation_cross_attn.1.\2",
|
||||
r"^net\.blocks\.(\d+)\.adaln_modulation_cross_attn\.2\.(.*)$":
|
||||
r"transformer_blocks.\1.adaln_modulation_cross_attn.2.\2",
|
||||
r"^net\.blocks\.(\d+)\.adaln_modulation_mlp\.1\.(.*)$":
|
||||
r"transformer_blocks.\1.adaln_modulation_mlp.1.\2",
|
||||
r"^net\.blocks\.(\d+)\.adaln_modulation_mlp\.2\.(.*)$":
|
||||
r"transformer_blocks.\1.adaln_modulation_mlp.2.\2",
|
||||
|
||||
# Layer norms: net.blocks.N.layer_norm_* -> transformer_blocks.N.norm*.norm
|
||||
r"^net\.blocks\.(\d+)\.layer_norm_self_attn\._extra_state$":
|
||||
r"transformer_blocks.\1.norm1.norm._extra_state",
|
||||
r"^net\.blocks\.(\d+)\.layer_norm_cross_attn\._extra_state$":
|
||||
r"transformer_blocks.\1.norm2.norm._extra_state",
|
||||
r"^net\.blocks\.(\d+)\.layer_norm_mlp\._extra_state$":
|
||||
r"transformer_blocks.\1.norm3.norm._extra_state",
|
||||
|
||||
# Final layer: net.final_layer.linear -> final_layer.proj_out
|
||||
r"^net\.final_layer\.linear\.(.*)$":
|
||||
r"final_layer.proj_out.\1",
|
||||
# Final layer AdaLN-LoRA: net.final_layer.adaln_modulation -> final_layer.linear_*
|
||||
r"^net\.final_layer\.adaln_modulation\.1\.(.*)$":
|
||||
r"final_layer.linear_1.\1",
|
||||
r"^net\.final_layer\.adaln_modulation\.2\.(.*)$":
|
||||
r"final_layer.linear_2.\1",
|
||||
|
||||
# Note: The following keys from official checkpoint are NOT mapped and can be safely ignored:
|
||||
# - net.pos_embedder.* (seq, dim_spatial_range, dim_temporal_range) - These are computed dynamically
|
||||
# in FastVideo's Cosmos25RotaryPosEmbed forward() method, so they don't need to be loaded.
|
||||
# - net.accum_* keys (training metadata) - These are skipped during checkpoint loading.
|
||||
})
|
||||
|
||||
lora_param_names_mapping: dict = field(
|
||||
default_factory=lambda: {
|
||||
r"^transformer_blocks\.(\d+)\.attn1\.to_q\.(.*)$":
|
||||
r"transformer_blocks.\1.attn1.to_q.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn1\.to_k\.(.*)$":
|
||||
r"transformer_blocks.\1.attn1.to_k.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn1\.to_v\.(.*)$":
|
||||
r"transformer_blocks.\1.attn1.to_v.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn1\.to_out\.(.*)$":
|
||||
r"transformer_blocks.\1.attn1.to_out.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn2\.to_q\.(.*)$":
|
||||
r"transformer_blocks.\1.attn2.to_q.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn2\.to_k\.(.*)$":
|
||||
r"transformer_blocks.\1.attn2.to_k.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn2\.to_v\.(.*)$":
|
||||
r"transformer_blocks.\1.attn2.to_v.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn2\.to_out\.(.*)$":
|
||||
r"transformer_blocks.\1.attn2.to_out.\2",
|
||||
r"^transformer_blocks\.(\d+)\.mlp\.(.*)$":
|
||||
r"transformer_blocks.\1.mlp.\2",
|
||||
})
|
||||
|
||||
# Cosmos 2.5 specific config parameters
|
||||
in_channels: int = 16
|
||||
out_channels: int = 16
|
||||
num_attention_heads: int = 16
|
||||
attention_head_dim: int = 128 # 2048 / 16
|
||||
num_layers: int = 28
|
||||
mlp_ratio: float = 4.0
|
||||
text_embed_dim: int = 1024
|
||||
adaln_lora_dim: int = 256
|
||||
use_adaln_lora: bool = True
|
||||
max_size: tuple[int, int, int] = (128, 240, 240)
|
||||
patch_size: tuple[int, int, int] = (1, 2, 2)
|
||||
rope_scale: tuple[float, float, float] = (1.0, 3.0, 3.0) # T, H, W scaling
|
||||
concat_padding_mask: bool = True
|
||||
extra_pos_embed_type: str | None = None # "learnable" or None
|
||||
# Note: Official checkpoint has use_crossattn_projection=True with 100K-dim input from Qwen 7B.
|
||||
# When enabled, must provide 100,352-dim embeddings to match the projection layer in checkpoint.
|
||||
use_crossattn_projection: bool = False
|
||||
crossattn_proj_in_channels: int = 100352 # Qwen 7B embedding dimension
|
||||
rope_enable_fps_modulation: bool = True
|
||||
qk_norm: str = "rms_norm"
|
||||
eps: float = 1e-6
|
||||
exclude_lora_layers: list[str] = field(default_factory=lambda: ["embedder"])
|
||||
|
||||
def __post_init__(self):
|
||||
super().__post_init__()
|
||||
self.out_channels = self.out_channels or self.in_channels
|
||||
self.hidden_size = self.num_attention_heads * self.attention_head_dim
|
||||
self.num_channels_latents = self.in_channels
|
||||
|
||||
|
||||
@dataclass
|
||||
class Cosmos25VideoConfig(DiTConfig):
|
||||
"""Configuration for Cosmos 2.5 video generation model."""
|
||||
arch_config: DiTArchConfig = field(default_factory=Cosmos25ArchConfig)
|
||||
prefix: str = "Cosmos25"
|
||||
@@ -5,10 +5,10 @@ from fastvideo.configs.models.encoders.base import (BaseEncoderOutput,
|
||||
from fastvideo.configs.models.encoders.clip import (
|
||||
CLIPTextConfig, CLIPVisionConfig, WAN2_1ControlCLIPVisionConfig)
|
||||
from fastvideo.configs.models.encoders.llama import LlamaConfig
|
||||
from fastvideo.configs.models.encoders.t5 import T5Config, T5LargeConfig
|
||||
from fastvideo.configs.models.encoders.t5 import T5Config
|
||||
|
||||
__all__ = [
|
||||
"EncoderConfig", "TextEncoderConfig", "ImageEncoderConfig",
|
||||
"BaseEncoderOutput", "CLIPTextConfig", "CLIPVisionConfig",
|
||||
"WAN2_1ControlCLIPVisionConfig", "LlamaConfig", "T5Config", "T5LargeConfig"
|
||||
"WAN2_1ControlCLIPVisionConfig", "LlamaConfig", "T5Config"
|
||||
]
|
||||
|
||||
@@ -70,31 +70,8 @@ class T5ArchConfig(TextEncoderArchConfig):
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class T5LargeArchConfig(T5ArchConfig):
|
||||
"""T5 Large architecture config with parameters for your specific model."""
|
||||
d_model: int = 1024
|
||||
d_kv: int = 128
|
||||
d_ff: int = 65536
|
||||
num_layers: int = 24
|
||||
num_decoder_layers: int | None = 24
|
||||
num_heads: int = 128
|
||||
decoder_start_token_id: int = 0
|
||||
n_positions: int = 512
|
||||
task_specific_params: dict | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class T5Config(TextEncoderConfig):
|
||||
arch_config: TextEncoderArchConfig = field(default_factory=T5ArchConfig)
|
||||
|
||||
prefix: str = "t5"
|
||||
|
||||
|
||||
@dataclass
|
||||
class T5LargeConfig(TextEncoderConfig):
|
||||
"""T5 Large configuration for your specific model."""
|
||||
arch_config: TextEncoderArchConfig = field(
|
||||
default_factory=T5LargeArchConfig)
|
||||
|
||||
prefix: str = "t5"
|
||||
|
||||
@@ -1,4 +1,3 @@
|
||||
from fastvideo.configs.models.vaes.cosmosvae import CosmosVAEConfig
|
||||
from fastvideo.configs.models.vaes.hunyuanvae import HunyuanVAEConfig
|
||||
from fastvideo.configs.models.vaes.stepvideovae import StepVideoVAEConfig
|
||||
from fastvideo.configs.models.vaes.wanvae import WanVAEConfig
|
||||
@@ -7,5 +6,4 @@ __all__ = [
|
||||
"HunyuanVAEConfig",
|
||||
"WanVAEConfig",
|
||||
"StepVideoVAEConfig",
|
||||
"CosmosVAEConfig",
|
||||
]
|
||||
|
||||
@@ -1,87 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.configs.models.vaes.base import VAEArchConfig, VAEConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class CosmosVAEArchConfig(VAEArchConfig):
|
||||
_name_or_path: str = ""
|
||||
base_dim: int = 96
|
||||
z_dim: int = 16
|
||||
dim_mult: tuple[int, ...] = (1, 2, 4, 4)
|
||||
num_res_blocks: int = 2
|
||||
attn_scales: tuple[float, ...] = ()
|
||||
temperal_downsample: tuple[bool, ...] = (False, True, True)
|
||||
dropout: float = 0.0
|
||||
decoder_base_dim: int | None = None
|
||||
is_residual: bool = False
|
||||
in_channels: int = 3
|
||||
out_channels: int = 3
|
||||
patch_size: int | None = None
|
||||
scale_factor_temporal: int = 4
|
||||
scale_factor_spatial: int = 8
|
||||
clip_output: bool = True
|
||||
latents_mean: tuple[float, ...] = (
|
||||
-0.7571,
|
||||
-0.7089,
|
||||
-0.9113,
|
||||
0.1075,
|
||||
-0.1745,
|
||||
0.9653,
|
||||
-0.1517,
|
||||
1.5508,
|
||||
0.4134,
|
||||
-0.0715,
|
||||
0.5517,
|
||||
-0.3632,
|
||||
-0.1922,
|
||||
-0.9497,
|
||||
0.2503,
|
||||
-0.2921,
|
||||
)
|
||||
latents_std: tuple[float, ...] = (
|
||||
2.8184,
|
||||
1.4541,
|
||||
2.3275,
|
||||
2.6558,
|
||||
1.2196,
|
||||
1.7708,
|
||||
2.6052,
|
||||
2.0743,
|
||||
3.2687,
|
||||
2.1526,
|
||||
2.8652,
|
||||
1.5579,
|
||||
1.6382,
|
||||
1.1253,
|
||||
2.8251,
|
||||
1.9160,
|
||||
)
|
||||
temporal_compression_ratio = 4
|
||||
spatial_compression_ratio = 8
|
||||
|
||||
def __post_init__(self):
|
||||
self.scaling_factor: torch.Tensor = 1.0 / torch.tensor(
|
||||
self.latents_std).view(1, self.z_dim, 1, 1, 1)
|
||||
self.shift_factor: torch.Tensor = torch.tensor(self.latents_mean).view(
|
||||
1, self.z_dim, 1, 1, 1)
|
||||
self.temporal_compression_ratio = self.scale_factor_temporal
|
||||
self.spatial_compression_ratio = self.scale_factor_spatial
|
||||
|
||||
|
||||
@dataclass
|
||||
class CosmosVAEConfig(VAEConfig):
|
||||
arch_config: CosmosVAEArchConfig = field(
|
||||
default_factory=CosmosVAEArchConfig)
|
||||
use_feature_cache: bool = True
|
||||
|
||||
use_tiling: bool = False
|
||||
use_temporal_tiling: bool = False
|
||||
use_parallel_tiling: bool = False
|
||||
|
||||
def __post_init__(self):
|
||||
self.blend_num_frames = (self.tile_sample_min_num_frames -
|
||||
self.tile_sample_stride_num_frames) * 2
|
||||
@@ -1,6 +1,5 @@
|
||||
from fastvideo.configs.pipelines.base import (PipelineConfig,
|
||||
SlidingTileAttnConfig)
|
||||
from fastvideo.configs.pipelines.cosmos import CosmosConfig
|
||||
from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
|
||||
from fastvideo.configs.pipelines.registry import (
|
||||
get_pipeline_config_cls_from_name)
|
||||
@@ -13,6 +12,5 @@ __all__ = [
|
||||
"HunyuanConfig", "FastHunyuanConfig", "PipelineConfig",
|
||||
"SlidingTileAttnConfig", "WanT2V480PConfig", "WanI2V480PConfig",
|
||||
"WanT2V720PConfig", "WanI2V720PConfig", "StepVideoT2VConfig",
|
||||
"SelfForcingWanT2V480PConfig", "CosmosConfig",
|
||||
"get_pipeline_config_cls_from_name"
|
||||
"SelfForcingWanT2V480PConfig", "get_pipeline_config_cls_from_name"
|
||||
]
|
||||
|
||||
@@ -45,7 +45,6 @@ class PipelineConfig:
|
||||
embedded_cfg_scale: float = 6.0
|
||||
flow_shift: float | None = None
|
||||
disable_autocast: bool = False
|
||||
is_causal: bool = False
|
||||
|
||||
# Model configuration
|
||||
dit_config: DiTConfig = field(default_factory=DiTConfig)
|
||||
|
||||
@@ -1,73 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.configs.models import DiTConfig, EncoderConfig, VAEConfig
|
||||
from fastvideo.configs.models.dits import CosmosVideoConfig
|
||||
from fastvideo.configs.models.encoders import BaseEncoderOutput, T5LargeConfig
|
||||
from fastvideo.configs.models.vaes import CosmosVAEConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
|
||||
|
||||
def t5_large_postprocess_text(outputs: BaseEncoderOutput) -> torch.Tensor:
|
||||
"""Postprocess T5 Large text encoder outputs for Cosmos pipeline.
|
||||
|
||||
Return raw last_hidden_state without truncation/padding.
|
||||
"""
|
||||
hidden_state = outputs.last_hidden_state
|
||||
|
||||
if hidden_state is None:
|
||||
raise ValueError("T5 Large outputs missing last_hidden_state")
|
||||
|
||||
nan_count = torch.isnan(hidden_state).sum()
|
||||
if nan_count > 0:
|
||||
hidden_state = hidden_state.masked_fill(torch.isnan(hidden_state), 0.0)
|
||||
|
||||
# Zero out embeddings beyond actual sequence length
|
||||
if outputs.attention_mask is not None:
|
||||
attention_mask = outputs.attention_mask
|
||||
lengths = attention_mask.sum(dim=1).cpu()
|
||||
for i, length in enumerate(lengths):
|
||||
hidden_state[i, length:] = 0
|
||||
|
||||
return hidden_state
|
||||
|
||||
|
||||
@dataclass
|
||||
class CosmosConfig(PipelineConfig):
|
||||
"""Configuration for Cosmos2 Video2World pipeline matching diffusers."""
|
||||
|
||||
dit_config: DiTConfig = field(default_factory=CosmosVideoConfig)
|
||||
|
||||
vae_config: VAEConfig = field(default_factory=CosmosVAEConfig)
|
||||
|
||||
text_encoder_configs: tuple[EncoderConfig, ...] = field(
|
||||
default_factory=lambda: (T5LargeConfig(), ))
|
||||
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor],
|
||||
...] = field(default_factory=lambda:
|
||||
(t5_large_postprocess_text, ))
|
||||
|
||||
dit_precision: str = "bf16"
|
||||
vae_precision: str = "fp16"
|
||||
text_encoder_precisions: tuple[str, ...] = field(
|
||||
default_factory=lambda: ("bf16", ))
|
||||
|
||||
conditioning_strategy: str = "frame_replace"
|
||||
min_num_conditional_frames: int = 1
|
||||
max_num_conditional_frames: int = 2
|
||||
sigma_conditional: float = 0.0001
|
||||
sigma_data: float = 1.0
|
||||
state_ch: int = 16
|
||||
state_t: int = 24
|
||||
text_encoder_class: str = "T5"
|
||||
|
||||
embedded_cfg_scale: int = 6
|
||||
flow_shift: float = 1.0
|
||||
|
||||
def __post_init__(self):
|
||||
self.vae_config.load_encoder = True
|
||||
self.vae_config.load_decoder = True
|
||||
|
||||
self._vae_latent_dim = 16
|
||||
@@ -5,7 +5,6 @@ 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.hunyuan import FastHunyuanConfig, HunyuanConfig
|
||||
from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
|
||||
|
||||
@@ -39,12 +38,9 @@ PIPE_NAME_TO_CONFIG: dict[str, type[PipelineConfig]] = {
|
||||
"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,
|
||||
# Add other specific weight variants
|
||||
}
|
||||
|
||||
@@ -56,7 +52,6 @@ PIPELINE_DETECTOR: dict[str, Callable[[str], bool]] = {
|
||||
"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(),
|
||||
# Add other pipeline architecture detectors
|
||||
}
|
||||
|
||||
|
||||
@@ -186,7 +186,3 @@ class SelfForcingWan2_2_T2V480PConfig(Wan2_2_T2V_A14B_Config):
|
||||
dmd_denoising_steps: list[int] | None = field(
|
||||
default_factory=lambda: [1000, 850, 700, 550, 350, 275, 200, 125])
|
||||
warp_denoising_step: bool = True
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self.vae_config.load_encoder = True
|
||||
self.vae_config.load_decoder = True
|
||||
|
||||
@@ -1,18 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
|
||||
|
||||
@dataclass
|
||||
class Cosmos_Predict2_2B_Video2World_SamplingParam(SamplingParam):
|
||||
# Video parameters
|
||||
height: int = 704
|
||||
width: int = 1280
|
||||
num_frames: int = 93
|
||||
fps: int = 16
|
||||
|
||||
# Denoising stage
|
||||
guidance_scale: float = 7.0
|
||||
negative_prompt: str = "The video captures a series of frames showing ugly scenes, static with no motion, motion blur, over-saturation, shaky footage, low resolution, grainy texture, pixelated images, poorly lit areas, underexposed and overexposed scenes, poor color balance, washed out colors, choppy sequences, jerky movements, low frame rate, artifacting, color banding, unnatural transitions, outdated special effects, fake elements, unconvincing visuals, poorly edited content, jump cuts, visual noise, and flickering. Overall, the video is of poor quality."
|
||||
num_inference_steps: int = 35
|
||||
@@ -7,8 +7,6 @@ from fastvideo.configs.sample.hunyuan import (FastHunyuanSamplingParam,
|
||||
HunyuanSamplingParam)
|
||||
from fastvideo.configs.sample.stepvideo import StepVideoT2VSamplingParam
|
||||
|
||||
from fastvideo.configs.sample.cosmos import Cosmos_Predict2_2B_Video2World_SamplingParam
|
||||
|
||||
# isort: off
|
||||
from fastvideo.configs.sample.wan import (
|
||||
FastWanT2V480P_SamplingParam,
|
||||
@@ -74,17 +72,8 @@ SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
|
||||
# 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,
|
||||
|
||||
# Add other specific weight variants
|
||||
}
|
||||
|
||||
|
||||
@@ -191,6 +191,8 @@ class SelfForcingWan2_1_T2V_1_3B_480P_SamplingParam(
|
||||
@dataclass
|
||||
class SelfForcingWan2_2_T2V_A14B_480P_SamplingParam(
|
||||
Wan2_2_T2V_A14B_SamplingParam):
|
||||
guidance_scale: float = 2.0
|
||||
guidance_scale_2: float = 2.0
|
||||
num_inference_steps: int = 8
|
||||
num_frames: int = 81
|
||||
height: int = 448
|
||||
|
||||
@@ -152,44 +152,3 @@ class TemporalRandomCrop:
|
||||
begin_index = random.randint(0, rand_end)
|
||||
end_index = min(begin_index + self.size, total_frames)
|
||||
return begin_index, end_index
|
||||
|
||||
|
||||
def best_output_size(
|
||||
width: int,
|
||||
height: int,
|
||||
width_stride: int,
|
||||
height_stride: int,
|
||||
max_area: int,
|
||||
) -> tuple[int, int]:
|
||||
"""
|
||||
Calculate the best output size (width, height) given the original dimensions, strides and max area.
|
||||
The aspect ratio is preserved as much as possible.
|
||||
|
||||
Args:
|
||||
width (int): Original width
|
||||
height (int): Original height
|
||||
width_stride (int): Width stride requirement
|
||||
height_stride (int): Height stride requirement
|
||||
max_area (int): Maximum allowed area (width * height)
|
||||
|
||||
Returns:
|
||||
tuple[int, int]: (new_width, new_height)
|
||||
"""
|
||||
aspect_ratio = width / height
|
||||
|
||||
# Scale dimensions if they exceed max_area
|
||||
current_area = width * height
|
||||
if current_area > max_area:
|
||||
scale = (max_area / current_area)**0.5
|
||||
width = int(width * scale)
|
||||
height = int(height * scale)
|
||||
|
||||
# Round to the nearest multiple of stride
|
||||
width = round(width / width_stride) * width_stride
|
||||
height = round(height / height_stride) * height_stride
|
||||
|
||||
# Ensure dimensions are at least one stride
|
||||
width = max(width, width_stride)
|
||||
height = max(height, height_stride)
|
||||
|
||||
return width, height
|
||||
|
||||
@@ -208,7 +208,7 @@ class VideoGenerator:
|
||||
|
||||
def _sanitize_filename_component(name: str) -> str:
|
||||
# Remove characters invalid on common filesystems, strip spaces/dots
|
||||
sanitized = re.sub(r'[\\/:*?"<>|]', '', name)
|
||||
sanitized = re.sub(r'[\/:*?"<>|]', '', name)
|
||||
sanitized = sanitized.strip().strip('.')
|
||||
sanitized = re.sub(r'\s+', ' ', sanitized)
|
||||
return sanitized or "video"
|
||||
|
||||
@@ -1,195 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Minimal image processing utilities for FastVideo.
|
||||
This module provides lightweight image preprocessing without external dependencies beyond PyTorch/NumPy/PIL.
|
||||
"""
|
||||
|
||||
import numpy as np
|
||||
import PIL.Image
|
||||
import torch
|
||||
|
||||
|
||||
class ImageProcessor:
|
||||
"""
|
||||
Minimal image processor for video frame preprocessing.
|
||||
|
||||
This is a lightweight alternative to diffusers.VideoProcessor that handles:
|
||||
- PIL image to tensor conversion
|
||||
- Resizing to specified dimensions
|
||||
- Normalization to [-1, 1] range
|
||||
|
||||
Args:
|
||||
vae_scale_factor: The VAE scale factor used to ensure dimensions are multiples of this value.
|
||||
"""
|
||||
|
||||
def __init__(self, vae_scale_factor: int = 8) -> None:
|
||||
self.vae_scale_factor = vae_scale_factor
|
||||
|
||||
def preprocess(
|
||||
self,
|
||||
image: PIL.Image.Image | np.ndarray | torch.Tensor,
|
||||
height: int | None = None,
|
||||
width: int | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Preprocess an image to a normalized torch tensor.
|
||||
|
||||
Args:
|
||||
image: Input image (PIL Image, NumPy array, or torch tensor)
|
||||
height: Target height. If None, uses image's original height.
|
||||
width: Target width. If None, uses image's original width.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Normalized tensor of shape (1, 3, height, width) or (1, 1, height, width) for grayscale,
|
||||
with values in range [-1, 1].
|
||||
"""
|
||||
# Handle different input types
|
||||
if isinstance(image, PIL.Image.Image):
|
||||
return self._preprocess_pil(image, height, width)
|
||||
elif isinstance(image, np.ndarray):
|
||||
return self._preprocess_numpy(image, height, width)
|
||||
elif isinstance(image, torch.Tensor):
|
||||
return self._preprocess_tensor(image, height, width)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Unsupported image type: {type(image)}. "
|
||||
"Supported types: PIL.Image.Image, np.ndarray, torch.Tensor")
|
||||
|
||||
def _preprocess_pil(
|
||||
self,
|
||||
image: PIL.Image.Image,
|
||||
height: int | None = None,
|
||||
width: int | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""Preprocess a PIL image."""
|
||||
if height is None:
|
||||
height = image.height
|
||||
if width is None:
|
||||
width = image.width
|
||||
|
||||
height = height - (height % self.vae_scale_factor)
|
||||
width = width - (width % self.vae_scale_factor)
|
||||
|
||||
image = image.resize((width, height),
|
||||
resample=PIL.Image.Resampling.LANCZOS)
|
||||
|
||||
image_np = np.array(image, dtype=np.float32) / 255.0
|
||||
|
||||
if image_np.ndim == 2: # Grayscale
|
||||
image_np = np.expand_dims(image_np, axis=-1)
|
||||
|
||||
return self._normalize_to_tensor(image_np)
|
||||
|
||||
def _preprocess_numpy(
|
||||
self,
|
||||
image: np.ndarray,
|
||||
height: int | None = None,
|
||||
width: int | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""Preprocess a numpy array."""
|
||||
# Determine target dimensions if not provided
|
||||
if image.ndim == 3:
|
||||
img_height, img_width = image.shape[:2]
|
||||
elif image.ndim == 2:
|
||||
img_height, img_width = image.shape
|
||||
else:
|
||||
raise ValueError(f"Expected 2D or 3D array, got {image.ndim}D")
|
||||
|
||||
if height is None:
|
||||
height = img_height
|
||||
if width is None:
|
||||
width = img_width
|
||||
|
||||
height = height - (height % self.vae_scale_factor)
|
||||
width = width - (width % self.vae_scale_factor)
|
||||
|
||||
if image.dtype == np.uint8:
|
||||
pil_image = PIL.Image.fromarray(image)
|
||||
else:
|
||||
# Assume normalized [0, 1] or similar
|
||||
if image.max() <= 1.0:
|
||||
image_uint8 = (image * 255).astype(np.uint8)
|
||||
else:
|
||||
image_uint8 = image.astype(np.uint8)
|
||||
pil_image = PIL.Image.fromarray(image_uint8)
|
||||
|
||||
pil_image = pil_image.resize((width, height),
|
||||
resample=PIL.Image.Resampling.LANCZOS)
|
||||
image_np = np.array(pil_image, dtype=np.float32) / 255.0
|
||||
|
||||
# Ensure 3D shape
|
||||
if image_np.ndim == 2:
|
||||
image_np = np.expand_dims(image_np, axis=-1)
|
||||
|
||||
return self._normalize_to_tensor(image_np)
|
||||
|
||||
def _preprocess_tensor(
|
||||
self,
|
||||
image: torch.Tensor,
|
||||
height: int | None = None,
|
||||
width: int | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""Preprocess a torch tensor."""
|
||||
# Determine target dimensions
|
||||
if image.ndim == 3: # (H, W, C) or (C, H, W)
|
||||
if image.shape[0] in (1, 3, 4): # Likely (C, H, W)
|
||||
img_height, img_width = image.shape[1], image.shape[2]
|
||||
else: # Likely (H, W, C)
|
||||
img_height, img_width = image.shape[0], image.shape[1]
|
||||
elif image.ndim == 2: # (H, W)
|
||||
img_height, img_width = image.shape
|
||||
else:
|
||||
raise ValueError(f"Expected 2D or 3D tensor, got {image.ndim}D")
|
||||
|
||||
if height is None:
|
||||
height = img_height
|
||||
if width is None:
|
||||
width = img_width
|
||||
|
||||
height = height - (height % self.vae_scale_factor)
|
||||
width = width - (width % self.vae_scale_factor)
|
||||
|
||||
if image.ndim == 2:
|
||||
image = image.unsqueeze(0).unsqueeze(0) # (1, 1, H, W)
|
||||
elif image.ndim == 3:
|
||||
if image.shape[0] in (1, 3, 4): # (C, H, W)
|
||||
image = image.unsqueeze(0) # (1, C, H, W)
|
||||
else: # (H, W, C) - need to rearrange
|
||||
image = image.permute(2, 0, 1).unsqueeze(0) # (1, C, H, W)
|
||||
|
||||
image = torch.nn.functional.interpolate(image,
|
||||
size=(height, width),
|
||||
mode="bilinear",
|
||||
align_corners=False)
|
||||
|
||||
if image.max() > 1.0: # Assume [0, 255] range
|
||||
image = image / 255.0
|
||||
|
||||
image = 2.0 * image - 1.0
|
||||
|
||||
return image
|
||||
|
||||
def _normalize_to_tensor(self, image_np: np.ndarray) -> torch.Tensor:
|
||||
"""
|
||||
Convert normalized numpy array [0, 1] to torch tensor [-1, 1].
|
||||
|
||||
Args:
|
||||
image_np: NumPy array with shape (H, W) or (H, W, C) with values in [0, 1]
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Shape (1, C, H, W) or (1, 1, H, W) with values in [-1, 1]
|
||||
"""
|
||||
# Convert to tensor
|
||||
if image_np.ndim == 2: # (H, W) - grayscale
|
||||
tensor = torch.from_numpy(image_np).unsqueeze(0).unsqueeze(
|
||||
0) # (1, 1, H, W)
|
||||
elif image_np.ndim == 3: # (H, W, C)
|
||||
tensor = torch.from_numpy(image_np).permute(2, 0, 1).unsqueeze(
|
||||
0) # (1, C, H, W)
|
||||
else:
|
||||
raise ValueError(f"Expected 2D or 3D array, got {image_np.ndim}D")
|
||||
|
||||
# Normalize to [-1, 1]
|
||||
tensor = 2.0 * tensor - 1.0
|
||||
|
||||
return tensor
|
||||
@@ -101,19 +101,12 @@ class BaseLayerWithLoRA(nn.Module):
|
||||
def set_lora_weights(self,
|
||||
A: torch.Tensor,
|
||||
B: torch.Tensor,
|
||||
lora_alpha: float | None = None,
|
||||
training_mode: bool = False,
|
||||
lora_path: str | None = None) -> None:
|
||||
self.lora_A = torch.nn.Parameter(
|
||||
A) # share storage with weights in the pipeline
|
||||
self.lora_B = torch.nn.Parameter(B)
|
||||
self.disable_lora = False
|
||||
|
||||
# Store rank and alpha directly
|
||||
rank = A.shape[0] # rank is the first dimension of A
|
||||
self.lora_rank = rank
|
||||
self.lora_alpha = int(lora_alpha) if lora_alpha is not None else rank
|
||||
|
||||
if not training_mode:
|
||||
self.merge_lora_weights()
|
||||
self.lora_path = lora_path
|
||||
@@ -141,13 +134,8 @@ class BaseLayerWithLoRA(nn.Module):
|
||||
current_device = self.base_layer.weight.data.device
|
||||
data = self.base_layer.weight.data.to(
|
||||
get_local_torch_device()).full_tensor()
|
||||
|
||||
# Apply LoRA with alpha scaling
|
||||
lora_delta = (self.slice_lora_b_weights(self.lora_B).to(data)
|
||||
@ self.slice_lora_a_weights(self.lora_A).to(data))
|
||||
if self.lora_alpha and self.lora_rank and self.lora_alpha != self.lora_rank:
|
||||
lora_delta *= (self.lora_alpha / self.lora_rank)
|
||||
data += lora_delta
|
||||
data += (self.slice_lora_b_weights(self.lora_B).to(data)
|
||||
@ self.slice_lora_a_weights(self.lora_A).to(data))
|
||||
unsharded_base_layer.weight = nn.Parameter(data.to(current_device))
|
||||
if isinstance(getattr(self.base_layer, "bias", None), DTensor):
|
||||
unsharded_base_layer.bias = nn.Parameter(
|
||||
@@ -166,13 +154,8 @@ class BaseLayerWithLoRA(nn.Module):
|
||||
else:
|
||||
current_device = self.base_layer.weight.data.device
|
||||
data = self.base_layer.weight.data.to(get_local_torch_device())
|
||||
|
||||
# Apply LoRA with alpha scaling
|
||||
lora_delta = (self.slice_lora_b_weights(self.lora_B.to(data))
|
||||
@ self.slice_lora_a_weights(self.lora_A.to(data)))
|
||||
if self.lora_alpha and self.lora_rank and self.lora_alpha != self.lora_rank:
|
||||
lora_delta *= (self.lora_alpha / self.lora_rank)
|
||||
data += lora_delta
|
||||
data += \
|
||||
(self.slice_lora_b_weights(self.lora_B.to(data)) @ self.slice_lora_a_weights(self.lora_A.to(data)))
|
||||
self.base_layer.weight.data = data.to(current_device,
|
||||
non_blocking=True)
|
||||
|
||||
|
||||
@@ -47,59 +47,6 @@ def _rotate_gptj(x: torch.Tensor) -> torch.Tensor:
|
||||
return x.flatten(-2)
|
||||
|
||||
|
||||
def apply_rotary_emb(
|
||||
x: torch.Tensor,
|
||||
freqs_cis: torch.Tensor | tuple[torch.Tensor, torch.Tensor],
|
||||
use_real: bool = True,
|
||||
use_real_unbind_dim: int = -1,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Apply rotary embeddings to input tensors using the given frequency tensor. This function applies rotary embeddings
|
||||
to the given query or key 'x' tensors using the provided frequency tensor 'freqs_cis'. The input tensors are
|
||||
reshaped as complex numbers, and the frequency tensor is reshaped for broadcasting compatibility. The resulting
|
||||
tensors contain rotary embeddings and are returned as real tensors.
|
||||
Args:
|
||||
x (`torch.Tensor`):
|
||||
Query or key tensor to apply rotary embeddings. [B, H, S, D] xk (torch.Tensor): Key tensor to apply
|
||||
freqs_cis (`Tuple[torch.Tensor]`): Precomputed frequency tensor for complex exponentials. ([S, D], [S, D],)
|
||||
Returns:
|
||||
Tuple[torch.Tensor, torch.Tensor]: Tuple of modified query tensor and key tensor with rotary embeddings.
|
||||
"""
|
||||
if use_real:
|
||||
cos, sin = freqs_cis # [S, D]
|
||||
# Match Diffusers broadcasting (sequence_dim=2 case)
|
||||
cos = cos[None, None, :, :]
|
||||
sin = sin[None, None, :, :]
|
||||
cos, sin = cos.to(x.device), sin.to(x.device)
|
||||
|
||||
if use_real_unbind_dim == -1:
|
||||
# Used for flux, cogvideox, hunyuan-dit
|
||||
x_real, x_imag = x.reshape(*x.shape[:-1], -1,
|
||||
2).unbind(-1) # [B, S, H, D//2]
|
||||
x_rotated = torch.stack([-x_imag, x_real], dim=-1).flatten(3)
|
||||
elif use_real_unbind_dim == -2:
|
||||
# Used for Stable Audio, OmniGen, CogView4 and Cosmos
|
||||
x_real, x_imag = x.reshape(*x.shape[:-1], 2,
|
||||
-1).unbind(-2) # [B, S, H, D//2]
|
||||
x_rotated = torch.cat([-x_imag, x_real], dim=-1)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"`use_real_unbind_dim={use_real_unbind_dim}` but should be -1 or -2."
|
||||
)
|
||||
|
||||
out = (x.float() * cos + x_rotated.float() * sin).to(x.dtype)
|
||||
|
||||
return out
|
||||
else:
|
||||
# used for lumina
|
||||
x_rotated = torch.view_as_complex(x.float().reshape(
|
||||
*x.shape[:-1], -1, 2))
|
||||
freqs_cis = freqs_cis.unsqueeze(2)
|
||||
x_out = torch.view_as_real(x_rotated * freqs_cis).flatten(3)
|
||||
|
||||
return x_out.type_as(x)
|
||||
|
||||
|
||||
def _apply_rotary_emb(
|
||||
x: torch.Tensor,
|
||||
cos: torch.Tensor,
|
||||
|
||||
@@ -177,79 +177,3 @@ def unpatchify(x, t, h, w, patch_size, channels) -> torch.Tensor:
|
||||
imgs = x.reshape(shape=(x.shape[0], c, t * pt, h * ph, w * pw))
|
||||
|
||||
return imgs
|
||||
|
||||
|
||||
def get_timestep_embedding(
|
||||
timesteps: torch.Tensor,
|
||||
embedding_dim: int,
|
||||
flip_sin_to_cos: bool = False,
|
||||
downscale_freq_shift: float = 1,
|
||||
scale: float = 1,
|
||||
max_period: int = 10000,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
This matches the implementation in Denoising Diffusion Probabilistic Models: Create sinusoidal timestep embeddings.
|
||||
Args
|
||||
timesteps (torch.Tensor):
|
||||
a 1-D Tensor of N indices, one per batch element. These may be fractional.
|
||||
embedding_dim (int):
|
||||
the dimension of the output.
|
||||
flip_sin_to_cos (bool):
|
||||
Whether the embedding order should be `cos, sin` (if True) or `sin, cos` (if False)
|
||||
downscale_freq_shift (float):
|
||||
Controls the delta between frequencies between dimensions
|
||||
scale (float):
|
||||
Scaling factor applied to the embeddings.
|
||||
max_period (int):
|
||||
Controls the maximum frequency of the embeddings
|
||||
Returns
|
||||
torch.Tensor: an [N x dim] Tensor of positional embeddings.
|
||||
"""
|
||||
assert len(timesteps.shape) == 1, "Timesteps should be a 1d-array"
|
||||
|
||||
half_dim = embedding_dim // 2
|
||||
exponent = -math.log(max_period) * torch.arange(
|
||||
start=0, end=half_dim, dtype=torch.float32, device=timesteps.device)
|
||||
exponent = exponent / (half_dim - downscale_freq_shift)
|
||||
|
||||
emb = torch.exp(exponent)
|
||||
emb = timesteps[:, None].float() * emb[None, :]
|
||||
|
||||
# scale embeddings
|
||||
emb = scale * emb
|
||||
|
||||
# concat sine and cosine embeddings
|
||||
emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=-1)
|
||||
|
||||
# flip sine and cosine embeddings
|
||||
if flip_sin_to_cos:
|
||||
emb = torch.cat([emb[:, half_dim:], emb[:, :half_dim]], dim=-1)
|
||||
|
||||
# zero pad
|
||||
if embedding_dim % 2 == 1:
|
||||
emb = torch.nn.functional.pad(emb, (0, 1, 0, 0))
|
||||
return emb
|
||||
|
||||
|
||||
class Timesteps(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
num_channels: int,
|
||||
flip_sin_to_cos: bool,
|
||||
downscale_freq_shift: float,
|
||||
scale: int = 1):
|
||||
super().__init__()
|
||||
self.num_channels = num_channels
|
||||
self.flip_sin_to_cos = flip_sin_to_cos
|
||||
self.downscale_freq_shift = downscale_freq_shift
|
||||
self.scale = scale
|
||||
|
||||
def forward(self, timesteps: torch.Tensor) -> torch.Tensor:
|
||||
t_emb = get_timestep_embedding(
|
||||
timesteps,
|
||||
self.num_channels,
|
||||
flip_sin_to_cos=self.flip_sin_to_cos,
|
||||
downscale_freq_shift=self.downscale_freq_shift,
|
||||
scale=self.scale,
|
||||
)
|
||||
return t_emb
|
||||
@@ -20,7 +20,7 @@ import fastvideo.envs as envs
|
||||
from fastvideo.attention import (DistributedAttention,
|
||||
LocalAttention)
|
||||
from fastvideo.configs.models.dits import WanVideoConfig
|
||||
from fastvideo.distributed.parallel_state import get_sp_world_size
|
||||
from fastvideo.distributed.parallel_state import get_sp_world_size, get_local_torch_device
|
||||
from fastvideo.forward_context import get_forward_context
|
||||
from fastvideo.layers.layernorm import (FP32LayerNorm, LayerNormScaleShift,
|
||||
RMSNorm, ScaleResidual,
|
||||
@@ -33,9 +33,29 @@ 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 CacheAppend(torch.autograd.Function):
|
||||
"""
|
||||
KV cache with shape [batch, seq_len, heads, head_dim].
|
||||
"""
|
||||
@staticmethod
|
||||
def forward(ctx, storage, active_cache, x, start, end):
|
||||
# Ensure storage has the same dtype as x
|
||||
storage.data[:, start:end] = x
|
||||
ctx.save_for_backward(storage.to(x.dtype))
|
||||
ctx.start = start
|
||||
ctx.end = end
|
||||
return storage[:, :end].to(x.dtype) # Ensure returned value has same dtype as input
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
start = ctx.start
|
||||
end = ctx.end
|
||||
return None, grad_output[:, :start], grad_output[:, start:end], None, None
|
||||
|
||||
class CausalWanSelfAttention(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
@@ -67,6 +87,10 @@ class CausalWanSelfAttention(nn.Module):
|
||||
causal=False,
|
||||
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
|
||||
AttentionBackendEnum.TORCH_SDPA))
|
||||
|
||||
self.k_cache = None
|
||||
self.v_cache = None
|
||||
self.counter = 0
|
||||
|
||||
def forward(self,
|
||||
q: torch.Tensor,
|
||||
@@ -84,6 +108,19 @@ class CausalWanSelfAttention(nn.Module):
|
||||
grid_sizes(Tensor): Shape [B, 3], the second dimension contains (F, H, W)
|
||||
freqs(Tensor): Rope freqs, shape [1024, C / num_heads / 2]
|
||||
"""
|
||||
if kv_cache is not None:
|
||||
# Then we are in _forward_inference mode
|
||||
if self.k_cache is None:
|
||||
assert self.counter == 0
|
||||
self.counter += 1
|
||||
del self.k_cache
|
||||
self.register_buffer("k_cache", torch.empty(1, self.max_attention_size, self.num_heads, self.head_dim, device=q.device, dtype=q.dtype), persistent=False)
|
||||
self.k_cache.requires_grad_(True)
|
||||
if self.v_cache is None:
|
||||
del self.v_cache
|
||||
self.register_buffer("v_cache", torch.empty(1, self.max_attention_size, self.num_heads, self.head_dim, device=v.device, dtype=v.dtype), persistent=False)
|
||||
self.v_cache.requires_grad_(True)
|
||||
|
||||
if cache_start is None:
|
||||
cache_start = current_start
|
||||
|
||||
@@ -128,6 +165,8 @@ class CausalWanSelfAttention(nn.Module):
|
||||
num_new_tokens = roped_query.shape[1]
|
||||
if self.local_attn_size != -1 and (current_end > kv_cache["global_end_index"].item()) and (
|
||||
num_new_tokens + kv_cache["local_end_index"].item() > kv_cache_size):
|
||||
raise Exception("Not implemented")
|
||||
# @TODO(Wei): This part has not been thoroughly tested yet. Use with caution.
|
||||
# Calculate the number of new tokens added in this step
|
||||
# Shift existing cache content left to discard oldest tokens
|
||||
# Clone the source slice to avoid overlapping memory error
|
||||
@@ -137,26 +176,44 @@ class CausalWanSelfAttention(nn.Module):
|
||||
kv_cache["k"][:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
|
||||
kv_cache["v"][:, sink_tokens:sink_tokens + num_rolled_tokens] = \
|
||||
kv_cache["v"][:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
|
||||
# self.k_cache[:, sink_tokens:sink_tokens + num_rolled_tokens] = \
|
||||
# self.k_cache[:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
|
||||
# self.v_cache[:, sink_tokens:sink_tokens + num_rolled_tokens] = \
|
||||
# self.v_cache[:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
|
||||
# Insert the new keys/values at the end
|
||||
local_end_index = kv_cache["local_end_index"].item() + current_end - \
|
||||
kv_cache["global_end_index"].item() - num_evicted_tokens
|
||||
local_start_index = local_end_index - num_new_tokens
|
||||
# local_k = CacheAppend.apply(self.k_cache, kv_cache["k"], roped_key, local_start_index, local_end_index)
|
||||
# local_v = CacheAppend.apply(self.v_cache, kv_cache["v"], v, local_start_index, local_end_index)
|
||||
kv_cache["k"][:, local_start_index:local_end_index] = roped_key
|
||||
kv_cache["v"][:, local_start_index:local_end_index] = v
|
||||
else:
|
||||
# Assign new keys/values directly up to current_end
|
||||
local_end_index = kv_cache["local_end_index"].item() + current_end - kv_cache["global_end_index"].item()
|
||||
local_start_index = local_end_index - num_new_tokens
|
||||
kv_cache["k"] = kv_cache["k"].detach()
|
||||
kv_cache["v"] = kv_cache["v"].detach()
|
||||
# logger.info("kv_cache['k'] is in comp graph: %s", kv_cache["k"].requires_grad or kv_cache["k"].grad_fn is not None)
|
||||
kv_cache["k"][:, local_start_index:local_end_index] = roped_key
|
||||
kv_cache["v"][:, local_start_index:local_end_index] = v
|
||||
# kv_cache["k"] = kv_cache["k"].detach()
|
||||
# kv_cache["v"] = kv_cache["v"].detach()
|
||||
# kv_cache["k"][:, local_start_index:local_end_index] = roped_key
|
||||
# kv_cache["v"][:, local_start_index:local_end_index] = v
|
||||
local_k = CacheAppend.apply(self.k_cache, kv_cache["k"], roped_key, local_start_index, local_end_index)
|
||||
# logger.info("Is local_k meta tensor: %s, %s", local_k.is_meta, local_k.shape)
|
||||
# logger.info("local_start_index: %d, local_end_index: %d, number of zeros in local k: %d", local_start_index, local_end_index, (local_k == 0).sum().item())
|
||||
local_v = CacheAppend.apply(self.v_cache, kv_cache["v"], v, local_start_index, local_end_index)
|
||||
|
||||
x = self.attn(
|
||||
roped_query,
|
||||
kv_cache["k"][:, max(0, local_end_index - self.max_attention_size):local_end_index],
|
||||
kv_cache["v"][:, max(0, local_end_index - self.max_attention_size):local_end_index]
|
||||
local_k,
|
||||
local_v
|
||||
)
|
||||
# x = self.attn(
|
||||
# roped_query,
|
||||
# kv_cache["k"][:, max(0, local_end_index - self.max_attention_size):local_end_index],
|
||||
# kv_cache["v"][:, max(0, local_end_index - self.max_attention_size):local_end_index]
|
||||
# )
|
||||
|
||||
kv_cache["k"] = local_k
|
||||
kv_cache["v"] = local_v
|
||||
kv_cache["global_end_index"].fill_(current_end)
|
||||
kv_cache["local_end_index"].fill_(local_end_index)
|
||||
|
||||
@@ -233,6 +290,9 @@ class CausalWanTransformerBlock(nn.Module):
|
||||
|
||||
self.scale_shift_table = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)
|
||||
|
||||
self.null_shift = torch.tensor([0], device=get_local_torch_device())
|
||||
self.null_scale = torch.tensor([0], device=get_local_torch_device())
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
@@ -283,9 +343,8 @@ class CausalWanTransformerBlock(nn.Module):
|
||||
attn_output, _ = self.to_out(attn_output)
|
||||
attn_output = attn_output.squeeze(1)
|
||||
|
||||
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)
|
||||
hidden_states, attn_output, gate_msa, self.null_shift, self.null_scale)
|
||||
|
||||
# 2. Cross-attention
|
||||
attn_output = self.attn2(norm_hidden_states,
|
||||
@@ -452,7 +511,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):
|
||||
|
||||
@@ -1,726 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import math
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from fastvideo.attention import DistributedAttention, LocalAttention
|
||||
from fastvideo.configs.models.dits.cosmos import CosmosVideoConfig
|
||||
from fastvideo.forward_context import get_forward_context
|
||||
from fastvideo.layers.layernorm import RMSNorm
|
||||
from fastvideo.layers.linear import ReplicatedLinear
|
||||
from fastvideo.layers.mlp import MLP
|
||||
from fastvideo.layers.rotary_embedding import apply_rotary_emb
|
||||
from fastvideo.layers.visual_embedding import Timesteps
|
||||
from fastvideo.models.dits.base import BaseDiT
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
class CosmosPatchEmbed(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
in_channels: int,
|
||||
out_channels: int,
|
||||
patch_size: tuple[int, int, int],
|
||||
bias: bool = True) -> None:
|
||||
super().__init__()
|
||||
self.patch_size = patch_size
|
||||
|
||||
self.proj = nn.Linear(in_channels * patch_size[0] * patch_size[1] *
|
||||
patch_size[2],
|
||||
out_channels,
|
||||
bias=bias)
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
batch_size, num_channels, num_frames, height, width = hidden_states.shape
|
||||
p_t, p_h, p_w = self.patch_size
|
||||
hidden_states = hidden_states.reshape(batch_size, num_channels,
|
||||
num_frames // p_t, p_t,
|
||||
height // p_h, p_h, width // p_w,
|
||||
p_w)
|
||||
hidden_states = hidden_states.permute(0, 2, 4, 6, 1, 3, 5,
|
||||
7).flatten(4, 7)
|
||||
hidden_states = self.proj(hidden_states)
|
||||
return hidden_states
|
||||
|
||||
|
||||
class CosmosTimestepEmbedding(nn.Module):
|
||||
|
||||
def __init__(self, in_features: int, out_features: int) -> None:
|
||||
super().__init__()
|
||||
self.linear_1 = nn.Linear(in_features, out_features, bias=False)
|
||||
self.activation = nn.SiLU()
|
||||
self.linear_2 = nn.Linear(out_features, 3 * out_features, bias=False)
|
||||
|
||||
def forward(self, timesteps: torch.Tensor) -> torch.Tensor:
|
||||
emb = self.linear_1(timesteps)
|
||||
emb = self.activation(emb)
|
||||
emb = self.linear_2(emb)
|
||||
return emb
|
||||
|
||||
|
||||
class CosmosEmbedding(nn.Module):
|
||||
|
||||
def __init__(self, embedding_dim: int, condition_dim: int) -> None:
|
||||
super().__init__()
|
||||
|
||||
self.time_proj = Timesteps(embedding_dim,
|
||||
flip_sin_to_cos=True,
|
||||
downscale_freq_shift=0.0)
|
||||
self.t_embedder = CosmosTimestepEmbedding(embedding_dim, condition_dim)
|
||||
self.norm = RMSNorm(embedding_dim, eps=1e-6)
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor,
|
||||
timestep: torch.LongTensor) -> torch.Tensor:
|
||||
timesteps_proj = self.time_proj(timestep).type_as(hidden_states)
|
||||
temb = self.t_embedder(timesteps_proj)
|
||||
embedded_timestep = self.norm(timesteps_proj)
|
||||
return temb, embedded_timestep
|
||||
|
||||
|
||||
class CosmosAdaLayerNorm(nn.Module):
|
||||
|
||||
def __init__(self, in_features: int, hidden_features: int) -> None:
|
||||
super().__init__()
|
||||
self.embedding_dim = in_features
|
||||
|
||||
self.activation = nn.SiLU()
|
||||
self.norm = nn.LayerNorm(in_features,
|
||||
elementwise_affine=False,
|
||||
eps=1e-6)
|
||||
self.linear_1 = nn.Linear(in_features, hidden_features, bias=False)
|
||||
self.linear_2 = nn.Linear(hidden_features, 2 * in_features, bias=False)
|
||||
|
||||
def forward(self,
|
||||
hidden_states: torch.Tensor,
|
||||
embedded_timestep: torch.Tensor,
|
||||
temb: torch.Tensor | None = None) -> torch.Tensor:
|
||||
embedded_timestep = self.activation(embedded_timestep)
|
||||
embedded_timestep = self.linear_1(embedded_timestep)
|
||||
embedded_timestep = self.linear_2(embedded_timestep)
|
||||
|
||||
if temb is not None:
|
||||
embedded_timestep = embedded_timestep + temb[..., :2 *
|
||||
self.embedding_dim]
|
||||
|
||||
shift, scale = embedded_timestep.chunk(2, dim=-1)
|
||||
with torch.autocast(device_type="cuda", enabled=False):
|
||||
hidden_states = self.norm(hidden_states)
|
||||
|
||||
if embedded_timestep.ndim == 2:
|
||||
shift, scale = (x.unsqueeze(1) for x in (shift, scale))
|
||||
|
||||
hidden_states = hidden_states * (1 + scale) + shift
|
||||
return hidden_states
|
||||
|
||||
|
||||
class CosmosAdaLayerNormZero(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
in_features: int,
|
||||
hidden_features: int | None = None) -> None:
|
||||
super().__init__()
|
||||
|
||||
self.norm = nn.LayerNorm(in_features,
|
||||
elementwise_affine=False,
|
||||
eps=1e-6)
|
||||
self.activation = nn.SiLU()
|
||||
|
||||
if hidden_features is None:
|
||||
self.linear_1 = nn.Identity()
|
||||
else:
|
||||
self.linear_1 = nn.Linear(in_features, hidden_features, bias=False)
|
||||
|
||||
self.linear_2 = nn.Linear(hidden_features, 3 * in_features, bias=False)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
embedded_timestep: torch.Tensor,
|
||||
temb: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
embedded_timestep = self.activation(embedded_timestep)
|
||||
embedded_timestep = self.linear_1(embedded_timestep)
|
||||
embedded_timestep = self.linear_2(embedded_timestep)
|
||||
|
||||
if temb is not None:
|
||||
embedded_timestep = embedded_timestep + temb
|
||||
|
||||
shift, scale, gate = embedded_timestep.chunk(3, dim=-1)
|
||||
|
||||
with torch.autocast(device_type="cuda", enabled=False):
|
||||
hidden_states = self.norm(hidden_states)
|
||||
|
||||
if embedded_timestep.ndim == 2:
|
||||
shift, scale, gate = (x.unsqueeze(1) for x in (shift, scale, gate))
|
||||
|
||||
hidden_states = hidden_states * (1 + scale) + shift
|
||||
return hidden_states, gate
|
||||
|
||||
|
||||
class CosmosSelfAttention(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
dim: int,
|
||||
num_heads: int,
|
||||
qk_norm=True,
|
||||
eps=1e-6,
|
||||
prefix: str = "") -> None:
|
||||
assert dim % num_heads == 0
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = dim // num_heads
|
||||
self.qk_norm = qk_norm
|
||||
self.eps = eps
|
||||
|
||||
# layers - use standard PyTorch layers when using torch backend
|
||||
self.to_q = nn.Linear(dim, dim, bias=False)
|
||||
self.to_k = nn.Linear(dim, dim, bias=False)
|
||||
self.to_v = nn.Linear(dim, dim, bias=False)
|
||||
self.to_out = nn.Linear(dim, dim, bias=False)
|
||||
self.dropout = nn.Dropout(0.0)
|
||||
|
||||
self.norm_q = RMSNorm(self.head_dim,
|
||||
eps=eps) if qk_norm else nn.Identity()
|
||||
self.norm_k = RMSNorm(self.head_dim,
|
||||
eps=eps) if qk_norm else nn.Identity()
|
||||
|
||||
def forward(self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor | None = None,
|
||||
attention_mask: torch.Tensor | None = None,
|
||||
image_rotary_emb: torch.Tensor | None = None) -> torch.Tensor:
|
||||
|
||||
if encoder_hidden_states is None:
|
||||
encoder_hidden_states = hidden_states
|
||||
|
||||
# Get QKV
|
||||
query = self.to_q(hidden_states)
|
||||
key = self.to_k(encoder_hidden_states)
|
||||
value = self.to_v(encoder_hidden_states)
|
||||
|
||||
# Reshape for multi-head attention
|
||||
query = query.unflatten(2, (self.num_heads, -1)).transpose(1, 2)
|
||||
key = key.unflatten(2, (self.num_heads, -1)).transpose(1, 2)
|
||||
value = value.unflatten(2, (self.num_heads, -1)).transpose(1, 2)
|
||||
|
||||
# Apply normalization
|
||||
if self.norm_q is not None:
|
||||
query = self.norm_q(query)
|
||||
if self.norm_k is not None:
|
||||
key = self.norm_k(key)
|
||||
|
||||
# Apply RoPE if provided
|
||||
if image_rotary_emb is not None:
|
||||
query = apply_rotary_emb(query,
|
||||
image_rotary_emb,
|
||||
use_real=True,
|
||||
use_real_unbind_dim=-2)
|
||||
key = apply_rotary_emb(key,
|
||||
image_rotary_emb,
|
||||
use_real=True,
|
||||
use_real_unbind_dim=-2)
|
||||
|
||||
# Prepare for GQA (Grouped Query Attention)
|
||||
if torch.onnx.is_in_onnx_export():
|
||||
query_idx = torch.tensor(query.size(3), device=query.device)
|
||||
key_idx = torch.tensor(key.size(3), device=key.device)
|
||||
value_idx = torch.tensor(value.size(3), device=value.device)
|
||||
else:
|
||||
query_idx = query.size(3)
|
||||
key_idx = key.size(3)
|
||||
value_idx = value.size(3)
|
||||
key = key.repeat_interleave(query_idx // key_idx, dim=3)
|
||||
value = value.repeat_interleave(query_idx // value_idx, dim=3)
|
||||
|
||||
# Attention computation
|
||||
attn_output = torch.nn.functional.scaled_dot_product_attention(
|
||||
query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False
|
||||
)
|
||||
attn_output = attn_output.transpose(1, 2).flatten(2, 3).type_as(query)
|
||||
|
||||
# Output projection
|
||||
attn_output = self.to_out(attn_output)
|
||||
attn_output = self.dropout(attn_output)
|
||||
|
||||
return attn_output
|
||||
|
||||
|
||||
class CosmosCrossAttention(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
dim: int,
|
||||
cross_attention_dim: int,
|
||||
num_heads: int,
|
||||
qk_norm=True,
|
||||
eps=1e-6,
|
||||
prefix: str = "") -> None:
|
||||
assert dim % num_heads == 0
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.cross_attention_dim = cross_attention_dim
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = dim // num_heads
|
||||
self.qk_norm = qk_norm
|
||||
self.eps = eps
|
||||
|
||||
self.to_q = nn.Linear(dim, dim, bias=False)
|
||||
self.to_k = nn.Linear(cross_attention_dim, dim, bias=False)
|
||||
self.to_v = nn.Linear(cross_attention_dim, dim, bias=False)
|
||||
self.to_out = nn.Linear(dim, dim, bias=False)
|
||||
self.dropout = nn.Dropout(0.0)
|
||||
|
||||
self.norm_q = RMSNorm(self.head_dim,
|
||||
eps=eps) if qk_norm else nn.Identity()
|
||||
self.norm_k = RMSNorm(self.head_dim,
|
||||
eps=eps) if qk_norm else nn.Identity()
|
||||
|
||||
def forward(self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
attention_mask: torch.Tensor | None = None) -> torch.Tensor:
|
||||
|
||||
# Get QKV
|
||||
query = self.to_q(hidden_states)
|
||||
key = self.to_k(encoder_hidden_states)
|
||||
value = self.to_v(encoder_hidden_states)
|
||||
|
||||
# Reshape for multi-head attention
|
||||
query = query.unflatten(2, (self.num_heads, -1)).transpose(1, 2)
|
||||
key = key.unflatten(2, (self.num_heads, -1)).transpose(1, 2)
|
||||
value = value.unflatten(2, (self.num_heads, -1)).transpose(1, 2)
|
||||
|
||||
# Apply normalization
|
||||
if self.norm_q is not None:
|
||||
query = self.norm_q(query)
|
||||
if self.norm_k is not None:
|
||||
key = self.norm_k(key)
|
||||
|
||||
# Prepare for GQA (Grouped Query Attention)
|
||||
if torch.onnx.is_in_onnx_export():
|
||||
query_idx = torch.tensor(query.size(3), device=query.device)
|
||||
key_idx = torch.tensor(key.size(3), device=key.device)
|
||||
value_idx = torch.tensor(value.size(3), device=value.device)
|
||||
else:
|
||||
query_idx = query.size(3)
|
||||
key_idx = key.size(3)
|
||||
value_idx = value.size(3)
|
||||
key = key.repeat_interleave(query_idx // key_idx, dim=3)
|
||||
value = value.repeat_interleave(query_idx // value_idx, dim=3)
|
||||
|
||||
# Attention computation
|
||||
attn_output = torch.nn.functional.scaled_dot_product_attention(
|
||||
query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False
|
||||
)
|
||||
attn_output = attn_output.transpose(1, 2).flatten(2, 3).type_as(query)
|
||||
|
||||
# Output projection
|
||||
attn_output = self.to_out(attn_output)
|
||||
attn_output = self.dropout(attn_output)
|
||||
|
||||
return attn_output
|
||||
|
||||
|
||||
class CosmosTransformerBlock(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
num_attention_heads: int,
|
||||
attention_head_dim: int,
|
||||
cross_attention_dim: int,
|
||||
mlp_ratio: float = 4.0,
|
||||
adaln_lora_dim: int = 256,
|
||||
qk_norm: str = "rms_norm",
|
||||
out_bias: bool = False,
|
||||
prefix: str = "",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
hidden_size = num_attention_heads * attention_head_dim
|
||||
|
||||
self.norm1 = CosmosAdaLayerNormZero(in_features=hidden_size,
|
||||
hidden_features=adaln_lora_dim)
|
||||
self.attn1 = CosmosSelfAttention(
|
||||
dim=hidden_size,
|
||||
num_heads=num_attention_heads,
|
||||
qk_norm=(qk_norm == "rms_norm"),
|
||||
eps=1e-5,
|
||||
prefix=f"{prefix}.attn1")
|
||||
|
||||
self.norm2 = CosmosAdaLayerNormZero(in_features=hidden_size,
|
||||
hidden_features=adaln_lora_dim)
|
||||
self.attn2 = CosmosCrossAttention(
|
||||
dim=hidden_size,
|
||||
cross_attention_dim=cross_attention_dim,
|
||||
num_heads=num_attention_heads,
|
||||
qk_norm=(qk_norm == "rms_norm"),
|
||||
eps=1e-5,
|
||||
prefix=f"{prefix}.attn2")
|
||||
|
||||
self.norm3 = CosmosAdaLayerNormZero(in_features=hidden_size,
|
||||
hidden_features=adaln_lora_dim)
|
||||
self.ff = MLP(hidden_size,
|
||||
int(hidden_size * mlp_ratio),
|
||||
act_type="gelu",
|
||||
bias=False)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
embedded_timestep: torch.Tensor,
|
||||
temb: torch.Tensor | None = None,
|
||||
image_rotary_emb: torch.Tensor | None = None,
|
||||
extra_pos_emb: torch.Tensor | None = None,
|
||||
attention_mask: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
if extra_pos_emb is not None:
|
||||
hidden_states = hidden_states + extra_pos_emb
|
||||
|
||||
norm_hidden_states, gate = self.norm1(hidden_states, embedded_timestep,
|
||||
temb)
|
||||
|
||||
attn_output = self.attn1(norm_hidden_states,
|
||||
image_rotary_emb=image_rotary_emb)
|
||||
hidden_states = hidden_states + gate * attn_output
|
||||
|
||||
norm_hidden_states, gate = self.norm2(hidden_states, embedded_timestep,
|
||||
temb)
|
||||
attn_output = self.attn2(norm_hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
attention_mask=attention_mask)
|
||||
|
||||
hidden_states = hidden_states + gate * attn_output
|
||||
|
||||
norm_hidden_states, gate = self.norm3(hidden_states, embedded_timestep,
|
||||
temb)
|
||||
ff_output = self.ff(norm_hidden_states)
|
||||
hidden_states = hidden_states + gate * ff_output
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class CosmosRotaryPosEmbed(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int,
|
||||
max_size: tuple[int, int, int] = (128, 240, 240),
|
||||
patch_size: tuple[int, int, int] = (1, 2, 2),
|
||||
base_fps: int = 24,
|
||||
rope_scale: tuple[float, float, float] = (2.0, 1.0, 1.0),
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
self.max_size = [
|
||||
size // patch
|
||||
for size, patch in zip(max_size, patch_size, strict=False)
|
||||
]
|
||||
self.patch_size = patch_size
|
||||
self.base_fps = base_fps
|
||||
|
||||
self.dim_h = hidden_size // 6 * 2
|
||||
self.dim_w = hidden_size // 6 * 2
|
||||
self.dim_t = hidden_size - self.dim_h - self.dim_w
|
||||
|
||||
self.h_ntk_factor = rope_scale[1]**(self.dim_h / (self.dim_h - 2))
|
||||
self.w_ntk_factor = rope_scale[2]**(self.dim_w / (self.dim_w - 2))
|
||||
self.t_ntk_factor = rope_scale[0]**(self.dim_t / (self.dim_t - 2))
|
||||
|
||||
|
||||
def forward(self,
|
||||
hidden_states: torch.Tensor,
|
||||
fps: int | None = None) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
fps = 16
|
||||
batch_size, num_channels, num_frames, height, width = hidden_states.shape
|
||||
pe_size = [
|
||||
num_frames // self.patch_size[0], height // self.patch_size[1],
|
||||
width // self.patch_size[2]
|
||||
]
|
||||
device = hidden_states.device
|
||||
|
||||
h_theta = 10000.0 * self.h_ntk_factor
|
||||
w_theta = 10000.0 * self.w_ntk_factor
|
||||
t_theta = 10000.0 * self.t_ntk_factor
|
||||
|
||||
seq = torch.arange(max(self.max_size),
|
||||
device=device,
|
||||
dtype=torch.float32)
|
||||
dim_h_range = (
|
||||
torch.arange(0, self.dim_h, 2, device=device,
|
||||
dtype=torch.float32)[:(self.dim_h // 2)] / self.dim_h)
|
||||
dim_w_range = (
|
||||
torch.arange(0, self.dim_w, 2, device=device,
|
||||
dtype=torch.float32)[:(self.dim_w // 2)] / self.dim_w)
|
||||
dim_t_range = (
|
||||
torch.arange(0, self.dim_t, 2, device=device,
|
||||
dtype=torch.float32)[:(self.dim_t // 2)] / self.dim_t)
|
||||
|
||||
h_spatial_freqs = 1.0 / (h_theta**dim_h_range)
|
||||
w_spatial_freqs = 1.0 / (w_theta**dim_w_range)
|
||||
temporal_freqs = 1.0 / (t_theta**dim_t_range)
|
||||
|
||||
emb_h = torch.outer(seq[:pe_size[1]],
|
||||
h_spatial_freqs)[None, :, None, :].repeat(
|
||||
pe_size[0], 1, pe_size[2], 1)
|
||||
emb_w = torch.outer(seq[:pe_size[2]],
|
||||
w_spatial_freqs)[None, None, :, :].repeat(
|
||||
pe_size[0], pe_size[1], 1, 1)
|
||||
|
||||
if fps is None:
|
||||
emb_t = torch.outer(seq[:pe_size[0]], temporal_freqs)
|
||||
else:
|
||||
temporal_scale = seq[:pe_size[0]] / fps * self.base_fps
|
||||
emb_t = torch.outer(temporal_scale,
|
||||
temporal_freqs)
|
||||
|
||||
emb_t = emb_t[:, None, None, :].repeat(1, pe_size[1], pe_size[2], 1)
|
||||
freqs = torch.cat([emb_t, emb_h, emb_w] * 2, dim=-1).flatten(0,
|
||||
2).float()
|
||||
cos = torch.cos(freqs)
|
||||
sin = torch.sin(freqs)
|
||||
return cos, sin
|
||||
|
||||
|
||||
class CosmosLearnablePositionalEmbed(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int,
|
||||
max_size: tuple[int, int, int],
|
||||
patch_size: tuple[int, int, int],
|
||||
eps: float = 1e-6,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
self.max_size = [
|
||||
size // patch
|
||||
for size, patch in zip(max_size, patch_size, strict=False)
|
||||
]
|
||||
self.patch_size = patch_size
|
||||
self.eps = eps
|
||||
|
||||
self.pos_emb_t = nn.Parameter(torch.zeros(self.max_size[0],
|
||||
hidden_size))
|
||||
self.pos_emb_h = nn.Parameter(torch.zeros(self.max_size[1],
|
||||
hidden_size))
|
||||
self.pos_emb_w = nn.Parameter(torch.zeros(self.max_size[2],
|
||||
hidden_size))
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
batch_size, num_channels, num_frames, height, width = hidden_states.shape
|
||||
pe_size = [
|
||||
num_frames // self.patch_size[0], height // self.patch_size[1],
|
||||
width // self.patch_size[2]
|
||||
]
|
||||
|
||||
emb_t = self.pos_emb_t[:pe_size[0]][None, :, None, None, :].repeat(
|
||||
batch_size, 1, pe_size[1], pe_size[2], 1)
|
||||
emb_h = self.pos_emb_h[:pe_size[1]][None, None, :, None, :].repeat(
|
||||
batch_size, pe_size[0], 1, pe_size[2], 1)
|
||||
emb_w = self.pos_emb_w[:pe_size[2]][None, None, None, :, :].repeat(
|
||||
batch_size, pe_size[0], pe_size[1], 1, 1)
|
||||
emb = emb_t + emb_h + emb_w
|
||||
emb = emb.flatten(1, 3)
|
||||
|
||||
norm = torch.linalg.vector_norm(emb,
|
||||
dim=-1,
|
||||
keepdim=True,
|
||||
dtype=torch.float32)
|
||||
norm = torch.add(self.eps,
|
||||
norm,
|
||||
alpha=np.sqrt(norm.numel() / emb.numel()))
|
||||
return (emb / norm).type_as(hidden_states)
|
||||
|
||||
|
||||
class CosmosTransformer3DModel(BaseDiT):
|
||||
_fsdp_shard_conditions = CosmosVideoConfig()._fsdp_shard_conditions
|
||||
_compile_conditions = CosmosVideoConfig()._compile_conditions
|
||||
# _supported_attention_backends = CosmosVideoConfig()._supported_attention_backends
|
||||
param_names_mapping = CosmosVideoConfig().param_names_mapping
|
||||
lora_param_names_mapping = CosmosVideoConfig().lora_param_names_mapping
|
||||
|
||||
def __init__(self, config: CosmosVideoConfig, hf_config: dict[str, Any]) -> None:
|
||||
super().__init__(config=config, hf_config=hf_config)
|
||||
|
||||
inner_dim = config.num_attention_heads * config.attention_head_dim
|
||||
self.hidden_size = config.hidden_size
|
||||
self.num_attention_heads = config.num_attention_heads
|
||||
self.in_channels = config.in_channels
|
||||
self.out_channels = config.out_channels
|
||||
self.num_channels_latents = config.num_channels_latents
|
||||
self.patch_size = config.patch_size
|
||||
self.max_size = config.max_size
|
||||
self.rope_scale = config.rope_scale
|
||||
self.concat_padding_mask = config.concat_padding_mask
|
||||
self.extra_pos_embed_type = config.extra_pos_embed_type
|
||||
|
||||
# 1. Patch Embedding
|
||||
patch_embed_in_channels = config.in_channels + 1 if config.concat_padding_mask else config.in_channels
|
||||
self.patch_embed = CosmosPatchEmbed(patch_embed_in_channels,
|
||||
inner_dim,
|
||||
config.patch_size,
|
||||
bias=False)
|
||||
|
||||
# 2. Positional Embedding
|
||||
self.rope = CosmosRotaryPosEmbed(hidden_size=config.attention_head_dim,
|
||||
max_size=config.max_size,
|
||||
patch_size=config.patch_size,
|
||||
rope_scale=config.rope_scale)
|
||||
|
||||
self.learnable_pos_embed = None
|
||||
if config.extra_pos_embed_type == "learnable":
|
||||
self.learnable_pos_embed = CosmosLearnablePositionalEmbed(
|
||||
hidden_size=inner_dim,
|
||||
max_size=config.max_size,
|
||||
patch_size=config.patch_size,
|
||||
)
|
||||
|
||||
# 3. Time Embedding
|
||||
self.time_embed = CosmosEmbedding(inner_dim, inner_dim)
|
||||
|
||||
# 4. Transformer Blocks
|
||||
self.transformer_blocks = nn.ModuleList([
|
||||
CosmosTransformerBlock(
|
||||
num_attention_heads=config.num_attention_heads,
|
||||
attention_head_dim=config.attention_head_dim,
|
||||
cross_attention_dim=config.text_embed_dim,
|
||||
mlp_ratio=config.mlp_ratio,
|
||||
adaln_lora_dim=config.adaln_lora_dim,
|
||||
qk_norm=config.qk_norm,
|
||||
out_bias=False,
|
||||
prefix=f"{config.prefix}.transformer_blocks.{i}",
|
||||
) for i in range(config.num_layers)
|
||||
])
|
||||
|
||||
# 5. Output norm & projection
|
||||
self.norm_out = CosmosAdaLayerNorm(inner_dim, config.adaln_lora_dim)
|
||||
self.proj_out = nn.Linear(inner_dim,
|
||||
config.out_channels *
|
||||
math.prod(config.patch_size),
|
||||
bias=False)
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
# For TeaCache
|
||||
self.previous_e0_even = None
|
||||
self.previous_e0_odd = None
|
||||
self.previous_residual_even = None
|
||||
self.previous_residual_odd = None
|
||||
self.is_even = True
|
||||
self.should_calc_even = True
|
||||
self.should_calc_odd = True
|
||||
self.accumulated_rel_l1_distance_even = 0
|
||||
self.accumulated_rel_l1_distance_odd = 0
|
||||
self.cnt = 0
|
||||
self.__post_init__()
|
||||
|
||||
def forward(self,
|
||||
hidden_states: torch.Tensor,
|
||||
timestep: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor | list[torch.Tensor],
|
||||
attention_mask: torch.Tensor | None = None,
|
||||
fps: int | None = None,
|
||||
condition_mask: torch.Tensor | None = None,
|
||||
padding_mask: torch.Tensor | None = None,
|
||||
**kwargs) -> torch.Tensor:
|
||||
forward_batch = get_forward_context().forward_batch
|
||||
enable_teacache = forward_batch is not None and forward_batch.enable_teacache
|
||||
|
||||
orig_dtype = hidden_states.dtype
|
||||
if not isinstance(encoder_hidden_states, torch.Tensor):
|
||||
encoder_hidden_states = encoder_hidden_states[0]
|
||||
|
||||
batch_size, num_channels, num_frames, height, width = hidden_states.shape
|
||||
|
||||
# 1. Concatenate padding mask if needed & prepare attention mask
|
||||
if condition_mask is not None:
|
||||
hidden_states = torch.cat([hidden_states, condition_mask], dim=1)
|
||||
|
||||
if self.concat_padding_mask:
|
||||
from torchvision import transforms
|
||||
padding_mask = transforms.functional.resize(
|
||||
padding_mask, list(hidden_states.shape[-2:]), interpolation=transforms.InterpolationMode.NEAREST
|
||||
)
|
||||
hidden_states = torch.cat(
|
||||
[hidden_states, padding_mask.unsqueeze(2).repeat(batch_size, 1, num_frames, 1, 1)], dim=1
|
||||
)
|
||||
|
||||
if attention_mask is not None:
|
||||
attention_mask = attention_mask.unsqueeze(1).unsqueeze(
|
||||
1) # [B, 1, 1, S]
|
||||
|
||||
# 2. Generate positional embeddings
|
||||
image_rotary_emb = self.rope(hidden_states, fps=fps)
|
||||
extra_pos_emb = self.learnable_pos_embed(
|
||||
hidden_states) if self.extra_pos_embed_type == "learnable" else None
|
||||
|
||||
# 3. Patchify input
|
||||
p_t, p_h, p_w = self.patch_size
|
||||
post_patch_num_frames = num_frames // p_t
|
||||
post_patch_height = height // p_h
|
||||
post_patch_width = width // p_w
|
||||
hidden_states = self.patch_embed(hidden_states)
|
||||
hidden_states = hidden_states.flatten(
|
||||
1, 3) # [B, T, H, W, C] -> [B, THW, C] codespell:ignore
|
||||
|
||||
# 4. Timestep embeddings
|
||||
if timestep.ndim == 1:
|
||||
temb, embedded_timestep = self.time_embed(hidden_states, timestep)
|
||||
elif timestep.ndim == 5:
|
||||
assert timestep.shape == (batch_size, 1, num_frames, 1, 1), (
|
||||
f"Expected timestep to have shape [B, 1, T, 1, 1], but got {timestep.shape}"
|
||||
)
|
||||
timestep = timestep.flatten()
|
||||
temb, embedded_timestep = self.time_embed(hidden_states, timestep)
|
||||
# We can do this because num_frames == post_patch_num_frames, as p_t is 1
|
||||
temb, embedded_timestep = (
|
||||
x.view(batch_size, post_patch_num_frames, 1, 1,
|
||||
-1).expand(-1, -1, post_patch_height, post_patch_width,
|
||||
-1).flatten(1, 3)
|
||||
for x in (temb, embedded_timestep)
|
||||
) # [BT, C] -> [B, T, 1, 1, C] -> [B, T, H, W, C] -> [B, THW, C] codespell:ignore
|
||||
else:
|
||||
raise ValueError(f"Unsupported timestep shape: {timestep.shape}")
|
||||
|
||||
# 6. Transformer blocks
|
||||
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
||||
for i, block in enumerate(self.transformer_blocks):
|
||||
hidden_states = self._gradient_checkpointing_func(
|
||||
block,
|
||||
hidden_states,
|
||||
encoder_hidden_states,
|
||||
embedded_timestep,
|
||||
temb,
|
||||
image_rotary_emb,
|
||||
extra_pos_emb,
|
||||
attention_mask,
|
||||
)
|
||||
else:
|
||||
for i, block in enumerate(self.transformer_blocks):
|
||||
hidden_states = block(
|
||||
hidden_states=hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
embedded_timestep=embedded_timestep,
|
||||
temb=temb,
|
||||
image_rotary_emb=image_rotary_emb,
|
||||
extra_pos_emb=extra_pos_emb,
|
||||
attention_mask=attention_mask,
|
||||
)
|
||||
|
||||
# 7. Output norm & projection & unpatchify
|
||||
hidden_states = self.norm_out(hidden_states, embedded_timestep, temb)
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
hidden_states = hidden_states.unflatten(2, (p_h, p_w, p_t, -1))
|
||||
hidden_states = hidden_states.unflatten(
|
||||
1, (post_patch_num_frames, post_patch_height, post_patch_width))
|
||||
# NOTE: The permutation order here is not the inverse operation of what happens when patching as usually expected.
|
||||
# It might be a source of confusion to the reader, but this is correct
|
||||
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
|
||||
@@ -1,961 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import math
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torchvision import transforms
|
||||
|
||||
from fastvideo.attention import DistributedAttention, LocalAttention
|
||||
from fastvideo.configs.models.dits.cosmos2_5 import Cosmos25VideoConfig
|
||||
from fastvideo.forward_context import get_forward_context
|
||||
from fastvideo.layers.layernorm import RMSNorm
|
||||
from fastvideo.layers.mlp import MLP
|
||||
from fastvideo.layers.rotary_embedding import apply_rotary_emb
|
||||
from fastvideo.layers.visual_embedding import Timesteps
|
||||
from fastvideo.models.dits.base import BaseDiT
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
|
||||
class Cosmos25PatchEmbed(nn.Module):
|
||||
"""
|
||||
COSMOS 2.5 patch embedding - converts video (B, C, T, H, W) to patches (B, T', H', W', D).
|
||||
Uses linear projection after rearranging patches.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int,
|
||||
out_channels: int,
|
||||
patch_size: tuple[int, int, int],
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.patch_size = patch_size
|
||||
self.dim = in_channels * patch_size[0] * patch_size[1] * patch_size[2]
|
||||
|
||||
self.proj = nn.Linear(self.dim, out_channels, bias=False)
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Args:
|
||||
hidden_states: (B, C, T, H, W)
|
||||
Returns:
|
||||
(B, T', H', W', D) where T'=T//pt, H'=H//ph, W'=W//pw
|
||||
"""
|
||||
batch_size, num_channels, num_frames, height, width = hidden_states.shape
|
||||
p_t, p_h, p_w = self.patch_size
|
||||
|
||||
# Rearrange: b c (t pt) (h ph) (w pw) -> b t h w (c pt ph pw)
|
||||
hidden_states = hidden_states.reshape(
|
||||
batch_size, num_channels,
|
||||
num_frames // p_t, p_t,
|
||||
height // p_h, p_h,
|
||||
width // p_w, p_w
|
||||
)
|
||||
hidden_states = hidden_states.permute(0, 2, 4, 6, 1, 3, 5, 7)
|
||||
hidden_states = hidden_states.flatten(4, 7) # Flatten patch dimensions
|
||||
|
||||
# Project to model dimension
|
||||
hidden_states = self.proj(hidden_states)
|
||||
return hidden_states
|
||||
|
||||
|
||||
class Cosmos25TimestepEmbedding(nn.Module):
|
||||
"""
|
||||
COSMOS 2.5 timestep embedding with AdaLN-LoRA support.
|
||||
Generates both standard embedding and AdaLN-LoRA parameters.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_features: int,
|
||||
out_features: int,
|
||||
use_adaln_lora: bool = True,
|
||||
adaln_lora_dim: int = 256,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.use_adaln_lora = use_adaln_lora
|
||||
|
||||
self.linear_1 = nn.Linear(in_features, out_features, bias=False)
|
||||
self.activation = nn.SiLU()
|
||||
|
||||
if use_adaln_lora:
|
||||
self.linear_2 = nn.Linear(out_features, 3 * out_features, bias=False)
|
||||
else:
|
||||
self.linear_2 = nn.Linear(out_features, out_features, bias=False)
|
||||
|
||||
def forward(self, sample: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor | None]:
|
||||
"""
|
||||
Returns:
|
||||
emb: Standard embedding (B, T, D)
|
||||
adaln_lora: AdaLN-LoRA parameters (B, T, 3D) or None
|
||||
"""
|
||||
emb = self.linear_1(sample)
|
||||
emb = self.activation(emb)
|
||||
emb = self.linear_2(emb)
|
||||
|
||||
if self.use_adaln_lora:
|
||||
adaln_lora = emb # (B, T, 3D)
|
||||
emb_standard = sample # Use input as standard embedding
|
||||
else:
|
||||
emb_standard = emb
|
||||
adaln_lora = None
|
||||
|
||||
return emb_standard, adaln_lora
|
||||
|
||||
|
||||
class Cosmos25Embedding(nn.Module):
|
||||
"""
|
||||
COSMOS 2.5 timestep conditioning embedding.
|
||||
Generates sinusoidal embeddings and processes them through MLP.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
embedding_dim: int,
|
||||
condition_dim: int,
|
||||
use_adaln_lora: bool = True,
|
||||
adaln_lora_dim: int = 256,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
self.time_proj = Timesteps(embedding_dim, flip_sin_to_cos=True, downscale_freq_shift=0.0)
|
||||
self.t_embedder = Cosmos25TimestepEmbedding(
|
||||
embedding_dim,
|
||||
condition_dim,
|
||||
use_adaln_lora=use_adaln_lora,
|
||||
adaln_lora_dim=adaln_lora_dim,
|
||||
)
|
||||
self.norm = RMSNorm(embedding_dim, eps=1e-6)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
timestep: torch.Tensor,
|
||||
) -> tuple[torch.Tensor, torch.Tensor | None]:
|
||||
"""
|
||||
Args:
|
||||
timestep: (B, T) tensor of timesteps
|
||||
|
||||
Returns:
|
||||
embedded_timestep: Normalized timestep embedding (B, T, D)
|
||||
adaln_lora: AdaLN-LoRA parameters (B, T, 3D) or None
|
||||
"""
|
||||
# Handle 2D timestep input (B, T) like the official model
|
||||
assert timestep.ndim == 2, f"Expected 2D timestep, got {timestep.ndim}D with shape {timestep.shape}"
|
||||
B, T = timestep.shape
|
||||
|
||||
# Flatten for Timesteps layer which expects 1D, then reshape back
|
||||
timestep_flat = timestep.flatten() # (B*T,)
|
||||
timesteps_proj = self.time_proj(timestep_flat).type_as(hidden_states) # (B*T, D)
|
||||
timesteps_proj = timesteps_proj.reshape(B, T, -1) # (B, T, D)
|
||||
|
||||
embedded_timestep, adaln_lora = self.t_embedder(timesteps_proj)
|
||||
embedded_timestep = self.norm(embedded_timestep)
|
||||
|
||||
return embedded_timestep, adaln_lora
|
||||
|
||||
|
||||
class Cosmos25AdaLayerNormZero(nn.Module):
|
||||
"""
|
||||
COSMOS 2.5 Adaptive Layer Normalization with zero initialization and gate.
|
||||
This is a simplified version that expects pre-computed shift/scale/gate parameters.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_features: int,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.norm = nn.LayerNorm(in_features, elementwise_affine=False, eps=1e-6)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
shift: torch.Tensor,
|
||||
scale: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Args:
|
||||
hidden_states: Input tensor
|
||||
shift: Shift parameter for modulation
|
||||
scale: Scale parameter for modulation
|
||||
|
||||
Returns:
|
||||
normalized_hidden_states: Modulated normalized hidden states
|
||||
"""
|
||||
# Apply layer norm and modulation
|
||||
hidden_states = self.norm(hidden_states)
|
||||
hidden_states = hidden_states * (1 + scale) + shift
|
||||
return hidden_states
|
||||
|
||||
|
||||
class Cosmos25SelfAttention(nn.Module):
|
||||
"""
|
||||
COSMOS 2.5 self-attention with QK normalization and RoPE.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
num_heads: int,
|
||||
qk_norm: bool = True,
|
||||
eps: float = 1e-6,
|
||||
supported_attention_backends: tuple[AttentionBackendEnum, ...] | None = None,
|
||||
) -> None:
|
||||
assert dim % num_heads == 0
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = dim // num_heads
|
||||
|
||||
self.to_q = nn.Linear(dim, dim, bias=False)
|
||||
self.to_k = nn.Linear(dim, dim, bias=False)
|
||||
self.to_v = nn.Linear(dim, dim, bias=False)
|
||||
self.to_out = nn.Linear(dim, dim, bias=False)
|
||||
|
||||
self.norm_q = RMSNorm(self.head_dim, eps=eps) if qk_norm else nn.Identity()
|
||||
self.norm_k = RMSNorm(self.head_dim, eps=eps) if qk_norm else nn.Identity()
|
||||
|
||||
# Use DistributedAttention for flexible backend support (torch SDPA / FlashAttention)
|
||||
# For single-GPU (non-distributed), use LocalAttention to avoid distributed requirements
|
||||
if supported_attention_backends is None:
|
||||
supported_attention_backends = (AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA)
|
||||
|
||||
# Always use DistributedAttention (requires distributed environment to be initialized)
|
||||
self.attn = DistributedAttention(
|
||||
num_heads=num_heads,
|
||||
head_size=self.head_dim,
|
||||
causal=False,
|
||||
supported_attention_backends=supported_attention_backends,
|
||||
prefix="self_attn"
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
rope_emb: tuple[torch.Tensor, torch.Tensor] | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Args:
|
||||
hidden_states: (B, S, D) where S = T*H*W
|
||||
rope_emb: Tuple of (cos, sin) for RoPE
|
||||
"""
|
||||
# Get QKV
|
||||
query = self.to_q(hidden_states)
|
||||
key = self.to_k(hidden_states)
|
||||
value = self.to_v(hidden_states)
|
||||
|
||||
# Reshape for multi-head attention: (B, S, D) -> (B, S, H, D_h) -> (B, H, S, D_h)
|
||||
query = query.unflatten(-1, (self.num_heads, self.head_dim)).transpose(1, 2)
|
||||
key = key.unflatten(-1, (self.num_heads, self.head_dim)).transpose(1, 2)
|
||||
value = value.unflatten(-1, (self.num_heads, self.head_dim)).transpose(1, 2)
|
||||
|
||||
# Apply QK normalization
|
||||
query = self.norm_q(query)
|
||||
key = self.norm_k(key)
|
||||
|
||||
# Apply RoPE if provided (query/key are now in (B, H, S, D_h) format)
|
||||
if rope_emb is not None:
|
||||
cos, sin = rope_emb
|
||||
query = apply_rotary_emb(query, (cos, sin), use_real=True, use_real_unbind_dim=-2)
|
||||
key = apply_rotary_emb(key, (cos, sin), use_real=True, use_real_unbind_dim=-2)
|
||||
|
||||
# Attention computation using DistributedAttention or LocalAttention
|
||||
# Both expect (B, S, H, D_h), so transpose first
|
||||
query = query.transpose(1, 2) # (B, H, S, D_h) -> (B, S, H, D_h)
|
||||
key = key.transpose(1, 2)
|
||||
value = value.transpose(1, 2)
|
||||
|
||||
|
||||
attn_output, _ = self.attn(query, key, value)
|
||||
# Reshape back: (B, S, H, D_h) -> (B, S, H*D_h)
|
||||
attn_output = attn_output.flatten(-2, -1)
|
||||
|
||||
# Output projection
|
||||
attn_output = self.to_out(attn_output)
|
||||
return attn_output
|
||||
|
||||
|
||||
class Cosmos25CrossAttention(nn.Module):
|
||||
"""
|
||||
COSMOS 2.5 cross-attention for text conditioning.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
cross_attention_dim: int,
|
||||
num_heads: int,
|
||||
qk_norm: bool = True,
|
||||
eps: float = 1e-6,
|
||||
supported_attention_backends: tuple[AttentionBackendEnum, ...] | None = None,
|
||||
) -> None:
|
||||
assert dim % num_heads == 0
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.cross_attention_dim = cross_attention_dim
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = dim // num_heads
|
||||
|
||||
self.to_q = nn.Linear(dim, dim, bias=False)
|
||||
self.to_k = nn.Linear(cross_attention_dim, dim, bias=False)
|
||||
self.to_v = nn.Linear(cross_attention_dim, dim, bias=False)
|
||||
self.to_out = nn.Linear(dim, dim, bias=False)
|
||||
|
||||
self.norm_q = RMSNorm(self.head_dim, eps=eps) if qk_norm else nn.Identity()
|
||||
self.norm_k = RMSNorm(self.head_dim, eps=eps) if qk_norm else nn.Identity()
|
||||
|
||||
if supported_attention_backends is None:
|
||||
supported_attention_backends = (AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA)
|
||||
|
||||
# Use LocalAttention for cross-attention since text embeddings are not sharded
|
||||
# in sequence parallelism (replicated across ranks)
|
||||
self.attn = LocalAttention(
|
||||
num_heads=num_heads,
|
||||
head_size=self.head_dim,
|
||||
causal=False,
|
||||
supported_attention_backends=supported_attention_backends,
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
attention_mask: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Args:
|
||||
hidden_states: (B, S, D)
|
||||
encoder_hidden_states: (B, N, D_text)
|
||||
"""
|
||||
# Get QKV
|
||||
query = self.to_q(hidden_states)
|
||||
key = self.to_k(encoder_hidden_states)
|
||||
value = self.to_v(encoder_hidden_states)
|
||||
|
||||
# Reshape for multi-head attention
|
||||
query = query.unflatten(-1, (self.num_heads, self.head_dim))
|
||||
key = key.unflatten(-1, (self.num_heads, self.head_dim))
|
||||
value = value.unflatten(-1, (self.num_heads, self.head_dim))
|
||||
|
||||
# Apply QK normalization
|
||||
query = self.norm_q(query)
|
||||
key = self.norm_k(key)
|
||||
|
||||
# LocalAttention expects (B, S, H, D_h), which is what we already have
|
||||
attn_output = self.attn(query, key, value)
|
||||
|
||||
# Reshape back: (B, S, H, D_h) -> (B, S, H*D_h)
|
||||
attn_output = attn_output.flatten(-2, -1)
|
||||
|
||||
# Output projection
|
||||
attn_output = self.to_out(attn_output)
|
||||
return attn_output
|
||||
|
||||
|
||||
class Cosmos25TransformerBlock(nn.Module):
|
||||
"""
|
||||
COSMOS 2.5 transformer block with self-attention, cross-attention, and MLP.
|
||||
Uses AdaLN-LoRA for conditioning.
|
||||
Matches the official architecture where modulation parameters are computed once per block.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
num_attention_heads: int,
|
||||
attention_head_dim: int,
|
||||
cross_attention_dim: int,
|
||||
mlp_ratio: float = 4.0,
|
||||
adaln_lora_dim: int = 256,
|
||||
use_adaln_lora: bool = True,
|
||||
qk_norm: bool = True,
|
||||
supported_attention_backends: tuple[AttentionBackendEnum, ...] | None = None,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
hidden_size = num_attention_heads * attention_head_dim
|
||||
self.use_adaln_lora = use_adaln_lora
|
||||
|
||||
# Layer norms (no modulation logic inside)
|
||||
self.norm1 = Cosmos25AdaLayerNormZero(hidden_size)
|
||||
self.norm2 = Cosmos25AdaLayerNormZero(hidden_size)
|
||||
self.norm3 = Cosmos25AdaLayerNormZero(hidden_size)
|
||||
|
||||
# Attention and MLP layers
|
||||
self.attn1 = Cosmos25SelfAttention(
|
||||
dim=hidden_size,
|
||||
num_heads=num_attention_heads,
|
||||
qk_norm=qk_norm,
|
||||
supported_attention_backends=supported_attention_backends,
|
||||
)
|
||||
self.attn2 = Cosmos25CrossAttention(
|
||||
dim=hidden_size,
|
||||
cross_attention_dim=cross_attention_dim,
|
||||
num_heads=num_attention_heads,
|
||||
qk_norm=qk_norm,
|
||||
supported_attention_backends=supported_attention_backends,
|
||||
)
|
||||
self.mlp = MLP(hidden_size, int(hidden_size * mlp_ratio), act_type="gelu", bias=False)
|
||||
|
||||
# AdaLN modulation layers (compute shift/scale/gate for each sub-layer)
|
||||
# These match the official model's adaln_modulation_* layers
|
||||
if use_adaln_lora:
|
||||
self.adaln_modulation_self_attn = nn.Sequential(
|
||||
nn.SiLU(),
|
||||
nn.Linear(hidden_size, adaln_lora_dim, bias=False),
|
||||
nn.Linear(adaln_lora_dim, 3 * hidden_size, bias=False),
|
||||
)
|
||||
self.adaln_modulation_cross_attn = nn.Sequential(
|
||||
nn.SiLU(),
|
||||
nn.Linear(hidden_size, adaln_lora_dim, bias=False),
|
||||
nn.Linear(adaln_lora_dim, 3 * hidden_size, bias=False),
|
||||
)
|
||||
self.adaln_modulation_mlp = nn.Sequential(
|
||||
nn.SiLU(),
|
||||
nn.Linear(hidden_size, adaln_lora_dim, bias=False),
|
||||
nn.Linear(adaln_lora_dim, 3 * hidden_size, bias=False),
|
||||
)
|
||||
else:
|
||||
self.adaln_modulation_self_attn = nn.Sequential(
|
||||
nn.SiLU(), nn.Linear(hidden_size, 3 * hidden_size, bias=False)
|
||||
)
|
||||
self.adaln_modulation_cross_attn = nn.Sequential(
|
||||
nn.SiLU(), nn.Linear(hidden_size, 3 * hidden_size, bias=False)
|
||||
)
|
||||
self.adaln_modulation_mlp = nn.Sequential(
|
||||
nn.SiLU(), nn.Linear(hidden_size, 3 * hidden_size, bias=False)
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
embedded_timestep: torch.Tensor,
|
||||
adaln_lora: torch.Tensor | None = None,
|
||||
rope_emb: tuple[torch.Tensor, torch.Tensor] | None = None,
|
||||
extra_pos_emb: torch.Tensor | None = None,
|
||||
attention_mask: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Args:
|
||||
hidden_states: (B, T, H, W, D)
|
||||
encoder_hidden_states: (B, N, D_text)
|
||||
embedded_timestep: (B, T, D)
|
||||
adaln_lora: (B, T, 3D) AdaLN-LoRA parameters
|
||||
rope_emb: Tuple of (cos, sin) for RoPE
|
||||
extra_pos_emb: Optional learnable positional embeddings
|
||||
"""
|
||||
# Add extra positional embeddings if provided
|
||||
if extra_pos_emb is not None:
|
||||
hidden_states = hidden_states + extra_pos_emb
|
||||
|
||||
B, T, H, W, D = hidden_states.shape
|
||||
|
||||
# Step 1: Compute ALL modulation parameters once (matches official model)
|
||||
if self.use_adaln_lora and adaln_lora is not None:
|
||||
shift_self_attn, scale_self_attn, gate_self_attn = (
|
||||
self.adaln_modulation_self_attn(embedded_timestep) + adaln_lora
|
||||
).chunk(3, dim=-1)
|
||||
shift_cross_attn, scale_cross_attn, gate_cross_attn = (
|
||||
self.adaln_modulation_cross_attn(embedded_timestep) + adaln_lora
|
||||
).chunk(3, dim=-1)
|
||||
shift_mlp, scale_mlp, gate_mlp = (
|
||||
self.adaln_modulation_mlp(embedded_timestep) + adaln_lora
|
||||
).chunk(3, dim=-1)
|
||||
else:
|
||||
shift_self_attn, scale_self_attn, gate_self_attn = self.adaln_modulation_self_attn(
|
||||
embedded_timestep
|
||||
).chunk(3, dim=-1)
|
||||
shift_cross_attn, scale_cross_attn, gate_cross_attn = self.adaln_modulation_cross_attn(
|
||||
embedded_timestep
|
||||
).chunk(3, dim=-1)
|
||||
shift_mlp, scale_mlp, gate_mlp = self.adaln_modulation_mlp(embedded_timestep).chunk(3, dim=-1)
|
||||
|
||||
# Reshape modulation parameters from (B, T, D) to (B, T, 1, 1, D) for broadcasting
|
||||
shift_self_attn = shift_self_attn.unsqueeze(2).unsqueeze(2).type_as(hidden_states)
|
||||
scale_self_attn = scale_self_attn.unsqueeze(2).unsqueeze(2).type_as(hidden_states)
|
||||
gate_self_attn = gate_self_attn.unsqueeze(2).unsqueeze(2).type_as(hidden_states)
|
||||
|
||||
shift_cross_attn = shift_cross_attn.unsqueeze(2).unsqueeze(2).type_as(hidden_states)
|
||||
scale_cross_attn = scale_cross_attn.unsqueeze(2).unsqueeze(2).type_as(hidden_states)
|
||||
gate_cross_attn = gate_cross_attn.unsqueeze(2).unsqueeze(2).type_as(hidden_states)
|
||||
|
||||
shift_mlp = shift_mlp.unsqueeze(2).unsqueeze(2).type_as(hidden_states)
|
||||
scale_mlp = scale_mlp.unsqueeze(2).unsqueeze(2).type_as(hidden_states)
|
||||
gate_mlp = gate_mlp.unsqueeze(2).unsqueeze(2).type_as(hidden_states)
|
||||
|
||||
# Step 2: Self-attention block
|
||||
norm_hidden_states = self.norm1(hidden_states, shift_self_attn, scale_self_attn)
|
||||
# Flatten for attention: (B, T, H, W, D) -> (B, THW, D)
|
||||
norm_hidden_states_flat = norm_hidden_states.flatten(1, 3)
|
||||
|
||||
attn_output = self.attn1(norm_hidden_states_flat, rope_emb=rope_emb)
|
||||
|
||||
# Reshape back and apply residual
|
||||
attn_output = attn_output.unflatten(1, (T, H, W)) # (B, T, H, W, D)
|
||||
hidden_states = hidden_states + gate_self_attn * attn_output
|
||||
|
||||
# Step 3: Cross-attention block
|
||||
norm_hidden_states = self.norm2(hidden_states, shift_cross_attn, scale_cross_attn)
|
||||
norm_hidden_states_flat = norm_hidden_states.flatten(1, 3)
|
||||
|
||||
attn_output = self.attn2(
|
||||
norm_hidden_states_flat,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
attention_mask=attention_mask,
|
||||
)
|
||||
|
||||
attn_output = attn_output.unflatten(1, (T, H, W))
|
||||
hidden_states = hidden_states + gate_cross_attn * attn_output
|
||||
|
||||
# Step 4: MLP block
|
||||
norm_hidden_states = self.norm3(hidden_states, shift_mlp, scale_mlp)
|
||||
norm_hidden_states_flat = norm_hidden_states.flatten(1, 3)
|
||||
|
||||
mlp_output = self.mlp(norm_hidden_states_flat)
|
||||
|
||||
mlp_output = mlp_output.unflatten(1, (T, H, W))
|
||||
hidden_states = hidden_states + gate_mlp * mlp_output
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class Cosmos25RotaryPosEmbed(nn.Module):
|
||||
"""
|
||||
COSMOS 2.5 3D Rotary Position Embedding with NTK-aware extrapolation.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int,
|
||||
max_size: tuple[int, int, int] = (128, 240, 240),
|
||||
patch_size: tuple[int, int, int] = (1, 2, 2),
|
||||
base_fps: int = 24,
|
||||
rope_scale: tuple[float, float, float] = (1.0, 1.0, 1.0),
|
||||
enable_fps_modulation: bool = True,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
self.max_size = [size // patch for size, patch in zip(max_size, patch_size, strict=True)]
|
||||
self.patch_size = patch_size
|
||||
self.base_fps = base_fps
|
||||
self.enable_fps_modulation = enable_fps_modulation
|
||||
|
||||
# Split dimensions: 1/3 for T, 1/3 for H, 1/3 for W
|
||||
self.dim_h = hidden_size // 6 * 2
|
||||
self.dim_w = hidden_size // 6 * 2
|
||||
self.dim_t = hidden_size - self.dim_h - self.dim_w
|
||||
|
||||
# NTK-aware extrapolation factors
|
||||
self.h_ntk_factor = rope_scale[1] ** (self.dim_h / (self.dim_h - 2))
|
||||
self.w_ntk_factor = rope_scale[2] ** (self.dim_w / (self.dim_w - 2))
|
||||
self.t_ntk_factor = rope_scale[0] ** (self.dim_t / (self.dim_t - 2))
|
||||
|
||||
def forward(
|
||||
self, hidden_states: torch.Tensor, fps: int | None = None
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Generate 3D RoPE embeddings.
|
||||
|
||||
Args:
|
||||
hidden_states: (B, T, H, W, D) - patch-embedded features
|
||||
fps: Frames per second for temporal scaling
|
||||
|
||||
Returns:
|
||||
cos, sin: RoPE embeddings (THW, D)
|
||||
"""
|
||||
batch_size, T, H, W, input_dim = hidden_states.shape
|
||||
device = hidden_states.device
|
||||
|
||||
# T, H, W are already patch dimensions after patch_embed
|
||||
# No need to divide by patch_size
|
||||
|
||||
# Generate frequency scales with NTK
|
||||
h_theta = 10000.0 * self.h_ntk_factor
|
||||
w_theta = 10000.0 * self.w_ntk_factor
|
||||
t_theta = 10000.0 * self.t_ntk_factor
|
||||
|
||||
seq = torch.arange(max(self.max_size), device=device, dtype=torch.float32)
|
||||
|
||||
# Use self.dim_h/w/t which were set during initialization
|
||||
dim_h_range = torch.arange(0, self.dim_h, 2, device=device, dtype=torch.float32)[: (self.dim_h // 2)] / self.dim_h
|
||||
dim_w_range = torch.arange(0, self.dim_w, 2, device=device, dtype=torch.float32)[: (self.dim_w // 2)] / self.dim_w
|
||||
dim_t_range = torch.arange(0, self.dim_t, 2, device=device, dtype=torch.float32)[: (self.dim_t // 2)] / self.dim_t
|
||||
|
||||
h_spatial_freqs = 1.0 / (h_theta ** dim_h_range)
|
||||
w_spatial_freqs = 1.0 / (w_theta ** dim_w_range)
|
||||
temporal_freqs = 1.0 / (t_theta ** dim_t_range)
|
||||
|
||||
# Generate positional embeddings
|
||||
half_emb_h = torch.outer(seq[:H], h_spatial_freqs)
|
||||
half_emb_w = torch.outer(seq[:W], w_spatial_freqs)
|
||||
|
||||
if self.enable_fps_modulation and fps is not None:
|
||||
# Apply FPS scaling
|
||||
half_emb_t = torch.outer(seq[:T] / fps * self.base_fps, temporal_freqs)
|
||||
else:
|
||||
half_emb_t = torch.outer(seq[:T], temporal_freqs)
|
||||
|
||||
# Broadcast and concatenate embeddings
|
||||
emb_t = half_emb_t[:, None, None, :].repeat(1, H, W, 1)
|
||||
emb_h = half_emb_h[None, :, None, :].repeat(T, 1, W, 1)
|
||||
emb_w = half_emb_w[None, None, :, :].repeat(T, H, 1, 1)
|
||||
|
||||
# Concatenate [t, h, w, t, h, w] for sin/cos pairs
|
||||
freqs = torch.cat([emb_t, emb_h, emb_w] * 2, dim=-1)
|
||||
freqs = freqs.flatten(0, 2).float() # (THW, D)
|
||||
|
||||
cos = torch.cos(freqs) # (THW, D)
|
||||
sin = torch.sin(freqs) # (THW, D)
|
||||
|
||||
return cos, sin
|
||||
|
||||
|
||||
class Cosmos25LearnablePositionalEmbed(nn.Module):
|
||||
"""
|
||||
COSMOS 2.5 learnable absolute positional embeddings (optional).
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int,
|
||||
max_size: tuple[int, int, int],
|
||||
patch_size: tuple[int, int, int],
|
||||
eps: float = 1e-6,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
self.max_size = [size // patch for size, patch in zip(max_size, patch_size, strict=True)]
|
||||
self.patch_size = patch_size
|
||||
self.eps = eps
|
||||
|
||||
self.pos_emb_t = nn.Parameter(torch.zeros(self.max_size[0], hidden_size))
|
||||
self.pos_emb_h = nn.Parameter(torch.zeros(self.max_size[1], hidden_size))
|
||||
self.pos_emb_w = nn.Parameter(torch.zeros(self.max_size[2], hidden_size))
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Args:
|
||||
hidden_states: (B, T, H, W, D)
|
||||
Returns:
|
||||
pos_emb: (B, T, H, W, D)
|
||||
"""
|
||||
B, T, H, W, D = hidden_states.shape
|
||||
|
||||
emb_t = self.pos_emb_t[:T][None, :, None, None, :].repeat(B, 1, H, W, 1)
|
||||
emb_h = self.pos_emb_h[:H][None, None, :, None, :].repeat(B, T, 1, W, 1)
|
||||
emb_w = self.pos_emb_w[:W][None, None, None, :, :].repeat(B, T, H, 1, 1)
|
||||
|
||||
emb = emb_t + emb_h + emb_w
|
||||
|
||||
# Normalize
|
||||
norm = torch.linalg.vector_norm(emb, dim=-1, keepdim=True, dtype=torch.float32)
|
||||
norm = torch.add(self.eps, norm, alpha=np.sqrt(norm.numel() / emb.numel()))
|
||||
return (emb / norm).type_as(hidden_states)
|
||||
|
||||
|
||||
class Cosmos25FinalLayer(nn.Module):
|
||||
"""
|
||||
COSMOS 2.5 final layer with AdaLN modulation and unpatchification.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int,
|
||||
out_channels: int,
|
||||
patch_size: tuple[int, int, int],
|
||||
adaln_lora_dim: int = 256,
|
||||
use_adaln_lora: bool = True,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.hidden_size = hidden_size
|
||||
self.use_adaln_lora = use_adaln_lora
|
||||
|
||||
self.norm = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
||||
self.activation = nn.SiLU()
|
||||
|
||||
if use_adaln_lora:
|
||||
self.linear_1 = nn.Linear(hidden_size, adaln_lora_dim, bias=False)
|
||||
self.linear_2 = nn.Linear(adaln_lora_dim, 2 * hidden_size, bias=False)
|
||||
else:
|
||||
self.linear_1 = nn.Identity()
|
||||
self.linear_2 = nn.Linear(hidden_size, 2 * hidden_size, bias=False)
|
||||
|
||||
# Output projection
|
||||
output_dim = out_channels * patch_size[0] * patch_size[1] * patch_size[2]
|
||||
self.proj_out = nn.Linear(hidden_size, output_dim, bias=False)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
embedded_timestep: torch.Tensor,
|
||||
adaln_lora: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Args:
|
||||
hidden_states: (B, T, H, W, D)
|
||||
embedded_timestep: (B, T, D) or (B, D)
|
||||
adaln_lora: (B, T, 3D) or None
|
||||
"""
|
||||
# Generate modulation parameters
|
||||
embedded_timestep = self.activation(embedded_timestep)
|
||||
embedded_timestep = self.linear_1(embedded_timestep)
|
||||
embedded_timestep = self.linear_2(embedded_timestep)
|
||||
|
||||
if self.use_adaln_lora and adaln_lora is not None:
|
||||
# Use first 2*hidden_size elements for shift/scale
|
||||
embedded_timestep = embedded_timestep + adaln_lora[..., : 2 * self.hidden_size]
|
||||
|
||||
shift, scale = embedded_timestep.chunk(2, dim=-1)
|
||||
|
||||
# Apply normalization and modulation
|
||||
hidden_states = self.norm(hidden_states)
|
||||
|
||||
# Reshape for broadcasting if needed
|
||||
if embedded_timestep.ndim == 2:
|
||||
shift, scale = (x.unsqueeze(1) for x in (shift, scale))
|
||||
elif embedded_timestep.ndim == 3 and hidden_states.ndim == 5:
|
||||
shift, scale = (x.unsqueeze(2).unsqueeze(2) for x in (shift, scale))
|
||||
|
||||
hidden_states = hidden_states * (1 + scale) + shift
|
||||
|
||||
# Project to output
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class Cosmos25Transformer3DModel(BaseDiT):
|
||||
"""
|
||||
COSMOS 2.5 DiT - MiniTrainDIT architecture adapted for FastVideo.
|
||||
|
||||
Key features:
|
||||
- AdaLN-LoRA conditioning
|
||||
- 3D RoPE with NTK-aware extrapolation
|
||||
- Optional learnable positional embeddings
|
||||
- QK normalization
|
||||
- Cross-attention projection (optional)
|
||||
"""
|
||||
|
||||
_fsdp_shard_conditions = Cosmos25VideoConfig()._fsdp_shard_conditions
|
||||
_compile_conditions = Cosmos25VideoConfig()._compile_conditions
|
||||
param_names_mapping = Cosmos25VideoConfig().param_names_mapping
|
||||
lora_param_names_mapping = Cosmos25VideoConfig().lora_param_names_mapping
|
||||
|
||||
def __init__(self, config: Cosmos25VideoConfig, hf_config: dict[str, Any]) -> None:
|
||||
super().__init__(config=config, hf_config=hf_config)
|
||||
|
||||
inner_dim = config.num_attention_heads * config.attention_head_dim
|
||||
self.hidden_size = inner_dim
|
||||
self.num_attention_heads = config.num_attention_heads
|
||||
self.in_channels = config.in_channels
|
||||
self.out_channels = config.out_channels
|
||||
self.num_channels_latents = config.num_channels_latents
|
||||
self.patch_size = config.patch_size
|
||||
self.max_size = config.max_size
|
||||
self.rope_scale = config.rope_scale
|
||||
self.concat_padding_mask = config.concat_padding_mask
|
||||
self.use_adaln_lora = getattr(config, "use_adaln_lora", True)
|
||||
self.adaln_lora_dim = getattr(config, "adaln_lora_dim", 256)
|
||||
self.extra_pos_embed_type = getattr(config, "extra_pos_embed_type", None)
|
||||
self.use_crossattn_projection = getattr(config, "use_crossattn_projection", False)
|
||||
|
||||
# 1. Patch Embedding
|
||||
# Account for: VAE channels + condition_mask (1) + padding_mask (1 if concat_padding_mask)
|
||||
patch_embed_in_channels = config.in_channels # Base VAE channels (16)
|
||||
patch_embed_in_channels += 1 # Always add 1 for condition_mask
|
||||
if config.concat_padding_mask:
|
||||
patch_embed_in_channels += 1 # Add 1 for padding_mask
|
||||
# Total: 16 + 1 + 1 = 18 (with concat_padding_mask=True)
|
||||
|
||||
self.patch_embed = Cosmos25PatchEmbed(
|
||||
patch_embed_in_channels, inner_dim, config.patch_size
|
||||
)
|
||||
|
||||
# 2. Positional Embeddings
|
||||
self.rope = Cosmos25RotaryPosEmbed(
|
||||
hidden_size=config.attention_head_dim,
|
||||
max_size=config.max_size,
|
||||
patch_size=config.patch_size,
|
||||
rope_scale=config.rope_scale,
|
||||
enable_fps_modulation=getattr(config, "rope_enable_fps_modulation", True),
|
||||
)
|
||||
|
||||
self.learnable_pos_embed = None
|
||||
if self.extra_pos_embed_type == "learnable":
|
||||
self.learnable_pos_embed = Cosmos25LearnablePositionalEmbed(
|
||||
hidden_size=inner_dim,
|
||||
max_size=config.max_size,
|
||||
patch_size=config.patch_size,
|
||||
)
|
||||
|
||||
# 3. Time Embedding
|
||||
self.time_embed = Cosmos25Embedding(
|
||||
inner_dim,
|
||||
inner_dim,
|
||||
use_adaln_lora=self.use_adaln_lora,
|
||||
adaln_lora_dim=self.adaln_lora_dim,
|
||||
)
|
||||
|
||||
# 4. Cross-attention projection (optional)
|
||||
if self.use_crossattn_projection:
|
||||
crossattn_proj_in_channels = getattr(config, "crossattn_proj_in_channels", config.text_embed_dim)
|
||||
self.crossattn_proj = nn.Sequential(
|
||||
nn.Linear(crossattn_proj_in_channels, config.text_embed_dim, bias=True),
|
||||
nn.GELU(),
|
||||
)
|
||||
|
||||
# 5. Transformer Blocks
|
||||
self.transformer_blocks = nn.ModuleList([
|
||||
Cosmos25TransformerBlock(
|
||||
num_attention_heads=config.num_attention_heads,
|
||||
attention_head_dim=config.attention_head_dim,
|
||||
cross_attention_dim=config.text_embed_dim,
|
||||
mlp_ratio=config.mlp_ratio,
|
||||
adaln_lora_dim=self.adaln_lora_dim,
|
||||
use_adaln_lora=self.use_adaln_lora,
|
||||
qk_norm=(config.qk_norm == "rms_norm"),
|
||||
supported_attention_backends=config._supported_attention_backends,
|
||||
)
|
||||
for i in range(config.num_layers)
|
||||
])
|
||||
|
||||
# 6. Final Layer
|
||||
self.final_layer = Cosmos25FinalLayer(
|
||||
hidden_size=inner_dim,
|
||||
out_channels=config.out_channels,
|
||||
patch_size=config.patch_size,
|
||||
adaln_lora_dim=self.adaln_lora_dim,
|
||||
use_adaln_lora=self.use_adaln_lora,
|
||||
)
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
self.__post_init__()
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
timestep: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor | list[torch.Tensor],
|
||||
attention_mask: torch.Tensor | None = None,
|
||||
fps: int | None = None,
|
||||
condition_mask: torch.Tensor | None = None,
|
||||
padding_mask: torch.Tensor | None = None,
|
||||
**kwargs,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Args:
|
||||
hidden_states: (B, C, T, H, W) latent video
|
||||
timestep: (B,) or (B, T) diffusion timesteps
|
||||
encoder_hidden_states: (B, N, D_text) text embeddings
|
||||
attention_mask: Optional attention mask
|
||||
fps: Frames per second
|
||||
condition_mask: (B, 1, T, H, W) conditioning mask
|
||||
padding_mask: (B, 1, H, W) padding mask
|
||||
"""
|
||||
if not isinstance(encoder_hidden_states, torch.Tensor):
|
||||
encoder_hidden_states = encoder_hidden_states[0]
|
||||
|
||||
batch_size, num_channels, num_frames, height, width = hidden_states.shape
|
||||
|
||||
# 1. Concatenate condition mask if provided
|
||||
if condition_mask is not None:
|
||||
hidden_states = torch.cat([hidden_states, condition_mask], dim=1)
|
||||
|
||||
# 2. Concatenate padding mask if needed
|
||||
if self.concat_padding_mask and padding_mask is not None:
|
||||
padding_mask = transforms.functional.resize(
|
||||
padding_mask,
|
||||
list(hidden_states.shape[-2:]),
|
||||
interpolation=transforms.InterpolationMode.NEAREST,
|
||||
)
|
||||
hidden_states = torch.cat(
|
||||
[hidden_states, padding_mask.unsqueeze(2).repeat(1, 1, num_frames, 1, 1)],
|
||||
dim=1,
|
||||
)
|
||||
|
||||
# 3. Patchify input
|
||||
p_t, p_h, p_w = self.patch_size
|
||||
post_patch_num_frames = num_frames // p_t
|
||||
post_patch_height = height // p_h
|
||||
post_patch_width = width // p_w
|
||||
|
||||
|
||||
hidden_states = self.patch_embed(hidden_states) # (B, T', H', W', D)
|
||||
|
||||
|
||||
# 4. Generate RoPE embeddings (after patchify, using patch dimensions)
|
||||
rope_emb = self.rope(hidden_states, fps=fps)
|
||||
|
||||
# 5. Generate learnable positional embeddings (if used)
|
||||
extra_pos_emb = None
|
||||
if self.learnable_pos_embed is not None:
|
||||
extra_pos_emb = self.learnable_pos_embed(hidden_states)
|
||||
|
||||
# 6. Timestep embeddings
|
||||
# Official model expects timestep in (B, T) format, so ensure it has 2D shape
|
||||
if timestep.ndim == 1:
|
||||
# Scalar timestep per sample: (B,) -> (B, 1)
|
||||
timestep = timestep.unsqueeze(1)
|
||||
elif timestep.ndim == 2:
|
||||
# Already in (B, T) format
|
||||
pass
|
||||
else:
|
||||
raise ValueError(f"Unsupported timestep shape: {timestep.shape}")
|
||||
|
||||
# Now timestep is always (B, T), pass directly to time_embed
|
||||
embedded_timestep, adaln_lora = self.time_embed(hidden_states, timestep)
|
||||
|
||||
# 7. Apply cross-attention projection (if used)
|
||||
if self.use_crossattn_projection:
|
||||
encoder_hidden_states = self.crossattn_proj(encoder_hidden_states)
|
||||
|
||||
|
||||
|
||||
|
||||
# Prepare attention mask
|
||||
if attention_mask is not None:
|
||||
attention_mask = attention_mask.unsqueeze(1).unsqueeze(1) # (B, 1, 1, N)
|
||||
|
||||
# 8. Transformer blocks
|
||||
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
||||
for block in self.transformer_blocks:
|
||||
hidden_states = self._gradient_checkpointing_func(
|
||||
block,
|
||||
hidden_states,
|
||||
encoder_hidden_states,
|
||||
embedded_timestep,
|
||||
adaln_lora,
|
||||
rope_emb,
|
||||
extra_pos_emb,
|
||||
attention_mask,
|
||||
)
|
||||
else:
|
||||
for i, block in enumerate(self.transformer_blocks):
|
||||
hidden_states = block(
|
||||
hidden_states=hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
embedded_timestep=embedded_timestep,
|
||||
adaln_lora=adaln_lora,
|
||||
rope_emb=rope_emb,
|
||||
extra_pos_emb=extra_pos_emb,
|
||||
attention_mask=attention_mask,
|
||||
)
|
||||
|
||||
# 9. Final layer - output norm & projection
|
||||
hidden_states = self.final_layer(hidden_states, embedded_timestep, adaln_lora)
|
||||
|
||||
# 10. Unpatchify: (B, T', H', W', P) -> (B, C, T, H, W)
|
||||
# After unflatten: (B, T', H', W', p_t, p_h, p_w, C) with dims [0,1,2,3,4,5,6,7]
|
||||
hidden_states = hidden_states.unflatten(-1, (p_t, p_h, p_w, self.out_channels))
|
||||
# Permute to: (B, C, T', p_t, H', p_h, W', p_w)
|
||||
hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6)
|
||||
# Flatten pairs to get (B, C, T, H, W)
|
||||
hidden_states = hidden_states.flatten(2, 3).flatten(3, 4).flatten(4, 5)
|
||||
|
||||
return hidden_states
|
||||
|
||||
@@ -180,7 +180,7 @@ class T5Attention(nn.Module):
|
||||
|
||||
self.qkv_proj = QKVParallelLinear(
|
||||
self.d_model,
|
||||
self.key_value_proj_dim,
|
||||
self.d_model // self.total_num_heads,
|
||||
self.total_num_heads,
|
||||
self.total_num_kv_heads,
|
||||
bias=False,
|
||||
@@ -198,7 +198,7 @@ class T5Attention(nn.Module):
|
||||
padding_size=self.relative_attention_num_buckets,
|
||||
quant_config=quant_config)
|
||||
self.o = RowParallelLinear(
|
||||
self.total_num_heads * self.key_value_proj_dim,
|
||||
self.d_model,
|
||||
self.d_model,
|
||||
bias=False,
|
||||
quant_config=quant_config,
|
||||
@@ -297,7 +297,7 @@ class T5Attention(nn.Module):
|
||||
) -> torch.Tensor:
|
||||
bs, seq_len, _ = hidden_states.shape
|
||||
num_seqs = bs
|
||||
n, c = self.n_heads, self.key_value_proj_dim
|
||||
n, c = self.n_heads, self.d_model // self.total_num_heads
|
||||
qkv, _ = self.qkv_proj(hidden_states)
|
||||
# Projection of 'own' hidden state (self-attention). No GQA here.
|
||||
q, k, v = qkv.split(self.inner_dim, dim=-1)
|
||||
@@ -540,10 +540,7 @@ class T5EncoderModel(TextEncoder):
|
||||
attn_metadata=attn_metadata,
|
||||
)
|
||||
|
||||
return BaseEncoderOutput(
|
||||
last_hidden_state=hidden_states,
|
||||
attention_mask=attention_mask,
|
||||
)
|
||||
return BaseEncoderOutput(last_hidden_state=hidden_states)
|
||||
|
||||
def load_weights(self, weights: Iterable[tuple[str,
|
||||
torch.Tensor]]) -> set[str]:
|
||||
|
||||
@@ -4,8 +4,6 @@
|
||||
# Copyright 2024 The TorchTune Authors.
|
||||
# Copyright 2025 The FastVideo Authors.
|
||||
|
||||
from __future__ import annotations
|
||||
import os
|
||||
import contextlib
|
||||
from collections.abc import Callable, Generator
|
||||
from itertools import chain
|
||||
@@ -199,11 +197,10 @@ def shard_model(
|
||||
Raises:
|
||||
ValueError: If no layer modules were sharded, indicating that no shard_condition was triggered.
|
||||
"""
|
||||
# Check if we should use size-based filtering
|
||||
use_size_filtering = os.environ.get("FASTVIDEO_FSDP2_AUTOWRAP", "0") == "1"
|
||||
|
||||
if not fsdp_shard_conditions:
|
||||
logger.warning("No FSDP shard conditions provided; nothing will be sharded.")
|
||||
if fsdp_shard_conditions is None or len(fsdp_shard_conditions) == 0:
|
||||
logger.warning(
|
||||
"The FSDP shard condition list is empty or None. No modules will be sharded in %s",
|
||||
type(model).__name__)
|
||||
return
|
||||
|
||||
fsdp_kwargs = {
|
||||
@@ -218,38 +215,20 @@ def shard_model(
|
||||
# iterating in reverse to start with
|
||||
# lowest-level modules first
|
||||
num_layers_sharded = 0
|
||||
|
||||
if use_size_filtering:
|
||||
# Size-based filtering mode
|
||||
min_params = int(os.environ.get("FASTVIDEO_FSDP2_MIN_PARAMS", "10000000"))
|
||||
logger.info("Using size-based filtering with threshold: %.2fM", min_params / 1e6)
|
||||
|
||||
for n, m in reversed(list(model.named_modules())):
|
||||
if any([shard_condition(n, m) for shard_condition in fsdp_shard_conditions]):
|
||||
# Count all parameters
|
||||
param_count = sum(p.numel() for p in m.parameters(recurse=True))
|
||||
|
||||
# Skip small modules
|
||||
if param_count < min_params:
|
||||
logger.info("Skipping module %s (%.2fM params < %.2fM threshold)",
|
||||
n, param_count / 1e6, min_params / 1e6)
|
||||
continue
|
||||
|
||||
# Shard this module
|
||||
logger.info("Sharding module %s (%.2fM params)", n, param_count / 1e6)
|
||||
fully_shard(m, **fsdp_kwargs)
|
||||
num_layers_sharded += 1
|
||||
else:
|
||||
# Shard all modules matching conditions
|
||||
for n, m in reversed(list(model.named_modules())):
|
||||
if any([shard_condition(n, m) for shard_condition in fsdp_shard_conditions]):
|
||||
fully_shard(m, **fsdp_kwargs)
|
||||
num_layers_sharded += 1
|
||||
|
||||
if num_layers_sharded == 0:
|
||||
raise ValueError(
|
||||
"No layer modules were sharded. Please check if shard conditions are working as expected."
|
||||
)
|
||||
# TODO(will): don't reshard after forward for the last layer to save on the
|
||||
# all-gather that will immediately happen Shard the model with FSDP,
|
||||
for n, m in reversed(list(model.named_modules())):
|
||||
if any([
|
||||
shard_condition(n, m)
|
||||
for shard_condition in fsdp_shard_conditions
|
||||
]):
|
||||
fully_shard(m, **fsdp_kwargs)
|
||||
num_layers_sharded += 1
|
||||
|
||||
if num_layers_sharded == 0:
|
||||
raise ValueError(
|
||||
"No layer modules were sharded. Please check if shard conditions are working as expected."
|
||||
)
|
||||
|
||||
# Finally shard the entire model to account for any stragglers
|
||||
fully_shard(model, **fsdp_kwargs)
|
||||
|
||||