Compare commits
25
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
c4a0f789da | ||
|
|
5ac14938d2 | ||
|
|
9d239e9f8b | ||
|
|
bdec816b31 | ||
|
|
2cd2e57d2e | ||
|
|
9370234294 | ||
|
|
50da62e722 | ||
|
|
4f3e8751db | ||
|
|
f4c58894d9 | ||
|
|
01c94ef385 | ||
|
|
2415226d25 | ||
|
|
404314d00f | ||
|
|
87489f0872 | ||
|
|
9ce7c8039e | ||
|
|
e1e25e95f9 | ||
|
|
490bde90e1 | ||
|
|
dc7596b973 | ||
|
|
335afa4457 | ||
|
|
3f77a6805a | ||
|
|
13d0aae706 | ||
|
|
404cbf4f3c | ||
|
|
958ffec844 | ||
|
|
31f000d1cc | ||
|
|
cd32b3e02f | ||
|
|
bf27908095 |
@@ -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"
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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 }}
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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}
|
||||
}
|
||||
|
||||
@@ -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 |
@@ -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 |
@@ -72,6 +72,15 @@ We give a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
|
||||
|
||||
https://github.com/user-attachments/assets/f3b6dd79-7b43-4b60-a0fa-3d6495ec5747
|
||||
|
||||
|
||||
## STA Configuration Logic
|
||||
Here is a diagram of how the window is configured and passed through the FastVideo pipeline:
|
||||
|
||||
<div align="center">
|
||||
<img src="../../../docs/source/_static/images/STA_configuration.png" width="80%"/>
|
||||
</div>
|
||||
|
||||
|
||||
## Why is STA Fast?
|
||||
2D/3D Sliding Window Attention (SWA) creates many mixed blocks in the attention map. Even though mixed blocks have less output value,a mixed block is significantly slower than a dense block due to the GPU-unfriendly masking operation.
|
||||
|
||||
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 98 KiB |
@@ -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.
|
||||
@@ -115,6 +115,7 @@ design/overview
|
||||
|
||||
contributing/overview
|
||||
contributing/developer_env/index
|
||||
contributing/profiling
|
||||
:::
|
||||
|
||||
:::{toctree}
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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,43 @@
|
||||
# NOTE: This is still a work in progress, and the checkpoints are not released yet.
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
# from fastvideo.configs.sample import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples_self_forcing_causal_wan2_2_14B_t2v"
|
||||
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(
|
||||
"rand0nmr/SFWan2.2-T2V-A14B-Diffusers",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
dit_cpu_offload=True, # DiT need to be offloaded for MoE
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
dmd_denoising_steps=[1000, 850, 700, 550, 350, 275, 200, 125],
|
||||
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
|
||||
pin_cpu_memory=True,
|
||||
init_weights_from_safetensors="/mnt/sharefs/users/hao.zhang/wei/SFwan2.2_distill_self_forcing_release_cfg2/checkpoint-246_weight_only/generator_inference_transformer/",
|
||||
init_weights_from_safetensors_2="/mnt/sharefs/users/hao.zhang/wei/SFwan2.2_distill_self_forcing_release_cfg2/checkpoint-246_weight_only/generator_2_inference_transformer/",
|
||||
num_frame_per_block=7,
|
||||
# image_encoder_cpu_offload=False,
|
||||
)
|
||||
|
||||
# sampling_param = SamplingParam.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
|
||||
# sampling_param.num_frames = 45
|
||||
# sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
|
||||
# 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."
|
||||
)
|
||||
_ = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, num_frames=81)
|
||||
|
||||
|
||||
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()
|
||||
@@ -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.
|
||||
@@ -62,7 +62,8 @@ validation_args=(
|
||||
optimizer_args=(
|
||||
--learning_rate 6e-6
|
||||
--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
|
||||
)
|
||||
|
||||
+136
@@ -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[@]}"
|
||||
@@ -1,14 +1,5 @@
|
||||
{
|
||||
"data": [
|
||||
{
|
||||
"caption": "A large metal cylinder is seen pressing down on a pile of Oreo cookies, flattening them as if they were under a hydraulic press.",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"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,
|
||||
@@ -28,7 +19,52 @@
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A watermelon wearing a helmet is crushed by a hydraulic press, causing it to flatten and burst open.",
|
||||
"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,
|
||||
|
||||
@@ -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
|
||||
)
|
||||
|
||||
@@ -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
|
||||
@@ -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(
|
||||
|
||||
@@ -15,13 +15,10 @@ 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
|
||||
|
||||
@@ -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"
|
||||
]
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -13,7 +13,7 @@ from fastvideo.configs.pipelines.wan import (
|
||||
FastWan2_1_T2V_480P_Config, FastWan2_2_TI2V_5B_Config,
|
||||
Wan2_2_I2V_A14B_Config, Wan2_2_T2V_A14B_Config, Wan2_2_TI2V_5B_Config,
|
||||
WanI2V480PConfig, WanI2V720PConfig, WanT2V480PConfig, WanT2V720PConfig,
|
||||
SelfForcingWanT2V480PConfig)
|
||||
SelfForcingWanT2V480PConfig, WANV2VConfig, SelfForcingWan2_2_T2V480PConfig)
|
||||
# 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,
|
||||
@@ -36,6 +37,7 @@ PIPE_NAME_TO_CONFIG: dict[str, type[PipelineConfig]] = {
|
||||
"FastVideo/stepvideo-t2v-diffusers": StepVideoT2VConfig,
|
||||
"FastVideo/Wan2.1-VSA-T2V-14B-720P-Diffusers": WanT2V720PConfig,
|
||||
"wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers": SelfForcingWanT2V480PConfig,
|
||||
"rand0nmr/SFWan2.2-T2V-A14B-Diffusers": SelfForcingWan2_2_T2V480PConfig,
|
||||
"Wan-AI/Wan2.2-TI2V-5B-Diffusers": Wan2_2_TI2V_5B_Config,
|
||||
"Wan-AI/Wan2.2-T2V-A14B-Diffusers": Wan2_2_T2V_A14B_Config,
|
||||
"Wan-AI/Wan2.2-I2V-A14B-Diffusers": Wan2_2_I2V_A14B_Config,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -157,3 +176,13 @@ class SelfForcingWanT2V480PConfig(WanT2V480PConfig):
|
||||
dmd_denoising_steps: list[int] | None = field(
|
||||
default_factory=lambda: [1000, 750, 500, 250])
|
||||
warp_denoising_step: bool = True
|
||||
|
||||
|
||||
@dataclass
|
||||
class SelfForcingWan2_2_T2V480PConfig(Wan2_2_T2V_A14B_Config):
|
||||
is_causal: bool = True
|
||||
flow_shift: float | None = 12.0
|
||||
boundary_ratio: float | None = 0.875
|
||||
dmd_denoising_steps: list[int] | None = field(
|
||||
default_factory=lambda: [1000, 850, 700, 550, 350, 275, 200, 125])
|
||||
warp_denoising_step: bool = True
|
||||
|
||||
@@ -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"
|
||||
@@ -200,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,
|
||||
|
||||
@@ -9,7 +9,7 @@ from fastvideo.configs.sample.stepvideo import StepVideoT2VSamplingParam
|
||||
|
||||
# isort: off
|
||||
from fastvideo.configs.sample.wan import (
|
||||
FastWanT2V480PConfig,
|
||||
FastWanT2V480P_SamplingParam,
|
||||
Wan2_1_Fun_1_3B_InP_SamplingParam,
|
||||
Wan2_2_I2V_A14B_SamplingParam,
|
||||
Wan2_2_T2V_A14B_SamplingParam,
|
||||
@@ -18,7 +18,9 @@ from fastvideo.configs.sample.wan import (
|
||||
WanI2V_14B_720P_SamplingParam,
|
||||
WanT2V_1_3B_SamplingParam,
|
||||
WanT2V_14B_SamplingParam,
|
||||
SelfForcingWanT2V480PConfig,
|
||||
Wan2_1_Fun_1_3B_Control_SamplingParam,
|
||||
SelfForcingWan2_1_T2V_1_3B_480P_SamplingParam,
|
||||
SelfForcingWan2_2_T2V_A14B_480P_SamplingParam,
|
||||
)
|
||||
# isort: on
|
||||
from fastvideo.logger import init_logger
|
||||
@@ -28,33 +30,50 @@ from fastvideo.utils import (maybe_download_model_index,
|
||||
logger = init_logger(__name__)
|
||||
# Registry maps specific model weights to their config classes
|
||||
SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
|
||||
"FastVideo/FastHunyuan-diffusers": FastHunyuanSamplingParam,
|
||||
"hunyuanvideo-community/HunyuanVideo": HunyuanSamplingParam,
|
||||
"FastVideo/stepvideo-t2v-diffusers": StepVideoT2VSamplingParam,
|
||||
"FastVideo/FastHunyuan-diffusers":
|
||||
FastHunyuanSamplingParam,
|
||||
"hunyuanvideo-community/HunyuanVideo":
|
||||
HunyuanSamplingParam,
|
||||
"FastVideo/stepvideo-t2v-diffusers":
|
||||
StepVideoT2VSamplingParam,
|
||||
|
||||
# Wan2.1
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WanT2V_1_3B_SamplingParam,
|
||||
"Wan-AI/Wan2.1-T2V-14B-Diffusers": WanT2V_14B_SamplingParam,
|
||||
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers": WanI2V_14B_480P_SamplingParam,
|
||||
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers": WanI2V_14B_720P_SamplingParam,
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers":
|
||||
WanT2V_1_3B_SamplingParam,
|
||||
"Wan-AI/Wan2.1-T2V-14B-Diffusers":
|
||||
WanT2V_14B_SamplingParam,
|
||||
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers":
|
||||
WanI2V_14B_480P_SamplingParam,
|
||||
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers":
|
||||
WanI2V_14B_720P_SamplingParam,
|
||||
"weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers":
|
||||
Wan2_1_Fun_1_3B_InP_SamplingParam,
|
||||
"IRMChen/Wan2.1-Fun-1.3B-Control-Diffusers":
|
||||
Wan2_1_Fun_1_3B_Control_SamplingParam,
|
||||
|
||||
# Wan2.2
|
||||
"Wan-AI/Wan2.2-TI2V-5B-Diffusers": Wan2_2_TI2V_5B_SamplingParam,
|
||||
"Wan-AI/Wan2.2-TI2V-5B-Diffusers":
|
||||
Wan2_2_TI2V_5B_SamplingParam,
|
||||
"FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers":
|
||||
Wan2_2_TI2V_5B_SamplingParam,
|
||||
"Wan-AI/Wan2.2-T2V-A14B-Diffusers": Wan2_2_T2V_A14B_SamplingParam,
|
||||
"Wan-AI/Wan2.2-I2V-A14B-Diffusers": Wan2_2_I2V_A14B_SamplingParam,
|
||||
"Wan-AI/Wan2.2-T2V-A14B-Diffusers":
|
||||
Wan2_2_T2V_A14B_SamplingParam,
|
||||
"Wan-AI/Wan2.2-I2V-A14B-Diffusers":
|
||||
Wan2_2_I2V_A14B_SamplingParam,
|
||||
|
||||
# FastWan2.1
|
||||
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers": FastWanT2V480PConfig,
|
||||
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers":
|
||||
FastWanT2V480P_SamplingParam,
|
||||
|
||||
# FastWan2.2
|
||||
"FastVideo/FastWan2.2-TI2V-5B-Diffusers": Wan2_2_TI2V_5B_SamplingParam,
|
||||
"FastVideo/FastWan2.2-TI2V-5B-Diffusers":
|
||||
Wan2_2_TI2V_5B_SamplingParam,
|
||||
|
||||
# Causal Self-Forcing Wan2.1
|
||||
"wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers": SelfForcingWanT2V480PConfig,
|
||||
"wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers":
|
||||
SelfForcingWan2_1_T2V_1_3B_480P_SamplingParam,
|
||||
"rand0nmr/SFWan2.2-T2V-A14B-Diffusers":
|
||||
SelfForcingWan2_2_T2V_A14B_480P_SamplingParam,
|
||||
# Add other specific weight variants
|
||||
}
|
||||
|
||||
@@ -64,6 +83,8 @@ SAMPLING_PARAM_DETECTOR: dict[str, Callable[[str], bool]] = {
|
||||
"wanpipeline": lambda id: "wanpipeline" in id.lower(),
|
||||
"wanimagetovideo": lambda id: "wanimagetovideo" in id.lower(),
|
||||
"stepvideo": lambda id: "stepvideo" in id.lower(),
|
||||
"wandmdpipeline": lambda id: "wandmdpipeline" in id.lower(),
|
||||
"wancausaldmdpipeline": lambda id: "wancausaldmdpipeline" in id.lower(),
|
||||
# Add other pipeline architecture detectors
|
||||
}
|
||||
|
||||
@@ -74,7 +95,9 @@ SAMPLING_FALLBACK_PARAM: dict[str, Any] = {
|
||||
"wanpipeline":
|
||||
WanT2V_1_3B_SamplingParam, # Base Wan config as fallback for any Wan variant
|
||||
"wanimagetovideo": WanI2V_14B_480P_SamplingParam,
|
||||
"stepvideo": StepVideoT2VSamplingParam
|
||||
"wandmdpipeline": FastWanT2V480P_SamplingParam,
|
||||
"wancausaldmdpipeline": SelfForcingWan2_1_T2V_1_3B_480P_SamplingParam,
|
||||
"stepvideo": StepVideoT2VSamplingParam,
|
||||
# Other fallbacks by architecture
|
||||
}
|
||||
|
||||
|
||||
@@ -97,7 +97,7 @@ class WanI2V_14B_720P_SamplingParam(WanT2V_14B_SamplingParam):
|
||||
|
||||
|
||||
@dataclass
|
||||
class FastWanT2V480PConfig(WanT2V_1_3B_SamplingParam):
|
||||
class FastWanT2V480P_SamplingParam(WanT2V_1_3B_SamplingParam):
|
||||
# DMD parameters
|
||||
# dmd_denoising_steps: list[int] | None = field(default_factory=lambda: [1000, 757, 522])
|
||||
num_inference_steps: int = 3
|
||||
@@ -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,9 +173,28 @@ 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 =============
|
||||
# =============================================
|
||||
@dataclass
|
||||
class SelfForcingWanT2V480PConfig(WanT2V_1_3B_SamplingParam):
|
||||
class SelfForcingWan2_1_T2V_1_3B_480P_SamplingParam(
|
||||
Wan2_1_Fun_1_3B_InP_SamplingParam):
|
||||
pass
|
||||
|
||||
|
||||
@dataclass
|
||||
class SelfForcingWan2_2_T2V_A14B_480P_SamplingParam(
|
||||
Wan2_2_T2V_A14B_SamplingParam):
|
||||
guidance_scale: float = 2.0
|
||||
guidance_scale_2: float = 2.0
|
||||
num_inference_steps: int = 8
|
||||
num_frames: int = 81
|
||||
height: int = 448
|
||||
width: int = 832
|
||||
fps: int = 16
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
+57
-3
@@ -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
@@ -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
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -20,7 +20,7 @@ import fastvideo.envs as envs
|
||||
from fastvideo.attention import (DistributedAttention,
|
||||
LocalAttention)
|
||||
from fastvideo.configs.models.dits import WanVideoConfig
|
||||
from fastvideo.distributed.parallel_state import get_sp_world_size
|
||||
from fastvideo.distributed.parallel_state import get_sp_world_size, get_local_torch_device
|
||||
from fastvideo.forward_context import get_forward_context
|
||||
from fastvideo.layers.layernorm import (FP32LayerNorm, LayerNormScaleShift,
|
||||
RMSNorm, ScaleResidual,
|
||||
@@ -36,6 +36,26 @@ from fastvideo.models.dits.wanvideo import WanT2VCrossAttention, WanTimeTextImag
|
||||
from fastvideo.platforms import AttentionBackendEnum, current_platform
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
class CacheAppend(torch.autograd.Function):
|
||||
"""
|
||||
KV cache with shape [batch, seq_len, heads, head_dim].
|
||||
"""
|
||||
@staticmethod
|
||||
def forward(ctx, storage, active_cache, x, start, end):
|
||||
# Ensure storage has the same dtype as x
|
||||
storage.data[:, start:end] = x
|
||||
ctx.save_for_backward(storage.to(x.dtype))
|
||||
ctx.start = start
|
||||
ctx.end = end
|
||||
return storage[:, :end].to(x.dtype) # Ensure returned value has same dtype as input
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
start = ctx.start
|
||||
end = ctx.end
|
||||
return None, grad_output[:, :start], grad_output[:, start:end], None, None
|
||||
|
||||
class CausalWanSelfAttention(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
@@ -67,6 +87,10 @@ class CausalWanSelfAttention(nn.Module):
|
||||
causal=False,
|
||||
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
|
||||
AttentionBackendEnum.TORCH_SDPA))
|
||||
|
||||
self.k_cache = None
|
||||
self.v_cache = None
|
||||
self.counter = 0
|
||||
|
||||
def forward(self,
|
||||
q: torch.Tensor,
|
||||
@@ -84,6 +108,19 @@ class CausalWanSelfAttention(nn.Module):
|
||||
grid_sizes(Tensor): Shape [B, 3], the second dimension contains (F, H, W)
|
||||
freqs(Tensor): Rope freqs, shape [1024, C / num_heads / 2]
|
||||
"""
|
||||
if kv_cache is not None:
|
||||
# Then we are in _forward_inference mode
|
||||
if self.k_cache is None:
|
||||
assert self.counter == 0
|
||||
self.counter += 1
|
||||
del self.k_cache
|
||||
self.register_buffer("k_cache", torch.empty(1, self.max_attention_size, self.num_heads, self.head_dim, device=q.device, dtype=q.dtype), persistent=False)
|
||||
self.k_cache.requires_grad_(True)
|
||||
if self.v_cache is None:
|
||||
del self.v_cache
|
||||
self.register_buffer("v_cache", torch.empty(1, self.max_attention_size, self.num_heads, self.head_dim, device=v.device, dtype=v.dtype), persistent=False)
|
||||
self.v_cache.requires_grad_(True)
|
||||
|
||||
if cache_start is None:
|
||||
cache_start = current_start
|
||||
|
||||
@@ -128,6 +165,8 @@ class CausalWanSelfAttention(nn.Module):
|
||||
num_new_tokens = roped_query.shape[1]
|
||||
if self.local_attn_size != -1 and (current_end > kv_cache["global_end_index"].item()) and (
|
||||
num_new_tokens + kv_cache["local_end_index"].item() > kv_cache_size):
|
||||
raise Exception("Not implemented")
|
||||
# @TODO(Wei): This part has not been thoroughly tested yet. Use with caution.
|
||||
# Calculate the number of new tokens added in this step
|
||||
# Shift existing cache content left to discard oldest tokens
|
||||
# Clone the source slice to avoid overlapping memory error
|
||||
@@ -137,23 +176,44 @@ class CausalWanSelfAttention(nn.Module):
|
||||
kv_cache["k"][:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
|
||||
kv_cache["v"][:, sink_tokens:sink_tokens + num_rolled_tokens] = \
|
||||
kv_cache["v"][:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
|
||||
# self.k_cache[:, sink_tokens:sink_tokens + num_rolled_tokens] = \
|
||||
# self.k_cache[:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
|
||||
# self.v_cache[:, sink_tokens:sink_tokens + num_rolled_tokens] = \
|
||||
# self.v_cache[:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
|
||||
# Insert the new keys/values at the end
|
||||
local_end_index = kv_cache["local_end_index"].item() + current_end - \
|
||||
kv_cache["global_end_index"].item() - num_evicted_tokens
|
||||
local_start_index = local_end_index - num_new_tokens
|
||||
# local_k = CacheAppend.apply(self.k_cache, kv_cache["k"], roped_key, local_start_index, local_end_index)
|
||||
# local_v = CacheAppend.apply(self.v_cache, kv_cache["v"], v, local_start_index, local_end_index)
|
||||
kv_cache["k"][:, local_start_index:local_end_index] = roped_key
|
||||
kv_cache["v"][:, local_start_index:local_end_index] = v
|
||||
else:
|
||||
# Assign new keys/values directly up to current_end
|
||||
local_end_index = kv_cache["local_end_index"].item() + current_end - kv_cache["global_end_index"].item()
|
||||
local_start_index = local_end_index - num_new_tokens
|
||||
kv_cache["k"][:, local_start_index:local_end_index] = roped_key
|
||||
kv_cache["v"][:, local_start_index:local_end_index] = v
|
||||
# kv_cache["k"] = kv_cache["k"].detach()
|
||||
# kv_cache["v"] = kv_cache["v"].detach()
|
||||
# kv_cache["k"][:, local_start_index:local_end_index] = roped_key
|
||||
# kv_cache["v"][:, local_start_index:local_end_index] = v
|
||||
local_k = CacheAppend.apply(self.k_cache, kv_cache["k"], roped_key, local_start_index, local_end_index)
|
||||
# logger.info("Is local_k meta tensor: %s, %s", local_k.is_meta, local_k.shape)
|
||||
# logger.info("local_start_index: %d, local_end_index: %d, number of zeros in local k: %d", local_start_index, local_end_index, (local_k == 0).sum().item())
|
||||
local_v = CacheAppend.apply(self.v_cache, kv_cache["v"], v, local_start_index, local_end_index)
|
||||
|
||||
x = self.attn(
|
||||
roped_query,
|
||||
kv_cache["k"][:, max(0, local_end_index - self.max_attention_size):local_end_index],
|
||||
kv_cache["v"][:, max(0, local_end_index - self.max_attention_size):local_end_index]
|
||||
local_k,
|
||||
local_v
|
||||
)
|
||||
# x = self.attn(
|
||||
# roped_query,
|
||||
# kv_cache["k"][:, max(0, local_end_index - self.max_attention_size):local_end_index],
|
||||
# kv_cache["v"][:, max(0, local_end_index - self.max_attention_size):local_end_index]
|
||||
# )
|
||||
|
||||
kv_cache["k"] = local_k
|
||||
kv_cache["v"] = local_v
|
||||
kv_cache["global_end_index"].fill_(current_end)
|
||||
kv_cache["local_end_index"].fill_(local_end_index)
|
||||
|
||||
@@ -176,7 +236,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 +269,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 +282,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")
|
||||
@@ -232,6 +290,9 @@ class CausalWanTransformerBlock(nn.Module):
|
||||
|
||||
self.scale_shift_table = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)
|
||||
|
||||
self.null_shift = torch.tensor([0], device=get_local_torch_device())
|
||||
self.null_scale = torch.tensor([0], device=get_local_torch_device())
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
@@ -249,29 +310,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))
|
||||
@@ -282,11 +343,8 @@ class CausalWanTransformerBlock(nn.Module):
|
||||
attn_output, _ = self.to_out(attn_output)
|
||||
attn_output = attn_output.squeeze(1)
|
||||
|
||||
null_shift = null_scale = torch.tensor([0], device=hidden_states.device)
|
||||
norm_hidden_states, hidden_states = self.self_attn_residual_norm(
|
||||
hidden_states, attn_output, gate_msa, null_shift, null_scale)
|
||||
norm_hidden_states, hidden_states = norm_hidden_states.to(
|
||||
orig_dtype), hidden_states.to(orig_dtype)
|
||||
hidden_states, attn_output, gate_msa, self.null_shift, self.null_scale)
|
||||
|
||||
# 2. Cross-attention
|
||||
attn_output = self.attn2(norm_hidden_states,
|
||||
@@ -295,13 +353,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 +419,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 +429,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__()
|
||||
@@ -487,12 +542,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 +598,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 +641,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 +655,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 +695,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 +708,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
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
@@ -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)
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
@@ -246,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)
|
||||
|
||||
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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 = {}
|
||||
@@ -213,9 +212,18 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
start_index = 0
|
||||
|
||||
# DMD loop in causal blocks
|
||||
# Optional per-block callback for streaming
|
||||
on_block = None
|
||||
try:
|
||||
on_block = getattr(batch, "extra",
|
||||
{}).get("on_block",
|
||||
None) # type: ignore[attr-defined]
|
||||
except Exception:
|
||||
on_block = None
|
||||
|
||||
with self.progress_bar(total=len(block_sizes) *
|
||||
len(timesteps)) as progress_bar:
|
||||
for current_num_frames in block_sizes:
|
||||
for block_idx, current_num_frames in enumerate(block_sizes):
|
||||
current_latents = latents[:, :, start_index:start_index +
|
||||
current_num_frames, :, :]
|
||||
# use BTCHW for DMD conversion routines
|
||||
@@ -223,6 +231,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 +285,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 +350,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,
|
||||
@@ -352,6 +364,20 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
)
|
||||
start_index += current_num_frames
|
||||
|
||||
# Invoke callback with block metadata (no large tensor transfer required)
|
||||
try:
|
||||
if callable(on_block):
|
||||
on_block(
|
||||
block_index=block_idx,
|
||||
total_blocks=len(block_sizes),
|
||||
start_index=start_index - current_num_frames,
|
||||
num_frames=current_num_frames,
|
||||
latents=current_latents,
|
||||
)
|
||||
except Exception as e:
|
||||
# Swallow callback errors so they don't break generation
|
||||
logger.warning("on_block callback failed: %s", str(e))
|
||||
|
||||
batch.latents = latents
|
||||
return batch
|
||||
|
||||
|
||||
@@ -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(
|
||||
@@ -282,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],
|
||||
@@ -568,7 +575,7 @@ class DenoisingStage(PipelineStage):
|
||||
fastvideo_args: The inference arguments.
|
||||
"""
|
||||
# TODO(kevin): STA mask search, currently only support Wan2.1 with 69x768x1280
|
||||
from fastvideo.STA_configuration import configure_sta
|
||||
from fastvideo.attention.backends.STA_configuration import configure_sta
|
||||
STA_mode = fastvideo_args.STA_mode
|
||||
skip_time_steps = fastvideo_args.skip_time_steps
|
||||
if batch.timesteps is None:
|
||||
@@ -675,7 +682,7 @@ class DenoisingStage(PipelineStage):
|
||||
raise NotImplementedError(
|
||||
"STA mask search is not supported for this resolution")
|
||||
|
||||
from fastvideo.STA_configuration import save_mask_search_results
|
||||
from fastvideo.attention.backends.STA_configuration import save_mask_search_results
|
||||
if batch.mask_search_final_result_pos is not None and batch.prompt is not None:
|
||||
save_mask_search_results(
|
||||
[
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -64,6 +64,29 @@ def mps_platform_plugin() -> str | None:
|
||||
return "fastvideo.platforms.mps.MpsPlatform" if is_mps else None
|
||||
|
||||
|
||||
def npu_platform_plugin() -> str | None:
|
||||
is_npu = False
|
||||
|
||||
try:
|
||||
import torch
|
||||
# 导入 torch_npu 以初始化 NPU 后端
|
||||
import torch_npu # noqa: F401
|
||||
if torch.npu.is_available():
|
||||
is_npu = True
|
||||
logger.info("NPU is available")
|
||||
except ImportError:
|
||||
logger.error(
|
||||
"NPU detection failed: PyTorch or PyTorch_NPU is not installed")
|
||||
except AttributeError:
|
||||
logger.error(
|
||||
"NPU detection failed: PyTorch has no 'npu' attribute (use Ascend-adapted PyTorch)"
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error("NPU detection failed: unknown error - %s", str(e))
|
||||
|
||||
return "fastvideo.platforms.npu.NPUPlatform" if is_npu else None
|
||||
|
||||
|
||||
def cpu_platform_plugin() -> str | None:
|
||||
"""Detect if CPU platform should be used."""
|
||||
# CPU is always available as a fallback
|
||||
@@ -93,6 +116,7 @@ builtin_platform_plugins = {
|
||||
'rocm': rocm_platform_plugin,
|
||||
'mps': mps_platform_plugin,
|
||||
'cpu': cpu_platform_plugin,
|
||||
'npu': npu_platform_plugin,
|
||||
}
|
||||
|
||||
|
||||
@@ -115,6 +139,11 @@ def resolve_current_platform_cls_qualname() -> str:
|
||||
if platform_cls_qualname is not None:
|
||||
return platform_cls_qualname
|
||||
|
||||
# Fall back to NPU
|
||||
platform_cls_qualname = npu_platform_plugin()
|
||||
if platform_cls_qualname is not None:
|
||||
return platform_cls_qualname
|
||||
|
||||
# Fall back to CPU as last resort
|
||||
platform_cls_qualname = cpu_platform_plugin()
|
||||
if platform_cls_qualname is not None:
|
||||
|
||||
@@ -67,6 +67,7 @@ class CudaPlatformBase(Platform):
|
||||
device_name: str = "cuda"
|
||||
device_type: str = "cuda"
|
||||
dispatch_key: str = "CUDA"
|
||||
ray_device_key: str = "GPU"
|
||||
device_control_env_var: str = "CUDA_VISIBLE_DEVICES"
|
||||
|
||||
@classmethod
|
||||
@@ -107,6 +108,13 @@ class CudaPlatformBase(Platform):
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
return float(torch.cuda.max_memory_allocated(device))
|
||||
|
||||
@classmethod
|
||||
def get_torch_device(cls):
|
||||
"""
|
||||
Return torch.cuda
|
||||
"""
|
||||
return torch.cuda
|
||||
|
||||
@classmethod
|
||||
def get_attn_backend_cls(cls, selected_backend: AttentionBackendEnum | None,
|
||||
head_size: int, dtype: torch.dtype) -> str:
|
||||
@@ -115,6 +123,7 @@ class CudaPlatformBase(Platform):
|
||||
|
||||
logger.info("Trying FASTVIDEO_ATTENTION_BACKEND=%s",
|
||||
envs.FASTVIDEO_ATTENTION_BACKEND)
|
||||
logger.info("Selected backend: %s", selected_backend)
|
||||
if selected_backend == AttentionBackendEnum.SLIDING_TILE_ATTN:
|
||||
try:
|
||||
from st_attn import sliding_tile_attention # noqa: F401
|
||||
@@ -144,6 +153,20 @@ class CudaPlatformBase(Platform):
|
||||
logger.info(
|
||||
"Sage Attention backend is not installed. Fall back to Flash Attention."
|
||||
)
|
||||
elif selected_backend == AttentionBackendEnum.SAGE_ATTN_THREE:
|
||||
try:
|
||||
from fastvideo.attention.backends.sageattn.api import sageattn_blackwell # noqa: F401
|
||||
|
||||
from fastvideo.attention.backends.sage_attn3 import ( # noqa: F401
|
||||
SageAttention3Backend)
|
||||
logger.info("Using Sage Attention 3 backend.")
|
||||
|
||||
return "fastvideo.attention.backends.sage_attn3.SageAttention3Backend"
|
||||
except ImportError as e:
|
||||
logger.info(e)
|
||||
logger.info(
|
||||
"Sage Attention 3 backend is not installed. Fall back to Flash Attention."
|
||||
)
|
||||
elif selected_backend == AttentionBackendEnum.VIDEO_SPARSE_ATTN:
|
||||
try:
|
||||
from vsa import block_sparse_attn # noqa: F401
|
||||
@@ -158,7 +181,11 @@ class CudaPlatformBase(Platform):
|
||||
"Failed to import Video Sparse Attention backend: %s",
|
||||
str(e))
|
||||
raise ImportError(
|
||||
"Video Sparse Attention backend is not installed. ") from e
|
||||
"The Video Sparse Attention backend is not installed. "
|
||||
"To install it, please follow the instructions at: "
|
||||
"https://hao-ai-lab.github.io/FastVideo/video_sparse_attention/installation.html "
|
||||
) from e
|
||||
|
||||
elif selected_backend == AttentionBackendEnum.VMOBA_ATTN:
|
||||
try:
|
||||
from csrc.attn.vmoba_attn.vmoba import ( # noqa: F401
|
||||
|
||||
@@ -1,6 +1,3 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/platforms/interface.py
|
||||
|
||||
import enum
|
||||
import random
|
||||
from typing import NamedTuple
|
||||
@@ -18,6 +15,7 @@ class AttentionBackendEnum(enum.Enum):
|
||||
SLIDING_TILE_ATTN = enum.auto()
|
||||
TORCH_SDPA = enum.auto()
|
||||
SAGE_ATTN = enum.auto()
|
||||
SAGE_ATTN_THREE = enum.auto()
|
||||
VIDEO_SPARSE_ATTN = enum.auto()
|
||||
VMOBA_ATTN = enum.auto()
|
||||
NO_ATTENTION = enum.auto()
|
||||
@@ -27,10 +25,12 @@ class PlatformEnum(enum.Enum):
|
||||
CUDA = enum.auto()
|
||||
ROCM = enum.auto()
|
||||
TPU = enum.auto()
|
||||
XPU = enum.auto()
|
||||
CPU = enum.auto()
|
||||
MPS = enum.auto()
|
||||
OOT = enum.auto()
|
||||
UNSPECIFIED = enum.auto()
|
||||
NPU = enum.auto()
|
||||
|
||||
|
||||
class CpuArchEnum(enum.Enum):
|
||||
@@ -61,11 +61,18 @@ class Platform:
|
||||
device_name: str
|
||||
device_type: str
|
||||
|
||||
# available dispatch keys:
|
||||
# check https://github.com/pytorch/pytorch/blob/313dac6c1ca0fa0cde32477509cce32089f8532a/torchgen/model.py#L134 # noqa
|
||||
# use "CPU" as a fallback for platforms not registered in PyTorch
|
||||
dispatch_key: str = "CPU"
|
||||
|
||||
# platform-agnostic way to specify the device control environment variable,
|
||||
# .e.g. CUDA_VISIBLE_DEVICES for CUDA.
|
||||
# hint: search for "get_visible_accelerator_ids_env_var" in
|
||||
# https://github.com/ray-project/ray/tree/master/python/ray/_private/accelerators # noqa
|
||||
device_control_env_var: str = "FASTVIDEO_DEVICE_CONTROL_ENV_VAR_PLACEHOLDER"
|
||||
|
||||
# available ray device keys:
|
||||
# https://github.com/ray-project/ray/blob/10ba5adadcc49c60af2c358a33bb943fb491a171/python/ray/_private/ray_constants.py#L438 # noqa
|
||||
# empty string means the device does not support ray
|
||||
ray_device_key: str = ""
|
||||
# The torch.compile backend for compiling simple and
|
||||
# standalone functions. The default value is "inductor" to keep
|
||||
# the same behavior as PyTorch.
|
||||
@@ -75,6 +82,8 @@ class Platform:
|
||||
|
||||
supported_quantization: list[str] = []
|
||||
|
||||
additional_env_vars: list[str] = []
|
||||
|
||||
def is_cuda(self) -> bool:
|
||||
return self._enum == PlatformEnum.CUDA
|
||||
|
||||
@@ -84,6 +93,9 @@ class Platform:
|
||||
def is_tpu(self) -> bool:
|
||||
return self._enum == PlatformEnum.TPU
|
||||
|
||||
def is_xpu(self) -> bool:
|
||||
return self._enum == PlatformEnum.XPU
|
||||
|
||||
def is_cpu(self) -> bool:
|
||||
return self._enum == PlatformEnum.CPU
|
||||
|
||||
@@ -97,6 +109,9 @@ class Platform:
|
||||
def is_mps(self) -> bool:
|
||||
return self._enum == PlatformEnum.MPS
|
||||
|
||||
def is_npu(self) -> bool:
|
||||
return self._enum == PlatformEnum.NPU
|
||||
|
||||
@classmethod
|
||||
def get_attn_backend_cls(cls, selected_backend: AttentionBackendEnum | None,
|
||||
head_size: int, dtype: torch.dtype) -> str:
|
||||
@@ -156,6 +171,13 @@ class Platform:
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
@classmethod
|
||||
def get_torch_device(cls):
|
||||
"""
|
||||
Check if the current platform supports torch device.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
@classmethod
|
||||
def inference_mode(cls):
|
||||
"""A device-specific wrapper of `torch.inference_mode`.
|
||||
|
||||
@@ -0,0 +1,131 @@
|
||||
import gc
|
||||
from datetime import timedelta
|
||||
|
||||
import torch
|
||||
from torch.distributed import ProcessGroup
|
||||
from torch.distributed.distributed_c10d import PrefixStore
|
||||
|
||||
import fastvideo.envs as envs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.platforms.interface import (AttentionBackendEnum, Platform,
|
||||
PlatformEnum)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class NPUPlatform(Platform):
|
||||
|
||||
_enum = PlatformEnum.NPU
|
||||
device_name: str = "npu"
|
||||
device_type: str = "npu"
|
||||
simple_compile_backend: str = "eager" # Disable torch.compile()
|
||||
ray_device_key: str = "NPU"
|
||||
device_control_env_var: str = "ASCEND_RT_VISIBLE_DEVICES"
|
||||
dispatch_key: str = "PrivateUse1"
|
||||
|
||||
def is_sleep_mode_available(self) -> bool:
|
||||
return True
|
||||
|
||||
@classmethod
|
||||
def get_device_capability(cls, device_id: int = 0):
|
||||
return None
|
||||
|
||||
@classmethod
|
||||
def get_device_name(cls, device_id: int = 0) -> str:
|
||||
return torch.npu.get_device_name(device_id)
|
||||
|
||||
@classmethod
|
||||
def is_async_output_supported(cls, enforce_eager: bool | None) -> bool:
|
||||
return True
|
||||
|
||||
@classmethod
|
||||
def inference_mode(cls):
|
||||
return torch.inference_mode()
|
||||
|
||||
@classmethod
|
||||
def set_device(cls, device: torch.device):
|
||||
torch.npu.set_device(device)
|
||||
|
||||
@classmethod
|
||||
def empty_cache(cls):
|
||||
torch.npu.empty_cache()
|
||||
|
||||
@classmethod
|
||||
def synchronize(cls):
|
||||
torch.npu.synchronize()
|
||||
|
||||
@classmethod
|
||||
def mem_get_info(cls) -> tuple[int, int]:
|
||||
return torch.npu.mem_get_info()
|
||||
|
||||
@classmethod
|
||||
def clear_npu_memory(cls):
|
||||
gc.collect()
|
||||
torch.npu.empty_cache()
|
||||
torch.npu.reset_peak_memory_stats()
|
||||
|
||||
@classmethod
|
||||
def get_attn_backend_cls(cls, selected_backend: AttentionBackendEnum | None,
|
||||
head_size: int, dtype: torch.dtype) -> str:
|
||||
logger.info("Trying FASTVIDEO_ATTENTION_BACKEND=%s",
|
||||
envs.FASTVIDEO_ATTENTION_BACKEND)
|
||||
if envs.FASTVIDEO_ATTENTION_BACKEND != "TORCH_SDPA":
|
||||
logger.info("Ascend NPU only supports the Torch SDPA backend.")
|
||||
else:
|
||||
logger.info("Using Torch SDPA backend.")
|
||||
return "fastvideo.attention.backends.sdpa.SDPABackend"
|
||||
|
||||
@classmethod
|
||||
def get_current_memory_usage(cls,
|
||||
device: torch.types.Device | None = None
|
||||
) -> float:
|
||||
torch.npu.reset_peak_memory_stats(device)
|
||||
return torch.npu.max_memory_allocated(device)
|
||||
|
||||
@classmethod
|
||||
def get_device_communicator_cls(cls) -> str:
|
||||
return "fastvideo.distributed.device_communicators.npu_communicator.NpuCommunicator"
|
||||
|
||||
@classmethod
|
||||
def is_pin_memory_available(cls):
|
||||
return True
|
||||
|
||||
@classmethod
|
||||
def get_torch_device(cls):
|
||||
"""
|
||||
Return torch.npu
|
||||
"""
|
||||
return torch.npu
|
||||
|
||||
@classmethod
|
||||
def stateless_init_device_torch_dist_pg(
|
||||
cls,
|
||||
backend: str,
|
||||
prefix_store: PrefixStore,
|
||||
group_rank: int,
|
||||
group_size: int,
|
||||
timeout: timedelta,
|
||||
) -> ProcessGroup:
|
||||
from torch.distributed import is_hccl_available
|
||||
from torch_npu._C._distributed_c10d import ProcessGroupHCCL
|
||||
|
||||
assert is_hccl_available()
|
||||
options = ProcessGroup.Options(backend=backend)
|
||||
pg: ProcessGroup = ProcessGroup(
|
||||
prefix_store,
|
||||
group_rank,
|
||||
group_size,
|
||||
options,
|
||||
)
|
||||
|
||||
backend_options = ProcessGroupHCCL.Options()
|
||||
backend_options._timeout = timeout
|
||||
|
||||
backend_class = ProcessGroupHCCL(prefix_store, group_rank, group_size,
|
||||
backend_options)
|
||||
device = torch.device("npu")
|
||||
backend_class._set_sequence_number_for_group()
|
||||
backend_type = ProcessGroup.BackendType.CUSTOM
|
||||
|
||||
pg._register_backend(device, backend_type, backend_class)
|
||||
return pg
|
||||
@@ -22,6 +22,7 @@ class RocmPlatform(Platform):
|
||||
device_name: str = "rocm"
|
||||
device_type: str = "cuda" # torch uses 'cuda' backend string
|
||||
dispatch_key: str = "CUDA"
|
||||
ray_device_key: str = "GPU"
|
||||
device_control_env_var: str = "CUDA_VISIBLE_DEVICES"
|
||||
|
||||
@classmethod
|
||||
|
||||
@@ -0,0 +1,451 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Utilities for managing the PyTorch profiler within FastVideo.
|
||||
|
||||
The profiler is shared across the process; this module adds a light-weight
|
||||
controller that gates collection based on named *regions*. Regions may be
|
||||
enabled through dedicated environment variables (e.g.
|
||||
``FASTVIDEO_TORCH_PROFILE_MODEL_LOADING=1``) or via the consolidated
|
||||
``FASTVIDEO_TORCH_PROFILE_REGIONS`` comma-separated list (e.g.
|
||||
``FASTVIDEO_TORCH_PROFILE_REGIONS=model_loading,training_dit``).
|
||||
|
||||
Typical usage from client code::
|
||||
|
||||
controller = TorchProfilerController(profiler, activities)
|
||||
with controller.region("training_dit"):
|
||||
run_training_step()
|
||||
|
||||
To introduce a new region, register it via :func:`register_profiler_region`
|
||||
and wrap the corresponding code in :meth:`TorchProfilerController.region`.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
from collections.abc import Callable
|
||||
import functools
|
||||
from collections.abc import Iterable
|
||||
|
||||
import torch
|
||||
|
||||
import fastvideo.envs as envs
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
_GLOBAL_PROFILER: torch.profiler.profile | None = None
|
||||
_GLOBAL_CONTROLLER: TorchProfilerController | None = None
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ProfilerRegion:
|
||||
"""Metadata describing a profiler region."""
|
||||
|
||||
name: str
|
||||
description: str
|
||||
default_enabled: bool = False
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if not self.name or self.name.strip() != self.name:
|
||||
raise ValueError(
|
||||
f"Profiler region name must be non-empty without surrounding whitespace: {self.name!r}"
|
||||
)
|
||||
if not self.name.islower():
|
||||
raise ValueError(
|
||||
f"Profiler region name must be lower-case: {self.name!r}")
|
||||
|
||||
|
||||
_REGISTERED_REGIONS: dict[str, ProfilerRegion] = {}
|
||||
|
||||
|
||||
def _normalize_token(token: str) -> str:
|
||||
return token.strip().lower()
|
||||
|
||||
|
||||
def register_profiler_region(
|
||||
name: str,
|
||||
description: str,
|
||||
*,
|
||||
default_enabled: bool = False,
|
||||
) -> None:
|
||||
"""Register a profiler region so configuration can validate inputs."""
|
||||
|
||||
canonical = _normalize_token(name)
|
||||
if canonical in _REGISTERED_REGIONS:
|
||||
raise ValueError(f"Profiler region {name!r} is already registered")
|
||||
|
||||
region = ProfilerRegion(
|
||||
name=canonical,
|
||||
description=description,
|
||||
default_enabled=bool(default_enabled),
|
||||
)
|
||||
_REGISTERED_REGIONS[canonical] = region
|
||||
|
||||
|
||||
def resolve_profiler_region(name: str) -> ProfilerRegion | None:
|
||||
"""Return the registered region matching ``name`` or ``None`` if absent."""
|
||||
|
||||
canonical = _normalize_token(name)
|
||||
return _REGISTERED_REGIONS.get(canonical)
|
||||
|
||||
|
||||
def list_profiler_regions() -> list[ProfilerRegion]:
|
||||
"""Return all registered profiler regions sorted by canonical name."""
|
||||
|
||||
return [_REGISTERED_REGIONS[name] for name in sorted(_REGISTERED_REGIONS)]
|
||||
|
||||
|
||||
_DEFAULT_ACTIVITIES: tuple[torch.profiler.ProfilerActivity, ...] = (
|
||||
torch.profiler.ProfilerActivity.CPU,
|
||||
torch.profiler.ProfilerActivity.CUDA,
|
||||
)
|
||||
|
||||
|
||||
def get_global_profiler() -> torch.profiler.profile | None:
|
||||
"""Return the global profiler instance if one was created."""
|
||||
|
||||
return _GLOBAL_PROFILER
|
||||
|
||||
|
||||
def set_global_profiler(profiler: torch.profiler.profile | None) -> None:
|
||||
global _GLOBAL_PROFILER
|
||||
_GLOBAL_PROFILER = profiler
|
||||
|
||||
|
||||
def get_global_controller() -> TorchProfilerController | None:
|
||||
return _GLOBAL_CONTROLLER
|
||||
|
||||
|
||||
def set_global_controller(controller: TorchProfilerController | None) -> None:
|
||||
global _GLOBAL_CONTROLLER
|
||||
_GLOBAL_CONTROLLER = controller
|
||||
|
||||
|
||||
register_profiler_region(
|
||||
name="profiler_region_model_loading",
|
||||
description="Module/model loading during pipeline initialization.",
|
||||
default_enabled=False,
|
||||
)
|
||||
# register_profiler_region(
|
||||
# name="profiler_region_inference_pre_denoising",
|
||||
# description="Pre-denoising inference steps (conditioning, preprocessing).",
|
||||
# )
|
||||
# register_profiler_region(
|
||||
# name="profiler_region_inference_denoising",
|
||||
# description="The main inference denoising loop.",
|
||||
# )
|
||||
# register_profiler_region(
|
||||
# name="profiler_region_inference_post_denoising",
|
||||
# description=
|
||||
# "Post-processing after denoising (decoder, conditioning restores).",
|
||||
# )
|
||||
register_profiler_region(
|
||||
name="profiler_region_training_save_checkpoint",
|
||||
description="Training save checkpoint operations.",
|
||||
)
|
||||
|
||||
# general training related regions
|
||||
register_profiler_region(
|
||||
name="profiler_region_training_validation",
|
||||
description="Validation loop during training.",
|
||||
)
|
||||
register_profiler_region(
|
||||
name="profiler_region_training_train_one_step",
|
||||
description="High-level step orchestration in the training loop.",
|
||||
)
|
||||
register_profiler_region(
|
||||
name="profiler_region_training_train",
|
||||
description="Single optimizer step including forward/backward passes.",
|
||||
)
|
||||
|
||||
# distillation specific regions
|
||||
register_profiler_region(
|
||||
name="profiler_region_distillation_teacher_forward",
|
||||
description="Teacher model forward pass in distillation pipelines.",
|
||||
)
|
||||
register_profiler_region(
|
||||
name="profiler_region_distillation_student_forward",
|
||||
description="Student model forward pass in distillation pipelines.",
|
||||
)
|
||||
register_profiler_region(
|
||||
name="profiler_region_distillation_loss",
|
||||
description="Distillation loss computation and aggregation.",
|
||||
)
|
||||
register_profiler_region(
|
||||
name="profiler_region_distillation_update",
|
||||
description="Parameter updates specific to distillation workflows.",
|
||||
)
|
||||
|
||||
|
||||
def get_or_create_profiler(trace_dir: str | None) -> TorchProfilerController:
|
||||
"""Create or reuse the process-wide torch profiler controller."""
|
||||
|
||||
existing = get_global_controller()
|
||||
if existing is not None:
|
||||
if trace_dir:
|
||||
logger.info("Reusing existing global torch profiler controller")
|
||||
return existing
|
||||
|
||||
if not trace_dir:
|
||||
logger.info("Torch profiler disabled; returning no-op controller")
|
||||
return TorchProfilerController(None, _DEFAULT_ACTIVITIES, disabled=True)
|
||||
|
||||
logger.info("Profiling enabled. Traces will be saved to: %s", trace_dir)
|
||||
logger.info(
|
||||
"Profiler config: record_shapes=%s, profile_memory=%s, with_stack=%s, with_flops=%s",
|
||||
envs.FASTVIDEO_TORCH_PROFILER_RECORD_SHAPES,
|
||||
envs.FASTVIDEO_TORCH_PROFILER_WITH_PROFILE_MEMORY,
|
||||
envs.FASTVIDEO_TORCH_PROFILER_WITH_STACK,
|
||||
envs.FASTVIDEO_TORCH_PROFILER_WITH_FLOPS,
|
||||
)
|
||||
logger.info("FASTVIDEO_TORCH_PROFILE_REGIONS=%s",
|
||||
envs.FASTVIDEO_TORCH_PROFILE_REGIONS)
|
||||
|
||||
profiler = torch.profiler.profile(
|
||||
activities=_DEFAULT_ACTIVITIES,
|
||||
record_shapes=envs.FASTVIDEO_TORCH_PROFILER_RECORD_SHAPES,
|
||||
profile_memory=envs.FASTVIDEO_TORCH_PROFILER_WITH_PROFILE_MEMORY,
|
||||
with_stack=envs.FASTVIDEO_TORCH_PROFILER_WITH_STACK,
|
||||
with_flops=envs.FASTVIDEO_TORCH_PROFILER_WITH_FLOPS,
|
||||
on_trace_ready=torch.profiler.tensorboard_trace_handler(trace_dir,
|
||||
use_gzip=True),
|
||||
)
|
||||
controller = TorchProfilerController(profiler, _DEFAULT_ACTIVITIES)
|
||||
controller.start()
|
||||
logger.info("Torch profiler started")
|
||||
return controller
|
||||
|
||||
|
||||
@dataclass
|
||||
class TorchProfilerConfig:
|
||||
"""Configuration for torch profiler region control.
|
||||
|
||||
Use :meth:`from_env` to construct an instance with defaults inherited from
|
||||
registered regions and optional overrides from the
|
||||
``FASTVIDEO_TORCH_PROFILE_REGIONS`` environment variable. The resulting
|
||||
``regions`` map is consumed by :class:`TorchProfilerController` to decide
|
||||
when collection should be enabled.
|
||||
"""
|
||||
|
||||
regions: dict[str, bool]
|
||||
|
||||
@classmethod
|
||||
def from_env(cls) -> TorchProfilerConfig:
|
||||
"""Build a configuration from process environment variables."""
|
||||
|
||||
requested_regions = {
|
||||
token.strip()
|
||||
for token in (getattr(envs, "FASTVIDEO_TORCH_PROFILE_REGIONS", "")
|
||||
or "").split(",") if token.strip()
|
||||
}
|
||||
|
||||
if not requested_regions:
|
||||
available = ", ".join(region.name
|
||||
for region in list_profiler_regions())
|
||||
raise ValueError(
|
||||
"FASTVIDEO_TORCH_PROFILE_REGIONS must list at least one region; "
|
||||
f"available regions: {available}")
|
||||
|
||||
regions: dict[str, bool] = {}
|
||||
available_regions = list_profiler_regions()
|
||||
available_names = ", ".join(region.name for region in available_regions)
|
||||
|
||||
for token in requested_regions:
|
||||
resolved = resolve_profiler_region(token)
|
||||
if resolved is None:
|
||||
logger.warning(
|
||||
"Unknown profiler region '%s'; available regions: %s",
|
||||
token, available_names)
|
||||
continue
|
||||
regions[resolved.name] = True
|
||||
|
||||
if not regions:
|
||||
raise ValueError(
|
||||
"FASTVIDEO_TORCH_PROFILE_REGIONS did not match any known regions; "
|
||||
f"requested={sorted(requested_regions)}, available={available_names}"
|
||||
)
|
||||
|
||||
return cls(regions=regions)
|
||||
|
||||
def __str__(self) -> str:
|
||||
return f"TorchProfilerConfig(regions={self.regions})"
|
||||
|
||||
|
||||
class TorchProfilerController:
|
||||
"""Helper that toggles torch profiler collection for named regions.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
profiler:
|
||||
The shared :class:`torch.profiler.profile` instance, or ``None`` if
|
||||
profiling is disabled.
|
||||
activities:
|
||||
Iterable of :class:`torch.profiler.ProfilerActivity` recorded by the
|
||||
profiler.
|
||||
config:
|
||||
Optional :class:`TorchProfilerConfig`. If omitted, :meth:`from_env`
|
||||
constructs one during initialization.
|
||||
|
||||
Examples
|
||||
--------
|
||||
Enabling an existing region from the command line::
|
||||
|
||||
FASTVIDEO_TORCH_PROFILE_REGIONS=model_loading,training_dit \
|
||||
python fastvideo/training/wan_training_pipeline.py ...
|
||||
|
||||
Wrapping a code block in a custom region::
|
||||
|
||||
controller = TorchProfilerController(profiler, activities)
|
||||
with controller.region("training_validation"):
|
||||
run_validation_epoch()
|
||||
|
||||
Adding a new region requires three steps:
|
||||
1. Define an env var in ``envs.py``.
|
||||
2. Add a default entry to ``register_profiler_region`` in this module.
|
||||
3. Wrap the target code in :meth:`region` using the new name.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
profiler: Any,
|
||||
activities: Iterable[torch.profiler.ProfilerActivity],
|
||||
config: TorchProfilerConfig | None = None,
|
||||
disabled: bool = False,
|
||||
) -> None:
|
||||
activities_tuple = tuple(activities)
|
||||
existing = get_global_controller()
|
||||
if existing is not None and not disabled:
|
||||
raise RuntimeError(
|
||||
"TorchProfilerController already initialized globally. Use get_global_controller()."
|
||||
)
|
||||
if disabled:
|
||||
self._profiler = None
|
||||
return
|
||||
|
||||
self._profiler = profiler
|
||||
self._activities = activities_tuple
|
||||
self._config = config or TorchProfilerConfig.from_env()
|
||||
self._collection_enabled = False
|
||||
self._active_region_depth = 0
|
||||
logger.info(
|
||||
"PROFILER: TorchProfilerController initialized with config: %s",
|
||||
self._config)
|
||||
set_global_profiler(self._profiler)
|
||||
set_global_controller(self)
|
||||
|
||||
@property
|
||||
def is_enabled(self) -> bool:
|
||||
"""Return ``True`` when the underlying profiler is collecting."""
|
||||
|
||||
if self._profiler is None:
|
||||
return False
|
||||
return self._collection_enabled
|
||||
|
||||
def is_region_enabled(self, region: str) -> bool:
|
||||
"""Return ``True`` if ``region`` should be collected."""
|
||||
|
||||
if self._profiler is None:
|
||||
return False
|
||||
return self._config.regions.get(region, False)
|
||||
|
||||
def _set_collection(self, enabled: bool) -> None:
|
||||
if self._profiler is None:
|
||||
return
|
||||
if self._collection_enabled == enabled:
|
||||
return
|
||||
event = ("fastvideo.profiler.enable_collection"
|
||||
if enabled else "fastvideo.profiler.disable_collection")
|
||||
with torch.profiler.record_function(event):
|
||||
self._profiler.toggle_collection_dynamic(enabled, self._activities)
|
||||
self._collection_enabled = enabled
|
||||
|
||||
@contextlib.contextmanager
|
||||
def region(self, region: str):
|
||||
"""Context manager that enables profiling for ``region`` if configured."""
|
||||
|
||||
if self._profiler is None:
|
||||
yield
|
||||
return
|
||||
|
||||
if not self.is_region_enabled(region):
|
||||
yield
|
||||
return
|
||||
|
||||
with torch.profiler.record_function(f"fastvideo.region::{region}"):
|
||||
self._active_region_depth += 1
|
||||
if self._active_region_depth == 1:
|
||||
logger.info(
|
||||
"PROFILER: Setting collection to True (depth=%s) for region %s",
|
||||
self._active_region_depth, region)
|
||||
self._set_collection(True)
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
self._active_region_depth -= 1
|
||||
logger.info("PROFILER: Decreasing active region depth to %s",
|
||||
self._active_region_depth)
|
||||
if self._active_region_depth == 0:
|
||||
logger.info(
|
||||
"PROFILER: Setting collection to False upon exiting region %s",
|
||||
region)
|
||||
self._set_collection(False)
|
||||
|
||||
def start(self) -> None:
|
||||
"""Start the profiler and pause collection until a region is entered."""
|
||||
|
||||
logger.info("PROFILER: Starting profiler...")
|
||||
if self._profiler is None:
|
||||
return
|
||||
self._profiler.start()
|
||||
logger.info("PROFILER: Profiler started")
|
||||
# Profiler starts with collection disabled by default.
|
||||
logger.info("PROFILER: Setting collection to False")
|
||||
self._set_collection(False)
|
||||
logger.info("PROFILER: Profiler started with collection disabled")
|
||||
|
||||
def stop(self) -> None:
|
||||
"""Stop the profiler after disabling collection and clearing state."""
|
||||
|
||||
if self._profiler is None:
|
||||
return
|
||||
|
||||
logger.info("PROFILER: Stopping profiler...")
|
||||
self._profiler.stop()
|
||||
logger.info("PROFILER: Profiler stopped")
|
||||
self._active_region_depth = 0
|
||||
set_global_profiler(None)
|
||||
set_global_controller(None)
|
||||
|
||||
@property
|
||||
def has_profiler(self) -> bool:
|
||||
"""Return ``True`` when a profiler instance is available."""
|
||||
|
||||
return self._profiler is not None
|
||||
|
||||
@property
|
||||
def activities(self) -> tuple[torch.profiler.ProfilerActivity, ...]:
|
||||
return tuple(self._activities)
|
||||
|
||||
@property
|
||||
def profiler(self) -> torch.profiler.profile | None:
|
||||
return self._profiler
|
||||
|
||||
|
||||
def profile_region(
|
||||
region: str) -> Callable[[Callable[..., Any]], Callable[..., Any]]:
|
||||
"""Wrap a bound method so it runs inside a profiler region if available."""
|
||||
|
||||
def decorator(fn: Callable[..., Any]) -> Callable[..., Any]:
|
||||
|
||||
@functools.wraps(fn)
|
||||
def wrapped(self, *args, **kwargs):
|
||||
controller = getattr(self, "profiler_controller", None)
|
||||
if controller is None or not controller.has_profiler:
|
||||
return fn(self, *args, **kwargs)
|
||||
with controller.region(region):
|
||||
return fn(self, *args, **kwargs)
|
||||
|
||||
return wrapped
|
||||
|
||||
return decorator
|
||||
@@ -0,0 +1,78 @@
|
||||
import os
|
||||
|
||||
from fastvideo.entrypoints.video_generator import VideoGenerator
|
||||
|
||||
|
||||
def _new_video_generator() -> VideoGenerator:
|
||||
# Bypass __init__ since we only test a pure helper method.
|
||||
return VideoGenerator.__new__(VideoGenerator)
|
||||
|
||||
|
||||
def test_prepare_output_path_file_sanitization(tmp_path):
|
||||
vg = _new_video_generator()
|
||||
target_dir = tmp_path / "dir"
|
||||
raw_path = target_dir / "inv:al*id?.mp4"
|
||||
|
||||
result = vg._prepare_output_path(str(raw_path), prompt="ignored")
|
||||
|
||||
assert os.path.dirname(result) == str(target_dir)
|
||||
assert os.path.basename(result) == "invalid.mp4"
|
||||
assert os.path.isdir(target_dir)
|
||||
|
||||
|
||||
def test_prepare_output_path_directory_prompt_derived(tmp_path):
|
||||
vg = _new_video_generator()
|
||||
out_dir = tmp_path / "outputs"
|
||||
prompt = "Hello:/\\*?<>| world"
|
||||
|
||||
result = vg._prepare_output_path(str(out_dir), prompt=prompt)
|
||||
|
||||
assert os.path.dirname(result) == str(out_dir)
|
||||
# spaces are preserved (collapsed) by sanitizer; here it becomes "Hello world.mp4"
|
||||
assert os.path.basename(result) == "Hello world.mp4"
|
||||
assert os.path.isdir(out_dir)
|
||||
|
||||
|
||||
def test_prepare_output_path_non_mp4_treated_as_dir(tmp_path):
|
||||
vg = _new_video_generator()
|
||||
weird_dir = tmp_path / "foo.gif"
|
||||
prompt = "My Video"
|
||||
|
||||
result = vg._prepare_output_path(str(weird_dir), prompt=prompt)
|
||||
|
||||
assert os.path.dirname(result) == str(weird_dir)
|
||||
assert os.path.basename(result) == "My Video.mp4"
|
||||
assert os.path.isdir(weird_dir)
|
||||
|
||||
|
||||
def test_prepare_output_path_uniqueness_suffix(tmp_path):
|
||||
vg = _new_video_generator()
|
||||
out_dir = tmp_path / "outputs"
|
||||
prompt = "Sample Name"
|
||||
|
||||
first = vg._prepare_output_path(str(out_dir), prompt=prompt)
|
||||
# simulate existing file
|
||||
os.makedirs(os.path.dirname(first), exist_ok=True)
|
||||
with open(first, "wb") as f:
|
||||
f.write(b"")
|
||||
|
||||
second = vg._prepare_output_path(str(out_dir), prompt=prompt)
|
||||
assert os.path.basename(second) == "Sample Name_1.mp4"
|
||||
|
||||
# simulate second existing file as well
|
||||
with open(second, "wb") as f:
|
||||
f.write(b"")
|
||||
third = vg._prepare_output_path(str(out_dir), prompt=prompt)
|
||||
assert os.path.basename(third) == "Sample Name_2.mp4"
|
||||
|
||||
|
||||
def test_prepare_output_path_empty_prompt_fallback(tmp_path):
|
||||
vg = _new_video_generator()
|
||||
out_dir = tmp_path / "outputs"
|
||||
bad_prompt = ":/\\*?<>| .." # sanitizes to empty, should fallback to "video"
|
||||
|
||||
result = vg._prepare_output_path(str(out_dir), prompt=bad_prompt)
|
||||
|
||||
assert os.path.dirname(result) == str(out_dir)
|
||||
assert os.path.basename(result) == "video.mp4"
|
||||
|
||||
@@ -62,7 +62,7 @@ def run_test(pytest_command: str):
|
||||
|
||||
sys.exit(result.returncode)
|
||||
|
||||
@app.function(gpu="L40S:1", image=image, timeout=900)
|
||||
@app.function(gpu="H100:1", image=image, timeout=900)
|
||||
def run_encoder_tests():
|
||||
run_test("pytest ./fastvideo/tests/encoders -vs")
|
||||
|
||||
@@ -118,6 +118,10 @@ def run_inference_lora_tests():
|
||||
def run_distill_dmd_tests():
|
||||
run_test("pytest ./fastvideo/tests/training/distill/test_distill_dmd.py -vs")
|
||||
|
||||
@app.function(gpu="L40S:2", image=image, timeout=900, secrets=[modal.Secret.from_dict({"WANDB_API_KEY": os.environ.get("WANDB_API_KEY", "")})])
|
||||
def run_self_forcing_tests():
|
||||
run_test("wandb login $WANDB_API_KEY && pytest ./fastvideo/tests/training/self-forcing/test_self_forcing.py -vs")
|
||||
|
||||
@app.function(gpu="L40S:1", image=image, timeout=900)
|
||||
def run_unit_test():
|
||||
run_test("pytest ./fastvideo/tests/dataset/ ./fastvideo/tests/workflow/ -vs")
|
||||
run_test("pytest ./fastvideo/tests/dataset/ ./fastvideo/tests/workflow/ ./fastvideo/tests/entrypoints/ -vs")
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user