Compare commits

..
Author SHA1 Message Date
Y-aang c1d7c51bc7 fixi2v demo bug 2025-11-19 20:53:43 +00:00
Y-aang 5de2d80ebd add full i2v demo 2025-11-19 15:40:04 +00:00
RandNMR73 6cd227349c add demo prompts 2025-11-19 05:49:01 +00:00
RandNMR73 7e15e5dab1 add demo images 2025-11-19 05:42:37 +00:00
RandNMR73 c9ee4f8bf8 add i2v demo 2025-11-18 09:40:05 +00:00
SolitaryThinker f764b43aaa nit 2025-11-16 23:10:47 +00:00
SolitaryThinker 592af8c954 fix lint 2025-11-16 23:09:04 +00:00
SolitaryThinker 4225a96a40 use new ckpt 2025-11-16 12:49:06 +00:00
JerryZhou54 733c14a2cb Change to fp32 inference 2025-11-16 04:53:57 +00:00
JerryZhou54 81b302409e Add i2v images 2025-11-16 04:51:14 +00:00
JerryZhou54 7fbd0cfec7 Fi lint 2025-11-15 23:50:04 +00:00
RandNMR73 85b9934079 Add inference for MoE SF 2025-11-15 23:44:21 +00:00
Y-aang c30779184f fix: incorrect dv in vsa Triton kernel causing test_vsa error (#879) 2025-11-14 22:00:39 -08:00
William Lin 9d188c0b6c [misc] update wechat and slack invite links (#875) 2025-11-12 23:03:56 -08:00
Mihir Jagtap 9dd7c54221 [docs] Update Home Readme.md with fixed links (#873) 2025-11-12 13:32:44 -08:00
William Lin 62b95d8287 [feat] prepare for wan2.2 SF (#861) 2025-11-04 18:06:48 -08:00
Kaiqin Kong fdf21702f5 [Docs] add diagrams to docs (#863) 2025-11-04 16:29:07 -08:00
Ohm-Rishabh 2972fc9449 Improve FSDP loading with size-based filtering (#853) 2025-11-04 15:31:07 -08:00
Mihir Jagtap 8f5712629f [docs] port to mkdocs (#855) 2025-11-04 14:31:56 -08:00
Kevin Lin 436c701b9f [bugfix] Add Cosmos2 sampling params to registry (#862) 2025-11-02 00:09:17 -07:00
Kevin Lin 543fea88e3 [Feature] Add Cosmos2 i2v pipeline (#837) 2025-10-30 20:03:57 -07:00
140 changed files with 4361 additions and 1781 deletions
+5 -5
View File
@@ -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)
WANDB_API_KEY=$(buildkite-agent secret get wandb_api_key)
HF_API_KEY=$(buildkite-agent secret get hf_api_key)
if [ -n "$MODAL_TOKEN_ID" ] && [ -n "$MODAL_TOKEN_SECRET" ]; then
log "Retrieved Modal credentials from Buildkite secrets"
@@ -63,15 +63,15 @@ MODAL_ENV="BUILDKITE_REPO=$BUILDKITE_REPO BUILDKITE_COMMIT=$BUILDKITE_COMMIT BUI
case "$TEST_TYPE" in
"encoder")
log "Running encoder tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_encoder_tests"
MODAL_COMMAND="$MODAL_ENV HF_API_KEY=$HF_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_encoder_tests"
;;
"vae")
log "Running VAE tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_vae_tests"
MODAL_COMMAND="$MODAL_ENV HF_API_KEY=$HF_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_vae_tests"
;;
"transformer")
log "Running transformer tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_transformer_tests"
MODAL_COMMAND="$MODAL_ENV HF_API_KEY=$HF_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_transformer_tests"
;;
"ssim")
log "Running SSIM tests..."
+18 -45
View File
@@ -1,82 +1,55 @@
# Sample workflow for building and deploying a Hugo site to GitHub Pages
name: Deploy FastVideo Docs to Pages
name: Deploy Documentation
on:
# Runs on pushes targeting the default branch
push:
branches:
- main
paths:
- "docs/**/*.md"
- "fastvideo/examples/**/*.py"
branches: [ main ]
pull_request:
branches:
- main
types: [opened, ready_for_review, synchronize, reopened]
paths:
- "docs/**/*.md"
- "fastvideo/examples/**/*.py"
branches: [ main ]
# 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 Pages
id: pages
uses: actions/configure-pages@v5
- name: Set up Python
- name: Setup Python
uses: actions/setup-python@v5
with:
python-version: "3.10"
python-version: '3.12'
- name: Install dependencies
run: |
cd docs
pip install -r requirements-docs.txt
- name: Build docs
run: |
cd docs
make clean
make html
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
- name: Upload artifact
uses: actions/upload-pages-artifact@v3
with:
path: ./docs/build/html
path: ./site
# 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
+10 -7
View File
@@ -14,6 +14,8 @@ wandb/
*.pt
cache_dir/
wandb/
venv/
.venv/
runs/
samples/
*validation/
@@ -37,12 +39,13 @@ dist/
eggs/
.eggs/
# Sphinx documentation
docs/_build/
docs/source/getting_started/examples/
docs/source/inference/examples/
docs/source/training/examples/
docs/source/distillation/examples/
# MkDocs documentation
site/
docs/getting_started/examples/
docs/inference/examples/
docs/training/examples/
docs/distillation/examples/
!requirements-mkdocs.txt
# VSCode
.vscode/
@@ -61,7 +64,7 @@ docs/source/distillation/examples/
!fastvideo/tests/ssim/reference_videos/**/*.mp4
# Static images
!docs/source/_static/images/**/*.png
!docs/assets/images/**/*.png
!comfyui/assets/**/*.png
!comfyui/assets/**/*.gif
+1
View File
@@ -10,6 +10,7 @@ exclude: |
demo/.*|
predict\.py|
scripts/.*|
prompts/.*|
fastvideo/data_preprocess/.*|
fastvideo/dataset/.*|
fastvideo/models/.*|
+9 -9
View File
@@ -7,7 +7,7 @@
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.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> |
| 🕹️ <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/XcY0Cpv" target="_blank"> <b> WeChat </b> </a> |
</p>
<div align="center">
@@ -49,10 +49,10 @@ conda activate fastvideo
pip install fastvideo
```
Please see our [docs](https://hao-ai-lab.github.io/FastVideo/getting_started/installation.html) for more detailed installation instructions.
Please see our [docs](https://hao-ai-lab.github.io/FastVideo/getting_started/installation/) for more detailed installation instructions.
## Sparse Distillation
For our sparse distillation techniques, please see our [distillation docs](https://hao-ai-lab.github.io/FastVideo/distillation/dmd.html) 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/) 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.html). 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/). Create a file called `example.py` with the following code:
```python
import os
@@ -100,15 +100,15 @@ 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.html).
For a more detailed guide, please see our [inference quick start](https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start/).
### Other docs:
- [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)
- [Design Overview](https://hao-ai-lab.github.io/FastVideo/design/overview/)
- [Contribution Guide](https://hao-ai-lab.github.io/FastVideo/getting_started/installation/)
## Distillation and Finetuning
- [Distillation Guide](https://hao-ai-lab.github.io/FastVideo/distillation/dmd.html)
- [Distillation Guide](https://hao-ai-lab.github.io/FastVideo/distillation/dmd/)
<!-- - [Finetuning Guide](https://hao-ai-lab.github.io/FastVideo/training/finetune.html) -->
## 📑 Development Plan
@@ -127,7 +127,7 @@ See details in [development roadmap](https://github.com/hao-ai-lab/FastVideo/iss
## 🤝 Contributing
We welcome all contributions. Please check out our guide [here](https://hao-ai-lab.github.io/FastVideo/contributing/overview.html)
We welcome all contributions. Please check out our guide [here](https://hao-ai-lab.github.io/FastVideo/contributing/overview/)
## Acknowledgement
We learned and reused code from the following projects:
+1 -1
View File
@@ -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/source/_static/images/STA_configuration.png" width="80%"/>
<img src="../../../docs/assets/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)
-26
View File
@@ -1,26 +0,0 @@
# 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"
+29 -10
View File
@@ -1,20 +1,39 @@
# FastVideo documents
# FastVideo Documentation
## Build the docs
This directory contains the FastVideo documentation built with MkDocs.
## Build the docs locally
```bash
# Install dependencies.
pip install -r requirements-docs.txt
# Install dependencies
pip install -r docs/requirements-mkdocs.txt
# Build the docs.
make clean
make html
# Serve docs with live reload (recommended for development)
mkdocs serve
# Or build static site
mkdocs build
```
## Open the docs with your browser
## View the docs
### Development server (with live reload)
```bash
python -m http.server -d build/html/
mkdocs serve
```
Launch your browser and open localhost:8000.
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.
+248
View File
@@ -0,0 +1,248 @@
# 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
+27
View File
@@ -0,0 +1,27 @@
# 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
+41
View File
@@ -0,0 +1,41 @@
.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;
}
Binary file not shown.

After

Width:  |  Height:  |  Size: 98 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 122 KiB

Before

Width:  |  Height:  |  Size: 194 KiB

After

Width:  |  Height:  |  Size: 194 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 378 KiB

Before

Width:  |  Height:  |  Size: 303 KiB

After

Width:  |  Height:  |  Size: 303 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 575 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

+6
View File
@@ -0,0 +1,6 @@
<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>

After

Width:  |  Height:  |  Size: 691 B

+18
View File
@@ -0,0 +1,18 @@
<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>

After

Width:  |  Height:  |  Size: 5.7 KiB

@@ -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,11 +3,3 @@
# 🧰 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,4 +1,3 @@
(runpod)=
# 📦 Developing FastVideo on RunPod
@@ -10,7 +9,7 @@ Choose a GPU that supports CUDA 12.8
Pick 1 or 2 L40S GPU(s)
![RunPod CUDA selection](../../_static/images/runpod_cuda.png)
![RunPod CUDA selection](../../assets/images/runpod_cuda.png)
When creating your pod template, use this image:
@@ -24,11 +23,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"
```
![RunPod template configuration](../../_static/images/runpod_template.png)
![RunPod template configuration](../../assets/images/runpod_template.png)
After deploying, the pod will take a few minutes to pull the image and start the SSH service.
![RunPod ssh](../../_static/images/runpod_ssh.png)
![RunPod ssh](../../assets/images/runpod_ssh.png)
## Working with the pod
@@ -1,4 +1,3 @@
(developer-overview)=
# 🛠️ Contributing to FastVideo
@@ -29,7 +29,6 @@ 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.
@@ -61,7 +60,6 @@ with set_current_fastvideo_args(fastvideo_args):
result = generate_video()
```
(design-pipeline-system)=
## Pipeline System
### `ComposedPipelineBase`
@@ -108,7 +106,8 @@ def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> Forward
return batch
```
(design-forwardbatch)=
![Pipeline execution and data flow](../assets/images/pipeline.png)
### ForwardBatch
Defined in `fastvideo/pipelines/pipeline_batch_info.py`, `ForwardBatch` encapsulates the data payload passed between pipeline stages. It typically holds:
@@ -120,12 +119,10 @@ 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:
@@ -152,7 +149,6 @@ 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:
@@ -170,7 +166,6 @@ 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:
@@ -188,7 +183,6 @@ FastVideo implements optimizations such as:
- Caching for common prompts
- Precision-tuned computation
(design-schedulers)=
### Schedulers
Schedulers manage the diffusion sampling process:
@@ -216,7 +210,10 @@ def step(
return prev_sample
```
(design-optimized-attention)=
This diagram shows how models are discovered, validated, and loaded across entrypoints, executors, pipelines, and model loaders.
![Model loading flow](../assets/images/load_models.png)
## Optimized Attention
The `fastvideo/attention/` directory contains optimized attention implementations crucial for efficient video diffusion:
@@ -240,17 +237,17 @@ self.attn = LocalAttention(
# export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
```
![Attention backend selector design](../assets/images/attention_backend.png)
### Attention Patterns
Supports various patterns with memory optimization techniques:
- **Cross/Self/Temporal/Global-Local Attention**
- Chunking, progressive computation, optimized masking
(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:
@@ -307,7 +304,6 @@ Efficient communication primitives minimize distributed overhead:
- **Tensor-Parallel AllReduce**: Combines partial results
- **Distributed Synchronization**: Coordinates execution
(design-forwardcontext)=
## Forward Context Management
### ForwardContext
@@ -330,7 +326,6 @@ 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:
@@ -357,7 +352,6 @@ 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:
@@ -388,7 +382,6 @@ 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,4 +1,3 @@
(v0-data-preprocess)=
# 🧱 Data Preprocess for Distillation
+12
View File
@@ -0,0 +1,12 @@
# 💡 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)
@@ -6,10 +6,11 @@ import re
from dataclasses import dataclass, field
from pathlib import Path
ROOT_DIR = Path(__file__).parent.parent.parent.resolve()
ROOT_DIR_RELATIVE = '../../../..'
ROOT_DIR = Path(__file__).parent.parent.resolve()
ROOT_DIR_RELATIVE = '../..'
EXAMPLE_DIR = ROOT_DIR / "examples"
EXAMPLE_DOC_DIR = ROOT_DIR / "docs/source/getting_started/examples"
EXAMPLE_DOC_DIR = ROOT_DIR / "docs/getting_started/examples"
GITHUB_REPO = "hao-ai-lab/FastVideo" # Update this to your repo
def fix_case(text: str) -> str:
@@ -71,9 +72,16 @@ class Index:
def generate(self) -> str:
content = f"# {self.title}\n\n{self.description}\n\n"
content += ":::{toctree}\n"
content += f":caption: {self.caption}\n:maxdepth: {self.maxdepth}\n"
content += "\n".join(self.documents) + "\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"
return content
@@ -142,30 +150,66 @@ class Example:
return fix_case(self.path.stem.replace("_", " ").title())
def generate(self) -> str:
# Convert the path to a relative path from __file__
make_relative = lambda path: ROOT_DIR_RELATIVE / path.relative_to(
ROOT_DIR)
# 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"
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":
# Add title for code files
if self.main_file.suffix != ".md":
content += f"# {self.title}\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"
# 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"
if not self.other_files:
return content
content += "## Example materials\n\n"
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'
}
for file in sorted(self.other_files):
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"
# 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
return content
@@ -195,7 +239,7 @@ class NestedStructure:
def create_category_indices() -> dict[str, Index]:
"""Create category indices with their respective configurations."""
main_index_dir = ROOT_DIR / "docs/source/examples"
main_index_dir = ROOT_DIR / "docs/examples"
if not main_index_dir.exists():
main_index_dir.mkdir(parents=True)
@@ -203,17 +247,16 @@ def create_category_indices() -> dict[str, Index]:
"inference":
Index(
path=ROOT_DIR /
"docs/source/inference/examples/examples_inference_index.md",
"docs/inference/examples/examples_inference_index.md",
title="🚀 Examples",
description=
"Inference examples demonstrate how to use FastVideo inference. We recommend starting with <project:basic.md>.",
"Inference examples demonstrate how to use FastVideo inference. We recommend starting with [basic.md](basic.md).",
caption="Examples",
maxdepth=1,
),
"training":
Index(
path=ROOT_DIR /
"docs/source/training/examples/examples_training_index.md",
path=ROOT_DIR / "docs/training/examples/examples_training_index.md",
title="🚀 Examples",
description=
"Training examples demonstrate how to use FastVideo training.",
@@ -223,7 +266,7 @@ def create_category_indices() -> dict[str, Index]:
"distillation":
Index(
path=ROOT_DIR /
"docs/source/distillation/examples/examples_distillation_index.md",
"docs/distillation/examples/examples_distillation_index.md",
title="🚀 Examples",
description=
"Distillation examples demonstrate how to use FastVideo distillation.",
@@ -246,9 +289,21 @@ 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:
category_dir = EXAMPLE_DIR / category
# 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
globs = [category_dir.glob(pattern) for pattern in glob_patterns]
for path in itertools.chain(*globs):
examples.append(Example(path, category))
@@ -279,11 +334,18 @@ 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
category_dir = EXAMPLE_DIR / example.category
# Use mapped directory name if available
dir_name = category_dir_mapping.get(example.category, example.category)
category_dir = EXAMPLE_DIR / dir_name
relative_path = example.path.relative_to(category_dir)
path_parts = relative_path.parts
@@ -415,7 +477,7 @@ def generate_nested_examples(nested_structures: dict[str, dict[str, dict[
category_index.documents.append(method)
def generate_examples(generate_main_index=False):
def generate_examples(generate_main_index: bool = False) -> None:
"""
Generate example documentation.
@@ -429,12 +491,14 @@ def generate_examples(generate_main_index=False):
# Create the main examples index only if requested
examples_index = None
if generate_main_index:
main_index_dir = ROOT_DIR / "docs/source/examples"
main_index_dir = ROOT_DIR / "docs/examples"
examples_index = Index(
path=main_index_dir / "examples_index.md",
title="💡 Examples",
description=
"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>.",
"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.",
caption="Examples",
maxdepth=2)
@@ -471,3 +535,19 @@ def generate_examples(generate_main_index=False):
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!")
+41
View File
@@ -0,0 +1,41 @@
# 🔧 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
@@ -30,17 +30,8 @@ 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,17 +31,8 @@ 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
+55
View File
@@ -0,0 +1,55 @@
# 🚀 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
+42
View File
@@ -0,0 +1,42 @@
# 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
+20 -101
View File
@@ -1,30 +1,24 @@
# Welcome to FastVideo
:::{figure} ../../assets/logos/logo.svg
:align: center
:alt: FastVideo
:class: no-scaled-link
:width: 60%
:::
<div style="text-align: center;">
<img src="assets/logos/logo.svg" alt="FastVideo" style="width: 60%;" />
</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;">
<strong>FastVideo is a unified inference and post-training framework for accelerated video generation.</strong>
</div>
<p style="text-align:center">
<div 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>
</p>
:::
</div>
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=_static/images/fastwan.png width="100%"/>
<img src="assets/images/fastwan.png" style="width: 100%;"/>
</div>
## Key Features
@@ -42,91 +36,16 @@ FastVideo has the following features:
## Documentation
% How to start using FastVideo?
Welcome to FastVideo! This documentation will help you get started with our unified inference and post-training framework for accelerated video generation.
:::{toctree}
:caption: Getting Started
:maxdepth: 1
Use the navigation menu on the left to explore different sections:
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`
- **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
@@ -1,4 +1,3 @@
(add-pipeline)=
# 🏗️ Adding a New Pipeline
@@ -1,4 +1,4 @@
(inference-configuration)=
# Configuration
## Multi-GPU Setup
@@ -1,4 +1,3 @@
(inference-optimizations)=
# Optimizations
@@ -16,8 +15,6 @@ This page describes the various options for speeding up generation times in Fast
- Caching Techniques
- [TeaCache](#optimizations-teacache)
(optimizations-backends)=
## Attention Backends
### Available Backends
@@ -49,8 +46,6 @@ 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`**
@@ -71,12 +66,6 @@ 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`**
@@ -87,8 +76,6 @@ 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`**
@@ -100,8 +87,6 @@ python setup_vsa.py install
Please see [this page](#vsa-installation) for more installation instructions.
(optimizations-sage)=
### Sage Attention
**`SAGE_ATTN`**
@@ -114,8 +99,6 @@ cd sageattention
python setup.py install # or pip install -e .
```
(optimizations-sage3)=
### Sage Attention 3
**`SAGE_ATTN_THREE`**
@@ -136,8 +119,6 @@ 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.
+66
View File
@@ -0,0 +1,66 @@
# 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.
-35
View File
@@ -1,35 +0,0 @@
@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
-15
View File
@@ -1,15 +0,0 @@
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
@@ -1,4 +1,3 @@
(sta-demo)=
# 🔍 Demo
This is is a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
@@ -1,4 +1,3 @@
(sta-installation)=
# 🔧 Installation
You can install the Sliding Tile Attention package using
-8
View File
@@ -1,8 +0,0 @@
.vertical-table-header th.head:not(.stub) {
writing-mode: sideways-lr;
white-space: nowrap;
max-width: 0;
p {
margin: 0;
}
}
@@ -1,39 +0,0 @@
<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> -->
-19
View File
@@ -1,19 +0,0 @@
# Summary
## Video Generator
```{autodoc2-summary}
fastvideo.VideoGenerator
```
## Initialization Configuration
```{autodoc2-summary}
fastvideo.configs.pipelines.PipelineConfig
```
## Sampling Configuration
```{autodoc2-summary}
fastvideo.configs.sample.SamplingParam
```
-22
View File
@@ -1,22 +0,0 @@
# 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
-275
View File
@@ -1,275 +0,0 @@
# 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,18 +0,0 @@
(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
-83
View File
@@ -1,83 +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.
```{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
```
-136
View File
@@ -1,136 +0,0 @@
(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,4 +1,3 @@
(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,4 +1,3 @@
(vsa-installation)=
# 🔧 Installation
You can install the Video Sparse Attention package using
@@ -0,0 +1,44 @@
# 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 = 73
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,8 +3,13 @@ import os
import requests
import base64
import time
import json
from pathlib import Path
import tempfile
from io import BytesIO
import gradio as gr
from PIL import Image
from fastvideo.configs.sample.base import SamplingParam
@@ -12,6 +17,7 @@ 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",
}
@@ -37,7 +43,7 @@ class RayServeClient:
f"{self.backend_url}/generate_video",
json=request_data,
headers=headers,
timeout=300
timeout=900 # 15 minutes timeout for longer video generation
)
round_trip_time = time.time() - start_time
@@ -81,49 +87,89 @@ def save_video_from_base64(video_data: str, output_dir: str, prompt: str) -> str
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"
def encode_image_to_base64(image_input) -> str:
"""Encode an image file path or in-memory image to a base64 string."""
if image_input is None:
return None
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>"""
mime_types = {
'.jpg': 'image/jpeg',
'.jpeg': 'image/jpeg',
'.png': 'image/png',
'.gif': 'image/gif',
'.webp': 'image/webp',
}
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>"
try:
if isinstance(image_input, str):
if not os.path.exists(image_input):
return None
with open(image_input, 'rb') as f:
image_bytes = f.read()
ext = os.path.splitext(image_input)[1].lower()
mime_type = mime_types.get(ext, 'image/jpeg')
elif isinstance(image_input, Image.Image):
buffer = BytesIO()
image_to_save = image_input.convert("RGB")
image_to_save.save(buffer, format="PNG")
image_bytes = buffer.getvalue()
mime_type = 'image/png'
else:
return None
image_base64 = base64.b64encode(image_bytes).decode('utf-8')
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>"
def load_example_prompts():
@@ -144,26 +190,83 @@ def load_example_prompts():
print(f"Warning: Could not read {filepath}: {e}")
return prompts, labels
examples, example_labels = load_from_file("prompts/prompts_final.txt")
# Load prompts from prompts.txt
examples, example_labels = load_from_file("examples/inference/gradio/serving/prompts.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"]
return examples, example_labels
# 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
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, seed, guidance_scale,
num_frames, height, width, randomize_seed, model_selection, progress
prompt, negative_prompt, use_negative_prompt, guidance_scale,
num_frames, height, width, model_selection, input_image, 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:
@@ -172,6 +275,15 @@ 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,
@@ -183,7 +295,7 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
"width": width,
"randomize_seed": randomize_seed,
"return_frames": False,
"image_path": None,
"image_data": image_data,
"model_path": MODEL_PATH_MAPPING.get(model_selection, "FastVideo/FastWan2.1-T2V-1.3B-Diffusers")
}
@@ -198,16 +310,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:
@@ -219,7 +331,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, timing_details
return video_path, used_seed, ""
else:
return None, "Failed to save video", ""
else:
@@ -228,7 +340,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 = load_example_prompts()
examples, example_labels, example_images = load_example_prompts()
theme = gr.themes.Base().set(
button_primary_background_fill="#2563eb",
@@ -239,33 +351,39 @@ 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,
'seed': params.seed,
}
# 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,
# }
return {
'height': 448,
'height': 480,
'width': 832,
'num_frames': 61,
'guidance_scale': 3.0,
'seed': 1024,
'num_frames': 73,
}
initial_values = get_default_values("FastWan2.1-T2V-1.3B")
# 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)
with gr.Blocks(title="FastWan", theme=theme) as demo:
# 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:
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_post_training/" 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_causalwan_preview/" target="_blank">Blog</a> | <a href="https://hao-ai-lab.github.io/FastVideo/" target="_blank">Docs</a> </p>
</div>
""")
@@ -280,8 +398,8 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
with gr.Row():
model_selection = gr.Dropdown(
choices=list(MODEL_PATH_MAPPING.keys()),
value="FastWan2.1-T2V-1.3B",
choices=available_models,
value=default_model,
label="Select Model",
interactive=True
)
@@ -312,69 +430,70 @@ 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=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.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="pil",
height=400,
)
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,
)
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,
)
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")
# randomize_seed = gr.Checkbox(label="Randomize seed", value=False)
seed_output = gr.Number(label="Used Seed", value=1000)
with gr.Column(scale=1, elem_classes="video-column"):
with gr.Column(scale=1):
result = gr.Video(
label="Generated Video",
show_label=True,
height=466,
width=600,
height=500,
container=True,
elem_classes="video-component"
autoplay=True,
)
gr.HTML("""
@@ -387,116 +506,10 @@ def create_gradio_interface(backend_url: str, default_params: dict[str, Sampling
}
.gradio-container {
max-width: 1200px !important;
max-width: 1400px !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,18 +524,30 @@ 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)
return examples[index]
return ""
selected_prompt = examples[index]
selected_image_path = example_images[index] if index < len(example_images) else None
if selected_image_path and os.path.exists(selected_image_path):
try:
with Image.open(selected_image_path) as img:
selected_image = img.convert("RGB")
except Exception:
selected_image = None
else:
selected_image = None
return selected_prompt, selected_image
return "", None
example_dropdown.change(
fn=on_example_select,
inputs=example_dropdown,
outputs=prompt,
outputs=[prompt, input_image],
)
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 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>
<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>
</div>
""")
@@ -537,6 +562,7 @@ 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]
@@ -545,29 +571,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(value=params.seed),
gr.update(visible=show_image_input),
)
return (
gr.update(value=448),
gr.update(value=832),
gr.update(value=61),
gr.update(value=20),
gr.update(value=3.0),
gr.update(value=1024),
gr.update(visible=show_image_input),
)
model_selection.change(
fn=on_model_selection_change,
inputs=model_selection,
outputs=[height, width, num_frames, guidance_scale, seed],
outputs=[height, width, num_frames, guidance_scale, image_tab],
)
def handle_generation(*args, progress=None, request: gr.Request = None):
model_selection, prompt, negative_prompt, use_negative_prompt, seed, guidance_scale, num_frames, height, width, randomize_seed = args
model_selection, prompt, negative_prompt, use_negative_prompt, guidance_scale, num_frames, height, width, input_image = args
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
result_path, seed_or_error, _ = generate_video(
prompt, negative_prompt, use_negative_prompt, guidance_scale,
num_frames, height, width, model_selection, input_image, progress
)
if result_path and os.path.exists(result_path):
@@ -575,14 +601,12 @@ 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(
@@ -592,14 +616,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,
# randomize_seed,
input_image,
],
outputs=[result, seed_output, error_output, timing_display],
outputs=[result, seed_output, error_output], # timing_display removed
concurrency_limit=20,
)
@@ -611,8 +635,11 @@ 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="FastVideo/FastWan2.1-T2V-1.3B-Diffusers,FastVideo/FastWan2.1-T2V-14B-Diffusers",
default="",
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,
@@ -621,8 +648,15 @@ def main():
args = parser.parse_args()
default_params = {}
model_paths = args.t2v_model_paths.split(",")
for model_path in model_paths:
# 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:
default_params[model_path] = SamplingParam.from_pretrained(model_path)
demo = create_gradio_interface(args.backend_url, default_params)
@@ -630,6 +664,8 @@ 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
@@ -674,23 +710,23 @@ def main():
<meta charset="UTF-8" />
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
<title>FastWan</title>
<meta name="title" content="FastWan">
<title>CausalWan</title>
<meta name="title" content="CausalWan">
<meta name="description" content="Make video generation go blurrrrrrr">
<meta name="keywords" content="FastVideo, video generation, AI, machine learning, FastWan">
<meta name="keywords" content="FastVideo, video generation, AI, machine learning, CausalWan">
<meta property="og:type" content="website">
<meta property="og:url" content="{base_url}/">
<meta property="og:title" content="FastWan">
<meta property="og:title" content="CausalWan">
<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="FastWan">
<meta property="og:site_name" content="CausalWan">
<meta property="twitter:card" content="summary_large_image">
<meta property="twitter:url" content="{base_url}/">
<meta property="twitter:title" content="FastWan">
<meta property="twitter:title" content="CausalWan">
<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">
@@ -720,7 +756,15 @@ def main():
app,
demo,
path="/gradio",
allowed_paths=[os.path.abspath("outputs"), os.path.abspath("fastvideo-logos")]
root_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")),
]
)
uvicorn.run(app, host=args.host, port=args.port)
@@ -0,0 +1,15 @@
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,6 +26,7 @@ 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 = {
@@ -42,6 +43,13 @@ 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,
}
}
@@ -58,6 +66,7 @@ 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):
@@ -91,11 +100,38 @@ 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"
# 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"
os.environ["FASTVIDEO_STAGE_LOGGING"] = "1"
@@ -157,22 +193,41 @@ 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,
)
@@ -185,6 +240,13 @@ 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,
@@ -200,7 +262,7 @@ class BaseModelDeployment:
@serve.deployment(
ray_actor_options={"num_cpus": 2, "num_gpus": 1, "runtime_env": {"conda": "fv"}},
ray_actor_options={"num_cpus": 2, "num_gpus": 1, "runtime_env": {"conda": "demo-fv"}},
)
class T2VModelDeployment(BaseModelDeployment):
def __init__(self, t2v_model_path: str, output_path: str = "outputs"):
@@ -210,7 +272,7 @@ class T2VModelDeployment(BaseModelDeployment):
@serve.deployment(
ray_actor_options={"num_cpus": 16, "num_gpus": 1, "runtime_env": {"conda": "fv"}},
ray_actor_options={"num_cpus": 16, "num_gpus": 1, "runtime_env": {"conda": "demo-fv"}},
)
class T2V14BModelDeployment(BaseModelDeployment):
def __init__(self, t2v_14b_model_path: str, output_path: str = "outputs"):
@@ -221,18 +283,32 @@ 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=50, ray_actor_options={"num_cpus": 2})
@serve.deployment(num_replicas=1, ray_actor_options={"num_cpus": 1})
@serve.ingress(app)
class FastVideoAPI:
def __init__(self, t2v_deployments: Dict[str, DeploymentHandle]):
def __init__(self, t2v_deployments: Dict[str, DeploymentHandle], i2v_deployments: Dict[str, DeploymentHandle] = None):
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'])
@@ -257,10 +333,10 @@ class FastVideoAPI:
model_name = self._get_model_name(video_request.model_path)
try:
if video_request.model_path not in self.t2v_deployments:
if video_request.model_path not in self.all_deployments:
raise ValueError(f"Model {video_request.model_path} not found")
response_ref = self.t2v_deployments[video_request.model_path].generate_video.remote(video_request)
response_ref = self.all_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)
@@ -291,18 +367,21 @@ 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"
assert model in SUPPORTED_MODELS, f"Model {model} not supported. Supported models: {SUPPORTED_MODELS}"
assert replica_count > 0, f"Replicas must be greater than 0"
def start_ray_serve(
*,
t2v_model_paths: str,
t2v_model_replicas: str,
t2v_model_paths: str = "",
t2v_model_replicas: str = "",
i2v_model_paths: str = "",
i2v_model_replicas: str = "",
output_path: str = "outputs",
host: str = "0.0.0.0",
port: int = 8000,
@@ -310,21 +389,39 @@ def start_ray_serve(
if not ray.is_initialized():
ray.init()
model_paths = t2v_model_paths.split(",")
replicas = [int(r) for r in t2v_model_replicas.split(",")]
validate_configuration(model_paths, replicas)
# 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)
# Create T2V deployments
t2v_deps = {}
for model_path, replica_count in zip(model_paths, replicas):
for model_path, replica_count in zip(t2v_paths, t2v_reps):
t2v_dep = T2VModelDeployment.options(num_replicas=replica_count).bind(model_path, output_path)
t2v_deps[model_path] = t2v_dep
api = FastVideoAPI.bind(t2v_deps)
# 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)
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(model_paths, replicas):
for model_path, replica_count in zip(t2v_paths, t2v_reps):
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")
@@ -340,12 +437,20 @@ if __name__ == "__main__":
parser = argparse.ArgumentParser(description="FastVideo Ray Serve Backend")
parser.add_argument("--t2v_model_paths",
type=str,
default="FastVideo/FastWan2.1-T2V-1.3B-Diffusers,FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers",
default="",
help="Comma separated list of paths to the T2V model(s)")
parser.add_argument("--t2v_model_replicas",
type=str,
default="4,4",
default="",
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",
@@ -361,13 +466,21 @@ if __name__ == "__main__":
args = parser.parse_args()
model_paths = args.t2v_model_paths.split(",")
replicas = [int(r) for r in args.t2v_model_replicas.split(",")]
validate_configuration(model_paths, replicas)
# 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)
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,
@@ -376,4 +489,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)
+5 -3
View File
@@ -1,3 +1,5 @@
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"
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"
@@ -20,8 +20,10 @@ DEFAULT_BACKEND_PORT = 8000
DEFAULT_FRONTEND_HOST = "0.0.0.0"
DEFAULT_FRONTEND_PORT = 7860
DEFAULT_OUTPUT_PATH = "outputs"
DEFAULT_T2V_MODELS = "FastVideo/FastWan2.1-T2V-1.3B-Diffusers,FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers"
DEFAULT_T2V_REPLICAS = "4,4"
DEFAULT_T2V_MODELS = ""
DEFAULT_T2V_REPLICAS = ""
DEFAULT_I2V_MODELS = "FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers"
DEFAULT_I2V_REPLICAS = "1"
HEALTH_CHECK_TIMEOUT = 5
HEALTH_CHECK_MAX_RETRIES = 100
@@ -100,6 +102,12 @@ 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
@@ -111,6 +119,10 @@ 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
@@ -173,6 +185,9 @@ 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}")
@@ -190,6 +205,14 @@ 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 -1
View File
@@ -1,5 +1,9 @@
from fastvideo.configs.models.dits.cosmos import CosmosVideoConfig
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"]
__all__ = [
"HunyuanVideoConfig", "WanVideoConfig", "StepVideoConfig",
"CosmosVideoConfig"
]
+104
View File
@@ -0,0 +1,104 @@
# 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"
@@ -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
from fastvideo.configs.models.encoders.t5 import T5Config, T5LargeConfig
__all__ = [
"EncoderConfig", "TextEncoderConfig", "ImageEncoderConfig",
"BaseEncoderOutput", "CLIPTextConfig", "CLIPVisionConfig",
"WAN2_1ControlCLIPVisionConfig", "LlamaConfig", "T5Config"
"WAN2_1ControlCLIPVisionConfig", "LlamaConfig", "T5Config", "T5LargeConfig"
]
+23
View File
@@ -70,8 +70,31 @@ 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,3 +1,4 @@
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
@@ -6,4 +7,5 @@ __all__ = [
"HunyuanVAEConfig",
"WanVAEConfig",
"StepVideoVAEConfig",
"CosmosVAEConfig",
]
@@ -0,0 +1,87 @@
# 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
+3 -1
View File
@@ -1,5 +1,6 @@
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)
@@ -12,5 +13,6 @@ __all__ = [
"HunyuanConfig", "FastHunyuanConfig", "PipelineConfig",
"SlidingTileAttnConfig", "WanT2V480PConfig", "WanI2V480PConfig",
"WanT2V720PConfig", "WanI2V720PConfig", "StepVideoT2VConfig",
"SelfForcingWanT2V480PConfig", "get_pipeline_config_cls_from_name"
"SelfForcingWanT2V480PConfig", "CosmosConfig",
"get_pipeline_config_cls_from_name"
]
+73
View File
@@ -0,0 +1,73 @@
# 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
View File
@@ -5,6 +5,7 @@ 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
@@ -38,9 +39,12 @@ 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
}
@@ -52,6 +56,7 @@ 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
}
+4
View File
@@ -186,3 +186,7 @@ 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
+18
View File
@@ -0,0 +1,18 @@
# 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
+11
View File
@@ -7,6 +7,8 @@ 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,
@@ -72,8 +74,17 @@ 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
}
-2
View File
@@ -191,8 +191,6 @@ 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
+1 -1
View File
@@ -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"
+195
View File
@@ -0,0 +1,195 @@
# 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
+53
View File
@@ -47,6 +47,59 @@ 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,
+76
View File
@@ -177,3 +177,79 @@ 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
+12 -70
View File
@@ -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, get_local_torch_device
from fastvideo.distributed.parallel_state import get_sp_world_size
from fastvideo.forward_context import get_forward_context
from fastvideo.layers.layernorm import (FP32LayerNorm, LayerNormScaleShift,
RMSNorm, ScaleResidual,
@@ -33,29 +33,9 @@ from fastvideo.layers.visual_embedding import (PatchEmbed)
from fastvideo.logger import init_logger
from fastvideo.models.dits.base import BaseDiT
from fastvideo.models.dits.wanvideo import WanT2VCrossAttention, WanTimeTextImageEmbedding
from fastvideo.platforms import AttentionBackendEnum, current_platform
from fastvideo.platforms import AttentionBackendEnum
logger = init_logger(__name__)
class 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,
@@ -87,10 +67,6 @@ 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,
@@ -108,19 +84,6 @@ 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
@@ -165,8 +128,6 @@ 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
@@ -176,44 +137,26 @@ 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()
# 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)
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
x = self.attn(
roped_query,
local_k,
local_v
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]
)
# 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)
@@ -290,9 +233,6 @@ 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,
@@ -343,8 +283,9 @@ 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, self.null_shift, self.null_scale)
hidden_states, attn_output, gate_msa, null_shift, null_scale)
# 2. Cross-attention
attn_output = self.attn2(norm_hidden_states,
@@ -511,6 +452,7 @@ class CausalWanTransformer3DModel(BaseDiT):
This function will be run for num_frame times.
Process the latent frames one by one (1560 tokens each)
"""
from fastvideo.platforms import current_platform
orig_dtype = hidden_states.dtype
if not isinstance(encoder_hidden_states, torch.Tensor):
+726
View File
@@ -0,0 +1,726 @@
# 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
+7 -4
View File
@@ -180,7 +180,7 @@ class T5Attention(nn.Module):
self.qkv_proj = QKVParallelLinear(
self.d_model,
self.d_model // self.total_num_heads,
self.key_value_proj_dim,
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.d_model,
self.total_num_heads * self.key_value_proj_dim,
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.d_model // self.total_num_heads
n, c = self.n_heads, self.key_value_proj_dim
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,7 +540,10 @@ class T5EncoderModel(TextEncoder):
attn_metadata=attn_metadata,
)
return BaseEncoderOutput(last_hidden_state=hidden_states)
return BaseEncoderOutput(
last_hidden_state=hidden_states,
attention_mask=attention_mask,
)
def load_weights(self, weights: Iterable[tuple[str,
torch.Tensor]]) -> set[str]:
+39 -18
View File
@@ -4,6 +4,8 @@
# 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
@@ -197,10 +199,11 @@ def shard_model(
Raises:
ValueError: If no layer modules were sharded, indicating that no shard_condition was triggered.
"""
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__)
# 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.")
return
fsdp_kwargs = {
@@ -215,20 +218,38 @@ def shard_model(
# iterating in reverse to start with
# lowest-level modules first
num_layers_sharded = 0
# 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."
)
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."
)
# Finally shard the entire model to account for any stragglers
fully_shard(model, **fsdp_kwargs)
+3 -2
View File
@@ -26,7 +26,8 @@ _TEXT_TO_VIDEO_DIT_MODELS = {
("dits", "hunyuanvideo", "HunyuanVideoTransformer3DModel"),
"WanTransformer3DModel": ("dits", "wanvideo", "WanTransformer3DModel"),
"CausalWanTransformer3DModel": ("dits", "causal_wanvideo", "CausalWanTransformer3DModel"),
"StepVideoModel": ("dits", "stepvideo", "StepVideoModel")
"StepVideoModel": ("dits", "stepvideo", "StepVideoModel"),
"CosmosTransformer3DModel": ("dits", "cosmos", "CosmosTransformer3DModel")
}
_IMAGE_TO_VIDEO_DIT_MODELS = {
@@ -39,6 +40,7 @@ _TEXT_ENCODER_MODELS = {
"CLIPTextModel": ("encoders", "clip", "CLIPTextModel"),
"LlamaModel": ("encoders", "llama", "LlamaModel"),
"UMT5EncoderModel": ("encoders", "t5", "UMT5EncoderModel"),
"T5EncoderModel": ("encoders", "t5", "T5EncoderModel"),
"STEP1TextEncoder": ("encoders", "stepllm", "STEP1TextEncoder"),
"BertModel": ("encoders", "clip", "CLIPTextModel"),
}
@@ -239,7 +241,6 @@ class _ModelRegistry:
def _raise_for_unsupported(self, architectures: list[str]) -> NoReturn:
all_supported_archs = self.get_supported_archs()
if any(arch in all_supported_archs for arch in architectures):
raise ValueError(
f"Model architectures {architectures} failed "
@@ -88,6 +88,14 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
The type of dynamic resolution-dependent timestep shifting to apply. Either "exponential" or "linear".
stochastic_sampling (`bool`, defaults to False):
Whether to use stochastic sampling.
final_sigmas_type (`str`, defaults to "sigma_min"):
The type of final sigmas to use. Either "sigma_min" or "zero".
sigma_max (`float`, *optional*):
The maximum sigma value for the noise schedule.
sigma_min (`float`, *optional*):
The minimum sigma value for the noise schedule.
sigma_data (`float`, *optional*):
The sigma data value for scaling.
"""
_compatibles: list[Any] = []
@@ -110,6 +118,10 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
use_beta_sigmas: bool | None = False,
time_shift_type: str = "exponential",
stochastic_sampling: bool = False,
final_sigmas_type: str = "sigma_min",
sigma_max: float | None = None,
sigma_min: float | None = None,
sigma_data: float | None = None,
):
if sum([
self.config.use_beta_sigmas, self.config.use_exponential_sigmas,
@@ -336,9 +348,9 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
sigmas_array: np.ndarray
if sigmas is None:
if timesteps_array is None:
timesteps_array = np.linspace(self._sigma_to_t(self.sigma_max),
self._sigma_to_t(self.sigma_min),
num_inference_steps)
t_max = self._sigma_to_t(self.sigma_max)
t_min = self._sigma_to_t(self.sigma_min)
timesteps_array = np.linspace(t_max, t_min, num_inference_steps)
sigmas_array = timesteps_array / self.config.num_train_timesteps
else:
sigmas_array = np.array(sigmas).astype(np.float32)
@@ -403,9 +415,7 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
[sigmas_tensor,
torch.ones(1, device=sigmas_tensor.device)])
else:
sigmas_tensor = torch.cat(
[sigmas_tensor,
torch.zeros(1, device=sigmas_tensor.device)])
sigmas_tensor = torch.cat([sigmas_tensor, torch.zeros(1, device=sigmas_tensor.device)])
self.timesteps = timesteps_tensor
self.sigmas = sigmas_tensor
@@ -505,7 +515,9 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
next_sigma = lower_sigmas[..., None]
dt = current_sigma - next_sigma
else:
assert self.step_index is not None, "step_index should not be None"
if self.step_index is None:
self._init_step_index(timestep)
sigma_idx = self.step_index
sigma = self.sigmas[sigma_idx]
sigma_next = self.sigmas[sigma_idx + 1]
@@ -522,7 +534,6 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
prev_sample = sample + dt * model_output
# upon completion increase step index by one
assert self._step_index is not None, "_step_index should not be None"
self._step_index += 1
if per_token_timesteps is None:
# Cast sample back to model compatible dtype
@@ -558,7 +569,7 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
min_inv_rho = sigma_min**(1 / rho)
max_inv_rho = sigma_max**(1 / rho)
sigmas = (max_inv_rho + ramp * (min_inv_rho - max_inv_rho))**rho
return sigmas
return torch.from_numpy(sigmas).to(dtype=in_sigmas.dtype, device=in_sigmas.device)
# Copied from diffusers.schedulers.scheduling_euler_discrete.EulerDiscreteScheduler._convert_to_exponential
def _convert_to_exponential(self, in_sigmas: torch.Tensor,
@@ -583,7 +594,7 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
sigmas = np.exp(
np.linspace(math.log(sigma_max), math.log(sigma_min),
num_inference_steps))
return sigmas
return torch.from_numpy(sigmas).to(dtype=in_sigmas.dtype, device=in_sigmas.device)
# Copied from diffusers.schedulers.scheduling_euler_discrete.EulerDiscreteScheduler._convert_to_beta
def _convert_to_beta(self,
@@ -614,7 +625,7 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
for timestep in 1 - np.linspace(0, 1, num_inference_steps)
]
])
return sigmas
return torch.from_numpy(sigmas).to(dtype=in_sigmas.dtype, device=in_sigmas.device)
def _time_shift_exponential(
self, mu: float, sigma: float,
@@ -86,6 +86,12 @@ class SelfForcingFlowMatchScheduler(BaseScheduler, ConfigMixin, SchedulerMixin):
return (prev_sample, )
return SelfForcingFlowMatchSchedulerOutput(prev_sample=prev_sample)
@staticmethod
def calculate_alpha_beta_high(sigma, sigma_bound):
alpha = (1 - sigma) / (1 - sigma_bound)
beta = torch.sqrt(sigma ** 2 - (alpha * sigma_bound) ** 2)
return alpha, beta
def add_noise(self, original_samples, noise, timestep):
"""
Diffusion forward corruption process.
@@ -105,6 +111,32 @@ class SelfForcingFlowMatchScheduler(BaseScheduler, ConfigMixin, SchedulerMixin):
sample = (1 - sigma) * original_samples + sigma * noise
return sample.type_as(noise)
def add_noise_high(self, original_samples, noise, timestep, boundary_timestep):
"""
Diffusion forward corruption process.
Input:
- clean_latent: the clean latent with shape [B*T, C, H, W]
- noise: the noise with shape [B*T, C, H, W]
- timestep: the timestep with shape [B*T]
Output: the corrupted latent with shape [B*T, C, H, W]
"""
if timestep.ndim == 2:
timestep = timestep.flatten(0, 1)
if boundary_timestep.ndim == 2:
boundary_timestep = boundary_timestep.flatten(0, 1)
self.sigmas = self.sigmas.to(noise.device)
self.timesteps = self.timesteps.to(noise.device)
timestep_id = torch.argmin(
(self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1)
boundary_timestep_id = torch.argmin(
(self.timesteps.unsqueeze(0) - boundary_timestep.unsqueeze(1)).abs(), dim=1)
sigma_boundary = self.sigmas[boundary_timestep_id].reshape(-1, 1, 1, 1)
alpha, beta = self.calculate_alpha_beta_high(sigma, sigma_boundary)
sample = alpha * original_samples + beta * noise
return sample.type_as(noise)
def training_target(self, sample, noise, timestep):
target = noise - sample
return target
+48
View File
@@ -180,3 +180,51 @@ def pred_noise_to_pred_video(pred_noise: torch.Tensor,
sigma_t = sigmas[timestep_id].reshape(-1, 1, 1, 1)
pred_video = noise_input_latent - sigma_t * pred_noise
return pred_video.to(dtype)
def pred_noise_to_x_bound(pred_noise: torch.Tensor,
noise_input_latent: torch.Tensor,
timestep: torch.Tensor,
boundary_timestep: torch.Tensor,
scheduler: Any) -> torch.Tensor:
"""
Convert predicted noise to clean latent.
Args:
pred_noise: the predicted noise with shape [B, C, H, W]
where B is batch_size or batch_size * num_frames
noise_input_latent: the noisy latent with shape [B, C, H, W],
timestep: the timestep with shape [1] or [bs * num_frames] or [bs, num_frames]
boundary_timestep: the boundary timestep with shape [1] or [bs * num_frames] or [bs, num_frames]
scheduler: the scheduler
Returns:
the predicted video with shape [B, C, H, W]
"""
# If timestep is [bs, num_frames]
if timestep.ndim == 2:
timestep = timestep.flatten(0, 1)
assert timestep.numel() == noise_input_latent.shape[0]
elif timestep.ndim == 1:
# If timestep is [1]
if timestep.shape[0] == 1:
timestep = timestep.expand(noise_input_latent.shape[0])
else:
assert timestep.numel() == noise_input_latent.shape[0]
else:
raise ValueError(f"[pred_noise_to_pred_video] Invalid timestep shape: {timestep.shape}")
# timestep shape should be [B]
dtype = pred_noise.dtype
device = pred_noise.device
pred_noise = pred_noise.double().to(device)
noise_input_latent = noise_input_latent.double().to(device)
sigmas = scheduler.sigmas.double().to(device)
timesteps = scheduler.timesteps.double().to(device)
timestep_id = torch.argmin(
(timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
sigma_t = sigmas[timestep_id].reshape(-1, 1, 1, 1)
boundary_timestep_id = torch.argmin(
(timesteps.unsqueeze(0) - boundary_timestep.unsqueeze(1)).abs(), dim=1)
sigma_t_boundary = sigmas[boundary_timestep_id].reshape(-1, 1, 1, 1)
pred_video = noise_input_latent - (sigma_t - sigma_t_boundary) * pred_noise
return pred_video.to(dtype)
@@ -0,0 +1,84 @@
# SPDX-License-Identifier: Apache-2.0
"""
Cosmos video diffusion pipeline implementation.
This module contains an implementation of the Cosmos video diffusion pipeline
using the modular pipeline architecture.
"""
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler)
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.pipelines.stages import (ConditioningStage, CosmosDenoisingStage,
CosmosLatentPreparationStage,
DecodingStage, InputValidationStage,
TextEncodingStage,
TimestepPreparationStage)
logger = init_logger(__name__)
class Cosmos2VideoToWorldPipeline(ComposedPipelineBase):
_required_config_modules = [
"text_encoder", "tokenizer", "vae", "transformer", "scheduler",
"safety_checker"
]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
self.modules["scheduler"] = FlowMatchEulerDiscreteScheduler(
shift=fastvideo_args.pipeline_config.flow_shift,
use_karras_sigmas=True)
sigma_max = 80.0
sigma_min = 0.002
sigma_data = 1.0
final_sigmas_type = "sigma_min"
if self.modules["scheduler"] is not None:
scheduler = self.modules["scheduler"]
scheduler.config.sigma_max = sigma_max
scheduler.config.sigma_min = sigma_min
scheduler.config.sigma_data = sigma_data
scheduler.config.final_sigmas_type = final_sigmas_type
scheduler.sigma_max = sigma_max
scheduler.sigma_min = sigma_min
scheduler.sigma_data = sigma_data
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
self.add_stage(stage_name="input_validation_stage",
stage=InputValidationStage())
self.add_stage(stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
))
self.add_stage(stage_name="conditioning_stage",
stage=ConditioningStage())
self.add_stage(stage_name="timestep_preparation_stage",
stage=TimestepPreparationStage(
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="latent_preparation_stage",
stage=CosmosLatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer"),
vae=self.get_module("vae")))
self.add_stage(stage_name="denoising_stage",
stage=CosmosDenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
EntryClass = Cosmos2VideoToWorldPipeline
@@ -50,7 +50,8 @@ class WanCausalDMDPipeline(LoRAPipeline, ComposedPipelineBase):
stage=CausalDMDDenosingStage(
transformer=self.get_module("transformer"),
transformer_2=self.get_module("transformer_2", None),
scheduler=self.get_module("scheduler")))
scheduler=self.get_module("scheduler"),
vae=self.get_module("vae")))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
+1
View File
@@ -25,6 +25,7 @@ _PIPELINE_NAME_TO_ARCHITECTURE_NAME: dict[str, str] = {
"WanCausalDMDPipeline": "wan",
"StepVideoPipeline": "stepvideo",
"HunyuanVideoPipeline": "hunyuan",
"Cosmos2VideoToWorldPipeline": "cosmos"
}
_PREPROCESS_WORKLOAD_TYPE_TO_PIPELINE_NAME: dict[WorkloadType, str] = {
+6 -2
View File
@@ -10,7 +10,8 @@ from fastvideo.pipelines.stages.base import PipelineStage
from fastvideo.pipelines.stages.causal_denoising import CausalDMDDenosingStage
from fastvideo.pipelines.stages.conditioning import ConditioningStage
from fastvideo.pipelines.stages.decoding import DecodingStage
from fastvideo.pipelines.stages.denoising import (DenoisingStage,
from fastvideo.pipelines.stages.denoising import (CosmosDenoisingStage,
DenoisingStage,
DmdDenoisingStage)
from fastvideo.pipelines.stages.encoding import EncodingStage
from fastvideo.pipelines.stages.image_encoding import (ImageEncodingStage,
@@ -18,7 +19,8 @@ from fastvideo.pipelines.stages.image_encoding import (ImageEncodingStage,
ImageVAEEncodingStage,
VideoVAEEncodingStage)
from fastvideo.pipelines.stages.input_validation import InputValidationStage
from fastvideo.pipelines.stages.latent_preparation import LatentPreparationStage
from fastvideo.pipelines.stages.latent_preparation import (
CosmosLatentPreparationStage, LatentPreparationStage)
from fastvideo.pipelines.stages.stepvideo_encoding import (
StepvideoPromptEncodingStage)
from fastvideo.pipelines.stages.text_encoding import TextEncodingStage
@@ -30,10 +32,12 @@ __all__ = [
"InputValidationStage",
"TimestepPreparationStage",
"LatentPreparationStage",
"CosmosLatentPreparationStage",
"ConditioningStage",
"DenoisingStage",
"DmdDenoisingStage",
"CausalDMDDenosingStage",
"CosmosDenoisingStage",
"EncodingStage",
"DecodingStage",
"ImageEncodingStage",
+164 -150
View File
@@ -4,7 +4,7 @@ from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
from fastvideo.models.utils import pred_noise_to_pred_video
from fastvideo.models.utils import pred_noise_to_pred_video, pred_noise_to_x_bound
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.denoising import DenoisingStage
from fastvideo.pipelines.stages.validators import StageValidators as V
@@ -34,13 +34,16 @@ class CausalDMDDenosingStage(DenoisingStage):
Denoising stage for causal diffusion.
"""
def __init__(self, transformer, scheduler, transformer_2=None) -> None:
def __init__(self,
transformer,
scheduler,
transformer_2=None,
vae=None) -> None:
super().__init__(transformer, scheduler, transformer_2)
# KV and cross-attention cache state (initialized on first forward)
self.transformer = transformer
self.transformer_2 = transformer_2
self.kv_cache1: list | None = None
self.crossattn_cache: list | None = None
self.vae = vae
# Model-dependent constants (aligned with causal_inference.py assumptions)
self.num_transformer_blocks = len(self.transformer.blocks)
self.num_frames_per_block = self.transformer.config.arch_config.num_frames_per_block
@@ -80,6 +83,13 @@ class CausalDMDDenosingStage(DenoisingStage):
timesteps = scheduler_timesteps[1000 - timesteps]
timesteps = timesteps.to(get_local_torch_device())
if fastvideo_args.pipeline_config.dit_config.boundary_ratio is not None:
boundary_timestep = fastvideo_args.pipeline_config.dit_config.boundary_ratio * self.scheduler.num_train_timesteps
else:
boundary_timestep = None
high_noise_timesteps = timesteps[timesteps >= boundary_timestep]
# Image kwargs (kept empty unless caller provides compatible args)
image_kwargs: dict = {}
@@ -103,127 +113,107 @@ class CausalDMDDenosingStage(DenoisingStage):
assert torch.isnan(prompt_embeds[0]).sum() == 0
# Initialize or reset caches
if self.kv_cache1 is None:
self._initialize_kv_cache(batch_size=latents.shape[0],
dtype=target_dtype,
device=latents.device)
self._initialize_crossattn_cache(
batch_size=latents.shape[0],
max_text_len=fastvideo_args.pipeline_config.
text_encoder_configs[0].arch_config.text_len,
dtype=target_dtype,
device=latents.device)
else:
assert self.crossattn_cache is not None
# reset cross-attention cache
for block_index in range(self.num_transformer_blocks):
self.crossattn_cache[block_index][
"is_init"] = False # type: ignore
# reset kv cache pointers
for block_index in range(len(self.kv_cache1)):
self.kv_cache1[block_index][
"global_end_index"] = torch.tensor( # type: ignore
[0],
dtype=torch.long,
device=latents.device)
self.kv_cache1[block_index][
"local_end_index"] = torch.tensor( # type: ignore
[0],
dtype=torch.long,
device=latents.device)
kv_cache1 = self._initialize_kv_cache(batch_size=latents.shape[0],
dtype=target_dtype,
device=latents.device)
kv_cache2 = None
if boundary_timestep is not None:
# Initialize the low noise kv cache
kv_cache2 = self._initialize_kv_cache(batch_size=latents.shape[0],
dtype=target_dtype,
device=latents.device)
# Optional: cache context features from provided image latents prior to generation
current_start_frame = 0
if getattr(batch, "image_latent", None) is not None:
image_latent = batch.image_latent
assert image_latent is not None
input_frames = image_latent.shape[2]
# timestep zero (or configured context noise) for cache warm-up
t_zero = torch.zeros([latents.shape[0]],
device=latents.device,
dtype=torch.long)
if independent_first_frame and input_frames >= 1:
# warm-up with the very first frame independently
image_first_btchw = image_latent[:, :, :1, :, :].to(
target_dtype).permute(0, 2, 1, 3, 4)
with torch.autocast(device_type="cuda",
dtype=target_dtype,
enabled=autocast_enabled):
_ = self.transformer(
image_first_btchw,
prompt_embeds,
t_zero,
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
current_start=current_start_frame *
self.frame_seq_length,
**image_kwargs,
**pos_cond_kwargs,
)
current_start_frame += 1
remaining_frames = input_frames - 1
else:
remaining_frames = input_frames
def _get_kv_cache(timestep: float) -> list[dict]:
if boundary_timestep is not None:
if timestep >= boundary_timestep:
return kv_cache1
else:
assert kv_cache2 is not None, "kv_cache2 is not initialized"
return kv_cache2
return kv_cache1
# process remaining input frames in blocks of num_frame_per_block
while remaining_frames > 0:
block = min(self.num_frames_per_block, remaining_frames)
ref_btchw = image_latent[:, :, current_start_frame:
current_start_frame +
block, :, :].to(target_dtype).permute(
0, 2, 1, 3, 4)
with torch.autocast(device_type="cuda",
dtype=target_dtype,
enabled=autocast_enabled):
_ = self.transformer(
ref_btchw,
prompt_embeds,
t_zero,
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
current_start=current_start_frame *
self.frame_seq_length,
**image_kwargs,
**pos_cond_kwargs,
)
current_start_frame += block
remaining_frames -= block
crossattn_cache = self._initialize_crossattn_cache(
batch_size=latents.shape[0],
max_text_len=fastvideo_args.pipeline_config.text_encoder_configs[0].
arch_config.text_len,
dtype=target_dtype,
device=latents.device)
# Base position offset from any cache warm-up
pos_start_base = current_start_frame
pos_start_base = 0
# Determine block sizes
if not independent_first_frame or (independent_first_frame
and batch.image_latent is not None):
if t % self.num_frames_per_block != 0:
raise ValueError(
"num_frames must be divisible by num_frames_per_block for causal DMD denoising"
block_sizes = [self.num_frames_per_block] * 7
block_sizes[0] = 1
start_index = 0
first_frame_latent = None
if batch.pil_image is not None:
# Causal video gen directly replaces the first frame of the latent with
# the image latent instead of appending along the channel dim
assert self.vae is not None, "VAE is not provided for causal video gen task"
self.vae = self.vae.to(get_local_torch_device())
first_frame_latent = self.vae.encode(batch.pil_image).mean.float()
if (hasattr(self.vae, "shift_factor")
and self.vae.shift_factor is not None):
if isinstance(self.vae.shift_factor, torch.Tensor):
first_frame_latent -= self.vae.shift_factor.to(
first_frame_latent.device, first_frame_latent.dtype)
else:
first_frame_latent -= self.vae.shift_factor
if isinstance(self.vae.scaling_factor, torch.Tensor):
first_frame_latent = first_frame_latent * self.vae.scaling_factor.to(
first_frame_latent.device, first_frame_latent.dtype)
else:
first_frame_latent = first_frame_latent * self.vae.scaling_factor
if fastvideo_args.vae_cpu_offload:
self.vae = self.vae.to("cpu")
# Fill the low noise and high noise kv cache with first_frame_latent and timestep 0
t_zero = torch.zeros([latents.shape[0], 1],
device=latents.device,
dtype=torch.long)
with torch.autocast(device_type="cuda",
dtype=target_dtype,
enabled=autocast_enabled), \
set_forward_context(current_timestep=0,
attn_metadata=None,
forward_batch=batch):
self.transformer(
first_frame_latent.to(target_dtype),
prompt_embeds,
t_zero,
kv_cache=kv_cache1,
crossattn_cache=crossattn_cache,
current_start=(pos_start_base + start_index) *
self.frame_seq_length,
start_frame=start_index,
**image_kwargs,
**pos_cond_kwargs,
)
num_blocks = t // self.num_frames_per_block
block_sizes = [self.num_frames_per_block] * num_blocks
start_index = 0
else:
if (t - 1) % self.num_frames_per_block != 0:
raise ValueError(
"(num_frames - 1) must be divisible by num_frame_per_block when independent_first_frame=True"
)
num_blocks = (t - 1) // self.num_frames_per_block
block_sizes = [1] + [self.num_frames_per_block] * num_blocks
start_index = 0
if boundary_timestep is not None:
self.transformer_2(
first_frame_latent.to(target_dtype),
prompt_embeds,
t_zero,
kv_cache=kv_cache2,
crossattn_cache=crossattn_cache,
current_start=(pos_start_base + start_index) *
self.frame_seq_length,
start_frame=start_index,
**image_kwargs,
**pos_cond_kwargs,
)
start_index += 1
block_sizes.pop(0)
latents[:, :, :1, :, :] = first_frame_latent
# DMD loop in causal blocks
# Optional per-block callback for streaming
on_block = None
try:
on_block = getattr(batch, "extra",
{}).get("on_block",
None) # type: ignore[attr-defined]
except Exception:
on_block = None
with self.progress_bar(total=len(block_sizes) *
len(timesteps)) as progress_bar:
for block_idx, current_num_frames in enumerate(block_sizes):
for current_num_frames in block_sizes:
current_latents = latents[:, :, start_index:start_index +
current_num_frames, :, :]
# use BTCHW for DMD conversion routines
@@ -231,7 +221,7 @@ class CausalDMDDenosingStage(DenoisingStage):
video_raw_latent_shape = noise_latents_btchw.shape
for i, t_cur in enumerate(timesteps):
if self.transformer_2 is not None and fastvideo_args.pipeline_config.dit_config.boundary_ratio is not None and t_cur < fastvideo_args.pipeline_config.dit_config.boundary_ratio * self.scheduler.num_train_timesteps:
if boundary_timestep is not None and t_cur < boundary_timestep:
current_model = self.transformer_2
else:
current_model = self.transformer
@@ -289,8 +279,8 @@ class CausalDMDDenosingStage(DenoisingStage):
latent_model_input,
prompt_embeds,
t_expanded_noise,
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
kv_cache=_get_kv_cache(t_cur),
crossattn_cache=crossattn_cache,
current_start=(pos_start_base + start_index) *
self.frame_seq_length,
start_frame=start_index,
@@ -299,12 +289,22 @@ class CausalDMDDenosingStage(DenoisingStage):
).permute(0, 2, 1, 3, 4)
# Convert pred noise to pred video with FM Euler scheduler utilities
pred_video_btchw = pred_noise_to_pred_video(
pred_noise=pred_noise_btchw.flatten(0, 1),
noise_input_latent=noise_latents.flatten(0, 1),
timestep=t_expand,
scheduler=self.scheduler).unflatten(
0, pred_noise_btchw.shape[:2])
if boundary_timestep is not None and t_cur >= boundary_timestep:
pred_video_btchw = pred_noise_to_x_bound(
pred_noise=pred_noise_btchw.flatten(0, 1),
noise_input_latent=noise_latents.flatten(0, 1),
timestep=t_expand,
boundary_timestep=torch.ones_like(t_expand) *
boundary_timestep,
scheduler=self.scheduler).unflatten(
0, pred_noise_btchw.shape[:2])
else:
pred_video_btchw = pred_noise_to_pred_video(
pred_noise=pred_noise_btchw.flatten(0, 1),
noise_input_latent=noise_latents.flatten(0, 1),
timestep=t_expand,
scheduler=self.scheduler).unflatten(
0, pred_noise_btchw.shape[:2])
if i < len(timesteps) - 1:
next_timestep = timesteps[i + 1] * torch.ones(
@@ -318,11 +318,23 @@ class CausalDMDDenosingStage(DenoisingStage):
batch.generator, list) else
batch.generator)).to(self.device)
noise_btchw = noise
noise_latents_btchw = self.scheduler.add_noise(
pred_video_btchw.flatten(0, 1),
noise_btchw.flatten(0, 1),
next_timestep).unflatten(0,
pred_video_btchw.shape[:2])
if boundary_timestep is not None and i < len(
high_noise_timesteps) - 1:
noise_latents_btchw = self.scheduler.add_noise_high(
pred_video_btchw.flatten(0, 1),
noise_btchw.flatten(0, 1), next_timestep,
torch.ones_like(next_timestep) *
boundary_timestep).unflatten(
0, pred_video_btchw.shape[:2])
elif boundary_timestep is not None and i == len(
high_noise_timesteps) - 1:
noise_latents_btchw = pred_video_btchw
else:
noise_latents_btchw = self.scheduler.add_noise(
pred_video_btchw.flatten(0, 1),
noise_btchw.flatten(0, 1),
next_timestep).unflatten(
0, pred_video_btchw.shape[:2])
current_latents = noise_latents_btchw.permute(
0, 2, 1, 3, 4)
else:
@@ -350,38 +362,40 @@ class CausalDMDDenosingStage(DenoisingStage):
attn_metadata=attn_metadata,
forward_batch=batch):
t_expanded_context = t_context.unsqueeze(1)
_ = current_model(
if boundary_timestep is not None:
self.transformer_2(
context_bcthw,
prompt_embeds,
t_expanded_context,
kv_cache=kv_cache2,
crossattn_cache=crossattn_cache,
current_start=(pos_start_base + start_index) *
self.frame_seq_length,
start_frame=start_index,
**image_kwargs,
**pos_cond_kwargs,
)
self.transformer(
context_bcthw,
prompt_embeds,
t_expanded_context,
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
kv_cache=kv_cache1,
crossattn_cache=crossattn_cache,
current_start=(pos_start_base + start_index) *
self.frame_seq_length,
start_frame=start_index,
**image_kwargs,
**pos_cond_kwargs,
)
start_index += current_num_frames
# Invoke callback with block metadata (no large tensor transfer required)
try:
if callable(on_block):
on_block(
block_index=block_idx,
total_blocks=len(block_sizes),
start_index=start_index - current_num_frames,
num_frames=current_num_frames,
latents=current_latents,
)
except Exception as e:
# Swallow callback errors so they don't break generation
logger.warning("on_block callback failed: %s", str(e))
start_index += current_num_frames
batch.latents = latents
return batch
def _initialize_kv_cache(self, batch_size, dtype, device) -> None:
def _initialize_kv_cache(self, batch_size, dtype, device) -> list[dict]:
"""
Initialize a Per-GPU KV cache aligned with the Wan model assumptions.
"""
@@ -415,10 +429,10 @@ class CausalDMDDenosingStage(DenoisingStage):
torch.tensor([0], dtype=torch.long, device=device),
})
self.kv_cache1 = kv_cache1
return kv_cache1
def _initialize_crossattn_cache(self, batch_size, max_text_len, dtype,
device) -> None:
device) -> list[dict]:
"""
Initialize a Per-GPU cross-attention cache aligned with the Wan model assumptions.
"""
@@ -444,7 +458,7 @@ class CausalDMDDenosingStage(DenoisingStage):
"is_init":
False,
})
self.crossattn_cache = crossattn_cache
return crossattn_cache
def verify_input(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
+6 -5
View File
@@ -76,11 +76,12 @@ class DecodingStage(PipelineStage):
vae_autocast_enabled = (
vae_dtype != torch.float32) and not fastvideo_args.disable_autocast
if isinstance(self.vae.scaling_factor, torch.Tensor):
latents = latents / self.vae.scaling_factor.to(
latents.device, latents.dtype)
else:
latents = latents / self.vae.scaling_factor
if hasattr(self.vae, 'scaling_factor'):
if isinstance(self.vae.scaling_factor, torch.Tensor):
latents = latents / self.vae.scaling_factor.to(
latents.device, latents.dtype)
else:
latents = latents / self.vae.scaling_factor
# Apply shifting if needed
if (hasattr(self.vae, "shift_factor")

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