Compare commits

...
Author SHA1 Message Date
SolitaryThinker 73bf0a0a2b fix 2025-10-28 10:00:55 +00:00
William Lin 2cd2e57d2e [ci] fix causal ssim test (#848) 2025-10-26 19:33:07 -07:00
William Linandainsley 9370234294 [feat] Add gradio local inference demo (#847)
Co-authored-by: ainsley <jzhang2765@wisc.edu>
2025-10-26 07:01:33 -07:00
Jinzhe Pan 50da62e722 [bugfix] always force spawn instead of fork (#852) 2025-10-23 16:36:50 -07:00
William Lin 4f3e8751db [bugfix] [misc] Use training_state_checkpointing_steps in scripts/ (#846) 2025-10-19 20:20:53 -07:00
Jinzhe PanandXingyu Long f4c58894d9 [Feat] add ray support (#838)
Co-authored-by: Xingyu Long <xingyulong97@gmail.com>
2025-10-16 23:17:54 -07:00
Ohm-Rishabh 01c94ef385 [feat] unified trainer logging (#841) 2025-10-16 23:16:16 -07:00
Zhang Peiyuan 2415226d25 Update WeChat Link 2025-10-13 21:02:46 -07:00
Jiali Chen 404314d00f [Feature]Add video-to-video (V2V) pipeline (#829) 2025-10-12 21:53:05 -07:00
zyang6andkiritorl 87489f0872 Add wan2.1 functionality support for Ascend NPU platform (#810)
Co-authored-by: kiritorl <1021709528@qq.com>
2025-10-09 16:25:08 -07:00
Zhang Peiyuan 9ce7c8039e Update Wechat link 2025-10-06 15:01:19 -07:00
William Lin e1e25e95f9 [feature] Add torch profiler (#827) 2025-10-06 07:59:46 -07:00
William Lin 490bde90e1 [bugfix] Allow overriding dit checkpoint for inference and Lower VSA LR in example scripts (#831) 2025-10-05 01:44:49 -07:00
dc7596b973 [self-forcing][8/n] Self-Forcing For Wan2.2-A14B + torch.compile training and distillation support (#818)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
2025-10-02 15:01:45 -07:00
William Lin 335afa4457 [bugfix] Use training_state_checkpointing_steps instead of checkpointing_steps (#821) 2025-09-28 15:22:43 -07:00
Yongqi Chen 3f77a6805a [Feature]Update count trainable param for FSDP2 (#820) 2025-09-28 15:22:04 -07:00
RandNMR73 13d0aae706 Add Sage Attention 3 Backend (#815) 2025-09-24 15:11:38 -07:00
William Lin 404cbf4f3c [self-forcing] [6/n] Add Ode Init training (#811) 2025-09-22 17:58:19 -07:00
William Lin 958ffec844 [bugfix] Update learning rates for sparse distillation recipe (#812) 2025-09-22 12:07:03 -07:00
31f000d1cc [self-forcing] [5/n] Add Self-Forcing distillation pipeline (#808)
Co-authored-by: RandNMR73 <notomatthew31@gmail.com>
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2025-09-20 19:32:10 -07:00
Yongqi Chen cd32b3e02f Update example files and readme (#809) 2025-09-20 18:15:59 -07:00
Zhang Peiyuan bf27908095 Update WeChat Link 2025-09-20 14:16:20 -07:00
William Lin c5f9ea53b2 [self-forcing] [4/n] Preprocessing for collecting ODE trajectory (#788) 2025-09-15 17:54:42 -07:00
William Lin d32a7184da [bugfix] Wan2.2 Boundary ratio (#804) 2025-09-15 11:17:35 -07:00
Wenxuan Tanandgemini-code-assist[bot] 2930abe456 [Bugfix] Fix VMoba requirements (#802)
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
2025-09-14 18:28:52 -07:00
151 changed files with 10452 additions and 923 deletions
+12
View File
@@ -104,6 +104,18 @@ steps:
- TEST_TYPE=distillation_dmd
agents:
queue: "default"
- path:
- "fastvideo/training/*self_forcing_distillation_pipeline.py"
- "fastvideo/tests/training/self-forcing/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Self-Forcing Tests"
env:
- TEST_TYPE=self_forcing
agents:
queue: "default"
- path:
- "fastvideo/**"
- "pyproject.toml"
+4
View File
@@ -110,6 +110,10 @@ case "$TEST_TYPE" in
MODAL_COMMAND="$MODAL_ENV WANDB_API_KEY=$WANDB_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_distill_dmd_tests"
;;
# run_inference_tests_vmoba
"self_forcing")
log "Running self-forcing tests..."
MODAL_COMMAND="$MODAL_ENV WANDB_API_KEY=$WANDB_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_self_forcing_tests"
;;
"inference_vmoba")
log "Running V-MoBA inference tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_inference_tests_vmoba"
+1 -1
View File
@@ -352,7 +352,7 @@ jobs:
volume_size: 100
disk_size: 100
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
test_command: "uv pip install -e .[test] && pytest ./fastvideo/dataset/ -vs && pytest ./fastvideo/workflow/ -vs"
test_command: "uv pip install -e .[test] && pytest ./fastvideo/dataset/ -vs && pytest ./fastvideo/workflow/ -vs && pytest ./fastvideo/entrypoints/ -vs"
timeout_minutes: 30
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
+4 -7
View File
@@ -12,9 +12,6 @@ exclude: |
scripts/.*|
fastvideo/data_preprocess/.*|
fastvideo/dataset/.*|
fastvideo/distill/.*|
fastvideo/distill\.py|
fastvideo/distill_adv\.py|
fastvideo/models/.*|
fastvideo/sample/.*|
fastvideo/train\.py|
@@ -44,10 +41,10 @@ repos:
- id: codespell
additional_dependencies: ['tomli']
args: ['--toml', 'pyproject.toml']
- repo: https://github.com/PyCQA/isort
rev: 6.0.1
hooks:
- id: isort
# - repo: https://github.com/PyCQA/isort
# rev: 6.0.1
# hooks:
# - id: isort
- repo: https://github.com/jackdewinter/pymarkdown
rev: v0.9.30
hooks:
+3 -3
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/S7HLCSTh" target="_blank"> <b> WeChat </b> </a> |
| 🕹️ <a href="https://fastwan.fastvideo.org/"<b>Online Demo</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start.html"><b> Quick Start</b></a> | 🤗 <a href="https://huggingface.co/collections/FastVideo/fastwan-6886a305d9799c8cd1496408" target="_blank"><b>FastWan</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-3csdw1isz-Euq8_Q8~baewG8hxjXs2gQ" target="_blank"> <b>Slack</b> </a> | 🟣💬 <a href="https://ibb.co/tMwknPLY" target="_blank"> <b> WeChat </b> </a> |
</p>
<div align="center">
@@ -155,8 +155,8 @@ If you find FastVideo useful, please considering citing our work:
}
@article{zhang2025vsa,
title={VSA: Faster Video Diffusion with Trainable Sparse Attention},
author={Zhang, Peiyuan and Huang, Haofeng and Chen, Yongqi and Lin, Will and Liu, Zhengzhong and Stoica, Ion and Xing, Eric and Zhang, Hao},
title={Vsa: Faster video diffusion with trainable sparse attention},
author={Zhang, Peiyuan and Chen, Yongqi and Huang, Haofeng and Lin, Will and Liu, Zhengzhong and Stoica, Ion and Xing, Eric and Zhang, Hao},
journal={arXiv preprint arXiv:2505.13389},
year={2025}
}
+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

+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

+3 -1
View File
@@ -20,5 +20,7 @@ setup(
"License :: OSI Approved :: Apache Software License",
],
python_requires='>=3.12',
install_requires=[]
install_requires=[
"flash-attn >= 2.7.1",
]
)
+10 -2
View File
@@ -6,8 +6,16 @@ import time
import os
import torch
from typing import Tuple
from flash_attn import flash_attn_varlen_func # Use the new flash attention function
from flash_attn.flash_attn_interface import _flash_attn_varlen_forward, _flash_attn_varlen_backward
try:
from flash_attn import flash_attn_varlen_func # Use the new flash attention function
from flash_attn.flash_attn_interface import _flash_attn_varlen_forward, _flash_attn_varlen_backward
except ImportError:
def _unsupported(*args, **kwargs):
raise ImportError("flash-attn is not installed. Please install it, e.g., `pip install flash-attn`.")
_flash_attn_varlen_forward = _unsupported
_flash_attn_varlen_backward = _unsupported
flash_attn_varlen_func = _unsupported
from functools import lru_cache
from einops import rearrange
+53
View File
@@ -0,0 +1,53 @@
# Profiling FastVideo
!!! warning
Profiling is only intended for FastVideo developers and maintainers to understand the proportion of time spent in different parts of the codebase. **FastVideo end-users should never turn on profiling** as it will significantly slow down the inference.
## Profiling with PyTorch
FastVideo exposes a process-wide torch profiler that you can enable via environment variables. Set `FASTVIDEO_TORCH_PROFILER_DIR` to an absolute directory path to start collecting traces, and specify the regions you want recorded with `FASTVIDEO_TORCH_PROFILE_REGIONS`:
```bash
FASTVIDEO_TORCH_PROFILER_DIR=/mnt/traces/fastvideo \
FASTVIDEO_TORCH_PROFILE_REGIONS="profiler_region_model_loading,profiler_region_training_step"
```
All profiled regions must be registered in `fastvideo.profiler`; the current list includes:
- `profiler_region_model_loading` — pipeline/module loading
- `profiler_region_inference_pre_denoising`
- `profiler_region_inference_denoising`
- `profiler_region_inference_post_denoising`
- `profiler_region_training_checkpoint_saving`
- `profiler_region_training_dit`
- `profiler_region_training_validation`
- `profiler_region_training_epoch`
- `profiler_region_training_step`
- `profiler_region_training_backward`
- `profiler_region_training_optimizer`
- `profiler_region_distillation_teacher_forward`
- `profiler_region_distillation_student_forward`
- `profiler_region_distillation_loss`
- `profiler_region_distillation_update`
While profiling is enabled, FastVideo records additional annotations:
- `fastvideo.region::<name>` spans are emitted when entering a region.
- `fastvideo.profiler.enable_collection` / `fastvideo.profiler.disable_collection` events mark when torch profiler collection is toggled on or off.
Only one profiler instance is created per process; subsequent pipelines reuse the same controller. If you set `FASTVIDEO_TORCH_PROFILE_REGIONS` incorrectly (e.g. misspelled name), FastVideo logs a warning and ignores that entry.
Additional knobs:
- `FASTVIDEO_TORCH_PROFILER_RECORD_SHAPES`
- `FASTVIDEO_TORCH_PROFILER_WITH_PROFILE_MEMORY`
- `FASTVIDEO_TORCH_PROFILER_WITH_STACK`
- `FASTVIDEO_TORCH_PROFILER_WITH_FLOPS`
Traces can be visualized using <https://ui.perfetto.dev/>.
### Best Practices
- Keep the profiled step count small; traces can be large and slow down job shutdown while the profiler flushes data.
- After profiling, clean up trace directories to avoid filling disks.
- When adding new regions, register them in `fastvideo.profiler` and wrap the corresponding code block with `with self.profiler_controller.region("your_region"):` or the `@profile_region` decorator.
+1
View File
@@ -115,6 +115,7 @@ design/overview
contributing/overview
contributing/developer_env/index
contributing/profiling
:::
:::{toctree}
+42 -2
View File
@@ -1,32 +1,40 @@
(inference-optimizations)=
# Optimizations
This page describes the various options for speeding up generation times in FastVideo.
## Table of Contents
- Optimized Attention Backends
- [Flash Attention](#optimizations-flash)
- [Sliding Tile Attention](#optimizations-sta)
- [Sage Attention](#optimizations-sage)
- [Sage Attention 3](#optimizations-sage3)
- Caching Techniques
- [TeaCache](#optimizations-teacache)
(optimizations-backends)=
## Attention Backends
### Available Backends
- Torch SDPA: `FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA`
- Flash Attention 2 and 3: `FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN`
- Sliding Tile Attention: `FASTVIDEO_ATTENTION_BACKEND=SLIDING_TILE_ATTN`
- Video Sparse Attention: `FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN`
- Sage Attention: `FASTVIDEO_ATTENTION_BACKEND=SAGE_ATTN`
- Sage Attention 3: `FASTVIDEO_ATTENTION_BACKEND=SAGE_ATTN_THREE`
### Configuring Backends
There are two ways to configure the attention backend in FastVideo.
#### 1. In Python
In python, set the `FASTVIDEO_ATTENTION_BACKEND` environment variable before instantiating `VideoGenerator` like this:
```python
@@ -34,6 +42,7 @@ os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "SLIDING_TILE_ATTN"
```
#### 2. In CLI
You can also set the environment variable on the command line:
```bash
@@ -41,6 +50,7 @@ FASTVIDEO_ATTENTION_BACKEND=SAGE_ATTN python example.py
```
(optimizations-flash)=
### Flash Attention
**`FLASH_ATTN`**
@@ -57,7 +67,7 @@ And if using a Hopper+ GPU (ie H100), installing [Flash Attention 3](https://git
git clone https://github.com/Dao-AILab/flash-attention.git && cd flash-attention
cd hopper
pip install ninja
pip install ninja
python setup.py install
```
@@ -66,7 +76,9 @@ FastVideo will automatically detect and use `FA3` if it is installed when using
:::
(optimizations-sta)=
### Sliding Tile Attention
**`SLIDING_TILE_ATTN`**
```bash
@@ -76,7 +88,9 @@ 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`**
```bash
@@ -87,19 +101,45 @@ python setup_vsa.py install
Please see [this page](#vsa-installation) for more installation instructions.
(optimizations-sage)=
### Sage Attention
**`SAGE_ATTN`**
To use [SageAttention](https://github.com/thu-ml/SageAttention) 2.1.1, please compile from source:
```bash
git clone https://github.com/thu-ml/SageAttention.git
cd sageattention
cd sageattention
python setup.py install # or pip install -e .
```
(optimizations-sage3)=
### Sage Attention 3
**`SAGE_ATTN_THREE`**
[SageAttention 3](https://huggingface.co/jt-zhang/SageAttention3) is an advanced attention mechanism that leverages FP4 quantization and Blackwell GPU Tensor Cores for significant performance improvements.
#### Hardware Requirements
- RTX5090
#### Installation
Note that Sage Attention 3 requires `python>=3.13`, `torch>=2.8.0`, `CUDA >=12.8`. If you are using `uv` and using `torch==2.8.0` make sure that `sentencepiece==0.2.1` in the pyproject.toml file.
To use Sage Attention 3 in FastVideo, first get access to the SageAttention3 code, then move `sageattn/` and `setup.py` to the directory `fastvideo/attention/backends`, then install from using:
```bash
python setup.py install
```
(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.
### What is TeaCache?
@@ -0,0 +1,140 @@
#!/bin/bash
#SBATCH --job-name=t2v
#SBATCH --partition=main
#SBATCH --nodes=1
#SBATCH --ntasks=1
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:1
#SBATCH --cpus-per-task=128
#SBATCH --mem=1440G
#SBATCH --output=dmd_t2v_output/t2v_%j.out
#SBATCH --error=dmd_t2v_output/t2v_%j.err
#SBATCH --exclusive
# Basic Info
export NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
# different cache dir for different processes
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
export MASTER_PORT=29503
export TOKENIZERS_PARALLELISM=false
export WANDB_API_KEY=your_wandb_api_key
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
# Configs
NUM_GPUS=1
# Model paths for Self-Forcing DMD distillation:
GENERATOR_MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-14B-Diffusers" # Teacher model
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Critic model
DATA_DIR=your_data_dir
VALIDATION_DATASET_FILE=your_validation_data_dir
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
training_args=(
--tracker_project_name SFwan_t2v_distill_self_forcing_dmd
--output_dir your_output_dir
--max_train_steps 4000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 21
--num_height 480
--num_width 832
--enable_gradient_checkpointing_type "full"
--log_visualization
--simulate_generator_forward
--num_frames 81
--num_frame_per_block 3 # Frame generation block size for self-forcing
--enable_gradient_masking
--gradient_mask_last_n_frames 21
)
parallel_args=(
--num_gpus $NUM_GPUS # 64
--sp_size 1
--tp_size 1
--hsdp_replicate_dim 1 # 64
--hsdp_shard_dim $NUM_GPUS
)
model_args=(
--model_path $GENERATOR_MODEL_PATH # TODO: check if you can remove this in this script
--pretrained_model_name_or_path $GENERATOR_MODEL_PATH
--real_score_model_path $REAL_SCORE_MODEL_PATH
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
)
dataset_args=(
--data_path "$DATA_DIR"
--dataloader_num_workers 4
)
validation_args=(
--log_validation
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 50
--validation_sampling_steps "4"
--validation_guidance_scale "6.0" # not used for dmd inference
)
optimizer_args=(
--learning_rate 1e-5
--mixed_precision "bf16"
--training_state_checkpointing_steps 500
--weight_only_checkpointing_steps 500
--weight_decay 0.01
--betas '0.0,0.999'
--max_grad_norm 1.0
)
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
--training_cfg_rate 0.0
--dit_precision "fp32"
--flow_shift 5
--seed 1000
--use_ema True
--ema_decay 0.99
--ema_start_step 100
--init_weights_from_safetensors your_ode_init_weights_path
)
dmd_args=(
--dmd_denoising_steps '1000,750,500,250'
--min_timestep_ratio 0.02
--max_timestep_ratio 0.98
--dfake_gen_update_ratio 5
--real_score_guidance_scale 3.0
--fake_score_learning_rate 8e-6
--fake_score_betas '0.0,0.999'
--warp_denoising_step
)
self_forcing_args=(
--independent_first_frame False # Whether to treat first frame independently
--same_step_across_blocks True # Whether to use same denoising step across all blocks
--last_step_only False # Whether to only use the last denoising step
--context_noise 0 # Amount of noise to add during context caching (0 = no noise)
)
torchrun \
--nnodes 1 \
--master_port $MASTER_PORT \
--nproc_per_node $NUM_GPUS \
fastvideo/training/wan_self_forcing_distillation_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_args[@]}" \
"${self_forcing_args[@]}"
@@ -0,0 +1,3 @@
#!/bin/bash
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
@@ -0,0 +1,24 @@
#!/bin/bash
GPU_NUM=1 # 2,4,8
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
MODEL_TYPE="wan"
DATA_MERGE_PATH="data/crush-smol/merge.txt"
OUTPUT_DIR="data/crush-smol_processed_t2v/"
torchrun --nproc_per_node=$GPU_NUM \
fastvideo/pipelines/preprocess/v1_preprocess.py \
--model_path $MODEL_PATH \
--data_merge_path $DATA_MERGE_PATH \
--preprocess_video_batch_size 8 \
--seed 42 \
--max_height 480 \
--max_width 832 \
--num_frames 81 \
--dataloader_num_workers 0 \
--output_dir=$OUTPUT_DIR \
--train_fps 16 \
--samples_per_file 8 \
--flush_frequency 8 \
--video_length_tolerance_range 5 \
--preprocess_task "t2v"
@@ -0,0 +1,157 @@
#!/bin/bash
#SBATCH --job-name=t2v
#SBATCH --partition=main
#SBATCH --nodes=4
#SBATCH --ntasks=4
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=128
#SBATCH --mem=1440G
#SBATCH --output=dmd_t2v_output/t2v_%j.out
#SBATCH --error=dmd_t2v_output/t2v_%j.err
#SBATCH --exclusive
# Basic Info
export NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
export NCCL_DEBUG_SUBSYS=INIT,NET
# different cache dir for different processes
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
export MASTER_PORT=29500
export NODE_RANK=$SLURM_PROCID
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
export MASTER_ADDR=${nodes[0]}
export TOKENIZERS_PARALLELISM=false
export WANDB_API_KEY="2f25ad37933894dbf0966c838c0b8494987f9f2f"
# export WANDB_API_KEY='your_wandb_api_key_here'
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
# Configs
NUM_GPUS=8
# Model paths for Self-Forcing DMD distillation with Wan2.2:
# GENERATOR_MODEL_PATH="Wan-AI/Wan2.2-T2V-A14B-Diffusers" # Updated to Wan2.2
# REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.2-T2V-A14B-Diffusers" # Teacher model
# FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.2-T2V-A14B-Diffusers" # Critic model
GENERATOR_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Updated to Wan2.2
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Teacher model
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Critic model
# DATA_DIR="data/test-text-preprocessing/Node_0_GPU_1_File_1/combined_parquet_dataset/"
DATA_DIR="/mnt/weka/home/hao.zhang/matthew/FastVideo/data/test-text-preprocessing"
# DATA_DIR=data/crush-smol_processed_t2v/combined_parquet_dataset
# DATA_DIR="/mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn-upload/latents_i2v/train/"
# VALIDATION_DATASET_FILE="data/crush-smol-single_processed_t2v/validation.json"
VALIDATION_DATASET_FILE="/mnt/weka/home/hao.zhang/wl/FastVideo/examples/distill/Wan2.2-TI2V-5B-Diffusers/Data-free/validation_64.json"
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
training_args=(
--tracker_project_name SFwan2.2_t2v_distill_self_forcing_dmd # Updated for Wan2.2
--output_dir "/mnt/sharefs/users/hao.zhang/SFwan2.2_t2v_finetune"
--override_transformer_cls_name "CausalWanTransformer3DModel"
--max_train_steps 4000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 21
--num_height 448 # Updated to match Wan2.2 config
--num_width 832 # Updated to match Wan2.2 config
--enable_gradient_checkpointing_type "full"
--simulate_generator_forward
# --log_visualization
--num_frames 81
--num_frame_per_block 3 # Frame generation block size for self-forcing
--enable_gradient_masking
--gradient_mask_last_n_frames 21
# --init_weights_from_safetensors /mnt/sharefs/users/hao.zhang/wl/models/sf_ode_init_wan22_checkpoints/high/3k/
# --init_weights_from_safetensors_2 /mnt/sharefs/users/hao.zhang/wl/models/sf_ode_init_wan22_checkpoints/low/3k/
)
parallel_args=(
--num_gpus 32 # 64
--sp_size 1
--tp_size 1
--hsdp_replicate_dim 1 # 64
--hsdp_shard_dim 32
)
model_args=(
--model_path $GENERATOR_MODEL_PATH
--pretrained_model_name_or_path $GENERATOR_MODEL_PATH
--real_score_model_path $REAL_SCORE_MODEL_PATH
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
)
dataset_args=(
--data_path "$DATA_DIR"
--dataloader_num_workers 4
)
validation_args=(
--log_validation
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 20
--validation_sampling_steps "4"
--validation_guidance_scale "6.0" # not used for dmd inference
)
optimizer_args=(
--learning_rate 1e-5
--mixed_precision "bf16"
--training_state_checkpointing_steps 500
--weight_only_checkpointing_steps 500
--weight_decay 0.01
--betas '0.0,0.999'
--max_grad_norm 1.0
)
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
--training_cfg_rate 0.0
--dit_precision "fp32"
--flow_shift 5
--seed 1000
--use_ema True
--ema_decay 0.99
--ema_start_step 100
)
dmd_args=(
--dmd_denoising_steps '1000,750,500,250'
--min_timestep_ratio 0.02
--max_timestep_ratio 0.98
--dfake_gen_update_ratio 5
--real_score_guidance_scale 3.0
--fake_score_learning_rate 8e-6
--fake_score_betas '0.0,0.999'
--warp_denoising_step
)
self_forcing_args=(
--independent_first_frame False # Whether to treat first frame independently
--same_step_across_blocks True # Whether to use same denoising step across all blocks
--last_step_only False # Whether to only use the last denoising step
--context_noise 0 # Amount of noise to add during context caching (0 = no noise)
)
srun torchrun \
--nnodes $SLURM_JOB_NUM_NODES \
--nproc_per_node $NUM_GPUS \
--node_rank $SLURM_PROCID \
--rdzv_backend=c10d \
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
fastvideo/training/wan_self_forcing_distillation_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_args[@]}" \
"${self_forcing_args[@]}"
@@ -39,15 +39,18 @@ echo "NODE_RANK: $NODE_RANK"
# Configs
NUM_GPUS=8
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATA_DIR=your_data_dir
VALIDATION_DATASET_FILE=your_validation_dataset_file
OUTPUT_DIR="checkpoints/wan_t2v_finetune"
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name wan_t2v_distill_dmd_VSA
--output_dir"checkpoints/wan_t2v_finetune"
--output_dir $OUTPUT_DIR
--max_train_steps 4000
--train_batch_size 1
--train_sp_batch_size 1
@@ -72,6 +75,8 @@ parallel_args=(
model_args=(
--model_path $MODEL_PATH
--pretrained_model_name_or_path $MODEL_PATH
--real_score_model_path $REAL_SCORE_MODEL_PATH
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
)
# Dataset arguments
@@ -91,7 +96,7 @@ validation_args=(
# Optimizer arguments
optimizer_args=(
--learning_rate 1e-5
--learning_rate 2e-6
--mixed_precision "bf16"
--training_state_checkpointing_steps 500
--weight_only_checkpointing_steps 500
@@ -134,4 +139,4 @@ srun torchrun \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_args[@]}"
"${dmd_args[@]}"
@@ -39,15 +39,18 @@ echo "NODE_RANK: $NODE_RANK"
# Configs
NUM_GPUS=8
MODEL_PATH="Wan-AI/Wan2.1-T2V-14B-Diffusers"
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-14B-Diffusers"
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-14B-Diffusers"
DATA_DIR=your_data_dir
VALIDATION_DATASET_FILE=your_validation_dataset_file
OUTPUT_DIR="checkpoints/wan_t2v_finetune"
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name wan_t2v_distill_dmd_VSA
--output_dir "checkpoints/wan_t2v_finetune"
--output_dir "$OUTPUT_DIR"
--max_train_steps 4000
--train_batch_size 1
--train_sp_batch_size 1
@@ -72,6 +75,8 @@ parallel_args=(
model_args=(
--model_path $MODEL_PATH
--pretrained_model_name_or_path $MODEL_PATH
--real_score_model_path $REAL_SCORE_MODEL_PATH
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
)
# Dataset arguments
@@ -91,7 +96,7 @@ validation_args=(
# Optimizer arguments
optimizer_args=(
--learning_rate 1e-5
--learning_rate 2e-6
--mixed_precision "bf16"
--training_state_checkpointing_steps 500
--weight_only_checkpointing_steps 500
@@ -134,4 +139,4 @@ srun torchrun \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_args[@]}"
"${dmd_args[@]}"
@@ -39,15 +39,18 @@ echo "NODE_RANK: $NODE_RANK"
# Configs
NUM_GPUS=8
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATA_DIR=your_data_dir
VALIDATION_DATASET_FILE=your_validation_dataset_file
OUTPUT_DIR="checkpoints/wan_t2v_finetune"
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name wan_t2v_distill_dmd
--output_dir "checkpoints/wan_t2v_finetune"
--output_dir "$OUTPUT_DIR"
--max_train_steps 4000
--train_batch_size 1
--train_sp_batch_size 1
@@ -72,6 +75,8 @@ parallel_args=(
model_args=(
--model_path $MODEL_PATH
--pretrained_model_name_or_path $MODEL_PATH
--real_score_model_path $REAL_SCORE_MODEL_PATH
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
)
# Dataset arguments
@@ -91,7 +96,7 @@ validation_args=(
# Optimizer arguments
optimizer_args=(
--learning_rate 1e-5
--learning_rate 2e-6
--mixed_precision "bf16"
--training_state_checkpointing_steps 500
--weight_only_checkpointing_steps 500
@@ -133,4 +138,4 @@ srun torchrun \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_args[@]}"
"${dmd_args[@]}"
@@ -1,3 +1,3 @@
#!/bin/bash
python scripts/huggingface/download_hf.py --repo_id "FastVideo/Wan-Syn_77x448x832_600k" --local_dir "FastVideo/Wan-Syn_77x448x832_600k" --repo_type "dataset"
python scripts/huggingface/download_hf.py --repo_id "FastVideo/Wan-Syn_77x448x832_600k" --local_dir "FastVideo/Wan-Syn_77x448x832_600k" --repo_type "dataset"
@@ -40,15 +40,18 @@ echo "NODE_RANK: $NODE_RANK"
# Configs
NUM_GPUS=8
MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
DATA_DIR=your_data_dir
VALIDATION_DIR=your_validation_path #(example:validation_64.json)
OUTPUT_DIR="checkpoints/wan_t2v_finetune"
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name Wan_distillation
--output_dir "your_output_dir"
--output_dir "$OUTPUT_DIR"
--max_train_steps 4000
--train_batch_size 1
--train_sp_batch_size 1
@@ -73,6 +76,8 @@ parallel_args=(
model_args=(
--model_path $MODEL_PATH
--pretrained_model_name_or_path $MODEL_PATH
--real_score_model_path $REAL_SCORE_MODEL_PATH
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
)
# Dataset arguments
@@ -92,11 +97,11 @@ validation_args=(
# Optimizer arguments
optimizer_args=(
--learning_rate 2e-5
--learning_rate 4e-6
--lr_scheduler "cosine_with_min_lr"
--min_lr_ratio 0.5
--lr_warmup_steps 100
--fake_score_learning_rate 1e-5
--fake_score_learning_rate 2e-6
--fake_score_lr_scheduler "cosine_with_min_lr"
--mixed_precision "bf16"
--training_state_checkpointing_steps 500
@@ -141,4 +146,4 @@ srun torchrun \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_args[@]}"
"${dmd_args[@]}"
@@ -40,6 +40,8 @@ echo "NODE_RANK: $NODE_RANK"
# Configs
NUM_GPUS=8
MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
DATA_DIR=your_data_dir
VALIDATION_DIR=your_validation_path #(example:validation_64.json)
# export CUDA_VISIBLE_DEVICES=4,5
@@ -73,6 +75,8 @@ parallel_args=(
model_args=(
--model_path $MODEL_PATH
--pretrained_model_name_or_path $MODEL_PATH
--real_score_model_path $REAL_SCORE_MODEL_PATH
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
)
# Dataset arguments
@@ -92,11 +96,11 @@ validation_args=(
# Optimizer arguments
optimizer_args=(
--learning_rate 2e-5
--learning_rate 4e-6
--lr_scheduler "cosine_with_min_lr"
--min_lr_ratio 0.5
--lr_warmup_steps 100
--fake_score_learning_rate 1e-5
--fake_score_learning_rate 2e-6
--fake_score_lr_scheduler "cosine_with_min_lr"
--mixed_precision "bf16"
--training_state_checkpointing_steps 500
@@ -142,4 +146,4 @@ srun torchrun \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_args[@]}"
"${dmd_args[@]}"
@@ -14,26 +14,29 @@ export FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN
# Configs
NUM_GPUS=1
MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
DATA_DIR="data/crush-smol_processed_ti2v/combined_parquet_dataset/"
VALIDATION_DATASET_FILE="examples/distill/Wan2.2-TI2V-5B-Diffusers/crush_smol/validation.json"
OUTPUT_DIR="checkpoints/wan_t2v_finetune"
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name wan_t2v_distill_dmd_VSA
--output_dir="checkpoints/wan_t2v_finetune"
--max_train_steps=4000
--train_batch_size=1
--output_dir "$OUTPUT_DIR"
--max_train_steps 4000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps=1
--gradient_accumulation_steps 1
--num_latent_t 31
--num_height 704
--num_width 1280
--num_frames 121
--enable_gradient_checkpointing_type "full"
--training_state_checkpointing_steps=500
--weight_only_checkpointing_steps=500
--training_state_checkpointing_steps 500
--weight_only_checkpointing_steps 500
)
# Parallel arguments
@@ -49,6 +52,8 @@ parallel_args=(
model_args=(
--model_path $MODEL_PATH
--pretrained_model_name_or_path $MODEL_PATH
--real_score_model_path $REAL_SCORE_MODEL_PATH
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
)
# Dataset arguments
@@ -68,8 +73,8 @@ validation_args=(
# Optimizer arguments
optimizer_args=(
--learning_rate=1e-5
--mixed_precision="bf16"
--learning_rate 2e-6
--mixed_precision "bf16"
--weight_decay 0.01
--max_grad_norm 1.0
)
@@ -107,4 +112,4 @@ torchrun \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_args[@]}"
"${dmd_args[@]}"
@@ -14,6 +14,8 @@ export FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN
# Configs
NUM_GPUS=1
MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
DATA_DIR="data/crush-smol_processed_ti2v/combined_parquet_dataset/"
VALIDATION_DATASET_FILE="examples/distill/Wan2.2-TI2V-5B-Diffusers/crush_smol/validation.json"
# export CUDA_VISIBLE_DEVICES=4,5
@@ -51,6 +53,8 @@ parallel_args=(
model_args=(
--model_path $MODEL_PATH
--pretrained_model_name_or_path $MODEL_PATH
--real_score_model_path $REAL_SCORE_MODEL_PATH
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
)
# Dataset arguments
@@ -109,4 +113,4 @@ torchrun \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_args[@]}"
"${dmd_args[@]}"
@@ -1,3 +1,3 @@
#!/bin/bash
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
@@ -21,4 +21,4 @@ torchrun --nproc_per_node=$GPU_NUM \
--samples_per_file 8 \
--flush_frequency 8 \
--video_length_tolerance_range 5 \
--preprocess_task "t2v"
--preprocess_task "t2v"
+44
View File
@@ -0,0 +1,44 @@
from fastvideo import VideoGenerator
# from fastvideo.configs.sample import SamplingParam
OUTPUT_PATH = "video_samples"
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(
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
# FastVideo will automatically handle distributed setup
num_gpus=2,
use_fsdp_inference=True,
dit_cpu_offload=False,
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
distributed_executor_backend="ray",
# image_encoder_cpu_offload=False,
)
# Generate videos with the same simple API, regardless of GPU count
prompt = (
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
)
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True)
# Generate another video with a different prompt, without reloading the
# model!
prompt2 = (
"A majestic lion strides across the golden savanna, its powerful frame "
"glistening under the warm afternoon sun. The tall grass ripples gently in "
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
"cinematic.")
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True)
if __name__ == "__main__":
main()
@@ -0,0 +1,36 @@
from fastvideo import VideoGenerator
# from fastvideo.configs.sample import SamplingParam
OUTPUT_PATH = "video_samples_wan2_1_Fun"
OUTPUT_NAME = "wan2.1_test"
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(
"IRMChen/Wan2.1-Fun-1.3B-Control-Diffusers",
# "alibaba-pai/Wan2.2-Fun-A14B-Control",
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=True,
dit_cpu_offload=True, # DiT need to be offloaded for MoE
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
# 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,
)
prompt = "一位年轻女性穿着一件粉色的连衣裙,裙子上有白色的装饰和粉色的纽扣。她的头发是紫色的,头上戴着一个红色的大蝴蝶结,显得非常可爱和精致。她还戴着一个红色的领结,整体造型充满了少女感和活力。她的表情温柔,双手轻轻交叉放在身前,姿态优雅。背景是简单的灰色,没有任何多余的装饰,使得人物更加突出。她的妆容清淡自然,突显了她的清新气质。整体画面给人一种甜美、梦幻的感觉,仿佛置身于童话世界中。"
negative_prompt = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
# prompt = "A young woman with beautiful, clear eyes and blonde hair stands in the forest, wearing a white dress and a crown. Her expression is serene, reminiscent of a movie star, with fair and youthful skin. Her brown long hair flows in the wind. The video quality is very high, with a clear view. High quality, masterpiece, best quality, high resolution, ultra-fine, fantastical."
# negative_prompt = "Twisted body, limb deformities, text captions, comic, static, ugly, error, messy code."
image_path = "https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset_Wan2_2/v1.0/8.png"
control_video_path = "https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset_Wan2_2/v1.0/pose.mp4"
video = generator.generate_video(prompt, negative_prompt=negative_prompt, image_path=image_path, video_path=control_video_path, output_path=OUTPUT_PATH, output_video_name=OUTPUT_NAME, save_video=True)
if __name__ == "__main__":
main()
+56
View File
@@ -0,0 +1,56 @@
# FastVideo Gradio Local Demo
This is a Gradio-based web interface for generating videos using the FastVideo framework. The demo allows users to create videos from text prompts with various customization options.
## Overview
The demo uses the FastVideo framework to generate videos based on text prompts. It provides a simple web interface built with Gradio that allows users to:
- Enter text prompts to generate videos
- Customize video parameters (dimensions, number of frames, etc.)
- Use negative prompts to guide the generation process
- Set or randomize seeds for reproducibility
---
## Usage
Run the demo with:
```bash
python examples/inference/gradio/local/gradio_local_demo.py
```
This will start a web server at `http://0.0.0.0:7860` where you can access the interface.
---
## Model Initialization
This demo initializes a `VideoGenerator` with the minimum required arguments for inference. Users can seamlessly adjust inference options between generations, including prompts, resolution, video length, *without ever needing to reload the model*.
## Video Generation
The core functionality is in the `generate_video` function, which:
1. Processes user inputs
2. Uses the FastVideo VideoGenerator from earlier to run inference (`generator.generate_video()`)
## Gradio Interface
The interface is built with several components:
- A text input for the prompt
- A video display for the result
- Inference options in a collapsible accordion:
- Height and width sliders
- Number of frames slider
- Guidance scale slider
- Negative prompt options
- Seed controls
### Inference Options
- **Height/Width**: Control the resolution of the generated video
- **Number of Frames**: Set how many frames to generate
- **Guidance Scale**: Control how closely the generation follows the prompt
- **Negative Prompt**: Specify what you don't want to see in the video
- **Seed**: Control randomness for reproducible results
@@ -0,0 +1,656 @@
import argparse
import os
import base64
import time
import gradio as gr
from fastvideo.entrypoints.video_generator import VideoGenerator
from fastvideo.configs.sample.base import SamplingParam
from copy import deepcopy
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",
}
def create_timing_display(inference_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;">N/A</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;">N/A</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 setup_model_environment(model_path: str) -> None:
if "fullattn" in model_path.lower():
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "FLASH_ATTN"
else:
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "VIDEO_SPARSE_ATTN"
os.environ["FASTVIDEO_STAGE_LOGGING"] = "1"
def load_example_prompts():
def contains_chinese(text):
return any('\u4e00' <= char <= '\u9fff' for char in text)
def load_from_file(filepath):
prompts, labels = [], []
try:
with open(filepath, "r", encoding='utf-8') as f:
for line in f:
line = line.strip()
if line and not contains_chinese(line):
label = line[:100] + "..." if len(line) > 100 else line
labels.append(label)
prompts.append(line)
except Exception as e:
print(f"Warning: Could not read {filepath}: {e}")
return prompts, labels
examples, example_labels = load_from_file("examples/inference/gradio/local/prompts_final.txt")
if not examples:
examples = ["A crowded rooftop bar buzzes with energy, the city skyline twinkling like a field of stars in the background."]
example_labels = ["Crowded rooftop bar at night"]
return examples, example_labels
def create_gradio_interface(default_params: dict[str, SamplingParam], generators: dict[str, VideoGenerator]):
def generate_video(
prompt, negative_prompt, use_negative_prompt, seed, guidance_scale,
num_frames, height, width, randomize_seed, model_selection, progress
):
model_path = MODEL_PATH_MAPPING.get(model_selection, "FastVideo/FastWan2.1-T2V-1.3B-Diffusers")
setup_model_environment(model_path)
try:
if progress:
progress(0.1, desc="Loading model for local inference...")
generator = generators[model_path]
params = deepcopy(default_params[model_path])
total_start_time = time.time()
if progress:
progress(0.2, desc="Configuring parameters...")
params.prompt = prompt
params.seed = int(seed)
params.guidance_scale = guidance_scale
params.num_frames = int(num_frames)
params.height = int(height)
params.width = int(width)
if randomize_seed:
params.seed = torch.randint(0, 1000000, (1, )).item()
if use_negative_prompt and negative_prompt:
params.negative_prompt = negative_prompt
else:
params.negative_prompt = default_params[model_path].negative_prompt
if progress:
progress(0.4, desc="Generating video locally...")
output_dir = "outputs/"
os.makedirs(output_dir, exist_ok=True)
start_time = time.time()
result = generator.generate_video(prompt=prompt, sampling_param=params, save_video=True, return_frames=False)
inference_time = time.time() - start_time
logging_info = result.get("logging_info", None)
if logging_info:
stage_names = logging_info.get_execution_order()
stage_execution_times = [
logging_info.get_stage_info(stage_name).get("execution_time", 0.0)
for stage_name in stage_names
]
else:
stage_names = []
stage_execution_times = []
total_time = time.time() - total_start_time
timing_details=create_timing_display(inference_time=inference_time, total_time=total_time, stage_execution_times=stage_execution_times, num_frames=params.num_frames)
safe_prompt = params.prompt[:100].replace(' ', '_').replace('/', '_').replace('\\', '_')
video_filename = f"{params.prompt[:100]}.mp4"
output_path = os.path.join(output_dir, video_filename)
if progress:
progress(1.0, desc="Generation complete!")
return output_path, params.seed, timing_details
except Exception as e:
print(f"An error occurred during local generation: {e}")
return None, f"Generation failed: {str(e)}", ""
examples, example_labels = load_example_prompts()
theme = gr.themes.Base().set(
button_primary_background_fill="#2563eb",
button_primary_background_fill_hover="#1d4ed8",
button_primary_text_color="white",
slider_color="#2563eb",
checkbox_background_color_selected="#2563eb",
)
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,
}
return {
'height': 448,
'width': 832,
'num_frames': 61,
'guidance_scale': 3.0,
'seed': 1024,
}
initial_values = get_default_values("FastWan2.1-T2V-1.3B")
with gr.Blocks(title="FastWan", theme=theme) as demo:
gr.Image("assets/full.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>
</div>
""")
with gr.Accordion("🎥 What Is FastVideo?", open=False):
gr.HTML("""
<div style="padding: 20px; line-height: 1.6;">
<p style="font-size: 16px; margin-bottom: 15px;">
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.
</p>
</div>
""")
with gr.Row():
model_selection = gr.Dropdown(
choices=list(MODEL_PATH_MAPPING.keys()),
value="FastWan2.1-T2V-1.3B",
label="Select Model",
interactive=True
)
with gr.Row():
example_dropdown = gr.Dropdown(
choices=example_labels,
label="Example Prompts",
value=None,
interactive=True,
allow_custom_value=False
)
with gr.Row():
with gr.Column(scale=6):
prompt = gr.Text(
label="Prompt",
show_label=False,
max_lines=3,
placeholder="Describe your scene...",
container=False,
lines=3,
autofocus=True,
)
with gr.Column(scale=1, min_width=120, elem_classes="center-button"):
run_button = gr.Button("Run", variant="primary", size="lg")
with gr.Row():
with gr.Column():
error_output = gr.Text(label="Error", 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():
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,
)
seed = gr.Slider(
label="Seed",
minimum=0,
maximum=1000000,
step=1,
value=initial_values['seed'],
)
randomize_seed = gr.Checkbox(label="Randomize seed", value=False)
seed_output = gr.Number(label="Used Seed")
with gr.Column(scale=1, elem_classes="video-column"):
result = gr.Video(
label="Generated Video",
show_label=True,
height=466,
width=600,
container=True,
elem_classes="video-component"
)
gr.HTML("""
<style>
.center-button {
display: flex !important;
justify-content: center !important;
height: 100% !important;
padding-top: 1.4em !important;
}
.gradio-container {
max-width: 1200px !important;
margin: 0 auto !important;
}
.main {
max-width: 1200px !important;
margin: 0 auto !important;
}
.gr-form, .gr-box, .gr-group {
max-width: 1200px !important;
}
.gr-video {
max-width: 500px !important;
margin: 0 auto !important;
}
.main-content-row {
display: flex !important;
align-items: flex-start !important;
min-height: 500px !important;
gap: 20px !important;
}
.advanced-options-column,
.video-column {
display: flex !important;
flex-direction: column !important;
flex: 1 !important;
min-height: 400px !important;
align-items: stretch !important;
}
.video-column > * {
margin-top: 0 !important;
}
.video-column .gr-video,
.video-component {
margin-top: 0 !important;
padding-top: 0 !important;
}
.video-column .gr-video .gr-form {
margin-top: 0 !important;
}
.advanced-options-column .gr-group,
.video-column .gr-video {
margin-top: 0 !important;
vertical-align: top !important;
}
.advanced-options-column > *:last-child,
.video-column > *:last-child {
flex-grow: 0 !important;
}
@media (max-width: 1400px) {
.main-content-row {
min-height: 600px !important;
}
.advanced-options-column,
.video-column {
min-height: 600px !important;
}
}
@media (max-width: 1200px) {
.main-content-row {
flex-direction: column !important;
align-items: stretch !important;
}
.advanced-options-column,
.video-column {
min-height: auto !important;
width: 100% !important;
}
}
.timing-card {
background: var(--background-fill-secondary) !important;
border: 1px solid var(--border-color-primary) !important;
color: var(--body-text-color) !important;
padding: 10px;
border-radius: 8px;
text-align: center;
min-height: 80px;
display: flex;
flex-direction: column;
justify-content: center;
}
.timing-card-highlight {
background: var(--background-fill-primary) !important;
border: 2px solid var(--color-accent) !important;
}
.performance-card {
background: var(--background-fill-secondary) !important;
border: 1px solid var(--border-color-primary) !important;
color: var(--body-text-color) !important;
padding: 10px;
border-radius: 6px;
text-align: center;
}
.gr-number input[readonly] {
background-color: var(--background-fill-secondary) !important;
border: 1px solid var(--border-color-primary) !important;
color: var(--body-text-color-subdued) !important;
cursor: default !important;
text-align: center !important;
font-weight: 500 !important;
}
</style>
""")
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 ""
example_dropdown.change(
fn=on_example_select,
inputs=example_dropdown,
outputs=prompt,
)
gr.HTML("""
<div style="text-align: center; margin-top: 10px; margin-bottom: 15px;">
<p style="font-size: 16px; margin: 0;">Note that this demo is meant to showcase FastWan's quality and that under a large number of requests, generation speed may be affected.</p>
</div>
""")
use_negative_prompt.change(
fn=lambda x: gr.update(visible=x),
inputs=use_negative_prompt,
outputs=negative_prompt,
)
def on_model_selection_change(selected_model):
if not selected_model:
selected_model = "FastWan2.1-T2V-1.3B"
model_path = MODEL_PATH_MAPPING.get(selected_model)
if model_path and model_path in default_params:
params = default_params[model_path]
return (
gr.update(value=params.height),
gr.update(value=params.width),
gr.update(value=params.num_frames),
gr.update(value=params.guidance_scale),
gr.update(value=params.seed),
)
return (
gr.update(value=448),
gr.update(value=832),
gr.update(value=61),
gr.update(value=3.0),
gr.update(value=1024),
)
model_selection.change(
fn=on_model_selection_change,
inputs=model_selection,
outputs=[height, width, num_frames, guidance_scale, seed],
)
def handle_generation(*args, progress=None, request: gr.Request = None):
model_selection, prompt, negative_prompt, use_negative_prompt, seed, guidance_scale, num_frames, height, width, randomize_seed = 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
)
if result_path and os.path.exists(result_path):
return (
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(
fn=handle_generation,
inputs=[
model_selection,
prompt,
negative_prompt,
use_negative_prompt,
seed,
guidance_scale,
num_frames,
height,
width,
randomize_seed,
],
outputs=[result, seed_output, error_output, timing_display],
concurrency_limit=20,
)
return demo
def main():
parser = argparse.ArgumentParser(description="FastVideo Gradio Local Demo")
parser.add_argument("--t2v_model_paths", type=str,
default="FastVideo/FastWan2.1-T2V-1.3B-Diffusers",
help="Comma separated list of paths to the T2V 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,
help="Port to bind to")
args = parser.parse_args()
generators = {}
default_params = {}
model_paths = args.t2v_model_paths.split(",")
for model_path in model_paths:
print(f"Loading model: {model_path}")
setup_model_environment(model_path)
generators[model_path] = VideoGenerator.from_pretrained(model_path)
default_params[model_path] = SamplingParam.from_pretrained(model_path)
demo = create_gradio_interface(default_params, generators)
print(f"Starting Gradio frontend at http://{args.host}:{args.port}")
print(f"T2V Models: {args.t2v_model_paths}")
from fastapi import FastAPI, Request, HTTPException
from fastapi.responses import HTMLResponse, FileResponse
import uvicorn
app = FastAPI()
@app.get("/logo.png")
def get_logo():
return FileResponse(
"assets/full.svg",
media_type="image/svg+xml",
headers={
"Cache-Control": "public, max-age=3600",
"Access-Control-Allow-Origin": "*"
}
)
@app.get("/favicon.ico")
def get_favicon():
favicon_path = "assets/icon-simple.svg"
if os.path.exists(favicon_path):
return FileResponse(
favicon_path,
media_type="image/svg+xml",
headers={
"Cache-Control": "public, max-age=3600",
"Access-Control-Allow-Origin": "*"
}
)
else:
raise HTTPException(status_code=404, detail="Favicon not found")
@app.get("/", response_class=HTMLResponse)
def index(request: Request):
base_url = str(request.base_url).rstrip('/')
return f"""
<!DOCTYPE html>
<html lang="en">
<head>
<meta charset="UTF-8" />
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
<title>FastWan</title>
<meta name="title" content="FastWan">
<meta name="description" content="Make video generation go blurrrrrrr">
<meta name="keywords" content="FastVideo, video generation, AI, machine learning, FastWan">
<meta property="og:type" content="website">
<meta property="og:url" content="{base_url}/">
<meta property="og:title" content="FastWan">
<meta property="og:description" content="Make video generation go blurrrrrrr">
<meta property="og:image" content="{base_url}/logo.png">
<meta property="og:image:width" content="1200">
<meta property="og:image:height" content="630">
<meta property="og:site_name" content="FastWan">
<meta property="twitter:card" content="summary_large_image">
<meta property="twitter:url" content="{base_url}/">
<meta property="twitter:title" content="FastWan">
<meta property="twitter:description" content="Make video generation go blurrrrrrr">
<meta property="twitter:image" content="{base_url}/logo.png">
<link rel="icon" type="image/png" sizes="32x32" href="/favicon.ico">
<link rel="icon" type="image/png" sizes="16x16" href="/favicon.ico">
<link rel="apple-touch-icon" href="/favicon.ico">
<style>
body, html {{
margin: 0;
padding: 0;
height: 100%;
overflow: hidden;
}}
iframe {{
width: 100%;
height: 100vh;
border: none;
}}
</style>
</head>
<body>
<iframe src="/gradio" width="100%" height="100%" style="border: none;"></iframe>
</body>
</html>
"""
app = gr.mount_gradio_app(
app,
demo,
path="/gradio",
allowed_paths=[os.path.abspath("outputs"), os.path.abspath("fastvideo-logos")]
)
uvicorn.run(app, host=args.host, port=args.port)
if __name__ == "__main__":
main()
@@ -0,0 +1,11 @@
A dynamic shot of a sleek black motorcycle accelerating down an empty highway at sunset. The bike's engine roars as it gains speed, smoke trailing from the tires. The rider, wearing a black leather jacket and helmet, leans forward with determination, gripping the handlebars tightly. The camera follows the motorcycle from a distance, capturing the dust kicked up behind it, then zooms in to show the intense focus on the rider's face. The background showcases the endless road stretching into the horizon with vibrant orange and pink hues of the setting sun. Medium shot transitioning to close-up.
A Jedi Master Yoda, recognizable by his green skin, large ears, and wise wrinkles, is performing on a small stage, strumming a guitar with great concentration. Yoda wears a casual robe and sits on a stool, his eyes closed as he plays, fully immersed in the music. The stage is dimly lit with spotlights highlighting Yoda, creating a mystical atmosphere. The background shows a live audience watching intently. Medium close-up shot focusing on Yoda's expressive face and hands moving gracefully over the guitar strings.
A cute, fluffy panda bear is preparing a meal in a cozy, modern kitchen. The panda is standing at a wooden countertop, wearing a white chef’s hat and apron. It skillfully stirs a pot on the stove with one hand while holding a spatula in the other. The kitchen is well-lit, with appliances and cabinets in pastel colors, creating a warm and inviting atmosphere. The panda moves gracefully, with a focused and determined expression, as steam rises from the pot. Medium shot focusing on the panda’s actions at the stove.
In a futuristic Tokyo rooftop during a heavy rainstorm, a robotic DJ stands behind a turntable, spinning vinyl records in a cyberpunk night setting. The robot has metallic, sleek body parts with glowing blue LED lights, and it moves gracefully with the beat. Raindrops create a shimmering effect as they hit the ground and the DJ. The surrounding environment features neon signs, towering skyscrapers, and a dark, misty atmosphere. The camera starts with a wide shot of the city skyline before zooming in on the DJ performing. Sci-fi, fantasy.
A realistic animated scene featuring a polar bear playing a guitar. The polar bear is standing upright, wearing a cozy fur vest and fingerless gloves. It holds the guitar with both hands, strumming the strings with one hand while plucking them with the other, showcasing natural, fluid motions. The polar bear's expressive face shows concentration and joy as it plays. The background is a snowy Arctic landscape with icebergs and a clear blue sky. The scene captures the bear from a mid-shot angle, focusing on its interaction with the guitar.
The scene opens to a breathtaking view of a tranquil ocean horizon at dusk, displaying a vibrant tapestry of oranges, pinks, and purples as the sun sets. In the foreground, tall, swaying palm trees frame the scene, their silhouettes stark against the colorful sky. The ocean itself shimmers with reflections of the sunset, creating a peaceful, almost ethereal atmosphere. A small boat can be seen in the distance, centered on the horizon, adding a sense of scale and solitude to the scene. The waves gently lap the shore, creating faint patterns on the sandy beach, which stretches across the foreground. Above, the sky is dotted with scattered clouds that catch the last light of the day, enhancing the drama and beauty of the scene. The overall mood is serene and contemplative, capturing a perfect moment of nature’s grandeur.
A large, modern semi-truck accelerating down an empty highway, gaining speed with each second. The truck's powerful engine roars as it moves forward, smoke billowing from the tires. The camera starts from a wide shot, capturing the truck in the distance, then smoothly zooms in to follow the vehicle as it speeds up. The truck's headlights illuminate the road ahead, casting a bright glow. The truck driver can be seen through the windshield, focused and determined. The background shows the vast openness of the highway stretching into the horizon under a clear blue sky. Medium to close-up shots of the truck as it accelerates.
Soft blue light pulses from the blade’s rune-etched hilt, illuminating nearby moss-covered roots and ferns. The surrounding trees are tall and gnarled, their branches curling like claws overhead. Fog swirls gently at ground level, parting slightly as a figure in a cloak approaches from the distance. Medium shot slowly zooming toward the sword, emphasizing its mystical aura.
The video opens with a tranquil scene in the heart of a dense forest, emphasizing two large, textured tree trunks in the foreground framing the view. Sunlight filters through the canopy above, casting intricate patterns of light and shadow on the trees and the ground. Between the tree trunks, a clear view of a calm, muddy river unfolds, its surface shimmering under the gentle sunlight. The riverbank is decorated with a variety of small bushes and vibrant foliage, subtly transitioning into the deep greens of tall, leafy plants. In the background, the dense forest looms, filled with dark, towering trees, their branches intertwining to form an intricate canopy. The scene is bathed in the soft glow of the sun, creating a serene and picturesque setting. Occasional sunbeams pierce through the foliage, adding a magical aura to the landscape. The vibrant reds and oranges of the smaller plants add contrast, bringing warmth to the earthy tones of the scenery. Overall, this harmonious blend of natural elements creates a peaceful and idyllic forest setting.
A lone figure stands on a large, moss-covered rock, surrounded by the soft rush of a nearby stream. The figure is wearing white sneakers and shorts, with a plaid shirt that hangs loosely in the breeze. The lighting creates dramatic shadows, enhancing the textures of the rock and the subtle movement of the water below. In the background, a waterfall cascades into the stream, completing this tranquil and serene nature scene.
In an industrial setting, a person leans casually against a railing, exuding a sense of confidence and composure. They are wearing a striking outfit, consisting of a vibrant, patterned jacket over a simple white crop top, creating a bold contrast. The atmosphere is infused with warm, ambient lighting that casts soft shadows on the concrete walls and metallic surfaces. Intricate wiring and pipes form an intricate backdrop, enhancing the urban aesthetic. Their relaxed posture and direct, engaging gaze suggest a sense of ease in this industrial environment. This scene encapsulates a blend of modern fashion and gritty, urban architecture, creating a visually compelling narrative.
@@ -0,0 +1,47 @@
A watermelon wearing a helmet is crushed by a hydraulic press, causing it to flatten and burst open.
The video shows a green and orange object being flattened as if it were under a hydraulic press, with the press moving down and compressing the object.
The video shows a cylindrical object with a cityscape image being flattened as if it were under a hydraulic press. The object is placed on a metal platform, and a large, striped cylinder presses down on it, causing it to collapse and release a liquid inside. The background features a green wall with a yellow and red warning sign.
A red toy car is being crushed by a large hydraulic press, which is flattening objects as if they were under a hydraulic press.
A large, cylindrical object is seen pressing down on a small orange ball, causing it to flatten as if it were under a hydraulic press. The background features a green wall with yellow and red warning signs.
The video shows a hydraulic press in action, flattening objects as if they were under a hydraulic press. The press is shown compressing a wooden object, which shatters into small pieces. The background features a green wall with a yellow sign displaying a lightning bolt.
A large metal cylinder is seen descending, flattening objects as if they were under a hydraulic press. The cylinder compresses a stack of matches and boxes, causing them to crumble into small pieces. The scene is set against a green background with yellow and red signs.
A large metal press is shown compressing a pile of colorful macarons, flattening them as if they were under a hydraulic press. The press moves down, crushing the macarons into a pile of crumbs and squishing the colorful filling out.
The video shows a metal press flattening objects as if they were under a hydraulic press. The press is pressing down on a pile of colorful gummy candies, squishing them into a pile of squiggly shapes. The press is made of metal and has a large base, and the gummy candies are of various colors, including red, green, and orange. The background is a green wall, and the press is placed on a metal surface.
A pile of colorful candies is being flattened by a hydraulic press, causing them to crumble into small pieces.
The video shows a stack of colorful sponges being flattened as if they were under a hydraulic press. The sponges, which are pink, white, blue, and green, are compressed into a smaller size, demonstrating the press's power. The background features a green wall with a yellow and red sign, adding context to the setting.
A bowling ball is placed on a metal platform, and a large metal cylinder descends from above, flattening the ball as if it were under a hydraulic press. The ball is crushed into a flat, round shape, leaving a pile of debris around it.
A large metal cylinder with yellow and black stripes is seen pressing down on a pile of popcorn, flattening the objects as if they were under a hydraulic press.
The video shows a close-up of an orange being flattened as if it were under a hydraulic press, with the press moving down and compressing the fruit until it is completely flattened.
The video shows a close-up of a metal cylinder pressing down on a yellow object, which is being flattened as if it were under a hydraulic press. The cylinder is positioned above the object, and the force is causing the object to compress and spread out, creating a visible deformation. The background is blurred, focusing attention on the action of the cylinder and the object being flattened.
A colorful puzzle ball is being crushed by a large metal cylinder, which flattens the objects as if they were under a hydraulic press.
The video shows a hydraulic press flattening objects as if they were under a hydraulic press. The press is shown in action, compressing two colorful objects that resemble sandwiches. The press is yellow and black striped, and the objects being flattened are placed on a metal plate. The background is green, and the press is moving down, compressing the objects.
The scene shows a metal press with a yellow and black striped pattern, holding a container filled with chocolate. A metal cylinder is descending, flattening the chocolate as if it were under a hydraulic press. The background is a green wall, and the press is mounted on a sturdy metal frame.
The video shows a colorful sponge being flattened as if it were under a hydraulic press, with the sponge being compressed and eventually flattened into a thin layer.
The video shows a hydraulic press in action, flattening objects as if they were under a hydraulic press. The press is pressing down on a stack of wooden blocks, causing them to crumble and break apart. The press is black and yellow striped, and the wooden blocks are small and rectangular. The background is green, and the press is sitting on a metal table.
A pile of colorful candies is being flattened by a hydraulic press, causing them to crumble into small pieces.
The video shows a stack of colorful sponges being flattened by a large, cylindrical object, which appears to be a hydraulic press. The sponges, which are pink, blue, white, and green, are compressed into a single layer, demonstrating the press's powerful force. The background features a green wall with a yellow and red sign, adding context to the industrial setting.
A bowling ball is placed on a metal platform, and a large metal cylinder descends from above, flattening the ball as if it were under a hydraulic press. The ball is crushed into a flat, round shape, demonstrating the immense pressure applied by the cylinder.
A large metal cylinder with yellow and black stripes is seen pressing down on a pile of popcorn, flattening the objects as if they were under a hydraulic press. The popcorn is crushed and scattered around the base of the cylinder, creating a satisfying visual effect.
The video shows a hydraulic press in action, flattening objects as if they were under a hydraulic press. The press is composed of a large, cylindrical metal cylinder with yellow and black stripes, and a metal base. The objects being flattened are two cylindrical blocks of cotton candy, one pink and one blue. The press is positioned on a metal table, and the background features a green wall with a yellow and red sign.
The video shows a large orange being flattened as if it were under a hydraulic press, with the press moving down and compressing the fruit until it is completely flattened.
The video shows a cylindrical object being pressed down onto a flat surface, causing the objects beneath it to be flattened as if they were under a hydraulic press. The objects being flattened appear to be yellow and are being crushed into a pile of debris. The background is a greenish-gray color, and the surface on which the objects are being flattened is metallic and shiny.
A green and blue object with a spiky texture is being flattened by a large, cylindrical metal press, demonstrating its resilience and durability.
The video shows a stack of caramelized sugar cubes being flattened as if they were under a hydraulic press, resulting in a messy pile of broken sugar on the table.
A large metal cylinder is seen pressing down on a pile of colorful jelly beans, flattening them as if they were under a hydraulic press.
The video shows a machine with a yellow and black striped cylinder pressing down on a stack of colorful sponges, flattening them as if they were under a hydraulic press. The machine is situated in a green-walled room with warning signs in the background.
The video shows a machine with a yellow and black striped cylinder, which is pressing down on two colorful objects, flattening them as if they were under a hydraulic press. The machine appears to be in a workshop or industrial setting, with a green wall in the background. The objects being flattened are green and orange, and the machine is covered in dirt and grime, indicating it has been used frequently.
The video shows a large, industrial press flattening objects as if they were under a hydraulic press. The press is shown in action, compressing a pile of pink objects into a pile of crumbs. The press is large and metallic, with a yellow and black striped pattern on its side. The background is a green wall with a yellow warning sign.
The video shows a pink, sparkly ball being crushed by a large, rusty cylinder, which flattens the objects as if they were under a hydraulic press.
A lime is being crushed by a hydraulic press, causing it to flatten and burst open, releasing its juice and segments.
The video shows a machine with a yellow and black striped cylinder, which is flattening objects as if they were under a hydraulic press. The machine is pressing down on two colorful objects, causing them to compress and flatten. The background is a green wall, and the machine appears to be in a workshop or industrial setting.
The video shows a large, yellow and black striped cylinder flattening objects as if they were under a hydraulic press. The objects being flattened are pink and are being crushed into small pieces. The background is a green wall with a yellow sign.
The video shows a machine with a yellow and black striped cylinder pressing down on two colorful objects, which are flattened as if they were under a hydraulic press. The machine is positioned on a metal platform, and the background is a green wall.
A green cube is being compressed by a hydraulic press, which flattens the object as if it were under a hydraulic press. The press is shown in action, with the cube being squeezed into a smaller shape.
A pink, sparkly ball is being crushed by a large, rusty cylinder, which flattens the objects as if they were under a hydraulic press.
A red cabbage is being crushed by a hydraulic press, which flattens the objects as if they were under a hydraulic press. The press is shown in action, compressing the cabbage into a smaller, more compact form.
A lime is being crushed by a hydraulic press, causing it to flatten and burst open, releasing its juice and pulp.
A large metal press is shown compressing a stack of burgers, causing them to be flattened and crushed into a pile of ground meat.
A pizza is being crushed by a hydraulic press, causing the toppings to spread out and the crust to crumble.
A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.
A large metal cylinder is seen pressing down on a pile of colorful candies, flattening them as if they were under a hydraulic press. The candies are crushed and broken into small pieces, creating a mess on the table.
A large metal cylinder is seen pressing down on a pile of Oreo cookies, flattening them as if they were under a hydraulic press.
@@ -0,0 +1,94 @@
#!/bin/bash
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export TOKENIZERS_PARALLELISM=false
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATA_DIR="data/crush-smol_processed_t2v_1_3b_ode_init/"
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
NUM_GPUS=1
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name "wan_ode_init"
--output_dir "wan_ode_init_crush_smol"
--override_transformer_cls_name "CausalWanTransformer3DModel"
--wandb_run_name "wan_ode_init_crush_smol"
--max_train_steps 6000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 21
--num_height 480
--num_width 832
--num_frames 77
--warp_denoising_step
--enable_gradient_checkpointing_type "full"
)
# Parallel arguments
parallel_args=(
--num_gpus $NUM_GPUS
--sp_size 1
--tp_size 1
--hsdp_replicate_dim 1
--hsdp_shard_dim 1
)
# Model arguments
model_args=(
--model_path $MODEL_PATH
--pretrained_model_name_or_path $MODEL_PATH
)
# Dataset arguments
dataset_args=(
--data_path "$DATA_DIR"
--dataloader_num_workers 1
)
# Validation arguments
validation_args=(
--log_validation
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 50
--validation_sampling_steps "50"
--validation_guidance_scale "6.0"
)
# Optimizer arguments
optimizer_args=(
--learning_rate 6e-6
--mixed_precision "bf16"
--weight_only_checkpointing_steps 1000
--training_state_checkpointing_steps 1000
--weight_decay 1e-4
--max_grad_norm 1.0
)
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
--training_cfg_rate 0.1
--multi_phased_distill_schedule "4000-1"
--not_apply_cfg_solver
--dit_precision "fp32"
--num_euler_timesteps 50
--ema_start_step 0
)
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
torchrun \
--nnodes 1 \
--nproc_per_node $NUM_GPUS \
fastvideo/training/ode_causal_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}"
@@ -0,0 +1,136 @@
#!/bin/bash
#SBATCH --job-name=2e6B8_16kFV_ode_vidprom
#SBATCH --partition=main
#SBATCH --nodes=1
#SBATCH --ntasks=1
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=128
#SBATCH --mem=1440G
#SBATCH --output=ode_vidprom16k/ode_vidprom8b16k_2e-6.out
#SBATCH --error=ode_vidprom16k/ode_vidprom8b16k_2e-6.err
#SBATCH --exclusive
set -e -x
# Environment Setup
source ~/conda/miniconda/bin/activate
conda activate your-conda-env
export WANDB_MODE="online"
export NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
# different cache dir for different processes
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
export MASTER_PORT=29500
export NODE_RANK=$SLURM_PROCID
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
export MASTER_ADDR=${nodes[0]}
export CUDA_VISIBLE_DEVICES=$SLURM_LOCALID
export TOKENIZERS_PARALLELISM=false
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_API_KEY=your-wandb-api-key
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
echo "MASTER_ADDR: $MASTER_ADDR"
echo "NODE_RANK: $NODE_RANK"
MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
DATA_DIR="your-data-dir"
VALIDATION_DATASET_FILE="examples/training/consistency_finetune/causal_ode_init/validation.json"
OUTPUT_DIR="your-output-dir"
INIT_WEIGHTS_FROM_SAFETENSORS="your-init-weights-from-safetensors" # bidirectional weights from Wan2.1-T2V-1.3B-Diffusers
NUM_GPUS=8
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name "wan_ode_init"
--output_dir $OUTPUT_DIR
--override_transformer_cls_name "CausalWanTransformer3DModel"
--wandb_run_name "vidprom_8b16k_ode_init_2e-6"
# --resume_from_checkpoint "ode_init_diffusers/"
--warp_denoising_step
--log_visualization
--max_train_steps 6001
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 21
--num_height 480
--num_width 832
--num_frames 81
--dmd_denoising_steps "1000,750,500,250"
--enable_gradient_checkpointing_type "full"
)
# Parallel arguments
parallel_args=(
--num_gpus $NUM_GPUS
--sp_size 1
--tp_size 1
--hsdp_replicate_dim $NUM_GPUS
--hsdp_shard_dim 1
)
# Model arguments
model_args=(
--model_path $MODEL_PATH
--pretrained_model_name_or_path $MODEL_PATH
)
# Dataset arguments
dataset_args=(
--data_path "$DATA_DIR"
--dataloader_num_workers 1
)
# Validation arguments
validation_args=(
--log_validation
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 50
--validation_sampling_steps "50"
--validation_guidance_scale "6.0"
--init_weights_from_safetensors $INIT_WEIGHTS_FROM_SAFETENSORS
)
# Optimizer arguments
optimizer_args=(
--learning_rate 2e-6
--mixed_precision "bf16"
--weight_only_checkpointing_steps 500
--training_state_checkpointing_steps 500
--weight_decay 1e-4
--max_grad_norm 1.0
)
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
--training_cfg_rate 0.1
--multi_phased_distill_schedule "4000-1"
--not_apply_cfg_solver
--dit_precision "fp32"
--num_euler_timesteps 50
--ema_start_step 0
# --enable_gradient_checkpointing_type "full"
)
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
srun torchrun \
--nnodes $SLURM_JOB_NUM_NODES \
--nproc_per_node $NUM_GPUS \
--node_rank $SLURM_PROCID \
--rdzv_backend=c10d \
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
fastvideo/training/ode_causal_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}"
@@ -0,0 +1,25 @@
#!/bin/bash
GPU_NUM=1 # 2,4,8
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
MODEL_TYPE="wan"
DATA_MERGE_PATH="$(dirname "$0")/crush_smol_prompts.txt"
OUTPUT_DIR="data/crush-smol_processed_t2v_1_3b_ode_init/"
torchrun --nproc_per_node=$GPU_NUM \
fastvideo/pipelines/preprocess/v1_preprocess.py \
--model_path $MODEL_PATH \
--data_merge_path $DATA_MERGE_PATH \
--preprocess_video_batch_size 1 \
--seed 42 \
--max_height 480 \
--max_width 832 \
--num_frames 81 \
--flow_shift 5.0 \
--dataloader_num_workers 0 \
--output_dir=$OUTPUT_DIR \
--train_fps 16 \
--samples_per_file 8 \
--flush_frequency 8 \
--video_length_tolerance_range 5 \
--preprocess_task "ode_trajectory"
@@ -0,0 +1,76 @@
{
"data": [
{
"caption": "A stylish woman walks down a Tokyo street filled with warm glowing neon and animated city signage. She wears a black leather jacket, a long red dress, and black boots, and carries a black purse. She wears sunglasses and red lipstick. She walks confidently and casually. The street is damp and reflective, creating a mirror effect of the colorful lights. Many pedestrians walk about.",
"image_path": null,
"video_path": null,
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "A white and orange tabby cat is seen happily darting through a dense garden, as if chasing something. Its eyes are wide and happy as it jogs forward, scanning the branches, flowers, and leaves as it walks. The path is narrow as it makes its way between all the plants. the scene is captured from a ground-level angle, following the cat closely, giving a low and intimate perspective. The image is cinematic with warm tones and a grainy texture. The scattered daylight between the leaves and plants above creates a warm contrast, accentuating the cat’s orange fur. The shot is clear and sharp, with a shallow depth of field.",
"image_path": null,
"video_path": null,
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "Elon Musk, dressed in a sleek white spacesuit with a reflective visor, walks confidently across the lunar surface. His posture is upright, and he moves steadily with purpose. The moon's rocky terrain and scattered boulders surround him, casting shadows under the dim sunlight. The background shows vast stretches of the moon's barren landscape with craters and dust clouds kicked up by his boots. The scene captures a wide shot, emphasizing the vastness and desolation of the lunar environment. ",
"image_path": null,
"video_path": null,
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "In a dynamic action-packed sequence set in the Marvel multiverse, Spider-Man and Venom engage in an intense battle. Spider-Man, in his classic red and blue suit, swings and dodges venomous attacks from the black symbiote-covered Venom. Both characters display a range of acrobatic moves and powerful strikes. The environment is a chaotic urban landscape with crumbling buildings and neon lights, reflecting the multiversal theme. The camera captures the epic fight from various angles, including wide shots to show the scale of destruction and close-ups to highlight their fierce expressions and physical combat. ",
"image_path": null,
"video_path": null,
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "A warm, family-oriented scene depicting a father getting ready to leave the house to buy milk. The father, a middle-aged man with a kind face and a casual outfit, picks up a jacket from the coat rack. His posture is upright as he bends down slightly to put on his shoes. In the background, there are glimpses of a cozy living room with a family photograph on the wall. The camera focuses closely on the father, capturing his gentle smile and reassuring nod towards the camera before he opens the front door and steps outside. Static medium close-up shot. ",
"image_path": null,
"video_path": null,
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "Close-up shot of a man with a prosthetic hand that functions as a rocket launcher. He looks at his new hand with a mix of amazement and concern, his facial expression showing a blend of curiosity and apprehension. The prosthetic hand is sleek and metallic, with intricate details that resemble a high-tech weapon. The background is a dimly lit laboratory with various scientific equipment and monitors displaying data. The man stands in a relaxed posture, his other hand resting on his hip, as he inspects his new limb. The scene is rendered in a realistic sci-fi style, emphasizing the futuristic technology and the man's emotional response to his new appendage. ",
"image_path": null,
"video_path": null,
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "Realistic CCTV footage style, Kim Taehyung from the band BTS is involved in a drug deal, caught on camera. Kim Taehyung appears nervous and cautious, wearing casual clothing typical of a public space. He exchanges items discreetly with another person, who is partially obscured. Both individuals maintain a watchful demeanor, occasionally glancing around to ensure no one is watching them. The lighting is dim, with flickering fluorescent lights casting shadows on their faces. The background shows a typical urban setting with blurred figures moving in the distance. Static camera angle, medium close-up shot focusing on the interaction between Taehyung and the other individual. ",
"image_path": null,
"video_path": null,
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "Photorealistic studio setup with professional lighting, showcasing detailed cubic dissections of experimental plastic and felt-like materials on a pristine white background. Each cube reveals intricate layers and textures of the materials, emphasizing their unique properties. The scene has a shallow depth of field initially, then slowly pulls out to reveal the full arrangement of cubes, maintaining a wide depth of field throughout the transition. ",
"image_path": null,
"video_path": null,
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
}
]
}
@@ -61,7 +61,8 @@ validation_args=(
optimizer_args=(
--learning_rate 2e-5
--mixed_precision "bf16"
--checkpointing_steps 2000
--weight_only_checkpointing_steps 2000
--training_state_checkpointing_steps 2000
--weight_decay 1e-4
--max_grad_norm 1.0
)
@@ -95,7 +95,8 @@ validation_args=(
optimizer_args=(
--learning_rate 1e-5
--mixed_precision "bf16"
--checkpointing_steps 1000
--weight_only_checkpointing_steps 1000
--training_state_checkpointing_steps 1000
--weight_decay 1e-4
--max_grad_norm 1.0
)
@@ -93,7 +93,8 @@ validation_args=(
optimizer_args=(
--learning_rate 1e-6
--mixed_precision "bf16"
--checkpointing_steps 1000
--weight_only_checkpointing_steps 1000
--training_state_checkpointing_steps 1000
--weight_decay 0.01
--max_grad_norm 1.0
)
@@ -0,0 +1,134 @@
#!/bin/bash
#SBATCH --job-name=t2v
#SBATCH --partition=main
#SBATCH --nodes=8
#SBATCH --ntasks=8
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=128
#SBATCH --mem=1440G
#SBATCH --output=VSA_t2v_output/t2v_%j.out
#SBATCH --error=VSA_t2v_output/t2v_%j.err
#SBATCH --exclusive
set -e -x
# Environment Setup
source ~/conda/miniconda/bin/activate
conda activate your_env
# Basic Info
export WANDB_MODE="online"
export NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
# different cache dir for different processes
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
export MASTER_PORT=29500
export NODE_RANK=$SLURM_PROCID
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
export MASTER_ADDR=${nodes[0]}
export CUDA_VISIBLE_DEVICES=$SLURM_LOCALID
export TOKENIZERS_PARALLELISM=false
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
echo "MASTER_ADDR: $MASTER_ADDR"
echo "NODE_RANK: $NODE_RANK"
# Configs
NUM_GPUS=8
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATA_DIR=your_data_dir
VALIDATION_DATASET_FILE=your_validation_dataset_file
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name wan_t2v_VSA
--output_dir "checkpoints/wan_t2v_finetune_VSA"
--max_train_steps 4000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 21
--num_height 480
--num_width 832
--num_frames 81
# --enable_gradient_checkpointing_type "full" # if OOM enable this
)
# Parallel arguments
parallel_args=(
--num_gpus 64
--sp_size 1
--tp_size 1
--hsdp_replicate_dim 64
--hsdp_shard_dim 1
)
# Model arguments
model_args=(
--model_path $MODEL_PATH
--pretrained_model_name_or_path $MODEL_PATH
)
# Dataset arguments
dataset_args=(
--data_path "$DATA_DIR"
--dataloader_num_workers 4
)
# Validation arguments
validation_args=(
--log_validation
--validation_dataset_file $VALIDATION_DATASET_FILE
--validation_steps 200
--validation_sampling_steps "50"
--validation_guidance_scale "5.0"
)
# Optimizer arguments
optimizer_args=(
--learning_rate 1e-6
--mixed_precision "bf16"
--weight_only_checkpointing_steps 1000
--training_state_checkpointing_steps 1000
--weight_decay 0.01
--max_grad_norm 1.0
)
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
--training_cfg_rate 0.1
--dit_precision "fp32"
--ema_start_step 0
--flow_shift 1
--seed 1000
)
# VSA arguments
vsa_args=(
--VSA_decay_rate 0.03 \
--VSA_decay_interval_steps 50 \
--VSA_sparsity 0.9 \
)
srun torchrun \
--nnodes $SLURM_JOB_NUM_NODES \
--nproc_per_node $NUM_GPUS \
--node_rank $SLURM_PROCID \
--rdzv_backend=c10d \
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
fastvideo/training/wan_training_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${vsa_args[@]}"
@@ -91,9 +91,10 @@ validation_args=(
# Optimizer arguments
optimizer_args=(
--learning_rate 1e-5
--learning_rate 1e-6
--mixed_precision "bf16"
--checkpointing_steps 1000
--weight_only_checkpointing_steps 1000
--training_state_checkpointing_steps 1000
--weight_decay 0.01
--max_grad_norm 1.0
)
@@ -61,7 +61,8 @@ validation_args=(
optimizer_args=(
--learning_rate 1e-5
--mixed_precision "bf16"
--checkpointing_steps 1000
--weight_only_checkpointing_steps 1000
--training_state_checkpointing_steps 1000
--weight_decay 1e-4
--max_grad_norm 1.0
)
@@ -95,7 +95,8 @@ validation_args=(
optimizer_args=(
--learning_rate 1e-5
--mixed_precision "bf16"
--checkpointing_steps 1000
--weight_only_checkpointing_steps 1000
--training_state_checkpointing_steps 1000
--weight_decay 1e-4
--max_grad_norm 1.0
)
@@ -61,7 +61,8 @@ validation_args=(
optimizer_args=(
--learning_rate 1e-5
--mixed_precision "bf16"
--checkpointing_steps 1000
--weight_only_checkpointing_steps 1000
--training_state_checkpointing_steps 1000
--weight_decay 1e-4
--max_grad_norm 1.0
)
@@ -61,7 +61,8 @@ validation_args=(
optimizer_args=(
--learning_rate 5e-5
--mixed_precision "bf16"
--checkpointing_steps 1000
--weight_only_checkpointing_steps 1000
--training_state_checkpointing_steps 1000
--weight_decay 1e-4
--max_grad_norm 1.0
)
@@ -92,7 +92,8 @@ validation_args=(
optimizer_args=(
--learning_rate 5e-5
--mixed_precision "bf16"
--checkpointing_steps 400
--weight_only_checkpointing_steps 400
--training_state_checkpointing_steps 400
--weight_decay 1e-4
--max_grad_norm 1.0
)
@@ -61,7 +61,8 @@ validation_args=(
optimizer_args=(
--learning_rate 5e-5
--mixed_precision "bf16"
--checkpointing_steps 400
--weight_only_checkpointing_steps 400
--training_state_checkpointing_steps 400
--weight_decay 1e-4
--max_grad_norm 1.0
)
@@ -28,4 +28,4 @@
"num_frames": 77
}
]
}
}
@@ -0,0 +1,72 @@
# SPDX-License-Identifier: Apache-2.0
import torch
from fastvideo.attention.backends.sageattn.api import sageattn_blackwell
from fastvideo.attention.backends.abstract import (AttentionBackend,
AttentionImpl,
AttentionMetadata,
AttentionMetadataBuilder)
from fastvideo.logger import init_logger
logger = init_logger(__name__)
class SageAttention3Backend(AttentionBackend):
accept_output_buffer: bool = True
@staticmethod
def get_supported_head_sizes() -> list[int]:
return [64, 128, 256]
@staticmethod
def get_name() -> str:
return "SAGE_ATTN_THREE"
@staticmethod
def get_impl_cls() -> type["SageAttention3Impl"]:
return SageAttention3Impl
@staticmethod
def get_metadata_cls() -> type["AttentionMetadata"]:
raise NotImplementedError
@staticmethod
def get_builder_cls() -> type["AttentionMetadataBuilder"]:
raise NotImplementedError
# @staticmethod
# def get_metadata_cls() -> Type["AttentionMetadata"]:
# return FlashAttentionMetadata
class SageAttention3Impl(AttentionImpl):
def __init__(
self,
num_heads: int,
head_size: int,
causal: bool,
softmax_scale: float,
num_kv_heads: int | None = None,
prefix: str = "",
**extra_impl_args,
) -> None:
self.causal = causal
self.softmax_scale = softmax_scale
self.dropout = extra_impl_args.get("dropout_p", 0.0)
def forward(
self,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
attn_metadata: AttentionMetadata,
) -> torch.Tensor:
query = query.transpose(1, 2)
key = key.transpose(1, 2)
value = value.transpose(1, 2)
output = sageattn_blackwell(query, key, value, is_causal=self.causal)
output = output.transpose(1, 2)
return output
+4 -4
View File
@@ -5,7 +5,6 @@ from dataclasses import dataclass
import torch
from einops import rearrange
from flash_attn.bert_padding import pad_input
from csrc.attn.vmoba_attn.vmoba import (moba_attn_varlen, process_moba_input,
process_moba_output)
@@ -134,6 +133,8 @@ class VMOBAAttentionImpl(AttentionImpl):
**extra_impl_args) -> None:
self.prefix = prefix
self.layer_idx = self._get_layer_idx(prefix)
from flash_attn.bert_padding import pad_input
self.pad_input = pad_input
def _get_layer_idx(self, prefix: str) -> int | None:
match = re.search(r"blocks\.(\d+)", prefix)
@@ -169,7 +170,6 @@ class VMOBAAttentionImpl(AttentionImpl):
moba_chunk_size = attn_metadata.st_chunk_size
moba_topk = attn_metadata.st_topk
# torch.distributed.breakpoint()
query, chunk_size = process_moba_input(query,
attn_metadata.patch_resolution,
moba_chunk_size)
@@ -205,8 +205,8 @@ class VMOBAAttentionImpl(AttentionImpl):
simsum_threshold=attn_metadata.moba_threshold,
threshold_type=attn_metadata.moba_threshold_type,
)
hidden_states = pad_input(hidden_states, indices_q, batch_size,
sequence_length)
hidden_states = self.pad_input(hidden_states, indices_q, batch_size,
sequence_length)
hidden_states = process_moba_output(hidden_states,
attn_metadata.patch_resolution,
moba_chunk_size)
+3 -1
View File
@@ -12,7 +12,7 @@ import torch
import fastvideo.envs as envs
from fastvideo.attention.backends.abstract import AttentionBackend
from fastvideo.logger import init_logger
from fastvideo.platforms import AttentionBackendEnum, current_platform
from fastvideo.platforms import AttentionBackendEnum
from fastvideo.utils import STR_BACKEND_ENV_VAR, resolve_obj_by_qualname
logger = init_logger(__name__)
@@ -117,6 +117,8 @@ def _cached_get_attn_backend(
selected_backend = backend_name_to_enum(backend_by_env_var)
# get device-specific attn_backend
from fastvideo.platforms import current_platform
if selected_backend not in supported_attention_backends:
selected_backend = None
attention_cls = current_platform.get_attn_backend_cls(
+5 -7
View File
@@ -15,18 +15,16 @@ class DiTArchConfig(ArchConfig):
reverse_param_names_mapping: dict = field(default_factory=dict)
lora_param_names_mapping: dict = field(default_factory=dict)
_supported_attention_backends: tuple[AttentionBackendEnum, ...] = (
AttentionBackendEnum.SLIDING_TILE_ATTN,
AttentionBackendEnum.SAGE_ATTN,
AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA,
AttentionBackendEnum.VIDEO_SPARSE_ATTN,
AttentionBackendEnum.VMOBA_ATTN,
)
AttentionBackendEnum.SLIDING_TILE_ATTN, AttentionBackendEnum.SAGE_ATTN,
AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA,
AttentionBackendEnum.VIDEO_SPARSE_ATTN, AttentionBackendEnum.VMOBA_ATTN,
AttentionBackendEnum.SAGE_ATTN_THREE)
hidden_size: int = 0
num_attention_heads: int = 0
num_channels_latents: int = 0
exclude_lora_layers: list[str] = field(default_factory=list)
boundary_ratio: float | None = None
def __post_init__(self) -> None:
if not self._compile_conditions:
@@ -2,13 +2,13 @@ from fastvideo.configs.models.encoders.base import (BaseEncoderOutput,
EncoderConfig,
ImageEncoderConfig,
TextEncoderConfig)
from fastvideo.configs.models.encoders.clip import (CLIPTextConfig,
CLIPVisionConfig)
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
__all__ = [
"EncoderConfig", "TextEncoderConfig", "ImageEncoderConfig",
"BaseEncoderOutput", "CLIPTextConfig", "CLIPVisionConfig", "LlamaConfig",
"T5Config"
"BaseEncoderOutput", "CLIPTextConfig", "CLIPVisionConfig",
"WAN2_1ControlCLIPVisionConfig", "LlamaConfig", "T5Config"
]
+12
View File
@@ -77,6 +77,8 @@ class CLIPTextConfig(TextEncoderConfig):
num_hidden_layers_override: int | None = None
require_post_norm: bool | None = None
enable_scale: bool = True
is_causal: bool = True
prefix: str = "clip"
@@ -87,4 +89,14 @@ class CLIPVisionConfig(ImageEncoderConfig):
num_hidden_layers_override: int | None = None
require_post_norm: bool | None = None
enable_scale: bool = True
is_causal: bool = True
prefix: str = "clip"
@dataclass
class WAN2_1ControlCLIPVisionConfig(CLIPVisionConfig):
num_hidden_layers_override: int | None = 31
require_post_norm: bool | None = False
enable_scale: bool = False
is_causal: bool = False
+3 -2
View File
@@ -4,12 +4,13 @@ from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
from fastvideo.configs.pipelines.registry import (
get_pipeline_config_cls_from_name)
from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
from fastvideo.configs.pipelines.wan import (WanI2V480PConfig, WanI2V720PConfig,
from fastvideo.configs.pipelines.wan import (SelfForcingWanT2V480PConfig,
WanI2V480PConfig, WanI2V720PConfig,
WanT2V480PConfig, WanT2V720PConfig)
__all__ = [
"HunyuanConfig", "FastHunyuanConfig", "PipelineConfig",
"SlidingTileAttnConfig", "WanT2V480PConfig", "WanI2V480PConfig",
"WanT2V720PConfig", "WanI2V720PConfig", "StepVideoT2VConfig",
"get_pipeline_config_cls_from_name"
"SelfForcingWanT2V480PConfig", "get_pipeline_config_cls_from_name"
]
+1
View File
@@ -87,6 +87,7 @@ class PipelineConfig:
# Wan2.2 TI2V parameters
ti2v_task: bool = False
boundary_ratio: float | None = None
# Compilation
# enable_torch_compile: bool = False
+6 -3
View File
@@ -11,9 +11,9 @@ from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
# isort: off
from fastvideo.configs.pipelines.wan import (
FastWan2_1_T2V_480P_Config, FastWan2_2_TI2V_5B_Config,
SelfForcingWanT2V480PConfig, Wan2_2_I2V_A14B_Config, Wan2_2_T2V_A14B_Config,
Wan2_2_TI2V_5B_Config, WanI2V480PConfig, WanI2V720PConfig, WanT2V480PConfig,
WanT2V720PConfig)
Wan2_2_I2V_A14B_Config, Wan2_2_T2V_A14B_Config, Wan2_2_TI2V_5B_Config,
WanI2V480PConfig, WanI2V720PConfig, WanT2V480PConfig, WanT2V720PConfig,
SelfForcingWanT2V480PConfig, WANV2VConfig)
# isort: on
from fastvideo.logger import init_logger
from fastvideo.utils import (maybe_download_model_index,
@@ -27,6 +27,7 @@ PIPE_NAME_TO_CONFIG: dict[str, type[PipelineConfig]] = {
"hunyuanvideo-community/HunyuanVideo": HunyuanConfig,
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WanT2V480PConfig,
"weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers": WanI2V480PConfig,
"IRMChen/Wan2.1-Fun-1.3B-Control-Diffusers": WANV2VConfig,
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers": WanI2V480PConfig,
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers": WanI2V720PConfig,
"Wan-AI/Wan2.1-T2V-14B-Diffusers": WanT2V720PConfig,
@@ -48,6 +49,7 @@ PIPELINE_DETECTOR: dict[str, Callable[[str], bool]] = {
"wanpipeline": lambda id: "wanpipeline" in id.lower(),
"wanimagetovideo": lambda id: "wanimagetovideo" in id.lower(),
"wandmdpipeline": lambda id: "wandmdpipeline" in id.lower(),
"wancausaldmdpipeline": lambda id: "wancausaldmdpipeline" in id.lower(),
"stepvideo": lambda id: "stepvideo" in id.lower(),
# Add other pipeline architecture detectors
}
@@ -60,6 +62,7 @@ PIPELINE_FALLBACK_CONFIG: dict[str, type[PipelineConfig]] = {
WanT2V480PConfig, # Base Wan config as fallback for any Wan variant
"wanimagetovideo": WanI2V480PConfig,
"wandmdpipeline": FastWan2_1_T2V_480P_Config,
"wancausaldmdpipeline": SelfForcingWanT2V480PConfig,
"stepvideo": StepVideoT2VConfig
# Other fallbacks by architecture
}
+20 -1
View File
@@ -7,7 +7,8 @@ import torch
from fastvideo.configs.models import DiTConfig, EncoderConfig, VAEConfig
from fastvideo.configs.models.dits import WanVideoConfig
from fastvideo.configs.models.encoders import (BaseEncoderOutput,
CLIPVisionConfig, T5Config)
CLIPVisionConfig, T5Config,
WAN2_1ControlCLIPVisionConfig)
from fastvideo.configs.models.vaes import WanVAEConfig
from fastvideo.configs.pipelines.base import PipelineConfig
@@ -54,6 +55,9 @@ class WanT2V480PConfig(PipelineConfig):
text_encoder_precisions: tuple[str, ...] = field(
default_factory=lambda: ("fp32", ))
# self-forcing params
warp_denoising_step: bool = True
# WanConfig-specific added parameters
def __post_init__(self):
@@ -97,6 +101,16 @@ class WanI2V720PConfig(WanI2V480PConfig):
flow_shift: float | None = 5.0
@dataclass
class WANV2VConfig(WanI2V480PConfig):
"""Configuration for WAN2.1 1.3B Control pipeline."""
image_encoder_config: EncoderConfig = field(
default_factory=WAN2_1ControlCLIPVisionConfig)
# CLIP encoder precision
image_encoder_precision: str = 'bf16'
@dataclass
class FastWan2_1_T2V_480P_Config(WanT2V480PConfig):
"""Base configuration for FastWan T2V 1.3B 480P pipeline architecture with DMD"""
@@ -133,6 +147,11 @@ class Wan2_2_T2V_A14B_Config(WanT2V480PConfig):
flow_shift: float | None = 12.0
boundary_ratio: float | None = 0.875
# self-forcing params
dmd_denoising_steps: list[int] | None = field(
default_factory=lambda: [1000, 750, 500, 250])
warp_denoising_step: bool = True
def __post_init__(self) -> None:
self.dit_config.boundary_ratio = self.boundary_ratio
+23
View File
@@ -18,6 +18,9 @@ class SamplingParam:
# Image inputs
image_path: str | None = None
# Video inputs
video_path: str | None = None
# Text inputs
prompt: str | list[str] | None = None
negative_prompt: str = "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"
@@ -48,6 +51,8 @@ class SamplingParam:
# Misc
save_video: bool = True
return_frames: bool = False
return_trajectory_latents: bool = False # returns all latents for each timestep
return_trajectory_decoded: bool = False # returns decoded latents for each timestep
def __post_init__(self) -> None:
self.data_type = "video" if self.num_frames > 1 else "image"
@@ -198,6 +203,12 @@ class SamplingParam:
default=SamplingParam.image_path,
help="Path to input image for image-to-video generation",
)
parser.add_argument(
"--video_path",
type=str,
default=SamplingParam.video_path,
help="Path to input video for video-to-video generation",
)
parser.add_argument(
"--moba-config-path",
type=str,
@@ -205,6 +216,18 @@ class SamplingParam:
help=
"Path to a JSON file containing V-MoBA specific configurations.",
)
parser.add_argument(
"--return-trajectory-latents",
action="store_true",
default=SamplingParam.return_trajectory_latents,
help="Whether to return the trajectory",
)
parser.add_argument(
"--return-trajectory-decoded",
action="store_true",
default=SamplingParam.return_trajectory_decoded,
help="Whether to return the decoded trajectory",
)
return parser
+3
View File
@@ -18,6 +18,7 @@ from fastvideo.configs.sample.wan import (
WanI2V_14B_720P_SamplingParam,
WanT2V_1_3B_SamplingParam,
WanT2V_14B_SamplingParam,
Wan2_1_Fun_1_3B_Control_SamplingParam,
SelfForcingWanT2V480PConfig,
)
# isort: on
@@ -39,6 +40,8 @@ SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers": WanI2V_14B_720P_SamplingParam,
"weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers":
Wan2_1_Fun_1_3B_InP_SamplingParam,
"IRMChen/Wan2.1-Fun-1.3B-Control-Diffusers":
Wan2_1_Fun_1_3B_Control_SamplingParam,
# Wan2.2
"Wan-AI/Wan2.2-TI2V-5B-Diffusers": Wan2_2_TI2V_5B_SamplingParam,
+17
View File
@@ -122,6 +122,17 @@ class Wan2_1_Fun_1_3B_InP_SamplingParam(SamplingParam):
num_inference_steps: int = 50
@dataclass
class Wan2_1_Fun_1_3B_Control_SamplingParam(SamplingParam):
fps: int = 16
num_frames: int = 49
height: int = 832
width: int = 480
guidance_scale: float = 6.0
teacache_params: WanTeaCacheParams = field(
default_factory=lambda: WanTeaCacheParams(teacache_thresh=0.1, ))
# =============================================
# ============= Wan2.2 TI2V Models =============
# =============================================
@@ -162,6 +173,12 @@ class Wan2_2_I2V_A14B_SamplingParam(Wan2_2_Base_SamplingParam):
# can be overridden during sampling
@dataclass
class Wan2_2_Fun_A14B_Control_SamplingParam(
Wan2_1_Fun_1_3B_Control_SamplingParam):
num_frames: int = 81
# =============================================
# ============= Causal Self-Forcing =============
# =============================================
+1 -1
View File
@@ -37,7 +37,7 @@ def getdataset(args) -> VideoCaptionMergedDataset:
temporal_sample=temporal_sample,
transform_topcrop=transform_topcrop,
seed=args.seed)
def gettextdataset(args) -> TextDataset:
return TextDataset(data_merge_path=args.data_merge_path,
@@ -12,6 +12,7 @@ import tqdm
# Dataset
from torch.utils.data import Dataset, Sampler
from torchdata.stateful_dataloader import StatefulDataLoader
from fastvideo.platforms import current_platform
from fastvideo.dataset.utils import collate_rows_from_parquet_schema
from fastvideo.distributed import (get_sp_world_size, get_world_group,
@@ -342,6 +343,7 @@ def build_parquet_map_style_dataloader(
collate_fn=passthrough,
num_workers=num_data_workers,
pin_memory=True,
pin_memory_device=current_platform.device_name,
persistent_workers=num_data_workers > 0,
)
return dataset, loader
@@ -6,7 +6,6 @@ import os
import torch
from torch.distributed import ProcessGroup
from fastvideo.platforms import current_platform
from fastvideo.platforms.interface import CpuArchEnum
from .base_device_communicator import DeviceCommunicatorBase
@@ -22,6 +21,8 @@ class CpuCommunicator(DeviceCommunicatorBase):
super().__init__(cpu_group, device, device_group, unique_name)
self.dist_module = torch.distributed
from fastvideo.platforms import current_platform
if (current_platform.get_cpu_architecture()
== CpuArchEnum.X86) and hasattr(
torch.ops._C,
@@ -0,0 +1,71 @@
import torch
from torch.distributed import ProcessGroup
from fastvideo.distributed.device_communicators.base_device_communicator import (
DeviceCommunicatorBase)
class NpuCommunicator(DeviceCommunicatorBase):
def __init__(self,
cpu_group: ProcessGroup,
device: torch.device | None = None,
device_group: ProcessGroup | None = None,
unique_name: str = ""):
super().__init__(cpu_group, device, device_group, unique_name)
from fastvideo.distributed.device_communicators.pyhccl import (
PyHcclCommunicator)
self.pyhccl_comm: PyHcclCommunicator | None = None
if self.world_size > 1:
self.pyhccl_comm = PyHcclCommunicator(
group=self.cpu_group,
device=self.device,
)
def all_reduce(self, input_, op: torch.distributed.ReduceOp | None = None):
pyhccl_comm = self.pyhccl_comm
assert pyhccl_comm is not None, "pyhccl_comm should not be None"
out = pyhccl_comm.all_reduce(input_, op=op)
if out is None:
# fall back to the default all-reduce using PyTorch.
# this usually happens during testing.
# when we run the model, allreduce only happens for the TP
# group, where we always have either custom allreduce or pyhccl.
out = input_.clone()
torch.distributed.all_reduce(out, group=self.device_group, op=op)
return out
def send(self, tensor: torch.Tensor, dst: int | None = None) -> None:
"""Sends a tensor to the destination rank in a non-blocking way"""
"""NOTE: `dst` is the local rank of the destination rank."""
if dst is None:
dst = (self.rank_in_group + 1) % self.world_size
pyhccl_comm = self.pyhccl_comm
if pyhccl_comm is not None and not pyhccl_comm.disabled:
pyhccl_comm.send(tensor, dst)
else:
torch.distributed.send(tensor, self.ranks[dst], self.device_group)
def recv(self,
size: torch.Size,
dtype: torch.dtype,
src: int | None = None) -> torch.Tensor:
"""Receives a tensor from the source rank."""
"""NOTE: `src` is the local rank of the source rank."""
if src is None:
src = (self.rank_in_group - 1) % self.world_size
tensor = torch.empty(size, dtype=dtype, device=self.device)
pyhccl_comm = self.pyhccl_comm
if pyhccl_comm is not None and not pyhccl_comm.disabled:
pyhccl_comm.recv(tensor, src)
else:
torch.distributed.recv(tensor, self.ranks[src], self.device_group)
return tensor
def destroy(self) -> None:
if self.pyhccl_comm is not None:
self.pyhccl_comm = None
@@ -0,0 +1,146 @@
import torch
import torch.distributed as dist
from torch.distributed import ProcessGroup, ReduceOp
from fastvideo.distributed.device_communicators.pyhccl_wrapper import (
HCCLLibrary, aclrtStream_t, buffer_type, hcclComm_t, hcclDataTypeEnum,
hcclRedOpTypeEnum, hcclUniqueId)
from fastvideo.distributed.utils import StatelessProcessGroup
from fastvideo.logger import init_logger
from fastvideo.utils import current_stream
logger = init_logger(__name__)
class PyHcclCommunicator:
def __init__(
self,
group: ProcessGroup | StatelessProcessGroup,
device: int | str | torch.device,
library_path: str | None = None,
):
"""
Args:
group: the process group to work on. If None, it will use the
default process group.
device: the device to bind the PyHcclCommunicator to. If None,
it will be bind to f"npu:{local_rank}".
library_path: the path to the HCCL library. If None, it will
use the default library path.
It is the caller's responsibility to make sure each communicator
is bind to a unique device.
"""
if not isinstance(group, StatelessProcessGroup):
assert dist.is_initialized()
assert dist.get_backend(group) != dist.Backend.HCCL, (
"PyHcclCommunicator should be attached to a non-HCCL group.")
# note: this rank is the rank in the group
self.rank = dist.get_rank(group)
self.world_size = dist.get_world_size(group)
else:
self.rank = group.rank
self.world_size = group.world_size
self.group = group
# if world_size == 1, no need to create communicator
if self.world_size == 1:
self.available = False
self.disabled = True
return
try:
self.hccl = HCCLLibrary(library_path)
except Exception:
logger.warning("disable hccl because of missing HCCL library")
# disable because of missing HCCL library
# e.g. in a non-NPU environment
self.available = False
self.disabled = True
return
self.available = True
self.disabled = False
logger.info("FastVideo is using pyhccl")
if isinstance(device, int):
device = torch.device(f"npu:{device}")
elif isinstance(device, str):
device = torch.device(device)
# now `device` is a `torch.device` object
assert isinstance(device, torch.device)
self.device = device
if self.rank == 0:
# get the unique id from HCCL
with torch.npu.device(device):
self.unique_id = self.hccl.hcclGetUniqueId()
else:
# construct an empty unique id
self.unique_id = hcclUniqueId()
if not isinstance(group, StatelessProcessGroup):
tensor = torch.ByteTensor(list(self.unique_id.internal))
ranks = dist.get_process_group_ranks(group)
# arg `src` in `broadcast` is the global rank
dist.broadcast(tensor, src=ranks[0], group=group)
byte_list = tensor.tolist()
for i, byte in enumerate(byte_list):
self.unique_id.internal[i] = byte
else:
self.unique_id = group.broadcast_obj(self.unique_id, src=0)
# hccl communicator and stream will use this device
# `torch.npu.device` is a context manager that changes the
# current npu device to the specified one
with torch.npu.device(device):
self.comm: hcclComm_t = self.hccl.hcclCommInitRank(
self.world_size, self.unique_id, self.rank)
stream = current_stream()
# A small all_reduce for warmup.
data = torch.zeros(1, device=device)
self.all_reduce(data)
stream.synchronize()
del data
def all_reduce(self,
in_tensor: torch.Tensor,
op: ReduceOp = ReduceOp.SUM,
stream=None) -> torch.Tensor:
if self.disabled:
return None
# hccl communicator created on a specific device
# will only work on tensors on the same device
# otherwise it will cause "illegal memory access"
assert in_tensor.device == self.device, (
f"this hccl communicator is created to work on {self.device}, "
f"but the input tensor is on {in_tensor.device}")
out_tensor = torch.empty_like(in_tensor)
if stream is None:
stream = current_stream()
self.hccl.hcclAllReduce(buffer_type(in_tensor.data_ptr()),
buffer_type(out_tensor.data_ptr()),
in_tensor.numel(),
hcclDataTypeEnum.from_torch(in_tensor.dtype),
hcclRedOpTypeEnum.from_torch(op), self.comm,
aclrtStream_t(stream.npu_stream))
return out_tensor
def broadcast(self, tensor: torch.Tensor, src: int, stream=None):
if self.disabled:
return
assert tensor.device == self.device, (
f"this hccl communicator is created to work on {self.device}, "
f"but the input tensor is on {tensor.device}")
if stream is None:
stream = current_stream()
buffer = buffer_type(tensor.data_ptr())
self.hccl.hcclBroadcast(buffer, tensor.numel(),
hcclDataTypeEnum.from_torch(tensor.dtype), src,
self.comm, aclrtStream_t(stream.npu_stream))
@@ -0,0 +1,208 @@
import ctypes
import platform
from dataclasses import dataclass
from typing import Any
import torch
from torch.distributed import ReduceOp
from fastvideo.logger import init_logger
from fastvideo.utils import find_hccl_library
logger = init_logger(__name__)
hcclResult_t = ctypes.c_int
hcclComm_t = ctypes.c_void_p
class hcclUniqueId(ctypes.Structure):
_fields_ = [("internal", ctypes.c_byte * 4108)]
aclrtStream_t = ctypes.c_void_p
buffer_type = ctypes.c_void_p
hcclDataType_t = ctypes.c_int
class hcclDataTypeEnum:
hcclInt8 = 0
hcclInt16 = 1
hcclInt32 = 2
hcclFloat16 = 3
hcclFloat32 = 4
hcclInt64 = 5
hcclUint64 = 6
hcclUint8 = 7
hcclUint16 = 8
hcclUint32 = 9
hcclFloat64 = 10
hcclBfloat16 = 11
hcclInt128 = 12
@classmethod
def from_torch(cls, dtype: torch.dtype) -> int:
if dtype == torch.int8:
return cls.hcclInt8
if dtype == torch.uint8:
return cls.hcclUint8
if dtype == torch.int32:
return cls.hcclInt32
if dtype == torch.int64:
return cls.hcclInt64
if dtype == torch.float16:
return cls.hcclFloat16
if dtype == torch.float32:
return cls.hcclFloat32
if dtype == torch.float64:
return cls.hcclFloat64
if dtype == torch.bfloat16:
return cls.hcclBfloat16
raise ValueError(f"Unsupported dtype: {dtype}")
hcclRedOp_t = ctypes.c_int
class hcclRedOpTypeEnum:
hcclSum = 0
hcclProd = 1
hcclMax = 2
hcclMin = 3
@classmethod
def from_torch(cls, op: ReduceOp) -> int:
if op == ReduceOp.SUM:
return cls.hcclSum
if op == ReduceOp.PRODUCT:
return cls.hcclProd
if op == ReduceOp.MAX:
return cls.hcclMax
if op == ReduceOp.MIN:
return cls.hcclMin
raise ValueError(f"Unsupported op: {op}")
@dataclass
class Function:
name: str
restype: Any
argtypes: list[Any]
class HCCLLibrary:
exported_functions = [
Function("HcclGetErrorString", ctypes.c_char_p, [hcclResult_t]),
Function("HcclGetRootInfo", hcclResult_t,
[ctypes.POINTER(hcclUniqueId)]),
Function("HcclCommInitRootInfo", hcclResult_t, [
ctypes.c_int,
ctypes.POINTER(hcclUniqueId),
ctypes.c_int,
ctypes.POINTER(hcclComm_t),
]),
Function("HcclAllReduce", hcclResult_t, [
buffer_type,
buffer_type,
ctypes.c_size_t,
hcclDataType_t,
hcclRedOp_t,
hcclComm_t,
aclrtStream_t,
]),
Function("HcclBroadcast", hcclResult_t, [
buffer_type,
ctypes.c_size_t,
hcclDataType_t,
ctypes.c_int,
hcclComm_t,
aclrtStream_t,
]),
Function("HcclCommDestroy", hcclResult_t, [hcclComm_t]),
]
# class attribute to store the mapping from the path to the library
# to avoid loading the same library multiple times
path_to_library_cache: dict[str, Any] = {}
# class attribute to store the mapping from library path
# to the correspongding directory
path_to_dict_mapping: dict[str, dict[str, Any]] = {}
def __init__(self, so_file: str | None = None):
so_file = so_file or find_hccl_library()
try:
if so_file not in HCCLLibrary.path_to_dict_mapping:
lib = ctypes.CDLL(so_file)
HCCLLibrary.path_to_library_cache[so_file] = lib
self.lib = HCCLLibrary.path_to_library_cache[so_file]
except Exception as e:
logger.error(
"Failed to load HCCL library from %s. "
"It is expected if you are not running on Ascend NPUs."
"Otherwise, the hccl library might not exist, be corrupted "
"or it does not support the current platform %s. "
"If you already have the library, please set the "
"environment variable HCCL_SO_PATH"
" to point to the correct hccl library path.", so_file,
platform.platform())
raise e
if so_file not in HCCLLibrary.path_to_dict_mapping:
_funcs: dict[str, Any] = {}
for func in HCCLLibrary.exported_functions:
f = getattr(self.lib, func.name)
f.restype = func.restype
f.argtypes = func.argtypes
_funcs[func.name] = f
HCCLLibrary.path_to_dict_mapping[so_file] = _funcs
self._funcs = HCCLLibrary.path_to_dict_mapping[so_file]
def hcclGetErrorString(self, result: hcclResult_t) -> str:
return self._funcs["HcclGetErrorString"](result).decode("utf-8")
def HCCL_CHECK(self, result: hcclResult_t) -> None:
if result != 0:
error_str = self.hcclGetErrorString(result)
raise RuntimeError(f"HCCL error: {error_str}")
def hcclGetUniqueId(self) -> hcclUniqueId:
unique_id = hcclUniqueId()
self.HCCL_CHECK(self._funcs["HcclGetRootInfo"](ctypes.byref(unique_id)))
return unique_id
def hcclCommInitRank(self, world_size: int, unique_id: hcclUniqueId,
rank: int) -> hcclComm_t:
comm = hcclComm_t()
self.HCCL_CHECK(self._funcs["HcclCommInitRootInfo"](
world_size, ctypes.byref(unique_id), rank, ctypes.byref(comm)))
return comm
def hcclAllReduce(self, sendbuff: buffer_type, recvbuff: buffer_type,
count: int, datatype: int, op: int, comm: hcclComm_t,
stream: aclrtStream_t) -> None:
self.HCCL_CHECK(self._funcs["HcclAllReduce"](sendbuff, recvbuff, count,
datatype, op, comm,
stream))
def hcclBroadcast(self, buf: buffer_type, count: int, datatype: int,
root: int, comm: hcclComm_t,
stream: aclrtStream_t) -> None:
self.HCCL_CHECK(self._funcs["HcclBroadcast"](buf, count, datatype, root,
comm, stream))
def hcclCommDestroy(self, comm: hcclComm_t) -> None:
self.HCCL_CHECK(self._funcs["HcclCommDestroy"](comm))
__all__ = [
"HCCLLibrary",
"hcclDataTypeEnum",
"hcclRedOpTypeEnum",
"hcclUniqueId",
"hcclComm_t",
"aclrtStream_t",
"buffer_type",
]
+37 -27
View File
@@ -45,7 +45,6 @@ from fastvideo.distributed.device_communicators.cpu_communicator import (
CpuCommunicator)
from fastvideo.distributed.utils import StatelessProcessGroup
from fastvideo.logger import init_logger
from fastvideo.platforms import current_platform
logger = init_logger(__name__)
@@ -190,7 +189,6 @@ class GroupCoordinator:
self.device = get_local_torch_device()
self.use_device_communicator = use_device_communicator
self.device_communicator: DeviceCommunicatorBase = None # type: ignore
if use_device_communicator and self.world_size > 1:
# Platform-aware device communicator selection
@@ -203,6 +201,15 @@ class GroupCoordinator:
device_group=self.device_group,
unique_name=self.unique_name,
)
elif current_platform.is_npu():
from fastvideo.distributed.device_communicators.npu_communicator import (
NpuCommunicator)
self.device_communicator = NpuCommunicator(
cpu_group=self.cpu_group,
device=self.device,
device_group=self.device_group,
unique_name=self.unique_name,
)
else:
# For MPS and CPU, use the CPU communicator
self.device_communicator = CpuCommunicator(
@@ -776,8 +783,13 @@ def init_distributed_environment(
):
# Determine the appropriate backend based on the platform
from fastvideo.platforms import current_platform
if backend == "nccl" and not current_platform.is_cuda_alike():
# Use gloo backend for non-CUDA platforms (MPS, CPU)
backend = "nccl"
if current_platform.is_cuda_alike():
logger.info("Using nccl backend for CUDA platform")
elif current_platform.is_npu():
backend = "hccl"
logger.info("Using hccl backend for NPU platform")
else:
backend = "gloo"
logger.info("Using gloo backend for %s platform",
current_platform.device_name)
@@ -791,21 +803,11 @@ def init_distributed_environment(
"distributed_init_method must be provided when initializing "
"distributed environment")
# For MPS, don't pass device_id as it doesn't support device indices
if current_platform.is_mps():
torch.distributed.init_process_group(
backend=backend,
init_method=distributed_init_method,
world_size=world_size,
rank=rank)
else:
# this backend is used for WORLD
torch.distributed.init_process_group(
backend=backend,
init_method=distributed_init_method,
world_size=world_size,
rank=rank,
device_id=device_id)
torch.distributed.init_process_group(
backend=backend,
init_method=distributed_init_method,
world_size=world_size,
rank=rank)
# set the local rank
# local_rank is not available in torch ProcessGroup,
# see https://github.com/pytorch/pytorch/issues/122816
@@ -948,9 +950,14 @@ def get_dp_rank() -> int:
def get_local_torch_device() -> torch.device:
"""Return the torch device for the current rank."""
return torch.device(f"cuda:{envs.LOCAL_RANK}"
) if current_platform.is_cuda_alike() else torch.device(
"mps")
from fastvideo.platforms import current_platform
if current_platform.is_npu():
device = torch.device(f"npu:{envs.LOCAL_RANK}")
elif current_platform.is_cuda_alike() or current_platform.is_cuda():
device = torch.device(f"cuda:{envs.LOCAL_RANK}")
else:
device = torch.device("mps")
return device
def maybe_init_distributed_environment_and_model_parallel(
@@ -968,7 +975,9 @@ def maybe_init_distributed_environment_and_model_parallel(
device = get_local_torch_device()
logger.info(
"Initializing distributed environment with world_size=%d, device=%s",
world_size, device)
world_size,
device,
local_main_process_only=False)
init_distributed_environment(
world_size=world_size,
@@ -979,10 +988,11 @@ def maybe_init_distributed_environment_and_model_parallel(
initialize_model_parallel(tensor_model_parallel_size=tp_size,
sequence_model_parallel_size=sp_size)
# Only set CUDA device if we're on a CUDA platform
if current_platform.is_cuda_alike():
device = torch.device(f"cuda:{local_rank}")
torch.cuda.set_device(device)
# set device if we're on a CUDA/NPU platform
from fastvideo.platforms import current_platform
device_type = current_platform.device_type
device = torch.device(f"{device_type}:{local_rank}")
current_platform.get_torch_device().set_device(device)
def model_parallel_is_initialized() -> bool:
+91 -38
View File
@@ -8,6 +8,7 @@ diffusion models.
import math
import os
import re
import time
from copy import deepcopy
from typing import Any
@@ -110,7 +111,7 @@ class VideoGenerator:
prompt: The prompt to use for generation (optional if prompt_txt is provided)
negative_prompt: The negative prompt to use (overrides the one in fastvideo_args)
output_path: Path to save the video (overrides the one in fastvideo_args)
output_video_name: Name of the video file to save. Default is the first 100 characters of the prompt.
prompt_path: Path to prompt file
save_video: Whether to save the video to disk
return_frames: Whether to return the raw frames
num_inference_steps: Number of denoising steps (overrides fastvideo_args)
@@ -127,8 +128,13 @@ class VideoGenerator:
Either the output dictionary, list of frames, or list of results for batch processing
"""
# Handle batch processing from text file
if self.fastvideo_args.prompt_txt is not None:
prompt_txt_path = self.fastvideo_args.prompt_txt
if sampling_param is None:
sampling_param = SamplingParam.from_pretrained(
self.fastvideo_args.model_path)
sampling_param.update(kwargs)
if self.fastvideo_args.prompt_txt is not None or sampling_param.prompt_path is not None:
prompt_txt_path = sampling_param.prompt_path or self.fastvideo_args.prompt_txt
if not os.path.exists(prompt_txt_path):
raise FileNotFoundError(
f"Prompt text file not found: {prompt_txt_path}")
@@ -142,22 +148,19 @@ class VideoGenerator:
logger.info("Found %d prompts in %s", len(prompts), prompt_txt_path)
if sampling_param is not None:
original_output_video_name = sampling_param.output_video_name
else:
original_output_video_name = None
results = []
for i, batch_prompt in enumerate(prompts):
logger.info("Processing prompt %d/%d: %s...", i + 1,
len(prompts), batch_prompt[:100])
try:
# Generate video for this prompt using the same logic below
if sampling_param is not None and original_output_video_name is not None:
sampling_param.output_video_name = original_output_video_name + f"_{i}"
output_path = self._prepare_output_path(
sampling_param.output_path, batch_prompt)
kwargs["output_path"] = output_path
result = self._generate_single_video(
batch_prompt, sampling_param, **kwargs)
prompt=batch_prompt,
sampling_param=sampling_param,
**kwargs)
# Add prompt info to result
if isinstance(result, dict):
@@ -181,8 +184,73 @@ class VideoGenerator:
# Single prompt generation (original behavior)
if prompt is None:
raise ValueError("Either prompt or prompt_txt must be provided")
output_path = self._prepare_output_path(sampling_param.output_path,
prompt)
kwargs["output_path"] = output_path
return self._generate_single_video(prompt=prompt,
sampling_param=sampling_param,
**kwargs)
return self._generate_single_video(prompt, sampling_param, **kwargs)
def _prepare_output_path(
self,
output_path: str,
prompt: str,
) -> str:
"""Build a unique, sanitized .mp4 output file path.
- If `output_path` ends with .mp4 (case-insensitive), treat it as a file path.
- Otherwise, treat `output_path` as a directory and derive the filename
from the prompt.
- Invalid filename characters are removed; if the name changes, a
warning is logged.
- If the target path already exists, a numeric suffix is appended.
"""
def _sanitize_filename_component(name: str) -> str:
# Remove characters invalid on common filesystems, strip spaces/dots
sanitized = re.sub(r'[\\/:*?"<>|]', '', name)
sanitized = sanitized.strip().strip('.')
sanitized = re.sub(r'\s+', ' ', sanitized)
return sanitized or "video"
base_path, extension = os.path.splitext(output_path)
extension_lower = extension.lower()
if extension_lower == ".mp4":
output_dir = os.path.dirname(output_path)
base_name = os.path.basename(
base_path) # filename without extension
sanitized_base = _sanitize_filename_component(base_name)
if sanitized_base != base_name:
logger.warning(
"The video name '%s' contained invalid characters. It has been renamed to '%s.mp4'",
os.path.basename(output_path),
sanitized_base,
)
video_name = f"{sanitized_base}.mp4"
else:
# Treat as directory; inform if an unexpected extension was provided.
if extension:
logger.info(
"Output path '%s' has non-mp4 extension '%s'; treating it as a directory and using a .mp4 filename derived from the prompt",
output_path,
extension,
)
output_dir = output_path
prompt_component = _sanitize_filename_component(prompt[:100])
video_name = f"{prompt_component}.mp4"
if output_dir:
os.makedirs(output_dir, exist_ok=True)
new_output_path = os.path.join(output_dir, video_name)
counter = 1
while os.path.exists(new_output_path):
name_part, ext_part = os.path.splitext(video_name)
new_video_name = f"{name_part}_{counter}{ext_part}"
new_output_path = os.path.join(output_dir, new_video_name)
counter += 1
return new_output_path
def _generate_single_video(
self,
@@ -200,15 +268,9 @@ class VideoGenerator:
raise TypeError(
f"`prompt` must be a string, but got {type(prompt)}")
prompt = prompt.strip()
if sampling_param is None:
sampling_param = SamplingParam.from_pretrained(
fastvideo_args.model_path)
else:
sampling_param = deepcopy(sampling_param)
kwargs["prompt"] = prompt
sampling_param.update(kwargs)
sampling_param = deepcopy(sampling_param)
output_path = kwargs["output_path"]
sampling_param.prompt = prompt
# Process negative prompt
if sampling_param.negative_prompt is not None:
sampling_param.negative_prompt = sampling_param.negative_prompt.strip(
@@ -277,7 +339,7 @@ class VideoGenerator:
height: {target_height}
width: {target_width}
video_length: {sampling_param.num_frames}
prompt: {prompt}
prompt: {sampling_param.prompt}
image_path: {sampling_param.image_path}
neg_prompt: {sampling_param.negative_prompt}
seed: {sampling_param.seed}
@@ -288,7 +350,7 @@ class VideoGenerator:
flow_shift: {fastvideo_args.pipeline_config.flow_shift}
embedded_guidance_scale: {fastvideo_args.pipeline_config.embedded_cfg_scale}
save_video: {sampling_param.save_video}
output_path: {sampling_param.output_path}
output_path: {output_path}
""" # type: ignore[attr-defined]
logger.info(debug_str)
@@ -298,13 +360,8 @@ class VideoGenerator:
eta=0.0,
n_tokens=n_tokens,
VSA_sparsity=fastvideo_args.VSA_sparsity,
extra={},
)
# Use prompt[:100] for video name
if batch.output_video_name is None:
batch.output_video_name = prompt[:100]
# Run inference
start_time = time.perf_counter()
output_batch = self.executor.execute_forward(batch, fastvideo_args)
@@ -324,15 +381,8 @@ class VideoGenerator:
# Save video if requested
if batch.save_video:
output_path = batch.output_path
if output_path:
os.makedirs(output_path, exist_ok=True)
video_path = os.path.join(output_path,
f"{batch.output_video_name}.mp4")
imageio.mimsave(video_path, frames, fps=batch.fps, format="mp4")
logger.info("Saved video to %s", video_path)
else:
logger.warning("No output path provided, video not saved")
imageio.mimsave(output_path, frames, fps=batch.fps, format="mp4")
logger.info("Saved video to %s", output_path)
if batch.return_frames:
return frames
@@ -344,6 +394,9 @@ class VideoGenerator:
"size": (target_height, target_width, batch.num_frames),
"generation_time": gen_time,
"logging_info": logging_info,
"trajectory": output_batch.trajectory_latents,
"trajectory_timesteps": output_batch.trajectory_timesteps,
"trajectory_decoded": output_batch.trajectory_decoded,
}
def set_lora_adapter(self,
+57 -3
View File
@@ -14,20 +14,29 @@ if TYPE_CHECKING:
FASTVIDEO_CACHE_ROOT: str = os.path.expanduser("~/.cache/fastvideo")
FASTVIDEO_CONFIG_ROOT: str = os.path.expanduser("~/.config/fastvideo")
FASTVIDEO_CONFIGURE_LOGGING: int = 1
FASTVIDEO_RAY_PER_WORKER_GPUS: float = 1.0
FASTVIDEO_LOGGING_LEVEL: str = "INFO"
FASTVIDEO_LOGGING_PREFIX: str = ""
FASTVIDEO_LOGGING_CONFIG_PATH: str | None = None
FASTVIDEO_TRACE_FUNCTION: int = 0
FASTVIDEO_ATTENTION_BACKEND: str | None = None
FASTVIDEO_ATTENTION_CONFIG: str | None = None
FASTVIDEO_WORKER_MULTIPROC_METHOD: str = "fork"
FASTVIDEO_WORKER_MULTIPROC_METHOD: str = "spawn"
FASTVIDEO_TARGET_DEVICE: str = "cuda"
MAX_JOBS: str | None = None
NVCC_THREADS: str | None = None
CMAKE_BUILD_TYPE: str | None = None
VERBOSE: bool = False
FASTVIDEO_TORCH_PROFILER_DIR: str | None = None
FASTVIDEO_TORCH_PROFILER_RECORD_SHAPES: bool = False
FASTVIDEO_TORCH_PROFILER_WITH_PROFILE_MEMORY: bool = False
FASTVIDEO_TORCH_PROFILER_WITH_STACK: bool = True
FASTVIDEO_TORCH_PROFILER_WITH_FLOPS: bool = False
FASTVIDEO_TORCH_PROFILE_REGIONS: str = ""
FASTVIDEO_SERVER_DEV_MODE: bool = False
FASTVIDEO_STAGE_LOGGING: bool = False
FASTVIDEO_HOST_IP: str = ""
FASTVIDEO_LOOPBACK_IP: str = ""
def get_default_cache_root() -> str:
@@ -113,6 +122,23 @@ environment_variables: dict[str, Callable[[], Any]] = {
os.path.join(get_default_cache_root(), "fastvideo"),
)),
# used in distributed environment to determine the ip address
# of the current node, when the node has multiple network interfaces.
# If you are using multi-node inference, you should set this differently
# on each node.
"FASTVIDEO_HOST_IP":
lambda: os.getenv("FASTVIDEO_HOST_IP", ""),
# Used to force set up loopback IP
"FASTVIDEO_LOOPBACK_IP":
lambda: os.getenv("FASTVIDEO_LOOPBACK_IP", ""),
# Number of GPUs per worker in Ray, if it is set to be a fraction,
# it allows ray to schedule multiple actors on a single GPU,
# so that users can colocate other actors on the same GPUs as FastVideo.
"FASTVIDEO_RAY_PER_WORKER_GPUS":
lambda: float(os.getenv("FASTVIDEO_RAY_PER_WORKER_GPUS", "1.0")),
# Interval in seconds to log a warning message when the ring buffer is full
"FASTVIDEO_RINGBUFFER_WARNING_INTERVAL":
lambda: int(os.environ.get("FASTVIDEO_RINGBUFFER_WARNING_INTERVAL", "60")),
@@ -175,6 +201,7 @@ environment_variables: dict[str, Callable[[], Any]] = {
# - "SLIDING_TILE_ATTN" : use Sliding Tile Attention
# - "VIDEO_SPARSE_ATTN": use Video Sparse Attention
# - "SAGE_ATTN": use Sage Attention
# - "SAGE_ATTN_THREE": use Sage Attention 3
"FASTVIDEO_ATTENTION_BACKEND":
lambda: os.getenv("FASTVIDEO_ATTENTION_BACKEND", None),
@@ -185,9 +212,8 @@ environment_variables: dict[str, Callable[[], Any]] = {
os.path.expanduser(os.getenv("FASTVIDEO_ATTENTION_CONFIG", "."))),
# Use dedicated multiprocess context for workers.
# Both spawn and fork work
"FASTVIDEO_WORKER_MULTIPROC_METHOD":
lambda: os.getenv("FASTVIDEO_WORKER_MULTIPROC_METHOD", "fork"),
lambda: os.getenv("FASTVIDEO_WORKER_MULTIPROC_METHOD", "spawn"),
# Enables torch profiler if set. Path to the directory where torch profiler
# traces are saved. Note that it must be an absolute path.
@@ -196,6 +222,34 @@ environment_variables: dict[str, Callable[[], Any]] = {
if os.getenv("FASTVIDEO_TORCH_PROFILER_DIR", None) is None else os.
path.expanduser(os.getenv("FASTVIDEO_TORCH_PROFILER_DIR", "."))),
# Enable torch profiler to record shapes if set
# FASTVIDEO_TORCH_PROFILER_RECORD_SHAPES=1. If not set, torch profiler will
# not record shapes.
"FASTVIDEO_TORCH_PROFILER_RECORD_SHAPES":
lambda: bool(
os.getenv("FASTVIDEO_TORCH_PROFILER_RECORD_SHAPES", "0") != "0"),
# Enable torch profiler to profile memory if set
# FASTVIDEO_TORCH_PROFILER_WITH_PROFILE_MEMORY=1. If not set, torch profiler
# will not profile memory.
"FASTVIDEO_TORCH_PROFILER_WITH_PROFILE_MEMORY":
lambda: bool(
os.getenv("FASTVIDEO_TORCH_PROFILER_WITH_PROFILE_MEMORY", "0") != "0"),
# Enable torch profiler to profile stack if set
# FASTVIDEO_TORCH_PROFILER_WITH_STACK=1. If not set, torch profiler WILL
# profile stack by default.
"FASTVIDEO_TORCH_PROFILER_WITH_STACK":
lambda: bool(os.getenv("FASTVIDEO_TORCH_PROFILER_WITH_STACK", "1") != "0"),
# Enable torch profiler to profile flops if set
# FASTVIDEO_TORCH_PROFILER_WITH_FLOPS=1. If not set, torch profiler will
# not profile flops.
"FASTVIDEO_TORCH_PROFILER_WITH_FLOPS":
lambda: bool(os.getenv("FASTVIDEO_TORCH_PROFILER_WITH_FLOPS", "0") != "0"),
"FASTVIDEO_TORCH_PROFILE_REGIONS":
lambda: os.getenv("FASTVIDEO_TORCH_PROFILE_REGIONS", ""),
# If set, fastvideo will run in development mode, which will enable
# some additional endpoints for developing and debugging,
# e.g. `/reset_prefix_cache`
+147 -8
View File
@@ -7,15 +7,21 @@ import json
from contextlib import contextmanager
from dataclasses import field
from enum import Enum
from typing import Any
from typing import Any, TYPE_CHECKING
from fastvideo.configs.configs import PreprocessConfig
from fastvideo.configs.pipelines.base import PipelineConfig, STA_Mode
from fastvideo.configs.utils import clean_cli_args
from fastvideo.logger import init_logger
from fastvideo.platforms import current_platform
from fastvideo.utils import FlexibleArgumentParser, StoreBoolean
if TYPE_CHECKING:
from ray.runtime_env import RuntimeEnv
from ray.util.placement_group import PlacementGroup
else:
RuntimeEnv = Any
PlacementGroup = Any
logger = init_logger(__name__)
@@ -91,6 +97,11 @@ class FastVideoArgs:
# Distributed executor backend
distributed_executor_backend: str = "mp"
# a few attributes for ray related
ray_placement_group: PlacementGroup | None = None
ray_runtime_env: RuntimeEnv | None = None
inference_mode: bool = True # if False == training mode
# HuggingFace specific parameters
@@ -133,6 +144,7 @@ class FastVideoArgs:
# Compilation
enable_torch_compile: bool = False
torch_compile_kwargs: dict[str, Any] = field(default_factory=dict)
disable_autocast: bool = False
@@ -158,12 +170,15 @@ class FastVideoArgs:
"transformer": True,
"vae": True,
})
override_transformer_cls_name: str | None = None
init_weights_from_safetensors: str = "" # path to safetensors file for initial weight loading
init_weights_from_safetensors_2: str = "" # path to safetensors file for initial weight loading for transformer_2
# # DMD parameters
# dmd_denoising_steps: List[int] | None = field(default=None)
# MoE parameters used by Wan2.2
boundary_ratio: float | None = None
boundary_ratio: float | None = 0.875
@property
def training_mode(self) -> bool:
@@ -329,6 +344,13 @@ class FastVideoArgs:
help="Use torch.compile to speed up DiT inference." +
"However, will likely cause precision drifts. See (https://github.com/pytorch/pytorch/issues/145213)",
)
parser.add_argument(
"--torch-compile-kwargs",
type=str,
default=None,
help=
"JSON string of kwargs to pass to torch.compile. Example: '{\"backend\":\"inductor\",\"mode\":\"reduce-overhead\"}'",
)
parser.add_argument(
"--dit-cpu-offload",
@@ -396,6 +418,21 @@ class FastVideoArgs:
default=FastVideoArgs.enable_stage_verification,
help="Enable input/output verification for pipeline stages",
)
parser.add_argument(
"--override-transformer-cls-name",
type=str,
default=FastVideoArgs.override_transformer_cls_name,
help="Override transformer cls name",
)
parser.add_argument(
"--init-weights-from-safetensors",
type=str,
help="Path to safetensors file for initial weight loading")
parser.add_argument(
"--init-weights-from-safetensors-2",
type=str,
help="Path to safetensors file for initial weight loading")
# Add pipeline configuration arguments
PipelineConfig.add_cli_args(parser)
@@ -424,6 +461,21 @@ class FastVideoArgs:
mode_value = getattr(args, attr, FastVideoArgs.mode.value)
kwargs['mode'] = ExecutionMode.from_string(
mode_value) if isinstance(mode_value, str) else mode_value
elif attr == 'torch_compile_kwargs':
# Parse JSON string for torch.compile kwargs
torch_compile_kwargs_str = getattr(args, 'torch_compile_kwargs',
None)
if torch_compile_kwargs_str:
try:
import json
kwargs['torch_compile_kwargs'] = json.loads(
torch_compile_kwargs_str)
except json.JSONDecodeError as e:
raise ValueError(
f"Invalid JSON for torch_compile_kwargs: {e}"
) from e
else:
kwargs['torch_compile_kwargs'] = {}
elif attr == 'workload_type':
# Convert string to WorkloadType enum
workload_type_value = getattr(args, 'workload_type',
@@ -465,6 +517,8 @@ class FastVideoArgs:
def check_fastvideo_args(self) -> None:
"""Validate inference arguments for consistency"""
from fastvideo.platforms import current_platform
if current_platform.is_mps():
self.use_fsdp_inference = False
@@ -603,7 +657,10 @@ class TrainingArgs(FastVideoArgs):
# text encoder & vae & diffusion model
pretrained_model_name_or_path: str = ""
dit_model_name_or_path: str = ""
# DMD model paths - separate paths for each network
real_score_model_path: str = "" # path for real score (teacher) model
fake_score_model_path: str = "" # path for fake score (critic) model
# diffusion setting
ema_decay: float = 0.0
@@ -618,6 +675,7 @@ class TrainingArgs(FastVideoArgs):
validation_guidance_scale: str = ""
validation_steps: float = 0.0
log_validation: bool = False
trackers: list[str] = dataclasses.field(default_factory=list)
tracker_project_name: str = ""
wandb_run_name: str = ""
seed: int | None = None
@@ -625,7 +683,6 @@ class TrainingArgs(FastVideoArgs):
# output
output_dir: str = ""
checkpoints_total_limit: int = 0
checkpointing_steps: int = 0
resume_from_checkpoint: str = "" # specify the checkpoint folder to resume from
# optimizer & scheduler
@@ -658,6 +715,7 @@ class TrainingArgs(FastVideoArgs):
linear_quadratic_threshold: float = 0.0
linear_range: float = 0.0
weight_decay: float = 0.0
betas: str = "0.9,0.999" # betas for optimizer, format: "beta1,beta2"
use_ema: bool = False
multi_phased_distill_schedule: str = ""
pred_decay_weight: float = 0.0
@@ -678,16 +736,28 @@ class TrainingArgs(FastVideoArgs):
# distillation args
generator_update_interval: int = 5
dfake_gen_update_ratio: int = 5 # self-forcing: how often to train generator vs critic
min_timestep_ratio: float = 0.2
max_timestep_ratio: float = 0.98
real_score_guidance_scale: float = 3.5
fake_score_learning_rate: float = 0.0 # separate learning rate for fake_score_transformer, if 0.0, use learning_rate
fake_score_lr_scheduler: str = "constant" # separate lr scheduler for fake_score_transformer, if not set, use lr_scheduler
fake_score_betas: str = "0.9,0.999" # betas for fake score optimizer, format: "beta1,beta2"
training_state_checkpointing_steps: int = 0 # for resuming training
weight_only_checkpointing_steps: int = 0 # for inference
log_visualization: bool = False
# simulate generator forward to match inference
simulate_generator_forward: bool = False
warp_denoising_step: bool = False
# Self-forcing specific arguments
num_frame_per_block: int = 3
independent_first_frame: bool = False
enable_gradient_masking: bool = True
gradient_mask_last_n_frames: int = 21
same_step_across_blocks: bool = False # Use same exit timestep for all blocks
last_step_only: bool = False # Only use the last timestep for training
context_noise: int = 0 # Context noise level for cache updates
@classmethod
def from_cli_args(cls, args: argparse.Namespace) -> "TrainingArgs":
@@ -789,6 +859,20 @@ class TrainingArgs(FastVideoArgs):
type=str,
help="Directory to cache models")
# DMD model paths - separate paths for each network
parser.add_argument(
"--generator-model-path",
type=str,
help="Path to generator (student) model for DMD distillation")
parser.add_argument(
"--real-score-model-path",
type=str,
help="Path to real score (teacher) model for DMD distillation")
parser.add_argument(
"--fake-score-model-path",
type=str,
help="Path to fake score (critic) model for DMD distillation")
# Diffusion settings
parser.add_argument("--ema-decay",
type=float,
@@ -844,9 +928,6 @@ class TrainingArgs(FastVideoArgs):
parser.add_argument("--checkpoints-total-limit",
type=int,
help="Maximum number of checkpoints to keep")
parser.add_argument("--checkpointing-steps",
type=int,
help="Steps between checkpoints")
parser.add_argument(
"--training-state-checkpointing-steps",
type=int,
@@ -963,6 +1044,10 @@ class TrainingArgs(FastVideoArgs):
help="Linear quadratic threshold")
parser.add_argument("--linear-range", type=float, help="Linear range")
parser.add_argument("--weight-decay", type=float, help="Weight decay")
parser.add_argument("--betas",
type=str,
default=TrainingArgs.betas,
help="Betas for optimizer (format: 'beta1,beta2')")
parser.add_argument("--use-ema",
action=StoreBoolean,
help="Whether to use EMA")
@@ -1013,6 +1098,13 @@ class TrainingArgs(FastVideoArgs):
type=int,
default=TrainingArgs.generator_update_interval,
help="Ratio of student updates to critic updates.")
parser.add_argument(
"--dfake-gen-update-ratio",
type=int,
default=TrainingArgs.dfake_gen_update_ratio,
help=
"Self-forcing: How often to train generator vs critic (train generator every N steps)."
)
parser.add_argument("--min-timestep-ratio",
type=float,
default=TrainingArgs.min_timestep_ratio,
@@ -1029,6 +1121,11 @@ class TrainingArgs(FastVideoArgs):
type=float,
default=TrainingArgs.fake_score_learning_rate,
help="Learning rate for fake score transformer")
parser.add_argument(
"--fake-score-betas",
type=str,
default=TrainingArgs.fake_score_betas,
help="Betas for fake score optimizer (format: 'beta1,beta2')")
parser.add_argument(
"--fake-score-lr-scheduler",
type=str,
@@ -1041,6 +1138,48 @@ class TrainingArgs(FastVideoArgs):
"--simulate-generator-forward",
action=StoreBoolean,
help="Whether to simulate generator forward to match inference")
parser.add_argument(
"--warp-denoising-step",
action=StoreBoolean,
help=
"Whether to warp denoising step according to the scheduler time shift"
)
# Self-forcing specific arguments
parser.add_argument(
"--num-frame-per-block",
type=int,
default=TrainingArgs.num_frame_per_block,
help="Number of frames per block for causal generation")
parser.add_argument(
"--independent-first-frame",
action=StoreBoolean,
help="Whether the first frame is independent in causal generation")
parser.add_argument(
"--enable-gradient-masking",
action=StoreBoolean,
help="Whether to enable frame-level gradient masking")
parser.add_argument(
"--gradient-mask-last-n-frames",
type=int,
default=TrainingArgs.gradient_mask_last_n_frames,
help="Number of last frames to enable gradients for")
parser.add_argument(
"--validate-cache-structure",
action=StoreBoolean,
help="Whether to validate KV cache structure (debug flag)")
parser.add_argument(
"--same-step-across-blocks",
action=StoreBoolean,
help="Whether to use the same exit timestep for all blocks")
parser.add_argument(
"--last-step-only",
action=StoreBoolean,
help="Whether to only use the last timestep for training")
parser.add_argument("--context-noise",
type=int,
default=TrainingArgs.context_noise,
help="Context noise level for cache updates")
return parser
+2 -1
View File
@@ -7,7 +7,6 @@ import torch.nn as nn
import torch.nn.functional as F
from fastvideo.layers.custom_op import CustomOp
from fastvideo.platforms import current_platform
@CustomOp.register("rms_norm")
@@ -34,6 +33,8 @@ class RMSNorm(CustomOp):
else var_hidden_size)
self.has_weight = has_weight
from fastvideo.platforms import current_platform
self.weight = torch.ones(hidden_size) if current_platform.is_cuda_alike(
) else torch.ones(hidden_size, dtype=dtype)
if self.has_weight:
+61 -39
View File
@@ -33,7 +33,7 @@ from fastvideo.layers.visual_embedding import (PatchEmbed)
from fastvideo.logger import init_logger
from fastvideo.models.dits.base import BaseDiT
from fastvideo.models.dits.wanvideo import WanT2VCrossAttention, WanTimeTextImageEmbedding
from fastvideo.platforms import AttentionBackendEnum, current_platform
from fastvideo.platforms import AttentionBackendEnum
logger = init_logger(__name__)
class CausalWanSelfAttention(nn.Module):
@@ -147,6 +147,9 @@ class CausalWanSelfAttention(nn.Module):
# Assign new keys/values directly up to current_end
local_end_index = kv_cache["local_end_index"].item() + current_end - kv_cache["global_end_index"].item()
local_start_index = local_end_index - num_new_tokens
kv_cache["k"] = kv_cache["k"].detach()
kv_cache["v"] = kv_cache["v"].detach()
# logger.info("kv_cache['k'] is in comp graph: %s", kv_cache["k"].requires_grad or kv_cache["k"].grad_fn is not None)
kv_cache["k"][:, local_start_index:local_end_index] = roped_key
kv_cache["v"][:, local_start_index:local_end_index] = v
x = self.attn(
@@ -176,7 +179,7 @@ class CausalWanTransformerBlock(nn.Module):
super().__init__()
# 1. Self-attention
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
self.norm1 = nn.LayerNorm(dim, eps, elementwise_affine=False)
self.to_q = ReplicatedLinear(dim, dim, bias=True)
self.to_k = ReplicatedLinear(dim, dim, bias=True)
self.to_v = ReplicatedLinear(dim, dim, bias=True)
@@ -209,8 +212,7 @@ class CausalWanTransformerBlock(nn.Module):
norm_type="layer",
eps=eps,
elementwise_affine=True,
dtype=torch.float32,
compute_dtype=torch.float32)
dtype=torch.float32)
# 2. Cross-attention
# Only T2V for now
@@ -223,8 +225,7 @@ class CausalWanTransformerBlock(nn.Module):
norm_type="layer",
eps=eps,
elementwise_affine=False,
dtype=torch.float32,
compute_dtype=torch.float32)
dtype=torch.float32)
# 3. Feed-forward
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
@@ -249,29 +250,29 @@ class CausalWanTransformerBlock(nn.Module):
if hidden_states.dim() == 4:
hidden_states = hidden_states.squeeze(1)
num_frames = temb.shape[1]
frame_seqlen = hidden_states.shape[1] // num_frames
frame_seqlen = hidden_states.shape[1] // num_frames
bs, seq_length, _ = hidden_states.shape
orig_dtype = hidden_states.dtype
# assert orig_dtype != torch.float32
e = self.scale_shift_table + temb.float()
e = self.scale_shift_table + temb
# e.shape: [batch_size, num_frames, 6, inner_dim]
assert e.shape == (bs, num_frames, 6, self.hidden_dim)
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
6, dim=2)
# *_msa.shape: [batch_size, num_frames, 1, inner_dim]
assert shift_msa.dtype == torch.float32
# assert shift_msa.dtype == torch.float32
# 1. Self-attention
norm_hidden_states = (self.norm1(hidden_states.float()).unflatten(dim=1, sizes=(num_frames, frame_seqlen)) *
(1 + scale_msa) + shift_msa).flatten(1, 2).to(orig_dtype)
norm_hidden_states = (self.norm1(hidden_states).unflatten(dim=1, sizes=(num_frames, frame_seqlen)) *
(1 + scale_msa) + shift_msa).flatten(1, 2)
query, _ = self.to_q(norm_hidden_states)
key, _ = self.to_k(norm_hidden_states)
value, _ = self.to_v(norm_hidden_states)
if self.norm_q is not None:
query = self.norm_q(query)
query = self.norm_q.forward_native(query)
if self.norm_k is not None:
key = self.norm_k(key)
key = self.norm_k.forward_native(key)
query = query.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
@@ -285,8 +286,6 @@ class CausalWanTransformerBlock(nn.Module):
null_shift = null_scale = torch.tensor([0], device=hidden_states.device)
norm_hidden_states, hidden_states = self.self_attn_residual_norm(
hidden_states, attn_output, gate_msa, null_shift, null_scale)
norm_hidden_states, hidden_states = norm_hidden_states.to(
orig_dtype), hidden_states.to(orig_dtype)
# 2. Cross-attention
attn_output = self.attn2(norm_hidden_states,
@@ -295,13 +294,10 @@ class CausalWanTransformerBlock(nn.Module):
crossattn_cache=crossattn_cache)
norm_hidden_states, hidden_states = self.cross_attn_residual_norm(
hidden_states, attn_output, 1, c_shift_msa, c_scale_msa)
norm_hidden_states, hidden_states = norm_hidden_states.to(
orig_dtype), hidden_states.to(orig_dtype)
# 3. Feed-forward
ff_output = self.ffn(norm_hidden_states)
hidden_states = self.mlp_residual(hidden_states, ff_output, c_gate_msa)
hidden_states = hidden_states.to(orig_dtype)
return hidden_states
@@ -364,8 +360,7 @@ class CausalWanTransformer3DModel(BaseDiT):
norm_type="layer",
eps=config.eps,
elementwise_affine=False,
dtype=torch.float32,
compute_dtype=torch.float32)
dtype=torch.float32)
self.proj_out = nn.Linear(
inner_dim, config.out_channels * math.prod(config.patch_size))
self.scale_shift_table = nn.Parameter(
@@ -375,7 +370,8 @@ class CausalWanTransformer3DModel(BaseDiT):
# Causal-specific
self.block_mask = None
self.num_frame_per_block = 1
self.num_frame_per_block = config.arch_config.num_frames_per_block
assert self.num_frame_per_block <= 3
self.independent_first_frame = False
self.__post_init__()
@@ -456,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):
@@ -487,12 +484,16 @@ class CausalWanTransformer3DModel(BaseDiT):
)
freqs_cos = freqs_cos.to(hidden_states.device)
freqs_sin = freqs_sin.to(hidden_states.device)
freqs_cis = (freqs_cos.float(),
freqs_sin.float()) if freqs_cos is not None else None
freqs_cis = (freqs_cos,
freqs_sin) if freqs_cos is not None else None
hidden_states = self.patch_embedding(hidden_states)
grid_sizes = torch.stack(
[torch.tensor(hidden_states[0].shape[1:], dtype=torch.long)])
hidden_states = hidden_states.flatten(2).transpose(1, 2)
encoder_hidden_states = torch.cat([encoder_hidden_states, encoder_hidden_states.new_zeros(1, self.text_len - encoder_hidden_states.size(1), encoder_hidden_states.size(2))], dim=1)
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
timestep.flatten(), encoder_hidden_states, encoder_hidden_states_image)
timestep_proj = timestep_proj.unflatten(1, (6, self.hidden_size)).unflatten(dim=0, sizes=timestep.shape)
@@ -539,14 +540,9 @@ class CausalWanTransformer3DModel(BaseDiT):
hidden_states = self.norm_out(hidden_states, shift, scale)
hidden_states = self.proj_out(hidden_states)
hidden_states = hidden_states.reshape(batch_size, post_patch_num_frames,
post_patch_height,
post_patch_width, p_t, p_h, p_w,
-1)
hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6)
output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
output = self.unpatchify(hidden_states, grid_sizes)
return output
return torch.stack(output)
def _forward_train(self,
hidden_states: torch.Tensor,
@@ -587,8 +583,8 @@ class CausalWanTransformer3DModel(BaseDiT):
)
freqs_cos = freqs_cos.to(hidden_states.device)
freqs_sin = freqs_sin.to(hidden_states.device)
freqs_cis = (freqs_cos.float(),
freqs_sin.float()) if freqs_cos is not None else None
freqs_cis = (freqs_cos,
freqs_sin) if freqs_cos is not None else None
# Construct blockwise causal attn mask
if self.block_mask is None:
@@ -601,8 +597,12 @@ class CausalWanTransformer3DModel(BaseDiT):
)
hidden_states = self.patch_embedding(hidden_states)
grid_sizes = torch.stack(
[torch.tensor(hidden_states[0].shape[1:], dtype=torch.long)])
hidden_states = hidden_states.flatten(2).transpose(1, 2)
encoder_hidden_states = torch.cat([encoder_hidden_states, encoder_hidden_states.new_zeros(1, self.text_len - encoder_hidden_states.size(1), encoder_hidden_states.size(2))], dim=1)
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
timestep.flatten(), encoder_hidden_states, encoder_hidden_states_image)
timestep_proj = timestep_proj.unflatten(1, (6, self.hidden_size)).unflatten(dim=0, sizes=timestep.shape)
@@ -637,14 +637,9 @@ class CausalWanTransformer3DModel(BaseDiT):
hidden_states = self.norm_out(hidden_states, shift, scale)
hidden_states = self.proj_out(hidden_states)
hidden_states = hidden_states.reshape(batch_size, post_patch_num_frames,
post_patch_height,
post_patch_width, p_t, p_h, p_w,
-1)
hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6)
output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
output = self.unpatchify(hidden_states, grid_sizes)
return output
return torch.stack(output)
def forward(
self,
@@ -655,3 +650,30 @@ class CausalWanTransformer3DModel(BaseDiT):
return self._forward_inference(*args, **kwargs)
else:
return self._forward_train(*args, **kwargs)
def unpatchify(self, x, grid_sizes):
r"""
Args:
x (List[Tensor]):
List of patchified features, each with shape [L, C_out * prod(patch_size)]
grid_sizes (Tensor):
Original spatial-temporal grid dimensions before patching,
Returns:
Tensor:
Reconstructed video tensors with shape [B, C_out, F, H / 8, W / 8]
"""
c = self.out_channels
out = []
for u, v in zip(x, grid_sizes.tolist()):
u = u[:math.prod(v)].view(*v, *self.patch_size, c)
u = u.permute(6, 0, 3, 1, 4, 2, 5)
# u = torch.einsum('fhwpqrc->cfphqwr', u.contiguous())
u = u.reshape(c, *[i * j for i, j in zip(v, self.patch_size)])
out.append(u)
return out
+4 -3
View File
@@ -685,9 +685,10 @@ class WanTransformer3DModel(CachableDiT):
encoder_hidden_states = torch.concat(
[encoder_hidden_states_image, encoder_hidden_states], dim=1)
encoder_hidden_states = encoder_hidden_states.to(
orig_dtype) if current_platform.is_mps(
) else encoder_hidden_states # cast to orig_dtype for MPS
if current_platform.is_mps() or current_platform.is_npu():
encoder_hidden_states = encoder_hidden_states.to(orig_dtype)
else:
encoder_hidden_states = encoder_hidden_states # cast to orig_dtype for MPS & NPU
assert encoder_hidden_states.dtype == orig_dtype
+2 -2
View File
@@ -140,7 +140,7 @@ class CLIPAttention(nn.Module):
"embed_dim must be divisible by num_heads "
f"(got `embed_dim`: {self.embed_dim} and `num_heads`:"
f" {self.num_heads}).")
self.scale = self.head_dim**-0.5
self.scale = self.head_dim**-0.5 if config.enable_scale else None
self.dropout = config.attention_dropout
self.qkv_proj = QKVParallelLinear(
@@ -166,7 +166,7 @@ class CLIPAttention(nn.Module):
self.head_dim,
self.num_heads_per_partition,
softmax_scale=self.scale,
causal=True,
causal=config.is_causal,
supported_attention_backends=config._supported_attention_backends)
def _shape(self, tensor: torch.Tensor, seq_len: int, bsz: int):
+2 -1
View File
@@ -37,7 +37,6 @@ from fastvideo.layers.quantization import QuantizationConfig
from fastvideo.layers.vocab_parallel_embedding import VocabParallelEmbedding
from fastvideo.models.encoders.base import TextEncoder
from fastvideo.models.loader.weight_utils import default_weight_loader
from fastvideo.platforms import current_platform
class AttentionType:
@@ -322,6 +321,8 @@ class T5Attention(nn.Module):
# Encoder/Decoder Self-Attention Layer, attn bias already cached.
assert attn_bias is not None
from fastvideo.platforms import current_platform
if attention_mask is not None:
attention_mask = attention_mask.view(
bs, 1, 1,
+50 -7
View File
@@ -29,7 +29,6 @@ from fastvideo.models.loader.weight_utils import (
filter_duplicate_safetensors_files, filter_files_not_needed_for_inference,
pt_weights_iterator, safetensors_weights_iterator)
from fastvideo.models.registry import ModelRegistry
from fastvideo.platforms import current_platform
from fastvideo.utils import PRECISION_TO_TYPE
logger = init_logger(__name__)
@@ -251,6 +250,8 @@ class TextEncoderLoader(ComponentLoader):
use_cpu_offload = fastvideo_args.text_encoder_cpu_offload and len(
getattr(model_config, "_fsdp_shard_conditions", [])) > 0
from fastvideo.platforms import current_platform
if fastvideo_args.text_encoder_cpu_offload:
target_device = torch.device(
"mps") if current_platform.is_mps() else torch.device("cpu")
@@ -274,12 +275,27 @@ class TextEncoderLoader(ComponentLoader):
# Explicitly move model to target device after loading weights
model = model.to(target_device)
from fastvideo.platforms import current_platform
if use_cpu_offload:
# Disable FSDP for MPS as it's not compatible
if current_platform.is_mps():
logger.info(
"Disabling FSDP sharding for MPS platform as it's not compatible"
)
elif current_platform.is_npu():
mesh = init_device_mesh(
"npu",
mesh_shape=(1, dist.get_world_size()),
mesh_dim_names=("offload", "replicate"),
)
shard_model(
model,
cpu_offload=True,
reshard_after_forward=True,
mesh=mesh["offload"],
fsdp_shard_conditions=model._fsdp_shard_conditions,
pin_cpu_memory=fastvideo_args.pin_cpu_memory)
else:
mesh = init_device_mesh(
"cuda",
@@ -325,6 +341,8 @@ class ImageEncoderLoader(TextEncoderLoader):
encoder_config = fastvideo_args.pipeline_config.image_encoder_config
encoder_config.update_model_arch(model_config)
from fastvideo.platforms import current_platform
if fastvideo_args.image_encoder_cpu_offload:
target_device = torch.device("mps") if current_platform.is_mps() else torch.device("cpu")
else:
@@ -379,6 +397,8 @@ class VAELoader(ComponentLoader):
vae_config = fastvideo_args.pipeline_config.vae_config
vae_config.update_model_arch(config)
from fastvideo.platforms import current_platform
if fastvideo_args.vae_cpu_offload:
target_device = torch.device("mps") if current_platform.is_mps() else torch.device("cpu")
else:
@@ -416,6 +436,11 @@ class TransformerLoader(ComponentLoader):
"Model config does not contain a _class_name attribute. "
"Only diffusers format is supported.")
logger.info("transformer cls_name: %s", cls_name)
if fastvideo_args.override_transformer_cls_name is not None:
cls_name = fastvideo_args.override_transformer_cls_name
logger.info("Overriding transformer cls_name to %s", cls_name)
fastvideo_args.model_paths["transformer"] = model_path
# Config from Diffusers supersedes fastvideo's model config
@@ -430,8 +455,24 @@ class TransformerLoader(ComponentLoader):
if not safetensors_list:
raise ValueError(f"No safetensors files found in {model_path}")
logger.info("Loading model from %s safetensors files in %s",
len(safetensors_list), model_path)
# Check if we should use custom initialization weights
custom_weights_path = getattr(fastvideo_args, 'init_weights_from_safetensors', None)
use_custom_weights = (custom_weights_path and os.path.exists(custom_weights_path) and
not hasattr(fastvideo_args, '_loading_teacher_critic_model'))
if use_custom_weights:
if 'transformer_2' in model_path:
custom_weights_path = getattr(fastvideo_args, 'init_weights_from_safetensors_2', None)
assert custom_weights_path is not None, "Custom initialization weights must be provided"
if os.path.isdir(custom_weights_path):
safetensors_list = glob.glob(
os.path.join(str(custom_weights_path), "*.safetensors"))
else:
assert custom_weights_path.endswith(".safetensors"), "Custom initialization weights must be a safetensors file"
safetensors_list = [custom_weights_path]
logger.info("Loading model from %s safetensors files: %s",
len(safetensors_list), safetensors_list)
default_dtype = PRECISION_TO_TYPE[
fastvideo_args.pipeline_config.dit_precision]
@@ -454,18 +495,20 @@ class TransformerLoader(ComponentLoader):
pin_cpu_memory=fastvideo_args.pin_cpu_memory,
fsdp_inference=fastvideo_args.use_fsdp_inference,
# TODO(will): make these configurable
default_dtype=default_dtype,
param_dtype=torch.bfloat16,
reduce_dtype=torch.float32,
output_dtype=None,
training_mode=fastvideo_args.training_mode)
training_mode=fastvideo_args.training_mode,
enable_torch_compile=fastvideo_args.enable_torch_compile,
torch_compile_kwargs=fastvideo_args.torch_compile_kwargs)
total_params = sum(p.numel() for p in model.parameters())
logger.info("Loaded model with %.2fB parameters", total_params / 1e9)
dtypes = set(param.dtype for param in model.parameters())
if len(dtypes) > 1:
model = model.to(default_dtype)
assert next(model.parameters()).dtype == default_dtype, "Model dtype does not match default dtype"
model = model.eval()
return model
+26 -5
View File
@@ -54,7 +54,7 @@ def set_default_dtype(dtype: torch.dtype) -> Generator[None, None, None]:
torch.set_default_dtype(old_dtype)
# TODO(PY): add compile option
# Supports optional torch.compile for FSDP-wrapped models during training
def maybe_load_fsdp_model(
model_cls: type[nn.Module],
init_params: dict[str, Any],
@@ -62,6 +62,7 @@ def maybe_load_fsdp_model(
device: torch.device,
hsdp_replicate_dim: int,
hsdp_shard_dim: int,
default_dtype: torch.dtype,
param_dtype: torch.dtype,
reduce_dtype: torch.dtype,
cpu_offload: bool = False,
@@ -69,6 +70,8 @@ def maybe_load_fsdp_model(
output_dtype: torch.dtype | None = None,
training_mode: bool = True,
pin_cpu_memory: bool = True,
enable_torch_compile: bool = False,
torch_compile_kwargs: dict[str, Any] | None = None,
) -> torch.nn.Module:
"""
Load the model with FSDP if is training, else load the model without FSDP.
@@ -87,7 +90,8 @@ def maybe_load_fsdp_model(
mp_policy=mp_policy,
)
with set_default_dtype(param_dtype), torch.device("meta"):
logger.info("Loading model with default_dtype: %s", default_dtype)
with set_default_dtype(default_dtype), torch.device("meta"):
model = model_cls(**init_params)
# Check if we should use FSDP
@@ -104,8 +108,17 @@ def maybe_load_fsdp_model(
if not training_mode and not fsdp_inference:
hsdp_replicate_dim = world_size
hsdp_shard_dim = 1
device_mesh = init_device_mesh(
if current_platform.is_npu():
with torch.device("cpu"):
device_mesh = init_device_mesh(
"npu",
# (Replicate(), Shard(dim=0))
mesh_shape=(hsdp_replicate_dim, hsdp_shard_dim),
mesh_dim_names=("replicate", "shard"),
)
else:
device_mesh = init_device_mesh(
"cuda",
# (Replicate(), Shard(dim=0))
mesh_shape=(hsdp_replicate_dim, hsdp_shard_dim),
@@ -125,7 +138,7 @@ def maybe_load_fsdp_model(
model,
weight_iterator,
device,
param_dtype,
default_dtype,
strict=True,
cpu_offload=cpu_offload,
param_names_mapping=param_names_mapping_fn,
@@ -137,6 +150,14 @@ def maybe_load_fsdp_model(
# Avoid unintended computation graph accumulation during inference
if isinstance(p, torch.nn.Parameter):
p.requires_grad = False
compile_in_loader = enable_torch_compile and training_mode
if compile_in_loader:
compile_kwargs = torch_compile_kwargs or {}
logger.info("Enabling torch.compile for FSDP training module with kwargs=%s",
compile_kwargs)
model = torch.compile(model, **compile_kwargs)
logger.info("torch.compile enabled for %s", type(model).__name__)
return model
@@ -635,8 +635,31 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
noise: torch.Tensor,
timestep: torch.IntTensor,
) -> torch.Tensor:
"""
Args:
clean_latent: the clean latent with shape [B, C, H, W],
where B is batch_size or batch_size * num_frames
noise: the noise with shape [B, C, H, W]
timestep: the timestep with shape [1] or [bs * num_frames] or [bs, num_frames]
Returns:
the corrupted latent with shape [B, C, H, W]
"""
# If timestep is [bs, num_frames]
if timestep.ndim == 2:
timestep = timestep.flatten(0, 1)
assert timestep.numel() == clean_latent.shape[0]
elif timestep.ndim == 1:
# If timestep is [1]
if timestep.shape[0] == 1:
timestep = timestep.expand(clean_latent.shape[0])
else:
assert timestep.numel() == clean_latent.shape[0]
else:
raise ValueError(f"[add_noise] Invalid timestep shape: {timestep.shape}")
# timestep shape should be [B]
self.sigmas = self.sigmas.to(noise.device)
timestep = timestep.expand(clean_latent.shape[0])
self.timesteps = self.timesteps.to(noise.device)
timestep_id = torch.argmin(
(self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
@@ -22,8 +22,10 @@ class SelfForcingFlowMatchSchedulerOutput(BaseOutput):
prev_sample: torch.FloatTensor
class SelfForcingFlowMatchScheduler(BaseScheduler, ConfigMixin, SchedulerMixin):
config_name = "scheduler_config.json"
order = 1
@register_to_config
def __init__(self, num_inference_steps=100, num_train_timesteps=1000, shift=3.0, sigma_max=1.0, sigma_min=0.003 / 1.002, inverse_timesteps=False, extra_one_step=False, reverse_sigmas=False, training=False):
self.num_train_timesteps = num_train_timesteps
self.shift = shift
@@ -62,8 +64,15 @@ class SelfForcingFlowMatchScheduler(BaseScheduler, ConfigMixin, SchedulerMixin):
def step(self, model_output: torch.FloatTensor, timestep: torch.FloatTensor, sample: torch.FloatTensor, to_final=False, return_dict=False, **kwargs):
if timestep.ndim == 2:
timestep = timestep.flatten(0, 1)
elif timestep.ndim == 0:
# handles the case where timestep is a scalar, this occurs when we
# use this scheduler for ODE trajectory
timestep = timestep.unsqueeze(0)
self.sigmas = self.sigmas.to(model_output.device)
self.timesteps = self.timesteps.to(model_output.device)
timestep = timestep.to(model_output.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)
+4 -4
View File
@@ -171,10 +171,10 @@ def pred_noise_to_pred_video(pred_noise: torch.Tensor,
# timestep shape should be [B]
dtype = pred_noise.dtype
device = pred_noise.device
pred_noise = pred_noise.float().to(device)
noise_input_latent = noise_input_latent.float().to(device)
sigmas = scheduler.sigmas.float().to(device)
timesteps = scheduler.timesteps.float().to(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)
+2 -1
View File
@@ -26,7 +26,6 @@ from fastvideo.configs.models.vaes import WanVAEConfig
from fastvideo.layers.activation import get_act_fn
from fastvideo.models.vaes.common import (DiagonalGaussianDistribution,
ParallelTiledVAE)
from fastvideo.platforms import current_platform
CACHE_T = 2
@@ -189,6 +188,8 @@ class WanCausalConv3d(nn.Conv3d):
self.padding = (0, 0, 0)
def forward(self, x, cache_x=None):
from fastvideo.platforms import current_platform
padding = list(self._padding)
if cache_x is not None and self._padding[4] > 0:
cache_x = cache_x.to(x.device)
+146 -29
View File
@@ -3,6 +3,7 @@
import os
import tempfile
from collections.abc import Callable
from typing import Any
from urllib.parse import unquote, urlparse
import imageio
@@ -11,6 +12,8 @@ import PIL.Image
import PIL.ImageOps
import requests
import torch
import torch.nn.functional as F
import torchvision.transforms.functional as TF
from packaging import version
if version.parse(version.parse(
@@ -131,12 +134,88 @@ def load_image(
return image
def _load_gif(gif_path: str) -> tuple[list[PIL.Image.Image], float | None]:
"""
Load frames from a GIF file.
Args:
gif_path: Path to the GIF file
Returns:
Tuple of (list of PIL images, original FPS or None)
"""
pil_images = []
original_fps = None
with PIL.Image.open(gif_path) as gif:
# Extract FPS from GIF metadata
if hasattr(gif, 'info') and 'duration' in gif.info:
duration_ms = gif.info['duration']
if duration_ms > 0:
original_fps = 1000.0 / duration_ms
# Extract all frames
try:
while True:
pil_images.append(gif.copy())
gif.seek(gif.tell() + 1)
except EOFError:
# End of GIF reached
pass
return pil_images, original_fps
def _load_video_with_ffmpeg(
video_path: str) -> tuple[list[PIL.Image.Image], float | None]:
"""
Load frames from a video file using ffmpeg.
Args:
video_path: Path to the video file
Returns:
Tuple of (list of PIL images, original FPS or None)
Raises:
AttributeError: If ffmpeg is not installed
"""
# Verify ffmpeg is available
try:
imageio.plugins.ffmpeg.get_exe()
except AttributeError as e:
raise AttributeError(
"Unable to find an ffmpeg installation on your machine. "
"Please install via `pip install imageio-ffmpeg`") from e
pil_images = []
original_fps = None
with imageio.get_reader(video_path) as reader:
# Try to extract FPS metadata
metadata = reader.get_meta_data()
original_fps = metadata.get('fps')
# Fallback: try format-specific metadata
if original_fps is None:
source_size = metadata.get('source_size', {})
if isinstance(source_size, dict):
original_fps = source_size.get('fps')
# Extract all frames
for frame in reader:
pil_images.append(PIL.Image.fromarray(frame))
return pil_images, original_fps
# adapted from diffusers.utils import load_video
def load_video(
video: str,
convert_method: Callable[[list[PIL.Image.Image]], list[PIL.Image.Image]]
| None = None,
) -> list[PIL.Image.Image]:
return_fps: bool = False,
) -> tuple[list[PIL.Image.Image], float | Any] | list[PIL.Image.Image]:
"""
Loads `video` to a list of PIL Image.
Args:
@@ -145,9 +224,12 @@ def load_video(
convert_method (Callable[[List[PIL.Image.Image]], List[PIL.Image.Image]], *optional*):
A conversion method to apply to the video after loading it. When set to `None` the images will be converted
to "RGB".
return_fps (`bool`, *optional*, defaults to `False`):
Whether to return the FPS of the video. If `True`, returns a tuple of (images, fps).
If `False`, returns only the list of images.
Returns:
`List[PIL.Image.Image]`:
The video as a list of PIL images.
`List[PIL.Image.Image]` or `Tuple[List[PIL.Image.Image], float | None]`:
The video as a list of PIL images. If `return_fps` is True, also returns the original FPS.
"""
is_url = video.startswith("http://") or video.startswith("https://")
is_file = os.path.isfile(video)
@@ -175,39 +257,27 @@ def load_video(
video_data = response.iter_content(chunk_size=8192)
for chunk in video_data:
temp_file.write(chunk)
video = video_path
was_tempfile_created = True
else:
video_path = video
pil_images = []
if video.endswith(".gif"):
gif = PIL.Image.open(video)
try:
while True:
pil_images.append(gif.copy())
gif.seek(gif.tell() + 1)
except EOFError:
pass
original_fps = None
else:
try:
imageio.plugins.ffmpeg.get_exe()
except AttributeError:
raise AttributeError(
"`Unable to find an ffmpeg installation on your machine. Please install via `pip install imageio-ffmpeg"
) from None
with imageio.get_reader(video) as reader:
# Read all frames
for frame in reader:
pil_images.append(PIL.Image.fromarray(frame))
if was_tempfile_created:
os.remove(video_path)
try:
if video_path.endswith(".gif"):
pil_images, original_fps = _load_gif(video_path)
else:
pil_images, original_fps = _load_video_with_ffmpeg(video_path)
finally:
# Clean up temporary file if it was created
if was_tempfile_created and os.path.exists(video_path):
os.remove(video_path)
if convert_method is not None:
pil_images = convert_method(pil_images)
return pil_images
return pil_images, original_fps if return_fps else pil_images
def get_default_height_width(
@@ -297,3 +367,50 @@ def resize(
else:
raise ValueError(f"resize_mode {resize_mode} is not supported")
return image
def create_default_image(width: int = 512, height: int = 512, color: tuple[int, int, int] = (0, 0, 0)) -> PIL.Image.Image:
"""
Create a default black PIL image.
Args:
width: Image width in pixels
height: Image height in pixels
color: RGB color tuple
Returns:
PIL.Image.Image: A new PIL image with specified dimensions and color
"""
return PIL.Image.new("RGB", (width, height), color=color)
def preprocess_reference_image_for_clip(image: PIL.Image.Image, device: torch.device) -> PIL.Image.Image:
"""
Preprocess reference image to match CLIP encoder requirements.
Applies normalization, resizing to 224x224, and denormalization to ensure
the image is in the correct format for CLIP processing.
Args:
image: Input PIL image
device: Target device for tensor operations
Returns:
Preprocessed PIL image ready for CLIP encoder
"""
# Convert PIL to tensor and normalize to [-1, 1] range
image_tensor = TF.to_tensor(image).sub_(0.5).div_(0.5).to(device)
# Resize to CLIP's expected input size (224x224) using bicubic interpolation
resized_tensor = F.interpolate(
image_tensor.unsqueeze(0),
size=(224, 224),
mode='bicubic',
align_corners=False
).squeeze(0)
# Denormalize back to [0, 1] range
denormalized_tensor = resized_tensor.mul_(0.5).add_(0.5)
return TF.to_pil_image(denormalized_tensor)
@@ -7,8 +7,6 @@ This module wires the causal DMD denoising stage into the modular pipeline.
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 import ComposedPipelineBase, LoRAPipeline
# isort: off
@@ -28,10 +26,6 @@ class WanCausalDMDPipeline(LoRAPipeline, ComposedPipelineBase):
"text_encoder", "tokenizer", "vae", "transformer", "scheduler"
]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
self.modules["scheduler"] = FlowMatchEulerDiscreteScheduler(
shift=fastvideo_args.pipeline_config.flow_shift)
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
"""Set up pipeline stages with proper dependency injection."""
@@ -55,6 +49,7 @@ class WanCausalDMDPipeline(LoRAPipeline, ComposedPipelineBase):
self.add_stage(stage_name="denoising_stage",
stage=CausalDMDDenosingStage(
transformer=self.get_module("transformer"),
transformer_2=self.get_module("transformer_2", None),
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="decoding_stage",
@@ -0,0 +1,83 @@
# SPDX-License-Identifier: Apache-2.0
"""
Wan video-to-video diffusion pipeline implementation.
This module contains an implementation of the Wan video-to-video diffusion pipeline
using the modular pipeline architecture.
"""
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.pipelines.lora_pipeline import LoRAPipeline
# isort: off
from fastvideo.pipelines.stages import (
RefImageEncodingStage, ConditioningStage, DecodingStage, DenoisingStage,
VideoVAEEncodingStage, InputValidationStage, LatentPreparationStage,
TextEncodingStage, TimestepPreparationStage)
# isort: on
from fastvideo.models.schedulers.scheduling_flow_unipc_multistep import (
FlowUniPCMultistepScheduler)
logger = init_logger(__name__)
class WanVideoToVideoPipeline(LoRAPipeline, ComposedPipelineBase):
_required_config_modules = [
"text_encoder", "tokenizer", "vae", "transformer", "scheduler", \
"image_encoder", "image_processor"
]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
shift=fastvideo_args.pipeline_config.flow_shift)
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")],
))
if (self.get_module("image_encoder") is not None
and self.get_module("image_processor") is not None):
self.add_stage(
stage_name="ref_image_encoding_stage",
stage=RefImageEncodingStage(
image_encoder=self.get_module("image_encoder"),
image_processor=self.get_module("image_processor"),
))
self.add_stage(stage_name="conditioning_stage",
stage=ConditioningStage())
self.add_stage(stage_name="timestep_preparation_stage",
stage=TimestepPreparationStage(
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer")))
self.add_stage(stage_name="video_latent_preparation_stage",
stage=VideoVAEEncodingStage(vae=self.get_module("vae")))
self.add_stage(stage_name="denoising_stage",
stage=DenoisingStage(
transformer=self.get_module("transformer"),
transformer_2=self.get_module("transformer_2"),
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
EntryClass = WanVideoToVideoPipeline
+50 -6
View File
@@ -14,12 +14,14 @@ import torch
from fastvideo.configs.pipelines import PipelineConfig
from fastvideo.distributed import (
maybe_init_distributed_environment_and_model_parallel)
maybe_init_distributed_environment_and_model_parallel, get_world_group)
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.logger import init_logger
from fastvideo.profiler import get_or_create_profiler
from fastvideo.models.loader.component_loader import PipelineComponentLoader
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages import PipelineStage
import fastvideo.envs as envs
from fastvideo.utils import (maybe_download_model,
verify_model_config_and_directory)
@@ -40,7 +42,7 @@ class ComposedPipelineBase(ABC):
_extra_config_module_map: dict[str, str] = {}
training_args: TrainingArgs | None = None
fastvideo_args: FastVideoArgs | TrainingArgs | None = None
modules: dict[str, torch.nn.Module] = {}
modules: dict[str, Any] = {}
post_init_called: bool = False
# TODO(will): args should support both inference args and training args
@@ -69,9 +71,18 @@ class ComposedPipelineBase(ABC):
maybe_init_distributed_environment_and_model_parallel(
fastvideo_args.tp_size, fastvideo_args.sp_size)
# Torch profiler. Enabled and configured through env vars:
# FASTVIDEO_TORCH_PROFILER_DIR=/path/to/save/trace
trace_dir = envs.FASTVIDEO_TORCH_PROFILER_DIR
self.profiler_controller = get_or_create_profiler(trace_dir)
self.profiler = self.profiler_controller.profiler
self.local_rank = get_world_group().local_rank
# Load modules directly in initialization
logger.info("Loading pipeline modules...")
self.modules = self.load_modules(fastvideo_args, loaded_modules)
with self.profiler_controller.region("profiler_region_model_loading"):
self.modules = self.load_modules(fastvideo_args, loaded_modules)
def set_trainable(self) -> None:
# Only train DiT
@@ -99,9 +110,30 @@ class ComposedPipelineBase(ABC):
self.initialize_pipeline(self.fastvideo_args)
if self.fastvideo_args.enable_torch_compile:
self.modules["transformer"] = torch.compile(
self.modules["transformer"])
logger.info("Torch Compile enabled for DiT")
transformer_module = self.modules["transformer"]
if self.fastvideo_args.training_mode:
logger.info(
"Torch Compile enabled via FSDP loader for training; skipping additional pipeline compile"
)
else:
fsdp_module_cls = None
try:
from torch.distributed.fsdp import FSDPModule # type: ignore
fsdp_module_cls = FSDPModule
except Exception: # pragma: no cover - FSDP not always available
fsdp_module_cls = None
if fsdp_module_cls is not None and isinstance(
transformer_module, fsdp_module_cls):
logger.info(
"Transformer is already FSDP-wrapped; skipping torch.compile in pipeline"
)
else:
compile_kwargs = self.fastvideo_args.torch_compile_kwargs or {}
logger.info("Enabling torch.compile for DiT with kwargs=%s",
compile_kwargs)
self.modules["transformer"] = torch.compile(
transformer_module, **compile_kwargs)
logger.info("Torch Compile enabled for DiT")
if not self.fastvideo_args.training_mode:
logger.info("Creating pipeline stages...")
@@ -331,6 +363,18 @@ class ComposedPipelineBase(ABC):
self._stage_name_mapping[stage_name] = stage
setattr(self, stage_name, stage)
def profile(self, is_start: bool = True):
if self.profiler is None:
raise RuntimeError("Profiler is not enabled.")
if is_start:
self.profiler.start()
else:
self.profiler.stop()
# only print profiler results on rank 0
if self.local_rank == 0:
print(self.profiler.key_averages().table(
sort_by="self_cuda_time_total"))
# TODO(will): don't hardcode no_grad
@torch.no_grad()
def forward(
+16 -1
View File
@@ -86,6 +86,11 @@ class ForwardBatch:
prompt_path: str | None = None
output_path: str = "outputs/"
output_video_name: str | None = None
# Video inputs
video_path: str | None = None
video_latent: torch.Tensor | None = None
# Primary encoder embeddings
prompt_embeds: list[torch.Tensor] = field(default_factory=list)
negative_prompt_embeds: list[torch.Tensor] | None = None
@@ -148,7 +153,12 @@ class ForwardBatch:
modules: dict[str, Any] = field(default_factory=dict)
# Final output (after pipeline completion)
output: Any = None
output: torch.Tensor | None = None
return_trajectory_latents: bool = False
return_trajectory_decoded: bool = False
trajectory_timesteps: list[torch.Tensor] | None = None
trajectory_latents: torch.Tensor | None = None
trajectory_decoded: list[torch.Tensor] | None = None
# Extra parameters that might be needed by specific pipeline implementations
extra: dict[str, Any] = field(default_factory=dict)
@@ -207,6 +217,10 @@ class TrainingBatch:
infos: list[dict[str, Any]] | None = None
mask_lat_size: torch.Tensor | None = None
# ODE trajectory supervision
trajectory_latents: torch.Tensor | None = None
trajectory_timesteps: torch.Tensor | None = None
# Transformer inputs
noisy_model_input: torch.Tensor | None = None
timesteps: torch.Tensor | None = None
@@ -237,6 +251,7 @@ class TrainingBatch:
fake_score_loss: float = 0.0
dmd_latent_vis_dict: dict[str, Any] = field(default_factory=dict)
latent_vis_dict: dict[str, Any] = field(default_factory=dict)
fake_score_latent_vis_dict: dict[str, Any] = field(default_factory=dict)
+1
View File
@@ -21,6 +21,7 @@ _PIPELINE_NAME_TO_ARCHITECTURE_NAME: dict[str, str] = {
"WanPipeline": "wan",
"WanDMDPipeline": "wan",
"WanImageToVideoPipeline": "wan",
"WanVideoToVideoPipeline": "wan",
"WanCausalDMDPipeline": "wan",
"StepVideoPipeline": "stepvideo",
"HunyuanVideoPipeline": "hunyuan",
@@ -0,0 +1,323 @@
# SPDX-License-Identifier: Apache-2.0
"""
ODE Trajectory Data Preprocessing pipeline implementation.
This module contains an implementation of the ODE Trajectory Data Preprocessing pipeline
using the modular pipeline architecture.
Sec 4.3 of CausVid paper: https://arxiv.org/pdf/2412.07772
"""
import os
from collections.abc import Iterator
from typing import Any
import pyarrow as pa
import torch
from torch.utils.data import DataLoader
from torchdata.stateful_dataloader import StatefulDataLoader
from tqdm import tqdm
from fastvideo.configs.sample import SamplingParam
from fastvideo.dataset import gettextdataset
from fastvideo.dataset.dataloader.parquet_io import (ParquetDatasetWriter,
records_to_table)
from fastvideo.dataset.dataloader.record_schema import (
ode_text_only_record_creator)
from fastvideo.dataset.dataloader.schema import (
pyarrow_schema_ode_trajectory_text_only)
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.models.schedulers.scheduling_self_forcing_flow_match import (
SelfForcingFlowMatchScheduler)
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.preprocess.preprocess_pipeline_base import (
BasePreprocessPipeline)
from fastvideo.pipelines.stages import (DecodingStage, DenoisingStage,
InputValidationStage,
LatentPreparationStage,
TextEncodingStage,
TimestepPreparationStage)
from fastvideo.utils import save_decoded_latents_as_video, shallow_asdict
logger = init_logger(__name__)
class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
"""ODE Trajectory preprocessing pipeline implementation."""
_required_config_modules = [
"text_encoder", "tokenizer", "vae", "transformer", "scheduler"
]
preprocess_dataloader: StatefulDataLoader
preprocess_loader_iter: Iterator[dict[str, Any]]
pbar: Any
num_processed_samples: int
def get_pyarrow_schema(self) -> pa.Schema:
"""Return the PyArrow schema for ODE Trajectory pipeline."""
return pyarrow_schema_ode_trajectory_text_only
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
assert fastvideo_args.pipeline_config.flow_shift == 5
self.modules["scheduler"] = SelfForcingFlowMatchScheduler(
shift=fastvideo_args.pipeline_config.flow_shift,
sigma_min=0.0,
extra_one_step=True)
self.modules["scheduler"].set_timesteps(num_inference_steps=48,
denoising_strength=1.0)
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="timestep_preparation_stage",
stage=TimestepPreparationStage(
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer", None)))
self.add_stage(stage_name="denoising_stage",
stage=DenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler"),
pipeline=self,
))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
def preprocess_text_and_trajectory(self, fastvideo_args: FastVideoArgs,
args):
"""Preprocess text-only data and generate trajectory information."""
for batch_idx, data in enumerate(self.pbar):
if data is None:
continue
with torch.inference_mode():
# For text-only processing, we only need text data
# Filter out samples without text
valid_indices = []
for i, text in enumerate(data["text"]):
if text and text.strip(): # Check if text is not empty
valid_indices.append(i)
self.num_processed_samples += len(valid_indices)
if not valid_indices:
continue
# Create new batch with only valid samples (text-only)
valid_data = {
"text": [data["text"][i] for i in valid_indices],
"path": [data["path"][i] for i in valid_indices],
}
# Add fps and duration if available in data
if "fps" in data:
valid_data["fps"] = [data["fps"][i] for i in valid_indices]
if "duration" in data:
valid_data["duration"] = [
data["duration"][i] for i in valid_indices
]
batch_captions = valid_data["text"]
# Encode text using the standalone TextEncodingStage API
prompt_embeds_list, prompt_masks_list = self.prompt_encoding_stage.encode_text(
batch_captions,
fastvideo_args,
encoder_index=[0],
return_attention_mask=True,
)
prompt_embeds = prompt_embeds_list[0]
prompt_attention_masks = prompt_masks_list[0]
assert prompt_embeds.shape[0] == prompt_attention_masks.shape[0]
sampling_params = SamplingParam.from_pretrained(args.model_path)
# encode negative prompt for trajectory collection
if sampling_params.guidance_scale > 1 and sampling_params.negative_prompt is not None:
negative_prompt_embeds_list, negative_prompt_masks_list = self.prompt_encoding_stage.encode_text(
sampling_params.negative_prompt,
fastvideo_args,
encoder_index=[0],
return_attention_mask=True,
)
negative_prompt_embed = negative_prompt_embeds_list[0][0]
negative_prompt_attention_mask = negative_prompt_masks_list[
0][0]
else:
negative_prompt_embed = None
negative_prompt_attention_mask = None
trajectory_latents = []
trajectory_timesteps = []
trajectory_decoded = []
for i, (prompt_embed, prompt_attention_mask) in enumerate(
zip(prompt_embeds, prompt_attention_masks,
strict=False)):
prompt_embed = prompt_embed.unsqueeze(0)
prompt_attention_mask = prompt_attention_mask.unsqueeze(0)
# Collect the trajectory data (text-to-video generation)
batch = ForwardBatch(**shallow_asdict(sampling_params), )
batch.prompt_embeds = [prompt_embed]
batch.prompt_attention_mask = [prompt_attention_mask]
batch.negative_prompt_embeds = [negative_prompt_embed]
batch.negative_attention_mask = [
negative_prompt_attention_mask
]
batch.num_inference_steps = 48
batch.return_trajectory_latents = True
# Enabling this will save the decoded trajectory videos.
# Used for debugging.
batch.return_trajectory_decoded = False
batch.height = args.max_height
batch.width = args.max_width
batch.fps = args.train_fps
batch.guidance_scale = 6.0
batch.do_classifier_free_guidance = True
result_batch = self.input_validation_stage(
batch, fastvideo_args)
result_batch = self.timestep_preparation_stage(
batch, fastvideo_args)
result_batch = self.latent_preparation_stage(
result_batch, fastvideo_args)
result_batch = self.denoising_stage(result_batch,
fastvideo_args)
result_batch = self.decoding_stage(result_batch,
fastvideo_args)
trajectory_latents.append(
result_batch.trajectory_latents.cpu())
trajectory_timesteps.append(
result_batch.trajectory_timesteps.cpu())
trajectory_decoded.append(result_batch.trajectory_decoded)
# Prepare extra features for text-only processing
extra_features = {
"trajectory_latents": trajectory_latents,
"trajectory_timesteps": trajectory_timesteps
}
if batch.return_trajectory_decoded:
for i, decoded_frames in enumerate(trajectory_decoded):
for j, decoded_frame in enumerate(decoded_frames):
save_decoded_latents_as_video(
decoded_frame,
f"decoded_videos/trajectory_decoded_{i}_{j}.mp4",
args.train_fps)
# Prepare batch data for Parquet dataset
batch_data: list[dict[str, Any]] = []
# Add progress bar for saving outputs
save_pbar = tqdm(enumerate(valid_data["path"]),
desc="Saving outputs",
unit="item",
leave=False)
for idx, video_path in save_pbar:
video_name = os.path.basename(video_path).split(".")[0]
# Convert tensors to numpy arrays
text_embedding = prompt_embeds[idx].cpu().numpy()
# Get extra features for this sample
sample_extra_features = {}
if extra_features:
for key, value in extra_features.items():
if isinstance(value, torch.Tensor):
sample_extra_features[key] = value[idx].cpu(
).numpy()
else:
assert isinstance(value, list)
if isinstance(value[idx], torch.Tensor):
sample_extra_features[key] = value[idx].cpu(
).float().numpy()
else:
sample_extra_features[key] = value[idx]
# Create record for Parquet dataset (text-only ODE schema)
record: dict[str, Any] = ode_text_only_record_creator(
video_name=video_name,
text_embedding=text_embedding,
caption=valid_data["text"][idx],
trajectory_latents=sample_extra_features[
"trajectory_latents"],
trajectory_timesteps=sample_extra_features[
"trajectory_timesteps"],
)
batch_data.append(record)
if batch_data:
write_pbar = tqdm(total=1,
desc="Writing to Parquet dataset",
unit="batch")
table = records_to_table(batch_data,
self.get_pyarrow_schema())
write_pbar.update(1)
write_pbar.close()
if not hasattr(self, 'dataset_writer'):
self.dataset_writer = ParquetDatasetWriter(
out_dir=self.combined_parquet_dir,
samples_per_file=args.samples_per_file,
)
self.dataset_writer.append_table(table)
logger.info("Collected batch with %s samples", len(table))
if self.num_processed_samples >= args.flush_frequency:
written = self.dataset_writer.flush()
logger.info("Flushed %s samples to parquet", written)
self.num_processed_samples = 0
# Final flush for any remaining samples
if hasattr(self, 'dataset_writer'):
written = self.dataset_writer.flush(write_remainder=True)
if written:
logger.info("Final flush wrote %s samples", written)
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs, args):
if not self.post_init_called:
self.post_init()
self.local_rank = int(os.getenv("RANK", 0))
os.makedirs(args.output_dir, exist_ok=True)
# Create directory for combined data
self.combined_parquet_dir = os.path.join(args.output_dir,
"combined_parquet_dataset")
os.makedirs(self.combined_parquet_dir, exist_ok=True)
# Loading dataset
train_dataset = gettextdataset(args)
self.preprocess_dataloader = DataLoader(
train_dataset,
batch_size=args.preprocess_video_batch_size,
num_workers=args.dataloader_num_workers,
)
self.preprocess_loader_iter = iter(self.preprocess_dataloader)
self.num_processed_samples = 0
# Add progress bar for video preprocessing
self.pbar = tqdm(self.preprocess_loader_iter,
desc="Processing videos",
unit="batch",
disable=self.local_rank != 0)
# Initialize class variables for data sharing
self.video_data: dict[str, Any] = {} # Store video metadata and paths
self.latent_data: dict[str, Any] = {} # Store latent tensors
self.preprocess_text_and_trajectory(fastvideo_args, args)
EntryClass = PreprocessPipeline_ODE_Trajectory
@@ -10,6 +10,8 @@ from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines.preprocess.preprocess_pipeline_i2v import (
PreprocessPipeline_I2V)
from fastvideo.pipelines.preprocess.preprocess_pipeline_ode_trajectory import (
PreprocessPipeline_ODE_Trajectory)
from fastvideo.pipelines.preprocess.preprocess_pipeline_t2v import (
PreprocessPipeline_T2V)
from fastvideo.pipelines.preprocess.preprocess_pipeline_text import (
@@ -48,13 +50,16 @@ def main(args) -> None:
text_encoder_cpu_offload=False,
pipeline_config=pipeline_config,
)
if args.preprocess_task == "t2v":
PreprocessPipeline = PreprocessPipeline_T2V
elif args.preprocess_task == "i2v":
PreprocessPipeline = PreprocessPipeline_I2V
elif args.preprocess_task == "text_only":
PreprocessPipeline = PreprocessPipeline_Text
elif args.preprocess_task == "ode_trajectory":
assert args.flow_shift is not None, "flow_shift is required for ode_trajectory"
fastvideo_args.pipeline_config.flow_shift = args.flow_shift
PreprocessPipeline = PreprocessPipeline_ODE_Trajectory
else:
raise ValueError(f"Invalid preprocess task: {args.preprocess_task}. "
f"Valid options: t2v, i2v, ode_trajectory, text_only")
@@ -100,10 +105,11 @@ if __name__ == "__main__":
parser.add_argument("--video_length_tolerance_range", type=int, default=2.0)
parser.add_argument("--group_frame", action="store_true") # TODO
parser.add_argument("--group_resolution", action="store_true") # TODO
parser.add_argument("--flow_shift", type=float, default=None)
parser.add_argument("--preprocess_task",
type=str,
default="t2v",
choices=["t2v", "i2v", "text_only"],
choices=["t2v", "i2v", "text_only", "ode_trajectory"],
help="Type of preprocessing task to run")
parser.add_argument("--train_fps", type=int, default=30)
parser.add_argument("--use_image_num", type=int, default=0)
+5 -1
View File
@@ -14,7 +14,9 @@ from fastvideo.pipelines.stages.denoising import (DenoisingStage,
DmdDenoisingStage)
from fastvideo.pipelines.stages.encoding import EncodingStage
from fastvideo.pipelines.stages.image_encoding import (ImageEncodingStage,
ImageVAEEncodingStage)
RefImageEncodingStage,
ImageVAEEncodingStage,
VideoVAEEncodingStage)
from fastvideo.pipelines.stages.input_validation import InputValidationStage
from fastvideo.pipelines.stages.latent_preparation import LatentPreparationStage
from fastvideo.pipelines.stages.stepvideo_encoding import (
@@ -35,7 +37,9 @@ __all__ = [
"EncodingStage",
"DecodingStage",
"ImageEncodingStage",
"RefImageEncodingStage",
"ImageVAEEncodingStage",
"VideoVAEEncodingStage",
"TextEncodingStage",
"StepvideoPromptEncodingStage",
]
+13 -10
View File
@@ -34,13 +34,15 @@ class CausalDMDDenosingStage(DenoisingStage):
Denoising stage for causal diffusion.
"""
def __init__(self, transformer, scheduler) -> None:
super().__init__(transformer, scheduler)
def __init__(self, transformer, scheduler, transformer_2=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
# Model-dependent constants (aligned with causal_inference.py assumptions)
self.num_transformer_blocks = self.transformer.config.arch_config.num_layers
self.num_transformer_blocks = len(self.transformer.blocks)
self.num_frames_per_block = self.transformer.config.arch_config.num_frames_per_block
self.sliding_window_num_frames = self.transformer.config.arch_config.sliding_window_num_frames
@@ -65,21 +67,18 @@ class CausalDMDDenosingStage(DenoisingStage):
-1] * self.transformer.config.arch_config.patch_size[-2]
self.frame_seq_length = latent_seq_length // patch_ratio
# TODO(will): make this a parameter once we add i2v support
independent_first_frame = self.transformer.independent_first_frame
independent_first_frame = self.transformer.independent_first_frame if hasattr(
self.transformer, 'independent_first_frame') else False
# Timesteps for DMD
timesteps = torch.tensor(
fastvideo_args.pipeline_config.dmd_denoising_steps,
dtype=torch.long).cpu()
if fastvideo_args.pipeline_config.warp_denoising_step:
logger.info("Warping timesteps...")
scheduler_timesteps = torch.cat((self.scheduler.timesteps.cpu(),
torch.tensor([0],
dtype=torch.float32)))
timesteps = scheduler_timesteps[1000 - timesteps]
timesteps = timesteps.to(get_local_torch_device())
logger.info("Using timesteps: %s", timesteps)
# Image kwargs (kept empty unless caller provides compatible args)
image_kwargs: dict = {}
@@ -223,6 +222,10 @@ 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:
current_model = self.transformer_2
else:
current_model = self.transformer
# Copy for pred conversion
noise_latents = noise_latents_btchw.clone()
latent_model_input = current_latents.to(target_dtype)
@@ -273,7 +276,7 @@ class CausalDMDDenosingStage(DenoisingStage):
(latent_model_input.shape[0], 1),
device=latent_model_input.device,
dtype=torch.long)
pred_noise_btchw = self.transformer(
pred_noise_btchw = current_model(
latent_model_input,
prompt_embeds,
t_expanded_noise,
@@ -338,7 +341,7 @@ class CausalDMDDenosingStage(DenoisingStage):
attn_metadata=attn_metadata,
forward_batch=batch):
t_expanded_context = t_context.unsqueeze(1)
_ = self.transformer(
_ = current_model(
context_bcthw,
prompt_embeds,
t_expanded_context,
+91 -48
View File
@@ -50,6 +50,63 @@ class DecodingStage(PipelineStage):
result.add_check("output", batch.output, [V.is_tensor, V.with_dims(5)])
return result
@torch.no_grad()
def decode(self, latents: torch.Tensor,
fastvideo_args: FastVideoArgs) -> torch.Tensor:
"""
Decode latent representations into pixel space using VAE.
Args:
latents: Input latent tensor with shape (batch, channels, frames, height_latents, width_latents)
fastvideo_args: Configuration containing:
- disable_autocast: Whether to disable automatic mixed precision (default: False)
- pipeline_config.vae_precision: VAE computation precision ("fp32", "fp16", "bf16")
- pipeline_config.vae_tiling: Whether to enable VAE tiling for memory efficiency
Returns:
Decoded video tensor with shape (batch, channels, frames, height, width),
normalized to [0, 1] range and moved to CPU as float32
"""
self.vae = self.vae.to(get_local_torch_device())
latents = latents.to(get_local_torch_device())
# Setup VAE precision
vae_dtype = PRECISION_TO_TYPE[
fastvideo_args.pipeline_config.vae_precision]
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
# Apply shifting if needed
if (hasattr(self.vae, "shift_factor")
and self.vae.shift_factor is not None):
if isinstance(self.vae.shift_factor, torch.Tensor):
latents += self.vae.shift_factor.to(latents.device,
latents.dtype)
else:
latents += self.vae.shift_factor
# Decode latents
with torch.autocast(device_type="cuda",
dtype=vae_dtype,
enabled=vae_autocast_enabled):
if fastvideo_args.pipeline_config.vae_tiling:
self.vae.enable_tiling()
# if fastvideo_args.vae_sp:
# self.vae.enable_parallel()
if not vae_autocast_enabled:
latents = latents.to(vae_dtype)
image = self.vae.decode(latents)
# Normalize image to [0, 1] range
image = (image / 2 + 0.5).clamp(0, 1)
return image
@torch.no_grad()
def forward(
self,
@@ -59,13 +116,28 @@ class DecodingStage(PipelineStage):
"""
Decode latent representations into pixel space.
This method processes the batch through the VAE decoder, converting latent
representations to pixel-space video/images. It also optionally decodes
trajectory latents for visualization purposes.
Args:
batch: The current batch information.
fastvideo_args: The inference arguments.
batch: The current batch containing:
- latents: Tensor to decode (batch, channels, frames, height_latents, width_latents)
- return_trajectory_decoded (optional): Flag to decode trajectory latents
- trajectory_latents (optional): Latents at different timesteps
- trajectory_timesteps (optional): Corresponding timesteps
fastvideo_args: Configuration containing:
- output_type: "latent" to skip decoding, otherwise decode to pixels
- vae_cpu_offload: Whether to offload VAE to CPU after decoding
- model_loaded: Track VAE loading state
- model_paths: Path to VAE model if loading needed
Returns:
The batch with decoded outputs.
Modified batch with:
- output: Decoded frames (batch, channels, frames, height, width) as CPU float32
- trajectory_decoded (if requested): List of decoded frames per timestep
"""
# load vae if not already loaded (used for memory constrained devices)
pipeline = self.pipeline() if self.pipeline else None
if not fastvideo_args.model_loaded["vae"]:
loader = VAELoader()
@@ -75,58 +147,29 @@ class DecodingStage(PipelineStage):
pipeline.add_module("vae", self.vae)
fastvideo_args.model_loaded["vae"] = True
self.vae = self.vae.to(get_local_torch_device())
latents = batch.latents
# TODO(will): remove this once we add input/output validation for stages
if latents is None:
raise ValueError("Latents must be provided")
# Skip decoding if output type is latent
if fastvideo_args.output_type == "latent":
image = latents
frames = batch.latents
else:
# Setup VAE precision
vae_dtype = PRECISION_TO_TYPE[
fastvideo_args.pipeline_config.vae_precision]
vae_autocast_enabled = (vae_dtype != torch.float32
) and not fastvideo_args.disable_autocast
frames = self.decode(batch.latents, fastvideo_args)
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")
and self.vae.shift_factor is not None):
if isinstance(self.vae.shift_factor, torch.Tensor):
latents += self.vae.shift_factor.to(latents.device,
latents.dtype)
else:
latents += self.vae.shift_factor
# Decode latents
with torch.autocast(device_type="cuda",
dtype=vae_dtype,
enabled=vae_autocast_enabled):
if fastvideo_args.pipeline_config.vae_tiling:
self.vae.enable_tiling()
# if fastvideo_args.vae_sp:
# self.vae.enable_parallel()
if not vae_autocast_enabled:
latents = latents.to(vae_dtype)
image = self.vae.decode(latents)
# Normalize image to [0, 1] range
image = (image / 2 + 0.5).clamp(0, 1)
# decode trajectory latents if needed
if batch.return_trajectory_decoded:
batch.trajectory_decoded = []
assert batch.trajectory_latents is not None, "batch should have trajectory latents"
for idx in range(batch.trajectory_latents.shape[1]):
# batch.trajectory_latents is [batch_size, timesteps, channels, frames, height, width]
cur_latent = batch.trajectory_latents[:, idx, :, :, :, :]
cur_timestep = batch.trajectory_timesteps[idx]
logger.info("decoding trajectory latent for timestep: %s",
cur_timestep)
decoded_frames = self.decode(cur_latent, fastvideo_args)
batch.trajectory_decoded.append(decoded_frames.cpu().float())
# Convert to CPU float32 for compatibility
image = image.cpu().float()
frames = frames.cpu().float()
# Update batch with decoded image
batch.output = image
batch.output = frames
# Offload models if needed
if hasattr(self, 'maybe_free_model_hooks'):
+46 -11
View File
@@ -85,8 +85,9 @@ class DenoisingStage(PipelineStage):
AttentionBackendEnum.SLIDING_TILE_ATTN,
AttentionBackendEnum.VIDEO_SPARSE_ATTN,
AttentionBackendEnum.VMOBA_ATTN,
AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA
) # hack
AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA,
AttentionBackendEnum.SAGE_ATTN_THREE) # hack
)
def forward(
@@ -204,14 +205,14 @@ class DenoisingStage(PipelineStage):
neg_prompt_embeds[0]).any(), "neg_prompt_embeds contains nan"
# (Wan2.2) Calculate timestep to switch from high noise expert to low noise expert
if fastvideo_args.pipeline_config.dit_config.boundary_ratio is not None:
boundary_timestep = fastvideo_args.pipeline_config.dit_config.boundary_ratio
if batch.boundary_timestep is not None:
logger.info("Overriding boundary timestep from %s to %s",
boundary_timestep, batch.boundary_timestep)
boundary_timestep = batch.boundary_timestep
boundary_ratio = fastvideo_args.pipeline_config.dit_config.boundary_ratio
if batch.boundary_ratio is not None:
logger.info("Overriding boundary ratio from %s to %s",
boundary_ratio, batch.boundary_ratio)
boundary_ratio = batch.boundary_ratio
boundary_timestep *= self.scheduler.num_train_timesteps
if boundary_ratio is not None:
boundary_timestep = boundary_ratio * self.scheduler.num_train_timesteps
else:
boundary_timestep = None
latent_model_input = latents.to(target_dtype)
@@ -253,6 +254,10 @@ class DenoisingStage(PipelineStage):
patch_size[2])
seq_len = int(math.ceil(seq_len / sp_world_size)) * sp_world_size
# Initialize lists for ODE trajectory
trajectory_timesteps: list[torch.Tensor] = []
trajectory_latents: list[torch.Tensor] = []
# Run denoising loop
with self.progress_bar(total=num_inference_steps) as progress_bar:
for i, t in enumerate(timesteps):
@@ -278,9 +283,15 @@ class DenoisingStage(PipelineStage):
current_guidance_scale = batch.guidance_scale_2
assert current_model is not None, "current_model is None"
# Expand latents for I2V
# Expand latents for V2V/I2V
latent_model_input = latents.to(target_dtype)
if batch.image_latent is not None:
if batch.video_latent is not None:
latent_model_input = torch.cat([
latent_model_input, batch.video_latent,
torch.zeros_like(latents)
],
dim=1).to(target_dtype)
elif batch.image_latent is not None:
assert not fastvideo_args.pipeline_config.ti2v_task, "image latents should not be provided for TI2V task"
latent_model_input = torch.cat(
[latent_model_input, batch.image_latent],
@@ -433,6 +444,11 @@ class DenoisingStage(PipelineStage):
latents = (1. - mask2[0]) * z + mask2[0] * latents
# latents = latents.unsqueeze(0)
# save trajectory latents if needed
if batch.return_trajectory_latents:
trajectory_timesteps.append(t)
trajectory_latents.append(latents)
# Update progress bar
if i == len(timesteps) - 1 or (
(i + 1) > num_warmup_steps and
@@ -440,9 +456,28 @@ class DenoisingStage(PipelineStage):
and progress_bar is not None):
progress_bar.update()
# Gather results if using sequence parallelism
trajectory_tensor: torch.Tensor | None = None
if trajectory_latents:
trajectory_tensor = torch.stack(trajectory_latents, dim=1)
trajectory_timesteps_tensor = torch.stack(trajectory_timesteps,
dim=0)
else:
trajectory_tensor = None
trajectory_timesteps_tensor = None
# Gather results if using sequence parallelism
if sp_group:
latents = sequence_model_parallel_all_gather(latents, dim=2)
if batch.return_trajectory_latents:
trajectory_tensor = trajectory_tensor.to(
get_local_torch_device())
trajectory_tensor = sequence_model_parallel_all_gather(
trajectory_tensor, dim=3)
if trajectory_tensor is not None and trajectory_timesteps_tensor is not None:
batch.trajectory_timesteps = trajectory_timesteps_tensor.cpu()
batch.trajectory_latents = trajectory_tensor.cpu()
# Update batch with final latents
batch.latents = latents
+238 -8
View File
@@ -1,8 +1,12 @@
# SPDX-License-Identifier: Apache-2.0
"""
Image encoding stages for I2V diffusion pipelines.
Image and video encoding stages for diffusion pipelines.
This module contains implementations of image encoding stages for diffusion pipelines.
This module contains implementations of encoding stages for diffusion pipelines:
- ImageEncodingStage: Encodes images using image encoders (e.g., CLIP)
- RefImageEncodingStage: Encodes reference image for Wan2.1 control pipeline
- ImageVAEEncodingStage: Encodes images to latent space using VAE for I2V generation
- VideoVAEEncodingStage: Encodes videos to latent space using VAE for V2V and control tasks
"""
import PIL
@@ -14,7 +18,9 @@ from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
from fastvideo.models.vaes.common import ParallelTiledVAE
from fastvideo.models.vision_utils import (get_default_height_width, normalize,
numpy_to_pt, pil_to_numpy, resize)
numpy_to_pt, pil_to_numpy, resize,
create_default_image,
preprocess_reference_image_for_clip)
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.base import PipelineStage
from fastvideo.pipelines.stages.validators import StageValidators as V
@@ -94,12 +100,60 @@ class ImageEncodingStage(PipelineStage):
return result
class RefImageEncodingStage(ImageEncodingStage):
"""
Stage for encoding reference image prompts into embeddings for Wan2.1 Control models.
This stage extends ImageEncodingStage with specialized preprocessing for reference images.
"""
@torch.no_grad()
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
"""
Encode the prompt into image encoder hidden states.
Args:
batch: The current batch information.
fastvideo_args: The inference arguments.
Returns:
The batch with encoded prompt embeddings.
"""
self.image_encoder = self.image_encoder.to(get_local_torch_device())
image = batch.pil_image
if image is None:
image = create_default_image()
# Preprocess reference image for CLIP encoder
image_tensor = preprocess_reference_image_for_clip(
image, get_local_torch_device())
image_inputs = self.image_processor(images=image_tensor,
return_tensors="pt").to(
get_local_torch_device())
with set_forward_context(current_timestep=0, attn_metadata=None):
outputs = self.image_encoder(**image_inputs)
image_embeds = outputs.last_hidden_state
batch.image_embeds.append(image_embeds)
if batch.pil_image is None:
batch.image_embeds = [
torch.zeros_like(x) for x in batch.image_embeds
]
return batch
class ImageVAEEncodingStage(PipelineStage):
"""
Stage for encoding pixel representations into latent space.
This stage handles the encoding of pixel representations into the final
input format (e.g., latents).
Stage for encoding image pixel representations into latent space.
This stage handles the encoding of image pixel representations into the final
input format (e.g., latents) for image-to-video generation.
"""
def __init__(self, vae: ParallelTiledVAE) -> None:
@@ -144,9 +198,9 @@ class ImageVAEEncodingStage(PipelineStage):
self.vae = self.vae.to(get_local_torch_device())
# Process single image for I2V
latent_height = height // self.vae.spatial_compression_ratio
latent_width = width // self.vae.spatial_compression_ratio
image = batch.pil_image
image = self.preprocess(
image,
@@ -296,3 +350,179 @@ class ImageVAEEncodingStage(PipelineStage):
result.add_check("image_latent", batch.image_latent,
[V.is_tensor, V.with_dims(5)])
return result
class VideoVAEEncodingStage(ImageVAEEncodingStage):
"""
Stage for encoding video pixel representations into latent space.
This stage handles the encoding of video pixel representations for video-to-video generation and control.
Inherits from ImageVAEEncodingStage to reuse common functionality.
"""
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
"""
Encode video pixel representations into latent space.
Args:
batch: The current batch information.
fastvideo_args: The inference arguments.
Returns:
The batch with encoded outputs.
"""
assert batch.video_latent is not None, "Video latent input is required for VideoVAEEncodingStage"
if fastvideo_args.mode == ExecutionMode.INFERENCE:
assert batch.height is not None and isinstance(batch.height, int)
assert batch.width is not None and isinstance(batch.width, int)
assert batch.num_frames is not None and isinstance(
batch.num_frames, int)
height = batch.height
width = batch.width
num_frames = batch.num_frames
elif fastvideo_args.mode == ExecutionMode.PREPROCESS:
assert batch.height is not None and isinstance(batch.height, list)
assert batch.width is not None and isinstance(batch.width, list)
assert batch.num_frames is not None and isinstance(
batch.num_frames, list)
num_frames = batch.num_frames[0]
height = batch.height[0]
width = batch.width[0]
self.vae = self.vae.to(get_local_torch_device())
# Prepare video tensor from control video
video_condition = self._prepare_control_video_tensor(
batch.video_latent, num_frames, height,
width).to(get_local_torch_device(), dtype=torch.float32)
# Setup VAE precision
vae_dtype = PRECISION_TO_TYPE[
fastvideo_args.pipeline_config.vae_precision]
vae_autocast_enabled = (
vae_dtype != torch.float32) and not fastvideo_args.disable_autocast
# Encode control video
with torch.autocast(device_type="cuda",
dtype=vae_dtype,
enabled=vae_autocast_enabled):
if fastvideo_args.pipeline_config.vae_tiling:
self.vae.enable_tiling()
if not vae_autocast_enabled:
video_condition = video_condition.to(vae_dtype)
encoder_output = self.vae.encode(video_condition)
generator = batch.generator
if generator is None:
raise ValueError("Generator must be provided")
latent_condition = self.retrieve_latents(encoder_output, generator)
if (hasattr(self.vae, "shift_factor")
and self.vae.shift_factor is not None):
if isinstance(self.vae.shift_factor, torch.Tensor):
latent_condition -= self.vae.shift_factor.to(
latent_condition.device, latent_condition.dtype)
else:
latent_condition -= self.vae.shift_factor
if isinstance(self.vae.scaling_factor, torch.Tensor):
latent_condition = latent_condition * self.vae.scaling_factor.to(
latent_condition.device, latent_condition.dtype)
else:
latent_condition = latent_condition * self.vae.scaling_factor
batch.video_latent = latent_condition
# Offload models if needed
if hasattr(self, 'maybe_free_model_hooks'):
self.maybe_free_model_hooks()
self.vae.to("cpu")
return batch
def _prepare_control_video_tensor(self, control_video, num_frames: int,
height: int, width: int) -> torch.Tensor:
"""
Prepare video tensor from control video input.
"""
if isinstance(control_video, list):
processed_frames = []
for i, frame in enumerate(control_video):
if i >= num_frames:
break
processed_frame = self.preprocess(
frame,
vae_scale_factor=self.vae.spatial_compression_ratio,
height=height,
width=width).to(get_local_torch_device(),
dtype=torch.float32)
processed_frames.append(processed_frame)
if processed_frames:
video_tensor = torch.cat(
[f.unsqueeze(2) for f in processed_frames], dim=2)
else:
video_tensor = torch.zeros(1,
3,
0,
height,
width,
device=get_local_torch_device(),
dtype=torch.float32)
elif isinstance(control_video, torch.Tensor):
# Handle tensor input [batch, channels, frames, height, width]
video_tensor = control_video.to(get_local_torch_device(),
dtype=torch.float32)
if video_tensor.shape[2] > num_frames:
video_tensor = video_tensor[:, :, :num_frames]
else:
raise ValueError(
f"Unsupported control_video type: {type(control_video)}. "
"Expected list of PIL Images or torch.Tensor.")
# Pad with zeros if we have fewer frames than required
current_frames = video_tensor.shape[2]
if current_frames < num_frames:
padding_frames = num_frames - current_frames
zero_padding = torch.zeros(video_tensor.shape[0],
video_tensor.shape[1],
padding_frames,
height,
width,
device=video_tensor.device,
dtype=video_tensor.dtype)
video_tensor = torch.cat([video_tensor, zero_padding], dim=2)
return video_tensor
def verify_input(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
"""Verify video encoding stage inputs."""
result = VerificationResult()
result.add_check("video_latent", batch.video_latent, V.not_none)
result.add_check("generator", batch.generator,
V.generator_or_list_generators)
if fastvideo_args.mode == ExecutionMode.PREPROCESS:
result.add_check("height", batch.height, V.list_not_empty)
result.add_check("width", batch.width, V.list_not_empty)
result.add_check("num_frames", batch.num_frames, V.list_not_empty)
else:
result.add_check("height", batch.height, V.positive_int)
result.add_check("width", batch.width, V.positive_int)
result.add_check("num_frames", batch.num_frames, V.positive_int)
return result
def verify_output(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
"""Verify video encoding stage outputs."""
result = VerificationResult()
result.add_check("video_latent", batch.video_latent,
[V.is_tensor, V.with_dims(5)])
return result
+50 -1
View File
@@ -9,7 +9,7 @@ from PIL import Image
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.models.vision_utils import load_image, load_video
from fastvideo.models.vision_utils import load_image, load_video, pil_to_numpy, numpy_to_pt, normalize, resize
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.base import PipelineStage
from fastvideo.pipelines.stages.validators import (StageValidators,
@@ -135,6 +135,55 @@ class InputValidationStage(PipelineStage):
batch.width = ow
batch.pil_image = img
# for v2v, get control video from video path
if batch.video_path is not None:
pil_images, original_fps = load_video(batch.video_path,
return_fps=True)
logger.info("Loaded video with %s frames, original FPS: %s",
len(pil_images), original_fps)
# Get target parameters from batch
target_fps = batch.fps
target_num_frames = batch.num_frames
target_height = batch.height
target_width = batch.width
if target_fps is not None and original_fps is not None:
frame_skip = max(1, int(original_fps // target_fps))
if frame_skip > 1:
pil_images = pil_images[::frame_skip]
effective_fps = original_fps / frame_skip
logger.info(
"Resampled video from %.1f fps to %.1f fps (skip=%s)",
original_fps, effective_fps, frame_skip)
# Limit to target number of frames
if target_num_frames is not None and len(
pil_images) > target_num_frames:
pil_images = pil_images[:target_num_frames]
logger.info("Limited video to %s frames (from %s total)",
target_num_frames, len(pil_images))
# Resize each PIL image to target dimensions
resized_images = []
for pil_img in pil_images:
resized_img = resize(pil_img,
target_height,
target_width,
resize_mode="default",
resample="lanczos")
resized_images.append(resized_img)
# Convert PIL images to numpy array
video_numpy = pil_to_numpy(resized_images)
video_numpy = normalize(video_numpy)
video_tensor = numpy_to_pt(video_numpy)
# Rearrange to [C, T, H, W] and add batch dimension -> [B, C, T, H, W]
input_video = video_tensor.permute(1, 0, 2, 3).unsqueeze(0)
batch.video_latent = input_video
return batch
def verify_input(self, batch: ForwardBatch,

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