Compare commits
30
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
d8bf1beabe | ||
|
|
30ea5bbbbc | ||
|
|
fd80ccbf88 | ||
|
|
93c30ce848 | ||
|
|
ea134fd785 | ||
|
|
4501ef4745 | ||
|
|
979e3e8d5b | ||
|
|
6b8f2d8228 | ||
|
|
3281151955 | ||
|
|
60f68d5b68 | ||
|
|
d20b8df607 | ||
|
|
51c94d6a73 | ||
|
|
dba38dbc03 | ||
|
|
2034cc3c4f | ||
|
|
c69afce2f6 | ||
|
|
b08e758eb3 | ||
|
|
3f3462d7ce | ||
|
|
c9c47dd89c | ||
|
|
9c4ef7c2f1 | ||
|
|
f25eb4b905 | ||
|
|
048d55ccbb | ||
|
|
a271c55fe4 | ||
|
|
f663ae0d8a | ||
|
|
c0911aa3dd | ||
|
|
5f59687ae7 | ||
|
|
5adbc81cdc | ||
|
|
f26d5c37c1 | ||
|
|
6a4ef42378 | ||
|
|
f1098c77dc | ||
|
|
0405b618f8 |
@@ -0,0 +1,106 @@
|
||||
name: Build Image Template
|
||||
|
||||
on:
|
||||
workflow_call:
|
||||
inputs:
|
||||
python_version:
|
||||
required: true
|
||||
type: string
|
||||
dockerfile_path:
|
||||
required: true
|
||||
type: string
|
||||
tag_suffix:
|
||||
required: true
|
||||
type: string
|
||||
|
||||
jobs:
|
||||
build-and-push:
|
||||
runs-on: ubuntu-latest
|
||||
permissions:
|
||||
contents: read
|
||||
packages: write
|
||||
|
||||
steps:
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Free up disk space
|
||||
run: |
|
||||
# Display initial space
|
||||
echo "Initial disk space:"
|
||||
df -h
|
||||
|
||||
# Remove large directories directly
|
||||
sudo rm -rf /usr/share/dotnet
|
||||
sudo rm -rf /usr/local/lib/android
|
||||
sudo rm -rf /opt/ghc
|
||||
sudo rm -rf /usr/local/share/boost
|
||||
sudo rm -rf /usr/share/swift
|
||||
sudo rm -rf /usr/local/lib/node_modules
|
||||
sudo rm -rf /usr/local/share/powershell
|
||||
sudo rm -rf /usr/share/rust
|
||||
sudo rm -rf /usr/local/.ghcup
|
||||
|
||||
# Remove cached files
|
||||
sudo rm -rf /var/lib/apt/lists/*
|
||||
sudo rm -rf /var/cache/apt/archives/*
|
||||
|
||||
# Clean Docker
|
||||
docker system prune -af --volumes
|
||||
|
||||
# Display available space after cleanup
|
||||
echo "Disk space after cleanup:"
|
||||
df -h
|
||||
|
||||
- name: Set up Docker Buildx
|
||||
uses: docker/setup-buildx-action@v3
|
||||
|
||||
- name: Login to GitHub Container Registry
|
||||
uses: docker/login-action@v3
|
||||
with:
|
||||
registry: ghcr.io
|
||||
username: ${{ github.repository_owner }}
|
||||
password: ${{ secrets.GITHUB_TOKEN }}
|
||||
|
||||
- name: Prepare tags
|
||||
id: prepare-tags
|
||||
run: |
|
||||
SHORT_SHA=$(echo ${{ github.sha }} | cut -c1-7)
|
||||
|
||||
TAGS="type=raw,value=${{ inputs.tag_suffix }}-latest"
|
||||
TAGS="${TAGS}\ntype=raw,value=${{ inputs.tag_suffix }}-sha-${SHORT_SHA}"
|
||||
|
||||
# Set Python 3.10 as the default image
|
||||
if [[ "${{ inputs.python_version }}" == "3.10" ]]; then
|
||||
TAGS="${TAGS}\ntype=raw,value=latest"
|
||||
fi
|
||||
|
||||
{
|
||||
echo "tags<<EOF"
|
||||
echo -e "$TAGS"
|
||||
echo "EOF"
|
||||
} >> $GITHUB_OUTPUT
|
||||
|
||||
- name: Extract metadata for Docker
|
||||
id: meta
|
||||
uses: docker/metadata-action@v5
|
||||
with:
|
||||
images: ghcr.io/${{ github.repository }}/fastvideo-dev
|
||||
tags: ${{ steps.prepare-tags.outputs.tags }}
|
||||
|
||||
- name: Build and push Docker image
|
||||
id: build-push
|
||||
uses: docker/build-push-action@v6
|
||||
with:
|
||||
context: .
|
||||
file: ${{ inputs.dockerfile_path }}
|
||||
push: true
|
||||
tags: ${{ steps.meta.outputs.tags }}
|
||||
labels: ${{ steps.meta.outputs.labels }}
|
||||
cache-from: type=gha
|
||||
cache-to: type=gha,mode=max
|
||||
|
||||
- name: Success message
|
||||
run: |
|
||||
echo "✅ Python ${{ inputs.python_version }} image successfully built and pushed to ghcr.io/${{ github.repository }}/fastvideo-dev:${{ inputs.tag_suffix }}-latest"
|
||||
echo "To run tests with this image, manually trigger the 'Run Tests' workflow."
|
||||
@@ -1,78 +1,52 @@
|
||||
name: Build and Push Docker Image
|
||||
name: Build and Push Docker Images
|
||||
|
||||
on:
|
||||
workflow_dispatch: # Only manual triggers
|
||||
workflow_dispatch:
|
||||
inputs:
|
||||
python_3_10:
|
||||
description: 'Build Python 3.10 image'
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
python_3_11:
|
||||
description: 'Build Python 3.11 image'
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
python_3_12:
|
||||
description: 'Build Python 3.12 image'
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
packages: write
|
||||
|
||||
jobs:
|
||||
build-and-push:
|
||||
runs-on: ubuntu-latest
|
||||
permissions:
|
||||
contents: read
|
||||
packages: write
|
||||
|
||||
steps:
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Free up disk space
|
||||
run: |
|
||||
# Display initial space
|
||||
echo "Initial disk space:"
|
||||
df -h
|
||||
|
||||
# Remove large directories directly
|
||||
sudo rm -rf /usr/share/dotnet
|
||||
sudo rm -rf /usr/local/lib/android
|
||||
sudo rm -rf /opt/ghc
|
||||
sudo rm -rf /usr/local/share/boost
|
||||
sudo rm -rf /usr/share/swift
|
||||
sudo rm -rf /usr/local/lib/node_modules
|
||||
sudo rm -rf /usr/local/share/powershell
|
||||
sudo rm -rf /usr/share/rust
|
||||
sudo rm -rf /usr/local/.ghcup
|
||||
|
||||
# Remove cached files
|
||||
sudo rm -rf /var/lib/apt/lists/*
|
||||
sudo rm -rf /var/cache/apt/archives/*
|
||||
|
||||
# Clean Docker
|
||||
docker system prune -af --volumes
|
||||
|
||||
# Display available space after cleanup
|
||||
echo "Disk space after cleanup:"
|
||||
df -h
|
||||
|
||||
- name: Set up Docker Buildx
|
||||
uses: docker/setup-buildx-action@v3
|
||||
|
||||
- name: Login to GitHub Container Registry
|
||||
uses: docker/login-action@v3
|
||||
with:
|
||||
registry: ghcr.io
|
||||
username: ${{ github.repository_owner }}
|
||||
password: ${{ secrets.GITHUB_TOKEN }}
|
||||
|
||||
- name: Extract metadata for Docker
|
||||
id: meta
|
||||
uses: docker/metadata-action@v5
|
||||
with:
|
||||
images: ghcr.io/${{ github.repository }}/fastvideo-dev
|
||||
tags: |
|
||||
type=raw,value=latest
|
||||
type=sha,format=short
|
||||
|
||||
- name: Build and push Docker image
|
||||
id: build-push
|
||||
uses: docker/build-push-action@v6
|
||||
with:
|
||||
context: .
|
||||
push: true
|
||||
tags: ${{ steps.meta.outputs.tags }}
|
||||
labels: ${{ steps.meta.outputs.labels }}
|
||||
cache-from: type=gha
|
||||
cache-to: type=gha,mode=max
|
||||
|
||||
- name: Success message
|
||||
run: |
|
||||
echo "✅ Image successfully built and pushed to ghcr.io/${{ github.repository }}/fastvideo-dev:latest"
|
||||
echo "To run tests with this image, manually trigger the 'Run Tests' workflow."
|
||||
build-python-3-10:
|
||||
if: ${{ github.event.inputs.python_3_10 == 'true' }}
|
||||
uses: ./.github/workflows/build-image-template.yml
|
||||
with:
|
||||
python_version: '3.10'
|
||||
dockerfile_path: docker/Dockerfile.python3.10
|
||||
tag_suffix: py3.10
|
||||
secrets: inherit
|
||||
|
||||
build-python-3-11:
|
||||
if: ${{ github.event.inputs.python_3_11 == 'true' }}
|
||||
uses: ./.github/workflows/build-image-template.yml
|
||||
with:
|
||||
python_version: '3.11'
|
||||
dockerfile_path: docker/Dockerfile.python3.11
|
||||
tag_suffix: py3.11
|
||||
secrets: inherit
|
||||
|
||||
build-python-3-12:
|
||||
if: ${{ github.event.inputs.python_3_12 == 'true' }}
|
||||
uses: ./.github/workflows/build-image-template.yml
|
||||
with:
|
||||
python_version: '3.12'
|
||||
dockerfile_path: docker/Dockerfile.python3.12
|
||||
tag_suffix: py3.12
|
||||
secrets: inherit
|
||||
+2
-1
@@ -40,6 +40,7 @@ eggs/
|
||||
# Sphinx documentation
|
||||
docs/_build/
|
||||
docs/source/getting_started/examples/
|
||||
docs/source/inference/examples/
|
||||
|
||||
# VSCode
|
||||
.vscode/
|
||||
@@ -58,4 +59,4 @@ docs/source/getting_started/examples/
|
||||
!fastvideo/v1/tests/ssim/reference_videos/**/*.mp4
|
||||
|
||||
# Static images
|
||||
!docs/source/_static/images/**/*.png
|
||||
!docs/source/_static/images/**/*.png
|
||||
|
||||
@@ -19,9 +19,11 @@ exclude: |
|
||||
fastvideo/sample/.*|
|
||||
fastvideo/train\.py|
|
||||
fastvideo/utils/.*|
|
||||
fastvideo/v1/examples/.*|
|
||||
examples/.*|
|
||||
fastvideo/v1/models/schedulers/scheduling_flow_match_euler_discrete.py|
|
||||
.github/workflows/fastvideo-publish.yml|
|
||||
.github/workflows/sta-publish.yml
|
||||
.github/workflows/sta-publish.yml|
|
||||
.github/workflows/build-image-template.yml
|
||||
)
|
||||
repos:
|
||||
- repo: https://github.com/google/yapf
|
||||
|
||||
@@ -12,7 +12,8 @@ https://github.com/user-attachments/assets/79af5fb8-707c-4263-b153-9ab2a01d3ac1
|
||||
|
||||
FastVideo currently offers: (with more to come)
|
||||
|
||||
- [NEW!] [Sliding Tile Attention](https://hao-ai-lab.github.io/blogs/sta/).
|
||||
- [NEW!] V1 inference API available. Full announcement coming soon!
|
||||
- [Sliding Tile Attention](https://hao-ai-lab.github.io/blogs/sta/).
|
||||
- FastHunyuan and FastMochi: consistency distilled video diffusion models for 8x inference speedup.
|
||||
- First open distillation recipes for video DiT, based on [PCM](https://github.com/G-U-N/Phased-Consistency-Model).
|
||||
- Support distilling/finetuning/inferencing state-of-the-art open video DiTs: 1. Mochi 2. Hunyuan.
|
||||
@@ -26,217 +27,55 @@ Dev in progress and highly experimental.
|
||||
- ```2025/02/18```: Release the inference code and kernel for [Sliding Tile Attention](https://hao-ai-lab.github.io/blogs/sta/).
|
||||
- ```2025/01/13```: Support Lora finetuning for HunyuanVideo.
|
||||
- ```2024/12/25```: Enable single 4090 inference for `FastHunyuan`, please rerun the installation steps to update the environment.
|
||||
- ```2024/12/17```: `FastVideo` v1.0 is released.
|
||||
- ```2024/12/17```: `FastVideo` v0.0.1 is released.
|
||||
|
||||
## 🔧 Installation from source
|
||||
The code is tested on Python 3.10-3.12, CUDA 12.4 and H100.
|
||||
## Getting Started
|
||||
|
||||
```
|
||||
# Clone FastVideo
|
||||
git clone https://github.com/hao-ai-lab/FastVideo.git && cd FastVideo
|
||||
- [Install FastVideo](https://hao-ai-lab.github.io/FastVideo/getting_started/installation.html)
|
||||
- [Design Overview](https://hao-ai-lab.github.io/FastVideo/design/overview.html)
|
||||
- [Contribution Guide](https://hao-ai-lab.github.io/FastVideo/getting_started/installation.html)
|
||||
|
||||
# Install FastVideo
|
||||
pip install -e .
|
||||
### Inference
|
||||
- [Quick Start](https://hao-ai-lab.github.io/FastVideo/inference/examples/basic.html)
|
||||
- V1 Inference API Guide (Coming soon!)
|
||||
|
||||
# Install Flash Attention (optional)
|
||||
pip install flash-attn==2.7.0.post2
|
||||
```
|
||||
### Distillation and Finetuning
|
||||
- [Distillation Guide](https://hao-ai-lab.github.io/FastVideo/training/distillation.html)
|
||||
- [Finetuning Guide](https://hao-ai-lab.github.io/FastVideo/training/finetuning.html)
|
||||
|
||||
To try Sliding Tile Attention (optional), please follow the instruction in [csrc/sliding_tile_attention/README.md](csrc/sliding_tile_attention/README.md) to install STA.
|
||||
|
||||
You can also install the Sliding Tile Attention package using
|
||||
|
||||
```
|
||||
pip install st_attn==0.0.4
|
||||
```
|
||||
|
||||
## 🚀 Inference
|
||||
### Inference StepVideo with Sliding Tile Attention
|
||||
First, download the model:
|
||||
|
||||
```
|
||||
python scripts/huggingface/download_hf.py --repo_id=stepfun-ai/stepvideo-t2v --local_dir=data/stepvideo-t2v --repo_type=model
|
||||
```
|
||||
|
||||
Use the following scripts to run inference for StepVideo. When using STA for inference, the generated videos will have dimensions of 204×768×768 (currently, this is the only supported shape).
|
||||
|
||||
```bash
|
||||
sh scripts/inference/inference_stepvideo_STA.sh # Inference stepvideo with STA
|
||||
sh scripts/inference/inference_stepvideo.sh # Inference original stepvideo
|
||||
```
|
||||
|
||||
### Inference HunyuanVideo with Sliding Tile Attention
|
||||
First, download the model:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/hunyuan --local_dir=data/hunyuan --repo_type=model
|
||||
```
|
||||
|
||||
We provide two examples in the following script to run inference with STA + [TeaCache](https://github.com/ali-vilab/TeaCache) and STA only.
|
||||
|
||||
```bash
|
||||
sh scripts/inference/inference_hunyuan_STA.sh
|
||||
```
|
||||
|
||||
### Video Demos using STA + Teacache
|
||||
Visit our [demo website](https://fast-video.github.io/) to explore our complete collection of examples. We shorten a single video generation process from 945s to 317s on H100.
|
||||
|
||||
### Inference FastHunyuan on single RTX4090
|
||||
We now support NF4 and LLM-INT8 quantized inference using BitsAndBytes for FastHunyuan. With NF4 quantization, inference can be performed on a single RTX 4090 GPU, requiring just 20GB of VRAM.
|
||||
|
||||
```bash
|
||||
# Download the model weight
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/FastHunyuan-diffusers --local_dir=data/FastHunyuan-diffusers --repo_type=model
|
||||
# CLI inference
|
||||
bash scripts/inference/inference_hunyuan_hf_quantization.sh
|
||||
```
|
||||
|
||||
For more information about the VRAM requirements for BitsAndBytes quantization, please refer to the table below (timing measured on an H100 GPU):
|
||||
|
||||
| Configuration | Memory to Init Transformer | Peak Memory After Init Pipeline (Denoise) | Diffusion Time | End-to-End Time |
|
||||
|--------------------------------|----------------------------|--------------------------------------------|----------------|-----------------|
|
||||
| BF16 + Pipeline CPU Offload | 23.883G | 33.744G | 81s | 121.5s |
|
||||
| INT8 + Pipeline CPU Offload | 13.911G | 27.979G | 88s | 116.7s |
|
||||
| NF4 + Pipeline CPU Offload | 9.453G | 19.26G | 78s | 114.5s |
|
||||
|
||||
For improved quality in generated videos, we recommend using a GPU with 80GB of memory to run the BF16 model with the original Hunyuan pipeline. To execute the inference, use the following section:
|
||||
|
||||
### FastHunyuan
|
||||
|
||||
```bash
|
||||
# Download the model weight
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/FastHunyuan --local_dir=data/FastHunyuan --repo_type=model
|
||||
# CLI inference
|
||||
bash scripts/inference/inference_hunyuan.sh
|
||||
```
|
||||
|
||||
You can also inference FastHunyuan in the [official Hunyuan github](https://github.com/Tencent/HunyuanVideo).
|
||||
|
||||
### FastMochi
|
||||
|
||||
```bash
|
||||
# Download the model weight
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/FastMochi-diffusers --local_dir=data/FastMochi-diffusers --repo_type=model
|
||||
# CLI inference
|
||||
bash scripts/inference/inference_mochi_sp.sh
|
||||
```
|
||||
|
||||
## 🎯 Distill
|
||||
Our distillation recipe is based on [Phased Consistency Model](https://github.com/G-U-N/Phased-Consistency-Model). We did not find significant improvement using multi-phase distillation, so we keep the one phase setup similar to the original latent consistency model's recipe.
|
||||
We use the [MixKit](https://huggingface.co/datasets/LanguageBind/Open-Sora-Plan-v1.1.0/tree/main/all_mixkit) dataset for distillation. To avoid running the text encoder and VAE during training, we preprocess all data to generate text embeddings and VAE latents.
|
||||
Preprocessing instructions can be found [data_preprocess.md](docs/data_preprocess.md). For convenience, we also provide preprocessed data that can be downloaded directly using the following command:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/HD-Mixkit-Finetune-Hunyuan --local_dir=data/HD-Mixkit-Finetune-Hunyuan --repo_type=dataset
|
||||
```
|
||||
|
||||
Next, download the original model weights with:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/hunyuan --local_dir=data/hunyuan --repo_type=model # original hunyuan
|
||||
python scripts/huggingface/download_hf.py --repo_id=genmo/mochi-1-preview --local_dir=data/mochi --repo_type=model # original mochi
|
||||
```
|
||||
|
||||
To launch the distillation process, use the following commands:
|
||||
|
||||
```
|
||||
bash scripts/distill/distill_hunyuan.sh # for hunyuan
|
||||
bash scripts/distill/distill_mochi.sh # for mochi
|
||||
```
|
||||
|
||||
We also provide an optional script for distillation with adversarial loss, located at `fastvideo/distill_adv.py`. Although we tried adversarial loss, we did not observe significant improvements.
|
||||
## Finetune
|
||||
### ⚡ Full Finetune
|
||||
Ensure your data is prepared and preprocessed in the format specified in [data_preprocess.md](docs/data_preprocess.md). For convenience, we also provide a mochi preprocessed Black Myth Wukong data that can be downloaded directly:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/Mochi-Black-Myth --local_dir=data/Mochi-Black-Myth --repo_type=dataset
|
||||
```
|
||||
|
||||
Download the original model weights as specified in [Distill Section](#-distill):
|
||||
|
||||
Then you can run the finetune with:
|
||||
|
||||
```
|
||||
bash scripts/finetune/finetune_mochi.sh # for mochi
|
||||
```
|
||||
|
||||
**Note that for finetuning, we did not tune the hyperparameters in the provided script.**
|
||||
### ⚡ Lora Finetune
|
||||
|
||||
Hunyuan supports Lora fine-tuning of videos up to 720p. Demos and prompts of Black-Myth-Wukong can be found in [here](https://huggingface.co/FastVideo/Hunyuan-Black-Myth-Wukong-lora-weight). You can download the Lora weight through:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/Hunyuan-Black-Myth-Wukong-lora-weight --local_dir=data/Hunyuan-Black-Myth-Wukong-lora-weight --repo_type=model
|
||||
```
|
||||
|
||||
#### Minimum Hardware Requirement
|
||||
- 40 GB GPU memory each for 2 GPUs with lora.
|
||||
- 30 GB GPU memory each for 2 GPUs with CPU offload and lora.
|
||||
|
||||
Currently, both Mochi and Hunyuan models support Lora finetuning through diffusers. To generate personalized videos from your own dataset, you'll need to follow three main steps: dataset preparation, finetuning, and inference.
|
||||
|
||||
#### Dataset Preparation
|
||||
We provide scripts to better help you get started to train on your own characters!
|
||||
You can run this to organize your dataset to get the videos2caption.json before preprocess. Specify your video folder and corresponding caption folder (caption files should be .txt files and have the same name with its video):
|
||||
|
||||
```
|
||||
python scripts/dataset_preparation/prepare_json_file.py --video_dir data/input_videos/ --prompt_dir data/captions/ --output_path data/output_folder/videos2caption.json --verbose
|
||||
```
|
||||
|
||||
Also, we provide script to resize your videos:
|
||||
|
||||
```
|
||||
python scripts/data_preprocess/resize_videos.py
|
||||
```
|
||||
|
||||
#### Finetuning
|
||||
After basic dataset preparation and preprocess, you can start to finetune your model using Lora:
|
||||
|
||||
```
|
||||
bash scripts/finetune/finetune_hunyuan_hf_lora.sh
|
||||
```
|
||||
|
||||
#### Inference
|
||||
For inference with Lora checkpoint, you can run the following scripts with additional parameter `--lora_checkpoint_dir`:
|
||||
|
||||
```
|
||||
bash scripts/inference/inference_hunyuan_hf.sh
|
||||
```
|
||||
|
||||
**We also provide scripts for Mochi in the same directory.**
|
||||
|
||||
#### Finetune with Both Image and Video
|
||||
Our codebase support finetuning with both image and video.
|
||||
|
||||
```bash
|
||||
bash scripts/finetune/finetune_hunyuan.sh
|
||||
bash scripts/finetune/finetune_mochi_lora_mix.sh
|
||||
```
|
||||
|
||||
For Image-Video Mixture Fine-tuning, make sure to enable the `--group_frame` option in your script.
|
||||
### Deprecated APIs
|
||||
- [V0 Inference (Deprecated)](https://hao-ai-lab.github.io/FastVideo/inference/v0_inference.html)
|
||||
|
||||
## 📑 Development Plan
|
||||
|
||||
- More distillation methods
|
||||
- [ ] Add Distribution Matching Distillation
|
||||
<!-- - More distillation methods -->
|
||||
<!-- - [ ] Add Distribution Matching Distillation -->
|
||||
- More models support
|
||||
- [ ] Add CogvideoX model
|
||||
- Code update
|
||||
- [ ] fp8 support
|
||||
- [ ] faster load model and save model support
|
||||
<!-- - [ ] Add CogvideoX model -->
|
||||
- [ ] Add StepVideo to V1
|
||||
- Optimization features
|
||||
- [ ] Teacache in V1
|
||||
- [ ] SageAttention in V1
|
||||
- Code updates
|
||||
- [ ] V1 Configuration API
|
||||
- [ ] Support Training in V1
|
||||
<!-- - [ ] fp8 support -->
|
||||
<!-- - [ ] faster load model and save model support -->
|
||||
|
||||
## 🤝 Contributing
|
||||
|
||||
We welcome all contributions. Please run `bash format.sh --all` before submitting a pull request.
|
||||
|
||||
## 🔧 Testing
|
||||
Run `pytest` to verify the data preprocessing, checkpoint saving, and sequence parallel pipelines. We recommend adding corresponding test cases in the `test` folder to support your contribution.
|
||||
We welcome all contributions. Please check out our guide [here](https://hao-ai-lab.github.io/FastVideo/developer_guide/overview.html)
|
||||
|
||||
## Acknowledgement
|
||||
We learned and reused code from the following projects: [PCM](https://github.com/G-U-N/Phased-Consistency-Model), [diffusers](https://github.com/huggingface/diffusers), [OpenSoraPlan](https://github.com/PKU-YuanGroup/Open-Sora-Plan), and [xDiT](https://github.com/xdit-project/xDiT).
|
||||
We learned and reused code from the following projects:
|
||||
- [PCM](https://github.com/G-U-N/Phased-Consistency-Model)
|
||||
- [diffusers](https://github.com/huggingface/diffusers)
|
||||
- [OpenSoraPlan](https://github.com/PKU-YuanGroup/Open-Sora-Plan)
|
||||
- [xDiT](https://github.com/xdit-project/xDiT)
|
||||
- [vLLM](https://github.com/vllm-project/vllm)
|
||||
- [SGLang](https://github.com/sgl-project/sglang)
|
||||
|
||||
We thank MBZUAI and Anyscale for their support throughout this project.
|
||||
We thank MBZUAI and [Anyscale](https://www.anyscale.com/) for their support throughout this project.
|
||||
|
||||
## Citation
|
||||
If you use FastVideo for your research, please cite our paper:
|
||||
|
||||
@@ -29,7 +29,7 @@ RUN echo "# Placeholder" > README.md
|
||||
|
||||
RUN conda run -n fastvideo-dev pip install --no-cache-dir --upgrade pip && \
|
||||
conda run -n fastvideo-dev pip install --no-cache-dir .[dev] && \
|
||||
conda run -n fastvideo-dev pip install --no-cache-dir flash-attn==2.7.0.post2 --no-build-isolation && \
|
||||
conda run -n fastvideo-dev pip install --no-cache-dir flash-attn==2.7.4.post1 --no-build-isolation && \
|
||||
conda clean -afy
|
||||
|
||||
COPY . .
|
||||
@@ -0,0 +1,48 @@
|
||||
FROM nvidia/cuda:12.4.1-devel-ubuntu20.04
|
||||
|
||||
ENV DEBIAN_FRONTEND=noninteractive
|
||||
|
||||
WORKDIR /FastVideo
|
||||
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
wget \
|
||||
git \
|
||||
ca-certificates \
|
||||
openssh-server \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
RUN wget https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh && \
|
||||
bash Miniconda3-latest-Linux-x86_64.sh -b -p /opt/conda && \
|
||||
rm Miniconda3-latest-Linux-x86_64.sh
|
||||
|
||||
ENV PATH=/opt/conda/bin:$PATH
|
||||
|
||||
RUN conda create --name fastvideo-dev python=3.11.11 -y
|
||||
|
||||
SHELL ["/bin/bash", "-c"]
|
||||
|
||||
# Copy just the pyproject.toml first to leverage Docker cache
|
||||
COPY pyproject.toml ./
|
||||
|
||||
# Create a dummy README to satisfy the installation
|
||||
RUN echo "# Placeholder" > README.md
|
||||
|
||||
RUN conda run -n fastvideo-dev pip install --no-cache-dir --upgrade pip && \
|
||||
conda run -n fastvideo-dev pip install --no-cache-dir .[dev] && \
|
||||
conda run -n fastvideo-dev pip install --no-cache-dir flash-attn==2.7.4.post1 --no-build-isolation && \
|
||||
conda clean -afy
|
||||
|
||||
COPY . .
|
||||
|
||||
RUN conda run -n fastvideo-dev pip install --no-cache-dir -e .[dev]
|
||||
|
||||
# Remove authentication headers
|
||||
RUN git config --unset-all http.https://github.com/.extraheader || true
|
||||
|
||||
# Set up automatic conda environment activation for all shells
|
||||
RUN echo 'source /opt/conda/etc/profile.d/conda.sh' >> /root/.bashrc && \
|
||||
echo 'conda activate fastvideo-dev' >> /root/.bashrc && \
|
||||
# Ensure .bashrc is sourced for SSH login shells
|
||||
echo 'if [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
|
||||
|
||||
EXPOSE 22
|
||||
@@ -0,0 +1,48 @@
|
||||
FROM nvidia/cuda:12.4.1-devel-ubuntu20.04
|
||||
|
||||
ENV DEBIAN_FRONTEND=noninteractive
|
||||
|
||||
WORKDIR /FastVideo
|
||||
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
wget \
|
||||
git \
|
||||
ca-certificates \
|
||||
openssh-server \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
RUN wget https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh && \
|
||||
bash Miniconda3-latest-Linux-x86_64.sh -b -p /opt/conda && \
|
||||
rm Miniconda3-latest-Linux-x86_64.sh
|
||||
|
||||
ENV PATH=/opt/conda/bin:$PATH
|
||||
|
||||
RUN conda create --name fastvideo-dev python=3.12.9 -y
|
||||
|
||||
SHELL ["/bin/bash", "-c"]
|
||||
|
||||
# Copy just the pyproject.toml first to leverage Docker cache
|
||||
COPY pyproject.toml ./
|
||||
|
||||
# Create a dummy README to satisfy the installation
|
||||
RUN echo "# Placeholder" > README.md
|
||||
|
||||
RUN conda run -n fastvideo-dev pip install --no-cache-dir --upgrade pip && \
|
||||
conda run -n fastvideo-dev pip install --no-cache-dir .[dev] && \
|
||||
conda run -n fastvideo-dev pip install --no-cache-dir flash-attn==2.7.4.post1 --no-build-isolation && \
|
||||
conda clean -afy
|
||||
|
||||
COPY . .
|
||||
|
||||
RUN conda run -n fastvideo-dev pip install --no-cache-dir -e .[dev]
|
||||
|
||||
# Remove authentication headers
|
||||
RUN git config --unset-all http.https://github.com/.extraheader || true
|
||||
|
||||
# Set up automatic conda environment activation for all shells
|
||||
RUN echo 'source /opt/conda/etc/profile.d/conda.sh' >> /root/.bashrc && \
|
||||
echo 'conda activate fastvideo-dev' >> /root/.bashrc && \
|
||||
# Ensure .bashrc is sourced for SSH login shells
|
||||
echo 'if [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
|
||||
|
||||
EXPOSE 22
|
||||
@@ -22,3 +22,4 @@ help:
|
||||
clean:
|
||||
@$(SPHINXBUILD) -M clean "$(SOURCEDIR)" "$(BUILDDIR)" $(SPHINXOPTS) $(O)
|
||||
rm -rf "$(SOURCEDIR)/getting_started/examples"
|
||||
rm -rf "$(SOURCEDIR)/inference/examples"
|
||||
|
||||
@@ -1,25 +1,15 @@
|
||||
sphinx==6.2.1
|
||||
sphinx-argparse==0.4.0
|
||||
sphinx-book-theme==1.0.1
|
||||
sphinx==7.4.7
|
||||
sphinx-argparse==0.5.2
|
||||
sphinx-autodoc2==0.5.0
|
||||
sphinx-book-theme==1.1.4
|
||||
sphinx-copybutton==0.5.2
|
||||
sphinx-design==0.6.1
|
||||
sphinx-togglebutton==0.3.2
|
||||
myst-parser==3.0.1
|
||||
msgspec
|
||||
cloudpickle
|
||||
commonmark # Required by sphinx-argparse when using :markdownhelp:
|
||||
|
||||
# packages to install to build the documentation
|
||||
cachetools
|
||||
pydantic >= 2.8
|
||||
-f https://download.pytorch.org/whl/cpu
|
||||
torch
|
||||
py-cpuinfo
|
||||
transformers
|
||||
mistral_common >= 1.5.4
|
||||
aiohttp
|
||||
starlette
|
||||
openai # Required by docs/source/serving/openai_compatible_server.md's vllm.entrypoints.openai.cli_args
|
||||
fastapi # Required by docs/source/serving/openai_compatible_server.md's vllm.entrypoints.openai.cli_args
|
||||
partial-json-parser # Required by docs/source/serving/openai_compatible_server.md's vllm.entrypoints.openai.cli_args
|
||||
requests
|
||||
zmq
|
||||
torch
|
||||
@@ -34,6 +34,6 @@
|
||||
}
|
||||
</style>
|
||||
|
||||
<div class="notification-bar">
|
||||
<!-- <div class="notification-bar">
|
||||
<p>You are viewing the latest developer preview docs. <a href="https://docs.vllm.ai/en/stable/">Click here</a> to view docs for the latest stable release.</p>
|
||||
</div>
|
||||
</div> -->
|
||||
|
||||
@@ -0,0 +1,19 @@
|
||||
# Summary
|
||||
|
||||
## Video Generator
|
||||
|
||||
```{autodoc2-summary}
|
||||
fastvideo.VideoGenerator
|
||||
```
|
||||
|
||||
## Initialization Configuration
|
||||
|
||||
```{autodoc2-summary}
|
||||
fastvideo.v1.configs.pipelines.PipelineConfig
|
||||
```
|
||||
|
||||
## Sampling Configuration
|
||||
|
||||
```{autodoc2-summary}
|
||||
fastvideo.v1.configs.sample.SamplingParam
|
||||
```
|
||||
@@ -0,0 +1,22 @@
|
||||
# type: ignore
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from docutils import nodes
|
||||
from myst_parser.parsers.sphinx_ import MystParser
|
||||
from sphinx.ext.napoleon import docstring
|
||||
|
||||
|
||||
class NapoleonParser(MystParser):
|
||||
|
||||
def parse(self, input_string: str, document: nodes.document) -> None:
|
||||
# Get the Sphinx configuration
|
||||
config = document.settings.env.config
|
||||
|
||||
parsed_content = str(
|
||||
docstring.GoogleDocstring(
|
||||
str(docstring.NumpyDocstring(input_string, config)),
|
||||
config,
|
||||
))
|
||||
return super().parse(parsed_content, document)
|
||||
|
||||
|
||||
Parser = NapoleonParser
|
||||
+62
-44
@@ -13,17 +13,19 @@
|
||||
# documentation root, use os.path.abspath to make it absolute, like shown here.
|
||||
|
||||
import datetime
|
||||
import inspect
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
import requests
|
||||
from sphinx.ext import autodoc
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
sys.path.append(os.path.abspath("../.."))
|
||||
REPO_ROOT = Path(__file__).resolve().parent.parent.parent
|
||||
print(os.path.abspath(REPO_ROOT))
|
||||
sys.path.append(os.path.abspath(REPO_ROOT))
|
||||
|
||||
# -- Project information -----------------------------------------------------
|
||||
|
||||
@@ -41,8 +43,7 @@ extensions = [
|
||||
"sphinx.ext.linkcode",
|
||||
"sphinx.ext.intersphinx",
|
||||
"sphinx_copybutton",
|
||||
"sphinx.ext.autodoc",
|
||||
"sphinx.ext.autosummary",
|
||||
"autodoc2",
|
||||
"myst_parser",
|
||||
"sphinxarg.ext",
|
||||
"sphinx_design",
|
||||
@@ -50,6 +51,31 @@ extensions = [
|
||||
]
|
||||
myst_enable_extensions = [
|
||||
"colon_fence",
|
||||
"fieldlist",
|
||||
]
|
||||
autodoc2_packages = [
|
||||
{
|
||||
"path": "../../fastvideo",
|
||||
"exclude_dirs": ["__pycache__", "third_party"],
|
||||
},
|
||||
]
|
||||
autodoc2_output_dir = "api"
|
||||
autodoc2_render_plugin = "myst"
|
||||
autodoc2_hidden_objects = ["dunder", "private", "inherited"]
|
||||
autodoc2_docstring_parser_regexes = [
|
||||
(".*", "docs.source.autodoc2_docstring_parser"),
|
||||
]
|
||||
autodoc2_sort_names = True
|
||||
autodoc2_index_template = None
|
||||
autodoc2_skip_module_regexes = [
|
||||
"fastvideo.dataset",
|
||||
"fastvideo.distill",
|
||||
"fastvideo.data_preprocess",
|
||||
"fastvideo.models",
|
||||
"fastvideo.sample",
|
||||
"fastvideo.utils",
|
||||
"fastvideo.distill_adv",
|
||||
"fastvideo.train",
|
||||
]
|
||||
|
||||
# Add any paths that contain templates here, relative to this directory.
|
||||
@@ -78,6 +104,11 @@ html_theme_options = {
|
||||
'repository_url': 'https://github.com/hao-ai-lab/FastVideo/',
|
||||
'use_repository_button': True,
|
||||
'use_edit_page_button': True,
|
||||
# Prevents the full API being added to the left sidebar of every page.
|
||||
# Reduces build time by 2.5x and reduces build size from ~225MB to ~95MB.
|
||||
'collapse_navbar': True,
|
||||
# Makes API visible in the right sidebar on API reference pages.
|
||||
'show_toc_level': 3,
|
||||
}
|
||||
# Add any paths that contain custom static files (such as style sheets) here,
|
||||
# relative to this directory. They are copied after the builtin static files,
|
||||
@@ -160,38 +191,38 @@ def linkcode_resolve(domain, info):
|
||||
return None
|
||||
if not info['module']:
|
||||
return None
|
||||
module = info['module']
|
||||
|
||||
# try to determine the correct file and line number to link to
|
||||
obj = sys.modules[module]
|
||||
# Get path from module name
|
||||
file = Path(f"{info['module'].replace('.', '/')}.py")
|
||||
path = REPO_ROOT / file
|
||||
if not path.exists():
|
||||
path = REPO_ROOT / file.with_suffix("") / "__init__.py"
|
||||
if not path.exists():
|
||||
return None
|
||||
|
||||
# get as specific as we can
|
||||
lineno: int = 0
|
||||
filename: str = ""
|
||||
try:
|
||||
for part in info['fullname'].split('.'):
|
||||
obj = getattr(obj, part)
|
||||
# Get the line number of the object
|
||||
with open(path) as f:
|
||||
lines = f.readlines()
|
||||
name = info['fullname'].split(".")[-1]
|
||||
pattern = fr"^( {{4}})*((def|class) )?{name}\b.*"
|
||||
for lineno, line in enumerate(lines, 1):
|
||||
if not line or line.startswith("#"):
|
||||
continue
|
||||
if re.match(pattern, line):
|
||||
break
|
||||
|
||||
if not (inspect.isclass(obj) or inspect.isfunction(obj)
|
||||
or inspect.ismethod(obj)):
|
||||
obj = obj.__class__ # type: ignore[assignment]
|
||||
# If the line number is not found, return None
|
||||
if lineno == len(lines):
|
||||
return None
|
||||
|
||||
lineno = inspect.getsourcelines(obj)[1]
|
||||
filename = (inspect.getsourcefile(obj)
|
||||
or f"{filename}.py").split("FastVideo/", 1)[1]
|
||||
except Exception:
|
||||
# For some things, like a class member, won't work, so
|
||||
# we'll use the line number of the parent (the class)
|
||||
pass
|
||||
|
||||
if filename.startswith("checkouts/"):
|
||||
# If the line number is found, create the URL
|
||||
filename = path.relative_to(REPO_ROOT)
|
||||
if "checkouts" in path.parts:
|
||||
# a PR build on readthedocs
|
||||
pr_number = filename.split("/")[1]
|
||||
filename = filename.split("/", 2)[2]
|
||||
pr_number = REPO_ROOT.name
|
||||
base, branch = get_repo_base_and_branch(pr_number)
|
||||
if base and branch:
|
||||
return f"https://github.com/{base}/blob/{branch}/{filename}#L{lineno}"
|
||||
|
||||
# Otherwise, link to the source file on the main branch
|
||||
return f"https://github.com/hao-ai-lab/FastVideo/blob/main/{filename}#L{lineno}"
|
||||
|
||||
@@ -203,6 +234,8 @@ autodoc_mock_imports = [
|
||||
"cpuinfo",
|
||||
"cv2",
|
||||
"torch",
|
||||
"huggingface_hub",
|
||||
"torchvision",
|
||||
"transformers",
|
||||
"psutil",
|
||||
"prometheus_client",
|
||||
@@ -231,18 +264,6 @@ for mock_target in autodoc_mock_imports:
|
||||
"been loaded into sys.modules when the sphinx build starts.",
|
||||
mock_target)
|
||||
|
||||
|
||||
class MockedClassDocumenter(autodoc.ClassDocumenter):
|
||||
"""Remove note about base class when a class is derived from object."""
|
||||
|
||||
def add_line(self, line: str, source: str, *lineno: int) -> None:
|
||||
if line == " Bases: :py:class:`object`":
|
||||
return
|
||||
super().add_line(line, source, *lineno)
|
||||
|
||||
|
||||
autodoc.ClassDocumenter = MockedClassDocumenter
|
||||
|
||||
intersphinx_mapping = {
|
||||
"python": ("https://docs.python.org/3", None),
|
||||
"typing_extensions":
|
||||
@@ -254,7 +275,4 @@ intersphinx_mapping = {
|
||||
"psutil": ("https://psutil.readthedocs.io/en/stable", None),
|
||||
}
|
||||
|
||||
autodoc_preserve_defaults = True
|
||||
autodoc_warningiserror = True
|
||||
|
||||
navigation_with_keys = False
|
||||
|
||||
@@ -0,0 +1,316 @@
|
||||
(add-pipeline)=
|
||||
|
||||
# 🏗️ Adding a New Diffusion Pipeline
|
||||
|
||||
This guide explains how to implement a custom diffusion pipeline in FastVideo, leveraging the framework's modular architecture for high-performance video generation.
|
||||
|
||||
## Implementation Process Overview
|
||||
|
||||
1. **Port Required Modules** - Identify and implement necessary model components
|
||||
2. **Create Directory Structure** - Set up pipeline files and folders
|
||||
3. **Implement Pipeline Class** - Build the pipeline using existing or custom stages
|
||||
4. **Register Your Pipeline** - Make it discoverable by the framework
|
||||
5. **Configure Your Pipeline** - (Coming soon)
|
||||
|
||||
Need help? Join our [Slack community](https://join.slack.com/t/fastvideo/shared_invite/zt-2zf6ru791-sRwI9lPIUJQq1mIeB_yjJg).
|
||||
|
||||
## Step 1: Pipeline Modules
|
||||
|
||||
### Identifying Required Modules
|
||||
|
||||
FastVideo uses the Hugging Face Diffusers format for model organization:
|
||||
|
||||
1. Examine the `model_index.json` in the HF model repository:
|
||||
|
||||
```json
|
||||
{
|
||||
"_class_name": "WanImageToVideoPipeline",
|
||||
"_diffusers_version": "0.33.0.dev0",
|
||||
"image_encoder": ["transformers", "CLIPVisionModelWithProjection"],
|
||||
"image_processor": ["transformers", "CLIPImageProcessor"],
|
||||
"scheduler": ["diffusers", "UniPCMultistepScheduler"],
|
||||
"text_encoder": ["transformers", "UMT5EncoderModel"],
|
||||
"tokenizer": ["transformers", "T5TokenizerFast"],
|
||||
"transformer": ["diffusers", "WanTransformer3DModel"],
|
||||
"vae": ["diffusers", "AutoencoderKLWan"]
|
||||
}
|
||||
```
|
||||
|
||||
1. For each component:
|
||||
- Note the originating library (`transformers` or `diffusers`)
|
||||
- Identify the class name
|
||||
- Check if it's already available in FastVideo
|
||||
|
||||
2. Review config files in each component's directory for architecture details
|
||||
|
||||
### Implementing Modules
|
||||
|
||||
Place new modules in the appropriate directories:
|
||||
- Encoders: `fastvideo/v1/models/encoders/`
|
||||
- VAEs: `fastvideo/v1/models/vaes/`
|
||||
- Transformer models: `fastvideo/v1/models/dits/`
|
||||
- Schedulers: `fastvideo/v1/models/schedulers/`
|
||||
|
||||
### Adapting Model Layers
|
||||
|
||||
#### Layer Replacements
|
||||
Replace standard PyTorch layers with FastVideo optimized versions:
|
||||
- nn.LayerNorm → fastvideo.v1.layers.layernorm.RMSNorm
|
||||
- Embedding layers → fastvideo.v1.layers.vocab_parallel_embedding modules
|
||||
- Activation functions → versions from fastvideo.v1.layers.activation
|
||||
|
||||
#### Distributed Linear Layers
|
||||
Use appropriate parallel layers for distribution:
|
||||
|
||||
```python
|
||||
# Output dimension parallelism
|
||||
from fastvideo.v1.layers.linear import ColumnParallelLinear
|
||||
self.q_proj = ColumnParallelLinear(
|
||||
input_size=hidden_size,
|
||||
output_size=head_size * num_heads,
|
||||
bias=bias,
|
||||
gather_output=False
|
||||
)
|
||||
|
||||
# Fused QKV projection
|
||||
from fastvideo.v1.layers.linear import QKVParallelLinear
|
||||
self.qkv_proj = QKVParallelLinear(
|
||||
hidden_size=hidden_size,
|
||||
head_size=attention_head_dim,
|
||||
total_num_heads=num_attention_heads,
|
||||
bias=True
|
||||
)
|
||||
|
||||
# Input dimension parallelism
|
||||
from fastvideo.v1.layers.linear import RowParallelLinear
|
||||
self.out_proj = RowParallelLinear(
|
||||
input_size=head_size * num_heads,
|
||||
output_size=hidden_size,
|
||||
bias=bias,
|
||||
input_is_parallel=True
|
||||
)
|
||||
```
|
||||
|
||||
### Attention Layers
|
||||
Replace standard attention with FastVideo's optimized attention:
|
||||
|
||||
```python
|
||||
# Local attention patterns
|
||||
from fastvideo.v1.attention import LocalAttention
|
||||
from fastvideo.v1.attention.backends.abstract import _Backend
|
||||
self.attn = LocalAttention(
|
||||
num_heads=num_heads,
|
||||
head_size=head_dim,
|
||||
dropout_rate=0.0,
|
||||
softmax_scale=None,
|
||||
causal=False,
|
||||
supported_attention_backends=(_Backend.FLASH_ATTN, _Backend.TORCH_SDPA)
|
||||
)
|
||||
|
||||
# Distributed attention for long sequences
|
||||
from fastvideo.v1.attention import DistributedAttention
|
||||
self.attn = DistributedAttention(
|
||||
num_heads=num_heads,
|
||||
head_size=head_dim,
|
||||
dropout_rate=0.0,
|
||||
softmax_scale=None,
|
||||
causal=False,
|
||||
supported_attention_backends=(_Backend.SLIDING_TILE_ATTN, _Backend.FLASH_ATTN, _Backend.TORCH_SDPA)
|
||||
)
|
||||
```
|
||||
|
||||
#### Define supported backend selection
|
||||
|
||||
```python
|
||||
_supported_attention_backends = (_Backend.FLASH_ATTN, _Backend.TORCH_SDPA)
|
||||
```
|
||||
|
||||
### Registering Models
|
||||
|
||||
Register implemented modules in the model registry:
|
||||
|
||||
```python
|
||||
# In fastvideo/v1/models/registry.py
|
||||
_TEXT_TO_VIDEO_DIT_MODELS = {
|
||||
"YourTransformerModel": ("dits", "yourmodule", "YourTransformerClass"),
|
||||
}
|
||||
|
||||
_VAE_MODELS = {
|
||||
"YourVAEModel": ("vaes", "yourvae", "YourVAEClass"),
|
||||
}
|
||||
```
|
||||
|
||||
## Step 2: Directory Structure
|
||||
|
||||
Create a new directory for your pipeline:
|
||||
|
||||
```
|
||||
fastvideo/v1/pipelines/
|
||||
├── your_pipeline/
|
||||
│ ├── __init__.py
|
||||
│ └── your_pipeline.py
|
||||
```
|
||||
|
||||
## Step 3: Implement Pipeline Class
|
||||
|
||||
Pipelines are composed of stages, each handling a specific part of the diffusion process:
|
||||
|
||||
- **InputValidationStage**: Validates input parameters
|
||||
- **Text Encoding Stages**: Handle text encoding (CLIP/Llama/T5)
|
||||
- **CLIPImageEncodingStage**: Processes image inputs
|
||||
- **TimestepPreparationStage**: Prepares diffusion timesteps
|
||||
- **LatentPreparationStage**: Manages latent representations
|
||||
- **ConditioningStage**: Processes conditioning inputs
|
||||
- **DenoisingStage**: Performs denoising diffusion
|
||||
- **DecodingStage**: Converts latents to pixels
|
||||
|
||||
### Creating Your Pipeline
|
||||
|
||||
```python
|
||||
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
from fastvideo.v1.pipelines.stages import (
|
||||
InputValidationStage, CLIPTextEncodingStage, TimestepPreparationStage,
|
||||
LatentPreparationStage, DenoisingStage, DecodingStage
|
||||
)
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
import torch
|
||||
|
||||
class MyCustomPipeline(ComposedPipelineBase):
|
||||
"""Custom diffusion pipeline implementation."""
|
||||
|
||||
# Define required model components from model_index.json
|
||||
_required_config_modules = [
|
||||
"text_encoder", "tokenizer", "vae", "transformer", "scheduler"
|
||||
]
|
||||
|
||||
@property
|
||||
def required_config_modules(self) -> List[str]:
|
||||
return self._required_config_modules
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
"""Initialize pipeline-specific components."""
|
||||
pass
|
||||
|
||||
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=CLIPTextEncodingStage(
|
||||
text_encoder=self.get_module("text_encoder"),
|
||||
tokenizer=self.get_module("tokenizer")
|
||||
)
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="timestep_preparation_stage",
|
||||
stage=TimestepPreparationStage(
|
||||
scheduler=self.get_module("scheduler")
|
||||
)
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="latent_preparation_stage",
|
||||
stage=LatentPreparationStage(
|
||||
scheduler=self.get_module("scheduler"),
|
||||
vae=self.get_module("vae")
|
||||
)
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="denoising_stage",
|
||||
stage=DenoisingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
scheduler=self.get_module("scheduler")
|
||||
)
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="decoding_stage",
|
||||
stage=DecodingStage(
|
||||
vae=self.get_module("vae")
|
||||
)
|
||||
)
|
||||
|
||||
# Register the pipeline class
|
||||
EntryClass = MyCustomPipeline
|
||||
```
|
||||
|
||||
### Creating Custom Stages (Optional)
|
||||
|
||||
If existing stages don't meet your needs, create custom ones:
|
||||
|
||||
```python
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
|
||||
class MyCustomStage(PipelineStage):
|
||||
"""Custom processing stage for the pipeline."""
|
||||
|
||||
def __init__(self, custom_module, other_param=None):
|
||||
super().__init__()
|
||||
self.custom_module = custom_module
|
||||
self.other_param = other_param
|
||||
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
# Access input data
|
||||
input_data = batch.some_attribute
|
||||
|
||||
# Validate inputs
|
||||
if input_data is None:
|
||||
raise ValueError("Required input is missing")
|
||||
|
||||
# Process with your module
|
||||
result = self.custom_module(input_data)
|
||||
|
||||
# Update batch with results
|
||||
batch.some_output = result
|
||||
|
||||
return batch
|
||||
```
|
||||
|
||||
Add your custom stage to the pipeline:
|
||||
|
||||
```python
|
||||
self.add_stage(
|
||||
stage_name="my_custom_stage",
|
||||
stage=MyCustomStage(
|
||||
custom_module=self.get_module("custom_module"),
|
||||
other_param="some_value"
|
||||
)
|
||||
)
|
||||
```
|
||||
|
||||
#### Stage Design Principles
|
||||
|
||||
1. **Single Responsibility**: Focus on one specific task
|
||||
2. **Functional Pattern**: Receive and return a `ForwardBatch` object
|
||||
3. **Dependency Injection**: Pass dependencies through constructor
|
||||
4. **Input Validation**: Validate inputs for clear error messages
|
||||
|
||||
## Step 4: Register Your Pipeline
|
||||
|
||||
Define `EntryClass` at the end of your pipeline file:
|
||||
|
||||
```python
|
||||
# Single pipeline class
|
||||
EntryClass = MyCustomPipeline
|
||||
|
||||
# Or multiple pipeline classes
|
||||
EntryClass = [MyCustomPipeline, MyOtherPipeline]
|
||||
```
|
||||
|
||||
The registry will automatically:
|
||||
1. Scan all packages under `fastvideo/v1/pipelines/`
|
||||
2. Look for `EntryClass` variables
|
||||
3. Register pipelines using their class names as identifiers
|
||||
|
||||
## Best Practices
|
||||
|
||||
- **Reuse Existing Components**: Leverage built-in stages and modules
|
||||
- **Follow Module Organization**: Place new modules in appropriate directories
|
||||
- **Match Model Patterns**: Follow existing code patterns and conventions
|
||||
@@ -0,0 +1,32 @@
|
||||
(docker)=
|
||||
# 🐳 Using the FastVideo Docker Image
|
||||
|
||||
If you prefer a containerized development environment or want to avoid managing dependencies manually, you can use our prebuilt Docker image:
|
||||
|
||||
**Image:** [`ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:latest`](https://ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev)
|
||||
|
||||
## Starting the container
|
||||
|
||||
```bash
|
||||
docker run --gpus all -it ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:latest
|
||||
```
|
||||
|
||||
This will:
|
||||
|
||||
- Start the container with GPU access
|
||||
- Drop you into a shell with the `fastvideo-dev` Conda environment preconfigured
|
||||
|
||||
## Using the container
|
||||
|
||||
```bash
|
||||
# Conda environment should already be active
|
||||
# FastVideo package installed in editable mode
|
||||
|
||||
# Pull the latest changes from remote
|
||||
cd /FastVideo
|
||||
git pull
|
||||
|
||||
# Run linters and tests
|
||||
pre-commit run --all-files
|
||||
pytest tests/
|
||||
```
|
||||
@@ -0,0 +1,13 @@
|
||||
(developer-env)
|
||||
|
||||
# 🧰 Developer Environment
|
||||
|
||||
Accelerate your FastVideo development workflow by leveraging Docker images and cloud GPUs for efficient experimentation and reproducible environments.
|
||||
|
||||
:::{toctree}
|
||||
:caption: Contents
|
||||
:maxdepth: 1
|
||||
|
||||
docker
|
||||
runpod
|
||||
:::
|
||||
@@ -0,0 +1,52 @@
|
||||
(runpod)=
|
||||
|
||||
# 📦 Developing FastVideo on RunPod
|
||||
|
||||
You can easily use the FastVideo Docker image as a custom container on [RunPod](https://www.runpod.io) for development or experimentation.
|
||||
|
||||
## Creating a new pod
|
||||
|
||||
Choose a GPU that supports CUDA 12.4
|
||||
|
||||

|
||||
|
||||
When creating your pod template, use this image:
|
||||
|
||||
```
|
||||
ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:latest
|
||||
```
|
||||
|
||||
Paste Container Start Command to support SSH ([RunPod Docs](https://docs.runpod.io/pods/configuration/use-ssh)):
|
||||
|
||||
```bash
|
||||
bash -c "apt update;DEBIAN_FRONTEND=noninteractive apt-get install openssh-server -y;mkdir -p ~/.ssh;cd $_;chmod 700 ~/.ssh;echo \"$PUBLIC_KEY\" >> authorized_keys;chmod 700 authorized_keys;service ssh start;sleep infinity"
|
||||
```
|
||||
|
||||

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

|
||||
|
||||
## Working with the pod
|
||||
|
||||
After SSH'ing into your pod, you'll find the `fastvideo-dev` Conda environment already activated.
|
||||
|
||||
To pull in the latest changes from the GitHub repo:
|
||||
|
||||
```bash
|
||||
cd /FastVideo
|
||||
git pull
|
||||
```
|
||||
|
||||
`If you have a persistent volume and want to keep your code changes, you can move /FastVideo to /workspace/FastVideo, or simply clone the repository there.`
|
||||
|
||||
Run your development workflows as usual:
|
||||
|
||||
```bash
|
||||
# Run linters
|
||||
pre-commit run --all-files
|
||||
|
||||
# Run tests
|
||||
pytest tests/
|
||||
```
|
||||
@@ -0,0 +1,52 @@
|
||||
(developer-overview)=
|
||||
|
||||
# 🛠️ Contributing to FastVideo
|
||||
|
||||
Thank you for your interest in contributing to FastVideo. We want to make the process as smooth for you as possible and this is a guide to help get you started!
|
||||
|
||||
Our community is open to everyone and welcomes any contributions no matter how large or small.
|
||||
|
||||
# Developer Environment:
|
||||
Do make sure you have CUDA 12.4 installed and supported. FastVideo currently only support Linux and CUDA GPUs, but we hope to support other platforms in the future.
|
||||
|
||||
We recommend using a fresh Python 3.10 Conda environment to develop FastVideo:
|
||||
|
||||
Install Miniconda:
|
||||
|
||||
```
|
||||
wget https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh
|
||||
bash Miniconda3-latest-Linux-x86_64.sh
|
||||
source ~/.bashrc
|
||||
```
|
||||
|
||||
Create and activate a Conda environment for FastVideo:
|
||||
|
||||
```
|
||||
conda create -n fastvideo python=3.10 -y
|
||||
conda activate fastvideo
|
||||
```
|
||||
|
||||
Clone the FastVideo repository and go to the FastVideo directory:
|
||||
|
||||
```
|
||||
git clone https://github.com/hao-ai-lab/FastVideo.git && cd FastVideo
|
||||
|
||||
```
|
||||
|
||||
Now you can install FastVideo and setup git hooks for running linting. By using `pre-commit`, the linters will run and have to pass before you'll be able to make a commit.
|
||||
|
||||
```bash
|
||||
pip install -e .[dev]
|
||||
|
||||
# Can also install flash-attn (optional)
|
||||
pip install flash-attn==2.7.4.post1 --no-build-isolation
|
||||
|
||||
# Linting, formatting and static type checking
|
||||
pre-commit install --hook-type pre-commit --hook-type commit-msg
|
||||
|
||||
# You can manually run pre-commit with
|
||||
pre-commit run --all-files
|
||||
|
||||
# Unit tests
|
||||
pytest tests/
|
||||
```
|
||||
@@ -0,0 +1,410 @@
|
||||
# 🔍 FastVideo Overview
|
||||
|
||||
This document outlines FastVideo's architecture for developers interested in framework internals or contributions. It serves as an onboarding guide for new contributors by providing an overview of the most important directories and files within the `fastvideo/v1/` codebase.
|
||||
|
||||
## Table of Contents - V1 Directory Structure and Files
|
||||
|
||||
- [`fastvideo/v1/pipelines/`](#design-pipeline-system) - Core diffusion pipeline components
|
||||
- [`fastvideo/v1/models/`](#design-model-components) - Model implementations
|
||||
- [`dits/`](#design-transformer-models) - Transformer-based diffusion models
|
||||
- [`vaes/`](#design-vae-variational-auto-encoder) - Variational autoencoders
|
||||
- [`encoders/`](#design-text-and-image-encoders) - Text and image encoders
|
||||
- [`schedulers/`](#design-schedulers) - Diffusion schedulers
|
||||
- [`fastvideo/v1/attention/`](#design-optimized-attention) - Optimized attention implementations
|
||||
- [`fastvideo/v1/distributed/`](#design-distributed-processing) - Distributed computing utilities
|
||||
- [`fastvideo/v1/layers/`](#design-tensor-parallelism) - Custom neural network layers
|
||||
- [`fastvideo/v1/platforms/`](#design-platforms) - Hardware platform abstractions
|
||||
- [`fastvideo/v1/worker/`](#design-executor-and-worker-abstractions) - Multi-GPU process management
|
||||
- [`fastvideo/v1/fastvideo_args.py`](#design-fastvideo-args) - Argument handling
|
||||
- [`fastvideo/v1/forward_context.py`](#design-forwardcontext) - Forward pass context management
|
||||
- `fastvideo/v1/utils.py` - Utility functions
|
||||
- [`fastvideo/v1/logger.py`](#design-logger) - Logging infrastructure
|
||||
|
||||
## Core Architecture
|
||||
|
||||
FastVideo separates model components from execution logic with these principles:
|
||||
- **Component Isolation**: Models (encoders, VAEs, transformers) are isolated from execution (pipelines, stages, distributed processing)
|
||||
- **Modular Design**: Components can be independently replaced
|
||||
- **Distributed Execution**: Supports various parallelism strategies (Tensor, Sequence)
|
||||
- **Custom Attention Backends**: Components can support and use different Attention implementations
|
||||
- **Pipeline Abstraction**: Consistent interface across diffusion models
|
||||
|
||||
(design-fastvideo-args)=
|
||||
## FastVideoArgs
|
||||
|
||||
The `FastVideoArgs` class in `fastvideo/v1/fastvideo_args.py` serves as the central configuration system for FastVideo. It contains all parameters needed to control model loading, inference configuration, performance optimization settings, and more.
|
||||
|
||||
Key features include:
|
||||
- **Command-line Interface**: Automatic conversion between CLI arguments and dataclass fields
|
||||
- **Configuration Groups**: Organized by functional areas (model loading, video params, optimization settings)
|
||||
- **Context Management**: Global access to current settings via `get_current_fastvideo_args()`
|
||||
- **Parameter Validation**: Ensures valid combinations of settings
|
||||
|
||||
Common configuration areas:
|
||||
- **Model paths and loading options**: `model_path`, `trust_remote_code`, `revision`
|
||||
- **Distributed execution settings**: `num_gpus`, `tp_size`, `sp_size`
|
||||
- **Video generation parameters**: `height`, `width`, `num_frames`, `num_inference_steps`
|
||||
- **Precision settings**: Control computation precision for different components
|
||||
|
||||
Example usage:
|
||||
|
||||
```python
|
||||
# Load arguments from command line
|
||||
fastvideo_args = prepare_fastvideo_args(sys.argv[1:])
|
||||
|
||||
# Access parameters
|
||||
model = load_model(fastvideo_args.model_path)
|
||||
|
||||
# Set as global context
|
||||
with set_current_fastvideo_args(fastvideo_args):
|
||||
# Code that requires access to these arguments
|
||||
result = generate_video()
|
||||
```
|
||||
|
||||
(design-pipeline-system)=
|
||||
## Pipeline System
|
||||
|
||||
### `ComposedPipelineBase`
|
||||
|
||||
This foundational class provides:
|
||||
|
||||
- **Model Loading**: Automatically loads components from HuggingFace-Diffusers-compatible model directories
|
||||
- **Stage Management**: Creates and orchestrates processing stages
|
||||
- **Data Flow Coordination**: Ensures proper state flow between stages
|
||||
|
||||
```python
|
||||
class MyCustomPipeline(ComposedPipelineBase):
|
||||
_required_config_modules = [
|
||||
"text_encoder", "tokenizer", "vae", "transformer", "scheduler"
|
||||
]
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
# Pipeline-specific initialization
|
||||
pass
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
self.add_stage("input_validation_stage", InputValidationStage())
|
||||
self.add_stage("text_encoding_stage", CLIPTextEncodingStage(
|
||||
text_encoder=self.get_module("text_encoder"),
|
||||
tokenizer=self.get_module("tokenizer")
|
||||
))
|
||||
# Additional stages...
|
||||
```
|
||||
|
||||
### Pipeline Stages
|
||||
Each stage handles a specific diffusion process component:
|
||||
- **Input Validation**: Parameter verification
|
||||
- **Text Encoding**: CLIP, LLaMA, or T5-based encoding
|
||||
- **Image Encoding**: Image input processing
|
||||
- **Timestep & Latent Preparation**: Setup for diffusion
|
||||
- **Denoising**: Core diffusion loop
|
||||
- **Decoding**: Latent-to-pixel conversion
|
||||
|
||||
Each stage implements a standard interface:
|
||||
|
||||
```python
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
# Process batch and update state
|
||||
return batch
|
||||
```
|
||||
|
||||
(design-forwardbatch)=
|
||||
### ForwardBatch
|
||||
|
||||
Defined in `fastvideo/v1/pipelines/pipeline_batch_info.py`, `ForwardBatch` encapsulates the data payload passed between pipeline stages. It typically holds:
|
||||
|
||||
- **Input Data**: Prompts, images, generation parameters
|
||||
- **Intermediate State**: Embeddings, latents, timesteps, accumulated during stage execution
|
||||
- **Output Storage**: Generated results and metadata
|
||||
- **Configuration**: Sampling parameters, precision settings
|
||||
|
||||
This structure facilitates clear state transitions between stages.
|
||||
|
||||
(design-model-components)=
|
||||
## Model Components
|
||||
|
||||
The `fastvideo/v1/models/` directory contains implementations of the core neural network models used in video diffusion:
|
||||
|
||||
(design-transformer-models)=
|
||||
### Transformer Models
|
||||
|
||||
Transformer networks perform the actual denoising during diffusion:
|
||||
|
||||
- **Location**: `fastvideo/v1/models/dits/`
|
||||
- **Examples**:
|
||||
- `WanTransformer3DModel`
|
||||
- `HunyuanVideoTransformer3DModel`
|
||||
|
||||
Features include:
|
||||
- Text/image conditioning
|
||||
- Standardized interface for model-specific optimizations
|
||||
|
||||
```python
|
||||
def forward(
|
||||
self,
|
||||
latents, # [B, T, C, H, W]
|
||||
encoder_hidden_states, # Text embeddings
|
||||
timestep, # Current diffusion timestep
|
||||
encoder_hidden_states_image=None, # Optional image embeddings
|
||||
**kwargs
|
||||
):
|
||||
# Perform denoising computation
|
||||
return noise_pred # Predicted noise residual
|
||||
```
|
||||
|
||||
(design-vae-variational-auto-encoder)=
|
||||
### VAE (Variational Auto-Encoder)
|
||||
|
||||
VAEs handle conversion between pixel space and latent space:
|
||||
|
||||
- **Location**: `fastvideo/v1/models/vaes/`
|
||||
- **Examples**:
|
||||
- `AutoencoderKLWan`
|
||||
- `AutoencoderKLHunyuanVideo`
|
||||
|
||||
These models compress image/video data to a more efficient latent representation (typically 4x-8x smaller in each dimension).
|
||||
|
||||
FastVideo's VAE implementations include:
|
||||
- Efficient video batch processing
|
||||
- Memory optimization
|
||||
- Optional tiling for large frames
|
||||
- Distributed weight support
|
||||
|
||||
(design-text-and-image-encoders)=
|
||||
### Text and Image Encoders
|
||||
|
||||
Encoders process conditioning inputs into embeddings:
|
||||
|
||||
- **Location**: `fastvideo/v1/models/encoders/`
|
||||
- **Text Encoders**:
|
||||
- `CLIPTextModel`
|
||||
- `LlamaModel`
|
||||
- `UMT5EncoderModel`
|
||||
- **Image Encoders**:
|
||||
- `CLIPVisionModel`
|
||||
|
||||
FastVideo implements optimizations such as:
|
||||
- Vocab parallelism for distributed processing
|
||||
- Caching for common prompts
|
||||
- Precision-tuned computation
|
||||
|
||||
(design-schedulers)=
|
||||
### Schedulers
|
||||
|
||||
Schedulers manage the diffusion sampling process:
|
||||
|
||||
- **Location**: `fastvideo/v1/models/schedulers/`
|
||||
- **Examples**:
|
||||
- `UniPCMultistepScheduler`
|
||||
- `FlowMatchEulerDiscreteScheduler`
|
||||
|
||||
These components control:
|
||||
- Diffusion timestep sequences
|
||||
- Noise prediction to latent update conversions
|
||||
- Quality/speed trade-offs
|
||||
|
||||
```python
|
||||
def step(
|
||||
self,
|
||||
model_output: torch.Tensor,
|
||||
timestep: torch.LongTensor,
|
||||
sample: torch.Tensor,
|
||||
**kwargs
|
||||
) -> torch.Tensor:
|
||||
# Process model output and update latents
|
||||
# Return updated latents
|
||||
return prev_sample
|
||||
```
|
||||
|
||||
(design-optimized-attention)=
|
||||
## Optimized Attention
|
||||
|
||||
The `fastvideo/v1/attention/` directory contains optimized attention implementations crucial for efficient video diffusion:
|
||||
|
||||
### Attention Backends
|
||||
Multiple implementations with automatic selection:
|
||||
- **FLASH_ATTN**: Optimized for supporting hardware
|
||||
- **TORCH_SDPA**: Built-in PyTorch scaled dot-product attention
|
||||
- **SLIDING_TILE_ATTN**: For very long sequences
|
||||
|
||||
```python
|
||||
# Configure available attention backends for this layer
|
||||
self.attn = LocalAttention(
|
||||
num_heads=num_heads,
|
||||
head_size=head_dim,
|
||||
causal=False,
|
||||
supported_attention_backends=(_Backend.FLASH_ATTN, _Backend.TORCH_SDPA)
|
||||
)
|
||||
|
||||
# Override via environment variable
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
|
||||
```
|
||||
|
||||
### Attention Patterns
|
||||
Supports various patterns with memory optimization techniques:
|
||||
- **Cross/Self/Temporal/Global-Local Attention**
|
||||
- Chunking, progressive computation, optimized masking
|
||||
|
||||
(design-distributed-processing)=
|
||||
## Distributed Processing
|
||||
|
||||
The `fastvideo/v1/distributed/` directory contains implementations for distributed model execution:
|
||||
|
||||
(design-tensor-parallelism)=
|
||||
### Tensor Parallelism
|
||||
|
||||
Tensor parallelism splits model weights across devices:
|
||||
|
||||
- **Implementation**: Through `RowParallelLinear` and `ColumnParallelLinear` layers
|
||||
- **Use cases**: Will be used by encoder models as their sequence lengths are shorter and enables efficient sharding.
|
||||
|
||||
```python
|
||||
# Tensor-parallel layers in a transformer block
|
||||
from fastvideo.v1.layers.linear import ColumnParallelLinear, RowParallelLinear
|
||||
|
||||
# Split along output dimension
|
||||
self.qkv_proj = ColumnParallelLinear(
|
||||
input_size=hidden_size,
|
||||
output_size=3 * hidden_size,
|
||||
bias=True,
|
||||
gather_output=False
|
||||
)
|
||||
|
||||
# Split along input dimension
|
||||
self.out_proj = RowParallelLinear(
|
||||
input_size=hidden_size,
|
||||
output_size=hidden_size,
|
||||
bias=True,
|
||||
input_is_parallel=True
|
||||
)
|
||||
```
|
||||
|
||||
### Sequence Parallelism
|
||||
|
||||
Sequence parallelism splits sequences across devices:
|
||||
|
||||
- **Implementation**: Through `DistributedAttention` and sequence splitting
|
||||
- **Use cases**: Long video sequences or high-resolution processing. Used by DiT models.
|
||||
|
||||
```python
|
||||
# Distributed attention for long sequences
|
||||
from fastvideo.v1.attention import DistributedAttention
|
||||
|
||||
self.attn = DistributedAttention(
|
||||
num_heads=num_heads,
|
||||
head_size=head_dim,
|
||||
causal=False,
|
||||
supported_attention_backends=(_Backend.SLIDING_TILE_ATTN, _Backend.FLASH_ATTN)
|
||||
)
|
||||
```
|
||||
|
||||
### Communication Primitives
|
||||
Efficient distributed operations via AllGather, AllReduce, and synchronization mechanisms.
|
||||
|
||||
Efficient communication primitives minimize distributed overhead:
|
||||
|
||||
- **Sequence-Parallel AllGather**: Collects sequence chunks
|
||||
- **Tensor-Parallel AllReduce**: Combines partial results
|
||||
- **Distributed Synchronization**: Coordinates execution
|
||||
|
||||
(design-forwardcontext)=
|
||||
## Forward Context Management
|
||||
|
||||
### ForwardContext
|
||||
|
||||
Defined in `fastvideo/v1/forward_context.py`, `ForwardContext` manages execution-specific state *within* a forward pass, particularly for low-level optimizations. It is accessed via `get_forward_context()`.
|
||||
|
||||
- **Attention Metadata**: Configuration for optimized attention kernels (`attn_metadata`)
|
||||
- **Profiling Data**: Potential hooks for performance metrics collection
|
||||
|
||||
This context-based approach enables:
|
||||
- Dynamic optimization based on execution state (e.g., attention backend selection)
|
||||
- Step-specific customizations within model components
|
||||
|
||||
Usage example:
|
||||
|
||||
```python
|
||||
with set_forward_context(current_timestep, attn_metadata, fastvideo_args):
|
||||
# During this forward pass, components can access context
|
||||
# through get_forward_context()
|
||||
output = model(inputs)
|
||||
```
|
||||
|
||||
(design-executor-and-worker-abstractions)=
|
||||
## Executor and Worker System
|
||||
|
||||
The `fastvideo/v1/worker/` directory contains the distributed execution framework:
|
||||
|
||||
### Executor Abstraction
|
||||
|
||||
FastVideo implements a flexible execution model for distributed processing:
|
||||
|
||||
- **Executor Base Class**: An abstract base class defining the interface for all executors
|
||||
- **MultiProcExecutor**: Primary implementation that spawns and manages worker processes
|
||||
- **GPU Workers**: Handle actual model execution on individual GPUs
|
||||
|
||||
The MultiProcExecutor implementation:
|
||||
1. Spawns worker processes for each GPU
|
||||
2. Establishes communication channels via pipes
|
||||
3. Coordinates distributed operations across workers
|
||||
4. Handles graceful startup and shutdown of the process group
|
||||
|
||||
Each GPU worker:
|
||||
1. Initializes the distributed environment
|
||||
2. Builds the pipeline for the specified model
|
||||
3. Executes requested operations on its assigned GPU
|
||||
4. Manages local resources and communicates results back to the executor
|
||||
|
||||
This design allows FastVideo to efficiently utilize multiple GPUs while providing a simple, unified interface for model execution.
|
||||
|
||||
(design-platforms)=
|
||||
## Platforms
|
||||
|
||||
The `fastvideo/v1/platforms/` directory provides hardware platform abstractions that enable FastVideo to run efficiently on different hardware configurations:
|
||||
|
||||
### Platform Abstraction
|
||||
|
||||
FastVideo's platform abstraction layer enables:
|
||||
- **Hardware Detection**: Automatic detection of available hardware
|
||||
- **Backend Selection**: Appropriate selection of compute kernels
|
||||
- **Memory Management**: Efficient utilization of hardware-specific memory features
|
||||
|
||||
The primary components include:
|
||||
- **Platform Interface**: Defines the common API for all platform implementations
|
||||
- **CUDA Platform**: Optimized implementation for NVIDIA GPUs
|
||||
- **Backend Enum**: Used throughout the codebase for feature selection
|
||||
|
||||
Usage example:
|
||||
|
||||
```python
|
||||
from fastvideo.v1.platforms import current_platform, _Backend
|
||||
|
||||
# Check hardware capabilities
|
||||
if current_platform.supports_backend(_Backend.FLASH_ATTN):
|
||||
# Use FlashAttention implementation
|
||||
else:
|
||||
# Fall back to standard implementation
|
||||
```
|
||||
|
||||
The platform system is designed to be extensible for future hardware targets.
|
||||
|
||||
(design-logger)=
|
||||
## Logger
|
||||
See [PR](https://github.com/hao-ai-lab/FastVideo/pull/356)
|
||||
|
||||
*TODO*: (help wanted) Add an environment variable that disables process-aware logging.
|
||||
|
||||
## Contributing to FastVideo
|
||||
|
||||
If you're a new contributor, here are some common areas to explore:
|
||||
|
||||
1. **Adding a new model**: Implement new model types in the appropriate subdirectory of `fastvideo/v1/models/`
|
||||
2. **Optimizing performance**: Look at attention implementations or memory management
|
||||
3. **Adding a new pipeline**: Create a new pipeline subclass in `fastvideo/v1/pipelines/`
|
||||
4. **Hardware support**: Extend the `platforms` module for new hardware targets
|
||||
|
||||
When adding code, follow these practices:
|
||||
- Use type hints for better code readability
|
||||
- Add appropriate docstrings
|
||||
- Maintain the separation between model components and execution logic
|
||||
- Follow existing patterns for distributed processing
|
||||
@@ -1,138 +0,0 @@
|
||||
(developer-guide)=
|
||||
|
||||
# Contributing to FastVideo
|
||||
|
||||
Thank you for your interest in contributing to FastVideo. We want to make the process as smooth for you as possible and this is a guide to help get you started!
|
||||
|
||||
Our community is open to everyone and welcomes any contributions no matter how large or small.
|
||||
|
||||
# Developer Environment:
|
||||
Do make sure you have CUDA 12.4 installed and supported. FastVideo currently only support Linux and CUDA GPUs, but we hope to support other platforms in the future.
|
||||
|
||||
We recommend using a fresh Python 3.10 Conda environment to develop FastVideo:
|
||||
|
||||
Install Miniconda:
|
||||
|
||||
```
|
||||
wget https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh
|
||||
bash Miniconda3-latest-Linux-x86_64.sh
|
||||
source ~/.bashrc
|
||||
```
|
||||
|
||||
Create and activate a Conda environment for FastVideo:
|
||||
|
||||
```
|
||||
conda create -n fastvideo python=3.10 -y
|
||||
conda activate fastvideo
|
||||
```
|
||||
|
||||
Clone the FastVideo repository and go to the FastVideo directory:
|
||||
|
||||
```
|
||||
git clone https://github.com/hao-ai-lab/FastVideo.git && cd FastVideo
|
||||
|
||||
```
|
||||
|
||||
Now you can install FastVideo and setup git hooks for running linting. By using `pre-commit`, the linters will run and have to pass before you'll be able to make a commit.
|
||||
|
||||
```bash
|
||||
pip install -e .[dev]
|
||||
|
||||
# Can also install flash-attn (optional)
|
||||
pip install flash-attn==2.7.0.post2 --no-build-isolation
|
||||
|
||||
# Linting, formatting and static type checking
|
||||
pre-commit install --hook-type pre-commit --hook-type commit-msg
|
||||
|
||||
# You can manually run pre-commit with
|
||||
pre-commit run --all-files
|
||||
|
||||
# Unit tests
|
||||
pytest tests/
|
||||
```
|
||||
|
||||
---
|
||||
## 🐳 Using the FastVideo Docker Image
|
||||
|
||||
If you prefer a containerized development environment or want to avoid managing dependencies manually, you can use our prebuilt Docker image:
|
||||
|
||||
**Image:** [`ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:latest`](https://ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev)
|
||||
|
||||
### Starting the container
|
||||
|
||||
```bash
|
||||
docker run --gpus all -it ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:latest
|
||||
```
|
||||
|
||||
This will:
|
||||
|
||||
- Start the container with GPU access
|
||||
- Drop you into a shell with the `fastvideo-dev` Conda environment preconfigured
|
||||
|
||||
### Using the container
|
||||
|
||||
```bash
|
||||
# Conda environment should already be active
|
||||
# FastVideo package installed in editable mode
|
||||
|
||||
# Pull the latest changes from remote
|
||||
cd /FastVideo
|
||||
git pull
|
||||
|
||||
# Run linters and tests
|
||||
pre-commit run --all-files
|
||||
pytest tests/
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 📦 Developing FastVideo on RunPod
|
||||
|
||||
You can easily use the FastVideo Docker image as a custom container on [RunPod](https://www.runpod.io) for development or experimentation.
|
||||
|
||||
### Creating a new pod
|
||||
|
||||
Choose a GPU that supports CUDA 12.4
|
||||
|
||||

|
||||
|
||||
When creating your pod template, use this image:
|
||||
|
||||
```
|
||||
ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:latest
|
||||
```
|
||||
|
||||
Paste Container Start Command to support SSH ([RunPod Docs](https://docs.runpod.io/pods/configuration/use-ssh)):
|
||||
|
||||
```bash
|
||||
bash -c "apt update;DEBIAN_FRONTEND=noninteractive apt-get install openssh-server -y;mkdir -p ~/.ssh;cd $_;chmod 700 ~/.ssh;echo \"$PUBLIC_KEY\" >> authorized_keys;chmod 700 authorized_keys;service ssh start;sleep infinity"
|
||||
```
|
||||
|
||||

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

|
||||
|
||||
### Working with the pod
|
||||
|
||||
After SSH'ing into your pod, you'll find the `fastvideo-dev` Conda environment already activated.
|
||||
|
||||
To pull in the latest changes from the GitHub repo:
|
||||
|
||||
```bash
|
||||
cd /FastVideo
|
||||
git pull
|
||||
```
|
||||
|
||||
`If you have a persistent volume and want to keep your code changes, you can move /FastVideo to /workspace/FastVideo, or simply clone the repository there.`
|
||||
|
||||
Run your development workflows as usual:
|
||||
|
||||
```bash
|
||||
# Run linters
|
||||
pre-commit run --all-files
|
||||
|
||||
# Run tests
|
||||
pytest tests/
|
||||
```
|
||||
@@ -9,7 +9,7 @@ from typing import Optional
|
||||
|
||||
ROOT_DIR = Path(__file__).parent.parent.parent.resolve()
|
||||
ROOT_DIR_RELATIVE = '../../../..'
|
||||
EXAMPLE_DIR = ROOT_DIR / "fastvideo/v1/examples"
|
||||
EXAMPLE_DIR = ROOT_DIR / "examples"
|
||||
EXAMPLE_DOC_DIR = ROOT_DIR / "docs/source/getting_started/examples"
|
||||
|
||||
|
||||
@@ -162,48 +162,49 @@ class Example:
|
||||
return content
|
||||
|
||||
|
||||
def generate_examples():
|
||||
# Create the EXAMPLE_DOC_DIR if it doesn't exist
|
||||
if not EXAMPLE_DOC_DIR.exists():
|
||||
EXAMPLE_DOC_DIR.mkdir(parents=True)
|
||||
def generate_examples(generate_main_index=False):
|
||||
"""
|
||||
Generate example documentation.
|
||||
|
||||
Args:
|
||||
generate_main_index (bool): Whether to generate the main examples index.
|
||||
If False, only category-specific indices will be generated.
|
||||
"""
|
||||
# Create empty indices with dynamic paths
|
||||
main_index_dir = ROOT_DIR / "docs/source/examples"
|
||||
if not main_index_dir.exists():
|
||||
main_index_dir.mkdir(parents=True)
|
||||
|
||||
# Create empty indices
|
||||
examples_index = Index(
|
||||
path=EXAMPLE_DOC_DIR / "examples_index.md",
|
||||
title="Examples",
|
||||
description=
|
||||
"A collection of examples demonstrating usage of FastVideo.\nAll documented examples are autogenerated using <gh-file:docs/source/generate_examples.py> from examples found in <gh-file:examples>.", # noqa: E501
|
||||
caption="Examples",
|
||||
maxdepth=2)
|
||||
# Category indices stored in reverse order because they are inserted into
|
||||
# examples_index.documents at index 0 in order
|
||||
# Create the main examples index only if requested
|
||||
examples_index = None
|
||||
if generate_main_index:
|
||||
examples_index = Index(
|
||||
path=main_index_dir / "examples_index.md",
|
||||
title="💡 Examples",
|
||||
description=
|
||||
"A collection of examples demonstrating usage of FastVideo.\nAll documented examples are autogenerated using <gh-file:docs/source/generate_examples.py> from examples found in <gh-file:examples>.", # noqa: E501
|
||||
caption="Examples",
|
||||
maxdepth=2)
|
||||
|
||||
# Category indices with dynamic paths based on category names
|
||||
category_indices = {
|
||||
# "other":
|
||||
# Index(
|
||||
# path=EXAMPLE_DOC_DIR / "examples_other_index.md",
|
||||
# title="Other",
|
||||
# description=
|
||||
# "Other examples that don't strongly fit into the online or offline serving categories.", # noqa: E501
|
||||
# caption="Examples",
|
||||
# ),
|
||||
# "online_serving":
|
||||
# Index(
|
||||
# path=EXAMPLE_DOC_DIR / "examples_online_serving_index.md",
|
||||
# title="Online Serving",
|
||||
# description=
|
||||
# "Online serving examples demonstrate how to use FastVideo in an online setting, where the model is queried for predictions in real-time.", # noqa: E501
|
||||
# caption="Examples",
|
||||
# ),
|
||||
"inference":
|
||||
Index(
|
||||
path=EXAMPLE_DOC_DIR / "examples_inference_index.md",
|
||||
title="Inference",
|
||||
path=ROOT_DIR /
|
||||
"docs/source/inference/examples/examples_inference_index.md",
|
||||
title="🚀 Examples",
|
||||
description=
|
||||
"Inference examples demonstrate how to use FastVideo in an offline setting, where the model is queried for predictions in batches. We recommend starting with <project:basic.md>.", # noqa: E501
|
||||
caption="Examples",
|
||||
),
|
||||
}
|
||||
|
||||
# Ensure all category doc directories exist
|
||||
for category, index in category_indices.items():
|
||||
category_dir = index.path.parent
|
||||
if not category_dir.exists():
|
||||
category_dir.mkdir(parents=True)
|
||||
|
||||
examples = []
|
||||
glob_patterns = ["*.py", "*.md", "*.sh"]
|
||||
# Find categorised examples
|
||||
@@ -216,34 +217,58 @@ def generate_examples():
|
||||
# Find examples in subdirectories
|
||||
for path in category_dir.glob("*/*.md"):
|
||||
examples.append(Example(path.parent, category))
|
||||
# Find uncategorised examples
|
||||
globs = [EXAMPLE_DIR.glob(pattern) for pattern in glob_patterns]
|
||||
for path in itertools.chain(*globs):
|
||||
examples.append(Example(path))
|
||||
# Find examples in subdirectories
|
||||
for path in EXAMPLE_DIR.glob("*/*.md"):
|
||||
# Skip categorised examples
|
||||
if path.parent.name in category_indices:
|
||||
continue
|
||||
examples.append(Example(path.parent))
|
||||
|
||||
# Generate the example documentation
|
||||
# Find uncategorised examples only if we're generating a main index
|
||||
if generate_main_index:
|
||||
globs = [EXAMPLE_DIR.glob(pattern) for pattern in glob_patterns]
|
||||
for path in itertools.chain(*globs):
|
||||
examples.append(Example(path))
|
||||
# Find examples in subdirectories
|
||||
for path in EXAMPLE_DIR.glob("*/*.md"):
|
||||
# Skip categorised examples
|
||||
if path.parent.name in category_indices:
|
||||
continue
|
||||
examples.append(Example(path.parent))
|
||||
|
||||
# Create document directories for each category based on category name and generate files
|
||||
for example in sorted(examples, key=lambda e: e.path.stem):
|
||||
print(example)
|
||||
doc_path = EXAMPLE_DOC_DIR / f"{example.path.stem}.md"
|
||||
|
||||
# Determine which index to use for this example
|
||||
if example.category is not None and example.category in category_indices:
|
||||
index = category_indices[example.category]
|
||||
elif generate_main_index:
|
||||
assert examples_index is not None
|
||||
index = examples_index # Default to main index if available
|
||||
else:
|
||||
# Skip examples without a category if no main index
|
||||
print(f"Skipping {example.path} (no category and no main index)")
|
||||
continue
|
||||
|
||||
# Place generated example markdown in the same directory as its index
|
||||
doc_path = index.path.parent / f"{example.path.stem}.md"
|
||||
with open(doc_path, "w+") as f:
|
||||
f.write(example.generate())
|
||||
# Add the example to the appropriate index
|
||||
assert example.category is not None
|
||||
index = category_indices.get(example.category, examples_index)
|
||||
# Add the example to the index
|
||||
index.documents.append(example.path.stem)
|
||||
|
||||
# Generate the index files
|
||||
# Generate the index files for categories
|
||||
for category_index in category_indices.values():
|
||||
if category_index.documents:
|
||||
examples_index.documents.insert(0, category_index.path.name)
|
||||
# Add to main index if it exists
|
||||
if generate_main_index:
|
||||
rel_path = category_index.path.relative_to(
|
||||
main_index_dir.parent)
|
||||
assert examples_index is not None
|
||||
examples_index.documents.insert(
|
||||
0,
|
||||
str(rel_path).replace(".md", ""))
|
||||
|
||||
# Write the category index file
|
||||
with open(category_index.path, "w+") as f:
|
||||
f.write(category_index.generate())
|
||||
|
||||
with open(examples_index.path, "w+") as f:
|
||||
f.write(examples_index.generate())
|
||||
# Write the main index file if requested
|
||||
if generate_main_index and examples_index:
|
||||
with open(examples_index.path, "w+") as f:
|
||||
f.write(examples_index.generate())
|
||||
|
||||
@@ -4,32 +4,33 @@
|
||||
|
||||
FastVideo currently only supports Linux and NVIDIA CUDA GPUs.
|
||||
|
||||
FastVideo has been tested on the following GPUs, but it should work on any GPUs that supports CUDA 12.4+, please create an issue if you discover any issues:
|
||||
- RTX 4090
|
||||
- A40
|
||||
- L40S
|
||||
- A100
|
||||
- H100
|
||||
|
||||
## Requirements
|
||||
|
||||
- OS: Linux
|
||||
- Python: 3.10-3.12
|
||||
- CUDA 12.4+ (Untested on CUDA < 12.4)
|
||||
- **OS: Linux**
|
||||
- **Python: 3.10-3.12**
|
||||
- **CUDA 12.4**
|
||||
- **At least 1 NVIDIA GPU**
|
||||
|
||||
## Installation Options
|
||||
|
||||
### Option 1: Quick Install
|
||||
## Quick start
|
||||
We recommend using an environment manager such as [Conda](#miniconda-recommended)
|
||||
|
||||
```bash
|
||||
pip install fastvideo
|
||||
```
|
||||
|
||||
### Option 2: Installation from Source
|
||||
Also optionally install flash-attn:
|
||||
|
||||
We recommend using a Python environment such as Conda.
|
||||
```bash
|
||||
pip install flash-attn==2.7.4.post1 --no-build-isolation
|
||||
```
|
||||
|
||||
#### 1. [Optional] Install Miniconda (if not already installed)
|
||||
[Click here to install from source](#installation-from-source)
|
||||
|
||||
## Environment Setup
|
||||
|
||||
We highly recommend using a `conda` or `virtualenv` environment:
|
||||
### Miniconda (Recommended)
|
||||
#### 1. Install Miniconda (if not already installed)
|
||||
|
||||
```bash
|
||||
wget https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh
|
||||
@@ -37,20 +38,22 @@ bash Miniconda3-latest-Linux-x86_64.sh
|
||||
source ~/.bashrc
|
||||
```
|
||||
|
||||
#### 2. [Optional] Create and activate a Conda environment for FastVideo
|
||||
#### 2. Create and activate a Conda environment for FastVideo
|
||||
|
||||
```bash
|
||||
conda create -n fastvideo python=3.10 -y
|
||||
conda activate fastvideo
|
||||
```
|
||||
|
||||
#### 3. Clone the FastVideo repository
|
||||
## Installation from Source
|
||||
|
||||
### 1. Clone the FastVideo repository
|
||||
|
||||
```bash
|
||||
git clone https://github.com/hao-ai-lab/FastVideo.git && cd FastVideo
|
||||
```
|
||||
|
||||
#### 4. Install FastVideo
|
||||
#### 2. Install FastVideo
|
||||
|
||||
Basic installation:
|
||||
|
||||
@@ -63,23 +66,25 @@ pip install -e .
|
||||
### Flash Attention
|
||||
|
||||
```bash
|
||||
pip install flash-attn==2.7.0.post2 --no-build-isolation
|
||||
pip install flash-attn==2.7.4.post1 --no-build-isolation
|
||||
```
|
||||
|
||||
### Sliding Tile Attention (STA) (Requires CUDA 12.4+ and H100)
|
||||
### Sliding Tile Attention (STA) (Requires CUDA 12.4 and H100)
|
||||
|
||||
To try Sliding Tile Attention (optional), please follow the instructions in [csrc/sliding_tile_attention/README.md](#sta-installation) to install STA.
|
||||
|
||||
## Docker Images
|
||||
We also have prebuilt docker images with FastVideo dependencies pre-installed:
|
||||
[Docker Images](#docker)
|
||||
## Development Environment Setup
|
||||
|
||||
If you're planning to contribute to FastVideo please see the following page:
|
||||
[Contributor Guide](#developer-guide)
|
||||
[Contributor Guide](#developer-overview)
|
||||
|
||||
## Hardware Requirements
|
||||
|
||||
### For Basic Inference
|
||||
- NVIDIA GPU with CUDA support
|
||||
- Minimum 20GB VRAM for quantized models (e.g., single RTX 4090)
|
||||
- NVIDIA GPU with CUDA 12.4support
|
||||
|
||||
### For Lora Finetuning
|
||||
- 40GB GPU memory each for 2 GPUs with lora
|
||||
|
||||
@@ -0,0 +1,143 @@
|
||||
# Quick Start
|
||||
|
||||
This page contains instructions and code to get you quickly started with video generation using FastVideo.
|
||||
|
||||
## Requirements
|
||||
- **OS: Linux**
|
||||
- **Python: 3.10-3.12**
|
||||
- **CUDA 12.4**
|
||||
- **At least 1 NVIDIA GPU**
|
||||
|
||||
## Installation
|
||||
|
||||
We recommend using a environment manager such as `Conda`.
|
||||
|
||||
```bash
|
||||
pip install fastvideo
|
||||
```
|
||||
|
||||
Also see the [Installation Guide](installation.md).
|
||||
|
||||
## Generating Your First Video
|
||||
Here's a minimal example to generate a video using the default settings. All of the following code snippets can be directly copied into a Python file and executed with
|
||||
|
||||
```bash
|
||||
python example.py
|
||||
```
|
||||
|
||||
```python
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
def main():
|
||||
# Create a video generator with a pre-trained model
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
num_gpus=1, # Adjust based on your hardware
|
||||
)
|
||||
|
||||
# Define a prompt for your video
|
||||
prompt = "A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes wide with interest."
|
||||
|
||||
# Generate the video
|
||||
video = generator.generate_video(prompt)
|
||||
if __name__ == '__main__':
|
||||
main()
|
||||
```
|
||||
|
||||
The generated video will be saved in the current directory under `outputs/` by default.
|
||||
|
||||
## Customizing Generation
|
||||
|
||||
You can customize various parameters when generating videos:
|
||||
|
||||
```python
|
||||
from fastvideo import VideoGenerator, SamplingParam
|
||||
|
||||
def main():
|
||||
# Create the generator
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
num_gpus=1,
|
||||
)
|
||||
|
||||
# Create and customize sampling parameters
|
||||
sampling_param = SamplingParam.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
|
||||
# How many frames to generate
|
||||
sampling_param.num_frames = 45
|
||||
# How many steps we denoise the video
|
||||
sampling_param.num_inference_steps = 30
|
||||
# How strongly does the video to conform to the prompt
|
||||
sampling_param.guidance_scale = 7.5
|
||||
|
||||
# Optional: provide an initial image for image-to-video generation
|
||||
sampling_param.image_path = "path/to/your/image.jpg" # Optional
|
||||
|
||||
# Generate video with custom parameters
|
||||
prompt = "A beautiful sunset over a calm ocean, with gentle waves."
|
||||
video = generator.generate_video(
|
||||
prompt,
|
||||
sampling_param=sampling_param,
|
||||
output_path="my_videos/", # controls where videos are saved
|
||||
return_frames=True # also return frames from this call (defaults to False)
|
||||
)
|
||||
|
||||
# `video` now contains frames
|
||||
if __name__ == '__main__':
|
||||
main()
|
||||
```
|
||||
|
||||
## Available Models
|
||||
|
||||
FastVideo supports various models for text-to-video generation:
|
||||
|
||||
- `Wan-AI/Wan2.1-T2V-1.3B-Diffusers` - A popular text-to-video model
|
||||
- `FastVideo/FastHunyuan-Diffusers` - A high-performance model for video generation
|
||||
|
||||
## Advanced Configuration
|
||||
|
||||
You can use PipelineConfig for more advanced customization:
|
||||
|
||||
```python
|
||||
from fastvideo import VideoGenerator, PipelineConfig
|
||||
|
||||
# Load the default configuration for a model
|
||||
config = PipelineConfig.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
|
||||
|
||||
# Modify configuration settings
|
||||
config.vae_config.scale_factor = 16
|
||||
config.vae_config.precision = "fp16"
|
||||
|
||||
# Create generator with custom config
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
num_gpus=4,
|
||||
pipeline_config=config,
|
||||
)
|
||||
|
||||
# Generate video
|
||||
prompt = "A futuristic cityscape with flying cars and neon signs."
|
||||
video = generator.generate_video(prompt)
|
||||
```
|
||||
|
||||
## Multi-GPU Setup
|
||||
|
||||
FastVideo automatically handles distributed setup when multiple GPUs are specified:
|
||||
|
||||
```python
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
num_gpus=4, # Will use 4 GPUs in parallel
|
||||
)
|
||||
```
|
||||
|
||||
## Hardware Requirements
|
||||
|
||||
- **Basic Inference**: Minimum 20GB VRAM for quantized models (e.g., single RTX 4090)
|
||||
- **LoRA Finetuning**: 40GB GPU memory each for 2 GPUs with LoRA
|
||||
- **Full Finetuning/Distillation**: Multiple high-memory GPUs recommended (e.g., H100)
|
||||
|
||||
## Next Steps
|
||||
|
||||
- Explore the [API Reference](../api/index.md) for more details
|
||||
- Learn about [Advanced Inference Options](../inference/overview_back.md)
|
||||
- See [Examples](../examples/index.md) for more usage scenarios
|
||||
@@ -0,0 +1,18 @@
|
||||
# V1 API
|
||||
|
||||
FastVideo's V1 API provides a streamlined interface for video generation tasks with powerful customization options. This page documents the primary components of the API.
|
||||
|
||||
## Video Generator
|
||||
|
||||
This class will be the primary Python API for generating videos and images.
|
||||
|
||||
```{autodoc2-summary}
|
||||
fastvideo.VideoGenerator
|
||||
```
|
||||
|
||||
VideoGenerator.from_pretrained() should be the primary way of creating a new video generator.
|
||||
|
||||
````{py:method} from_pretrained(model_path: str, device: typing.optional[str] = none, torch_dtype: typing.optional[torch.dtype] = none, pipeline_config: typing.optional[typing.union[str | fastvideo.v1.configs.pipelines.pipelineconfig]] = none, **kwargs) -> v1.entrypoints.video_generator.videogenerator
|
||||
:canonical: v1.entrypoints.video_generator.videogenerator.from_pretrained
|
||||
:classmethod:
|
||||
```
|
||||
+38
-14
@@ -32,7 +32,8 @@ FastVideo is a lightweight framework for accelerating large video diffusion mode
|
||||
|
||||
FastVideo currently offers: (with more to come)
|
||||
|
||||
- [NEW!] [Sliding Tile Attention](https://hao-ai-lab.github.io/blogs/sta/).
|
||||
- [NEW!] V1 inference API available. Full announcement coming soon!
|
||||
- [Sliding Tile Attention](https://hao-ai-lab.github.io/blogs/sta/).
|
||||
- FastHunyuan and FastMochi: consistency distilled video diffusion models for 8x inference speedup.
|
||||
- First open distillation recipes for video DiT, based on [PCM](https://github.com/G-U-N/Phased-Consistency-Model).
|
||||
- Support distilling/finetuning/inferencing state-of-the-art open video DiTs: 1. Mochi 2. Hunyuan.
|
||||
@@ -43,14 +44,34 @@ Dev in progress and highly experimental.
|
||||
|
||||
## Documentation
|
||||
|
||||
% How to start using vLLM?
|
||||
% How to start using FastVideo?
|
||||
|
||||
:::{toctree}
|
||||
:caption: Getting Started
|
||||
:maxdepth: 1
|
||||
|
||||
getting_started/installation
|
||||
getting_started/examples/examples_index
|
||||
getting_started/quick_start
|
||||
getting_started/v1_api
|
||||
:::
|
||||
|
||||
:::{toctree}
|
||||
:caption: Inference
|
||||
:maxdepth: 1
|
||||
|
||||
inference/examples/basic
|
||||
inference/examples/gradio
|
||||
inference/examples/optimizations
|
||||
inference/v0_inference
|
||||
:::
|
||||
|
||||
:::{toctree}
|
||||
:caption: Training
|
||||
:maxdepth: 1
|
||||
|
||||
training/data_preprocess
|
||||
training/distillation
|
||||
training/finetune
|
||||
:::
|
||||
|
||||
% What is STA Kernel?
|
||||
@@ -60,27 +81,30 @@ getting_started/examples/examples_index
|
||||
:maxdepth: 1
|
||||
|
||||
sliding_tile_attention/installation
|
||||
sliding_tile_attention/usage
|
||||
sliding_tile_attention/test
|
||||
sliding_tile_attention/demo
|
||||
:::
|
||||
|
||||
:::{toctree}
|
||||
:caption: Inference
|
||||
:caption: Design
|
||||
:maxdepth: 1
|
||||
|
||||
inference/wanvideo
|
||||
inference/stepvideo
|
||||
inference/hunyuanvideo
|
||||
inference/fasthunyuan
|
||||
inference/fastmochi
|
||||
design/overview
|
||||
:::
|
||||
|
||||
:::{toctree}
|
||||
:caption: Developer Guide
|
||||
:maxdepth: 1
|
||||
:maxdepth: 2
|
||||
|
||||
developer_guide/overview
|
||||
contributing/overview
|
||||
contributing/developer_env/index
|
||||
contributing/add_pipeline
|
||||
:::
|
||||
|
||||
:::{toctree}
|
||||
:caption: API Reference
|
||||
:maxdepth: 2
|
||||
|
||||
api/summary
|
||||
api/fastvideo/fastvideo
|
||||
:::
|
||||
|
||||
## Indices and tables
|
||||
|
||||
@@ -0,0 +1,74 @@
|
||||
(v0-inference)=
|
||||
|
||||
# [Deprecated] V0 Inference
|
||||
The following commands and APIs are deprecated but still supported until V1's API can completely replace all the features in this page.
|
||||
|
||||
## Inference StepVideo with Sliding Tile Attention
|
||||
First, download the model:
|
||||
|
||||
```
|
||||
python scripts/huggingface/download_hf.py --repo_id=stepfun-ai/stepvideo-t2v --local_dir=data/stepvideo-t2v --repo_type=model
|
||||
```
|
||||
|
||||
Use the following scripts to run inference for StepVideo. When using STA for inference, the generated videos will have dimensions of 204×768×768 (currently, this is the only supported shape).
|
||||
|
||||
```bash
|
||||
sh scripts/inference/inference_stepvideo_STA.sh # Inference stepvideo with STA
|
||||
sh scripts/inference/inference_stepvideo.sh # Inference original stepvideo
|
||||
```
|
||||
|
||||
## Inference HunyuanVideo with Sliding Tile Attention
|
||||
First, download the model:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/hunyuan --local_dir=data/hunyuan --repo_type=model
|
||||
```
|
||||
|
||||
We provide two examples in the following script to run inference with STA + [TeaCache](https://github.com/ali-vilab/TeaCache) and STA only.
|
||||
|
||||
```bash
|
||||
sh scripts/inference/inference_hunyuan_STA.sh
|
||||
```
|
||||
|
||||
## Video Demos using STA + Teacache
|
||||
Visit our [demo website](https://fast-video.github.io/) to explore our complete collection of examples. We shorten a single video generation process from 945s to 317s on H100.
|
||||
|
||||
## Inference FastHunyuan on single RTX4090
|
||||
We now support NF4 and LLM-INT8 quantized inference using BitsAndBytes for FastHunyuan. With NF4 quantization, inference can be performed on a single RTX 4090 GPU, requiring just 20GB of VRAM.
|
||||
|
||||
```bash
|
||||
# Download the model weight
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/FastHunyuan-diffusers --local_dir=data/FastHunyuan-diffusers --repo_type=model
|
||||
# CLI inference
|
||||
bash scripts/inference/inference_hunyuan_hf_quantization.sh
|
||||
```
|
||||
|
||||
For more information about the VRAM requirements for BitsAndBytes quantization, please refer to the table below (timing measured on an H100 GPU):
|
||||
|
||||
| Configuration | Memory to Init Transformer | Peak Memory After Init Pipeline (Denoise) | Diffusion Time | End-to-End Time |
|
||||
|--------------------------------|----------------------------|--------------------------------------------|----------------|-----------------|
|
||||
| BF16 + Pipeline CPU Offload | 23.883G | 33.744G | 81s | 121.5s |
|
||||
| INT8 + Pipeline CPU Offload | 13.911G | 27.979G | 88s | 116.7s |
|
||||
| NF4 + Pipeline CPU Offload | 9.453G | 19.26G | 78s | 114.5s |
|
||||
|
||||
For improved quality in generated videos, we recommend using a GPU with 80GB of memory to run the BF16 model with the original Hunyuan pipeline. To execute the inference, use the following section:
|
||||
|
||||
## FastHunyuan
|
||||
|
||||
```bash
|
||||
# Download the model weight
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/FastHunyuan --local_dir=data/FastHunyuan --repo_type=model
|
||||
# CLI inference
|
||||
bash scripts/inference/inference_hunyuan.sh
|
||||
```
|
||||
|
||||
You can also inference FastHunyuan in the [official Hunyuan github](https://github.com/Tencent/HunyuanVideo).
|
||||
|
||||
## FastMochi
|
||||
|
||||
```bash
|
||||
# Download the model weight
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/FastMochi-diffusers --local_dir=data/FastMochi-diffusers --repo_type=model
|
||||
# CLI inference
|
||||
bash scripts/inference/inference_mochi_sp.sh
|
||||
```
|
||||
@@ -1,6 +1,6 @@
|
||||
(sta-demo)=
|
||||
|
||||
# Demo
|
||||
# 🔍 Demo
|
||||
There is a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
|
||||
|
||||
<div style="text-align: center;">
|
||||
|
||||
@@ -1,10 +1,18 @@
|
||||
(sta-installation)=
|
||||
|
||||
# Installation
|
||||
# 🔧 Installation
|
||||
You can install the Sliding Tile Attention package using
|
||||
|
||||
```
|
||||
pip install st_attn==0.0.4
|
||||
```
|
||||
|
||||
# Building from Source
|
||||
We test our code on Pytorch 2.5.0 and CUDA>=12.4. Currently we only have implementation on H100.
|
||||
First, install C++20 for ThunderKittens:
|
||||
|
||||
```bash
|
||||
cd csrc/sliding_tile_attention/
|
||||
sudo apt update
|
||||
sudo apt install gcc-11 g++-11
|
||||
|
||||
@@ -23,3 +31,25 @@ export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
git submodule update --init --recursive
|
||||
python setup.py install
|
||||
```
|
||||
|
||||
# 🧪 Test
|
||||
|
||||
```bash
|
||||
python test/test_sta.py
|
||||
```
|
||||
|
||||
# 📋 Usage
|
||||
|
||||
```python
|
||||
from st_attn import sliding_tile_attention
|
||||
# assuming video size (T, H, W) = (30, 48, 80), text tokens = 256 with padding.
|
||||
# q, k, v: [batch_size, num_heads, seq_length, head_dim], seq_length = T*H*W + 256
|
||||
# a tile is a cube of size (6, 8, 8)
|
||||
# window_size in tiles: [(window_t, window_h, window_w), (..)...]. For example, window size (3, 3, 3) means a query can attend to (3x6, 3x8, 3x8) = (18, 24, 24) tokens out of the total 30x48x80 video.
|
||||
# text_length: int ranging from 0 to 256
|
||||
# If your attention contains text token (Hunyuan)
|
||||
out = sliding_tile_attention(q, k, v, window_size, text_length)
|
||||
# If your attention does not contain text token (StepVideo)
|
||||
out = sliding_tile_attention(q, k, v, window_size, 0, False)
|
||||
|
||||
```
|
||||
|
||||
@@ -1,7 +0,0 @@
|
||||
(sta-test)=
|
||||
|
||||
# Test
|
||||
|
||||
```bash
|
||||
python test/test_sta.py
|
||||
```
|
||||
@@ -1,17 +0,0 @@
|
||||
(sta-usage)=
|
||||
|
||||
# Usage
|
||||
|
||||
```python
|
||||
from st_attn import sliding_tile_attention
|
||||
# assuming video size (T, H, W) = (30, 48, 80), text tokens = 256 with padding.
|
||||
# q, k, v: [batch_size, num_heads, seq_length, head_dim], seq_length = T*H*W + 256
|
||||
# a tile is a cube of size (6, 8, 8)
|
||||
# window_size in tiles: [(window_t, window_h, window_w), (..)...]. For example, window size (3, 3, 3) means a query can attend to (3x6, 3x8, 3x8) = (18, 24, 24) tokens out of the total 30x48x80 video.
|
||||
# text_length: int ranging from 0 to 256
|
||||
# If your attention contains text token (Hunyuan)
|
||||
out = sliding_tile_attention(q, k, v, window_size, text_length)
|
||||
# If your attention does not contain text token (StepVideo)
|
||||
out = sliding_tile_attention(q, k, v, window_size, 0, False)
|
||||
|
||||
```
|
||||
@@ -1,5 +1,6 @@
|
||||
(v0-data-preprocess)=
|
||||
|
||||
## 🧱 Data Preprocess
|
||||
# 🧱 Data Preprocess
|
||||
|
||||
To save GPU memory, we precompute text embeddings and VAE latents to eliminate the need to load the text encoder and VAE during training.
|
||||
|
||||
@@ -18,10 +19,11 @@ bash scripts/preprocess/preprocess_hunyuan_data.sh # for hunyuan
|
||||
|
||||
The preprocessed dataset will be stored in `Image-Vid-Finetune-Mochi` or `Image-Vid-Finetune-HunYuan` correspondingly.
|
||||
|
||||
### Process your own dataset
|
||||
## Process your own dataset
|
||||
|
||||
If you wish to create your own dataset for finetuning or distillation, please structure you video dataset in the following format:
|
||||
|
||||
```
|
||||
path_to_dataset_folder/
|
||||
├── media/
|
||||
│ ├── 0.jpg
|
||||
@@ -29,6 +31,7 @@ path_to_dataset_folder/
|
||||
│ ├── 2.jpg
|
||||
├── video2caption.json
|
||||
└── merge.txt
|
||||
```
|
||||
|
||||
Format the JSON file as a list, where each item represents a media source:
|
||||
|
||||
@@ -0,0 +1,25 @@
|
||||
(v0-distill)=
|
||||
# 🎯 Distill
|
||||
Our distillation recipe is based on [Phased Consistency Model](https://github.com/G-U-N/Phased-Consistency-Model). We did not find significant improvement using multi-phase distillation, so we keep the one phase setup similar to the original latent consistency model's recipe.
|
||||
We use the [MixKit](https://huggingface.co/datasets/LanguageBind/Open-Sora-Plan-v1.1.0/tree/main/all_mixkit) dataset for distillation. To avoid running the text encoder and VAE during training, we prprocess all data to generate text embeddings and VAE latents.
|
||||
Preprocessing instructions can be found [data_preprocess.md](#v0-data-preprocess). For convenience, we also provide preprocessed data that can be downloaded directly using the following command:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/HD-Mixkit-Finetune-Hunyuan --local_dir=data/HD-Mixkit-Finetune-Hunyuan --repo_type=dataset
|
||||
```
|
||||
|
||||
Next, download the original model weights with:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/hunyuan --local_dir=data/hunyuan --repo_type=model # original hunyuan
|
||||
python scripts/huggingface/download_hf.py --repo_id=genmo/mochi-1-preview --local_dir=data/mochi --repo_type=model # original mochi
|
||||
```
|
||||
|
||||
To launch the distillation process, use the following commands:
|
||||
|
||||
```
|
||||
bash scripts/distill/distill_hunyuan.sh # for hunyuan
|
||||
bash scripts/distill/distill_mochi.sh # for mochi
|
||||
```
|
||||
|
||||
We also provide an optional script for distillation with adversarial loss, located at `fastvideo/distill_adv.py`. Although we tried adversarial loss, we did not observe significant improvements.
|
||||
@@ -0,0 +1,71 @@
|
||||
(v0-finetune)=
|
||||
# 🧠 Finetune
|
||||
## ⚡ Full Finetune
|
||||
Ensure your data is prepared and preprocessed in the format specified in [data_preprocess.md](#v0-data-preprocess). For convenience, we also provide a mochi preprocessed Black Myth Wukong data that can be downloaded directly:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/Mochi-Black-Myth --local_dir=data/Mochi-Black-Myth --repo_type=dataset
|
||||
```
|
||||
|
||||
Download the original model weights as specified in [Distill Section](#v0-distill):
|
||||
|
||||
Then you can run the finetune with:
|
||||
|
||||
```
|
||||
bash scripts/finetune/finetune_mochi.sh # for mochi
|
||||
```
|
||||
|
||||
**Note that for finetuning, we did not tune the hyperparameters in the provided script.**
|
||||
## ⚡ Lora Finetune
|
||||
|
||||
Hunyuan supports Lora fine-tuning of videos up to 720p. Demos and prompts of Black-Myth-Wukong can be found in [here](https://huggingface.co/FastVideo/Hunyuan-Black-Myth-Wukong-lora-weight). You can download the Lora weight through:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/Hunyuan-Black-Myth-Wukong-lora-weight --local_dir=data/Hunyuan-Black-Myth-Wukong-lora-weight --repo_type=model
|
||||
```
|
||||
|
||||
### Minimum Hardware Requirement
|
||||
- 40 GB GPU memory each for 2 GPUs with lora.
|
||||
- 30 GB GPU memory each for 2 GPUs with CPU offload and lora.
|
||||
|
||||
Currently, both Mochi and Hunyuan models support Lora finetuning through diffusers. To generate personalized videos from your own dataset, you'll need to follow three main steps: dataset preparation, finetuning, and inference.
|
||||
|
||||
### Dataset Preparation
|
||||
We provide scripts to better help you get started to train on your own characters!
|
||||
You can run this to organize your dataset to get the videos2caption.json before preprocess. Specify your video folder and corresponding caption folder (caption files should be .txt files and have the same name with its video):
|
||||
|
||||
```
|
||||
python scripts/dataset_preparation/prepare_json_file.py --video_dir data/input_videos/ --prompt_dir data/captions/ --output_path data/output_folder/videos2caption.json --verbose
|
||||
```
|
||||
|
||||
Also, we provide script to resize your videos:
|
||||
|
||||
```
|
||||
python scripts/data_preprocess/resize_videos.py
|
||||
```
|
||||
|
||||
### Finetuning
|
||||
After basic dataset preparation and preprocess, you can start to finetune your model using Lora:
|
||||
|
||||
```
|
||||
bash scripts/finetune/finetune_hunyuan_hf_lora.sh
|
||||
```
|
||||
|
||||
### Inference
|
||||
For inference with Lora checkpoint, you can run the following scripts with additional parameter `--lora_checkpoint_dir`:
|
||||
|
||||
```
|
||||
bash scripts/inference/inference_hunyuan_hf.sh
|
||||
```
|
||||
|
||||
**We also provide scripts for Mochi in the same directory.**
|
||||
|
||||
### Finetune with Both Image and Video
|
||||
Our codebase support finetuning with both image and video.
|
||||
|
||||
```bash
|
||||
bash scripts/finetune/finetune_hunyuan.sh
|
||||
bash scripts/finetune/finetune_mochi_lora_mix.sh
|
||||
```
|
||||
|
||||
For Image-Video Mixture Fine-tuning, make sure to enable the `--group_frame` option in your script.
|
||||
@@ -1,3 +1,41 @@
|
||||
# Basic
|
||||
# Basic Video Generation Tutorial
|
||||
The `VideoGenerator` class provides the primary Python interface for doing offline video generation, which is interacting with a diffusion pipeline without using a separate inference api server.
|
||||
|
||||
The class provides the main python interface for using FastVideo's inference pipeline.
|
||||
## Requirements
|
||||
- At least a single NVIDIA GPU with CUDA 12.4.
|
||||
- Python 3.10-3.12
|
||||
|
||||
## Installation
|
||||
If you have not installed FastVideo, please following these [instructions](https://hao-ai-lab.github.io/FastVideo/getting_started/installation.html) first.
|
||||
|
||||
## Usage
|
||||
The first script in this example shows the most basic usage of FastVideo. If you are new to Python and FastVideo, you should start here.
|
||||
|
||||
```bash
|
||||
# if you have not cloned the directory:
|
||||
git clone https://github.com/hao-ai-lab/FastVideo.git && cd FastVideo
|
||||
|
||||
python fastvideo/v1/examples/inference/basic/basic.py
|
||||
```
|
||||
|
||||
## Basic Walkthrough
|
||||
|
||||
All you need to generate videos using multi-gpus from state-of-the-art diffusion pipelines is the following few lines!
|
||||
|
||||
```python
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
def main():
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
num_gpus=1,
|
||||
)
|
||||
|
||||
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)
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
```
|
||||
|
||||
@@ -1 +1,41 @@
|
||||
print('Hello, world!')
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
# from fastvideo.v1.configs.sample import SamplingParam
|
||||
|
||||
|
||||
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",
|
||||
# if num_gpus > 1, FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
)
|
||||
|
||||
# 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."
|
||||
)
|
||||
video = generator.generate_video(prompt)
|
||||
# video = generator.generate_video(prompt, sampling_param=sampling_param, output_path="wan_t2v_videos/")
|
||||
|
||||
# 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)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
@@ -0,0 +1,62 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
|
||||
def main():
|
||||
|
||||
# This is the config class for the model initialization
|
||||
config = PipelineConfig.from_pretrained("FastVideo/FastHunyuan-Diffusers")
|
||||
# can be used to dump the config to a yaml file
|
||||
config.dump_to_yaml("config.yaml")
|
||||
print(config)
|
||||
# {
|
||||
# 'vae_config': {
|
||||
# 'scale_factor': 8,
|
||||
# 'sp': True,
|
||||
# 'tiling': True,
|
||||
# 'precision': 'fp16'
|
||||
# },
|
||||
# 'text_encoder_config': {
|
||||
# 'precision': 'fp16'
|
||||
# },
|
||||
# 'dit_config': {
|
||||
# 'precision': 'fp16'
|
||||
# },
|
||||
# 'inference_args': {
|
||||
# 'guidance_scale': 7.5,
|
||||
# 'num_inference_steps': 5,
|
||||
# 'seed': 1024,
|
||||
# 'guidance_rescale': 0.0,
|
||||
# 'flow_shift': 17,
|
||||
# 'num_inference_steps': 5,
|
||||
# }
|
||||
# }
|
||||
|
||||
config.vae_config.scale_factor = 16
|
||||
|
||||
# FastVideo will automatically used the optimal default arguments for the model
|
||||
# If a local path is provided, FastVideo will make a best effort attempt to
|
||||
# identify the optimal arguments.
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/FastHunyuan-Diffusers",
|
||||
num_gpus=4,
|
||||
config=config,
|
||||
# or
|
||||
config_path="config.yaml",
|
||||
)
|
||||
|
||||
sampling_param = SamplingParam.from_pretrained(
|
||||
"FastVideo/FastHunyuan-Diffusers")
|
||||
sampling_param.num_inference_steps = 5
|
||||
|
||||
# Generate videos with the same simple API, regardless of GPU count
|
||||
prompt = "A beautiful woman in a red dress walking down a street"
|
||||
video = generator.generate_video(prompt,
|
||||
sampling_param=sampling_param,
|
||||
num_inference_steps=6)
|
||||
|
||||
video2 = generator.generate_video(prompt2)
|
||||
prompt2 = "A beautiful woman in a blue dress walking down a street"
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
+4
-12
@@ -1,4 +1,4 @@
|
||||
# FastVideo VideoGenerator Gradio Demo
|
||||
# FastVideo Gradio 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.
|
||||
|
||||
@@ -16,7 +16,8 @@ The demo uses the FastVideo framework to generate videos based on text prompts.
|
||||
## Requirements
|
||||
|
||||
- Linux-based OS
|
||||
- Python 3.10
|
||||
- Python 3.10-3.12
|
||||
- NVIDIA GPU
|
||||
- Cuda 12.4
|
||||
- FastVideo
|
||||
|
||||
@@ -40,15 +41,6 @@ This will start a web server at `http://0.0.0.0:7860` where you can access the i
|
||||
|
||||
## Model Initialization
|
||||
|
||||
```python
|
||||
args = FastVideoArgs(model_path="FastVideo/FastHunyuan-Diffusers", num_gpus=2)
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
model_path=args.model_path,
|
||||
num_gpus=args.num_gpus
|
||||
)
|
||||
```
|
||||
|
||||
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, or even the number of inference steps, *without ever needing to reload the model*.
|
||||
|
||||
## Video Generation
|
||||
@@ -78,4 +70,4 @@ The interface is built with several components:
|
||||
- **Guidance Scale**: Control how closely the generation follows the prompt
|
||||
- **Inference Steps**: More steps can improve quality but take longer
|
||||
- **Negative Prompt**: Specify what you don't want to see in the video
|
||||
- **Seed**: Control randomness for reproducible results
|
||||
- **Seed**: Control randomness for reproducible results
|
||||
@@ -0,0 +1,169 @@
|
||||
import argparse
|
||||
import os
|
||||
from copy import deepcopy
|
||||
|
||||
import gradio as gr
|
||||
import torch
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.v1.configs.sample.base import SamplingParam
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser(description="FastVideo Gradio Demo")
|
||||
parser.add_argument("--model_path",
|
||||
type=str,
|
||||
default="FastVideo/FastHunyuan-diffusers",
|
||||
help="Path to the model")
|
||||
parser.add_argument("--num_gpus",
|
||||
type=int,
|
||||
default=1,
|
||||
help="Number of GPUs to use")
|
||||
parser.add_argument("--output_path",
|
||||
type=str,
|
||||
default="outputs",
|
||||
help="Path to save generated videos")
|
||||
parsed_args = parser.parse_args()
|
||||
|
||||
# args = FastVideoArgs(model_path="FastVideo/FastHunyuan-Diffusers", num_gpus=2)
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
model_path=parsed_args.model_path, num_gpus=parsed_args.num_gpus)
|
||||
|
||||
default_params = SamplingParam.from_pretrained(parsed_args.model_path)
|
||||
|
||||
def generate_video(
|
||||
prompt,
|
||||
negative_prompt,
|
||||
use_negative_prompt,
|
||||
seed,
|
||||
guidance_scale,
|
||||
num_frames,
|
||||
height,
|
||||
width,
|
||||
num_inference_steps,
|
||||
randomize_seed=False,
|
||||
):
|
||||
params = deepcopy(default_params)
|
||||
params.prompt = prompt
|
||||
params.negative_prompt = negative_prompt
|
||||
params.seed = seed
|
||||
params.guidance_scale = guidance_scale
|
||||
params.num_frames = num_frames
|
||||
params.height = height
|
||||
params.width = width
|
||||
params.num_inference_steps = num_inference_steps
|
||||
|
||||
if randomize_seed:
|
||||
params.seed = torch.randint(0, 1000000, (1, )).item()
|
||||
|
||||
if not use_negative_prompt:
|
||||
params.negative_prompt = None
|
||||
|
||||
generator.generate_video(prompt=prompt, sampling_param=params)
|
||||
|
||||
output_path = os.path.join(parsed_args.output_path,
|
||||
f"{params.prompt[:100]}.mp4")
|
||||
|
||||
return output_path, params.seed
|
||||
|
||||
examples = [
|
||||
"A hand enters the frame, pulling a sheet of plastic wrap over three balls of dough placed on a wooden surface. The plastic wrap is stretched to cover the dough more securely. The hand adjusts the wrap, ensuring that it is tight and smooth over the dough. The scene focuses on the hand’s movements as it secures the edges of the plastic wrap. No new objects appear, and the camera remains stationary, focusing on the action of covering the dough.",
|
||||
"A vintage train snakes through the mountains, its plume of white steam rising dramatically against the jagged peaks. The cars glint in the late afternoon sun, their deep crimson and gold accents lending a touch of elegance. The tracks carve a precarious path along the cliffside, revealing glimpses of a roaring river far below. Inside, passengers peer out the large windows, their faces lit with awe as the landscape unfolds.",
|
||||
"A crowded rooftop bar buzzes with energy, the city skyline twinkling like a field of stars in the background. Strings of fairy lights hang above, casting a warm, golden glow over the scene. Groups of people gather around high tables, their laughter blending with the soft rhythm of live jazz. The aroma of freshly mixed cocktails and charred appetizers wafts through the air, mingling with the cool night breeze.",
|
||||
]
|
||||
|
||||
with gr.Blocks() as demo:
|
||||
gr.Markdown("# FastVideo Inference Demo")
|
||||
|
||||
with gr.Group():
|
||||
with gr.Row():
|
||||
prompt = gr.Text(
|
||||
label="Prompt",
|
||||
show_label=False,
|
||||
max_lines=1,
|
||||
placeholder="Enter your prompt",
|
||||
container=False,
|
||||
)
|
||||
run_button = gr.Button("Run", scale=0)
|
||||
result = gr.Video(label="Result", show_label=False)
|
||||
|
||||
with gr.Accordion("Advanced options", open=False):
|
||||
with gr.Group():
|
||||
with gr.Row():
|
||||
height = gr.Slider(
|
||||
label="Height",
|
||||
minimum=256,
|
||||
maximum=1024,
|
||||
step=32,
|
||||
value=default_params.height,
|
||||
)
|
||||
width = gr.Slider(label="Width",
|
||||
minimum=256,
|
||||
maximum=1024,
|
||||
step=32,
|
||||
value=default_params.width)
|
||||
|
||||
with gr.Row():
|
||||
num_frames = gr.Slider(
|
||||
label="Number of Frames",
|
||||
minimum=21,
|
||||
maximum=163,
|
||||
value=default_params.num_frames,
|
||||
)
|
||||
guidance_scale = gr.Slider(
|
||||
label="Guidance Scale",
|
||||
minimum=1,
|
||||
maximum=12,
|
||||
value=default_params.guidance_scale,
|
||||
)
|
||||
num_inference_steps = gr.Slider(
|
||||
label="Inference Steps",
|
||||
minimum=4,
|
||||
maximum=100,
|
||||
value=default_params.num_inference_steps,
|
||||
)
|
||||
|
||||
with gr.Row():
|
||||
use_negative_prompt = gr.Checkbox(
|
||||
label="Use negative prompt", value=False)
|
||||
negative_prompt = gr.Text(
|
||||
label="Negative prompt",
|
||||
max_lines=1,
|
||||
placeholder="Enter a negative prompt",
|
||||
visible=False,
|
||||
)
|
||||
|
||||
seed = gr.Slider(label="Seed",
|
||||
minimum=0,
|
||||
maximum=1000000,
|
||||
step=1,
|
||||
value=default_params.seed)
|
||||
randomize_seed = gr.Checkbox(label="Randomize seed", value=True)
|
||||
seed_output = gr.Number(label="Used Seed")
|
||||
|
||||
gr.Examples(examples=examples, inputs=prompt)
|
||||
|
||||
use_negative_prompt.change(
|
||||
fn=lambda x: gr.update(visible=x),
|
||||
inputs=use_negative_prompt,
|
||||
outputs=default_params.negative_prompt,
|
||||
)
|
||||
|
||||
run_button.click(
|
||||
fn=generate_video,
|
||||
inputs=[
|
||||
prompt,
|
||||
negative_prompt,
|
||||
use_negative_prompt,
|
||||
seed,
|
||||
guidance_scale,
|
||||
num_frames,
|
||||
height,
|
||||
width,
|
||||
num_inference_steps,
|
||||
randomize_seed,
|
||||
],
|
||||
outputs=[result, seed_output],
|
||||
)
|
||||
|
||||
demo.queue(max_size=20).launch(server_name="0.0.0.0", server_port=7860)
|
||||
@@ -0,0 +1,104 @@
|
||||
# Speeding Up Generation
|
||||
|
||||
This page describes the various options for speeding up generation times.
|
||||
|
||||
## Table of Contents
|
||||
- Optimized Attention Backends
|
||||
- [Flash Attention](#optimizations-flash)
|
||||
- [Sliding Tile Attention](#optimizations-sta)
|
||||
- [Sage Attention](#optimizations-sage)
|
||||
|
||||
- Caching Techniques
|
||||
- [TeaCache](#optimizations-teacache)
|
||||
|
||||
|
||||
(optimizations-backends)=
|
||||
## Attention Backends
|
||||
|
||||
`attention_example.py` shows how to set `FASTVIDEO_ATTENTION_BACKEND` env var to change attention backends. To run this example:
|
||||
```bash
|
||||
python examples/inference/optimizations/attention_example.py
|
||||
```
|
||||
|
||||
In python, set the `FASTVIDEO_ATTENTION_BACKEND` before instantiating `VideoGenerator` like this:
|
||||
|
||||
```python
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "SLIDING_TILE_ATTN"
|
||||
```
|
||||
|
||||
You can also set the env var when running any of the other example like this:
|
||||
```bash
|
||||
FASTVIDEO_ATTENTION_BACKEND=SAGE_ATTN python example.py
|
||||
```
|
||||
|
||||
(optimizations-flash)=
|
||||
### Flash Attention
|
||||
|
||||
(optimizations-sta)=
|
||||
### Sliding Tile Attention
|
||||
|
||||
(optimizations-sage)=
|
||||
### Sage Attention
|
||||
|
||||
|
||||
(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?
|
||||
|
||||
See the official [TeaCache](https://github.com/ali-vilab/TeaCache) repo and their paper for more details.
|
||||
|
||||
|
||||
### How to Enable TeaCache
|
||||
|
||||
Enabling TeaCache is straightforward - simply add the `enable_teacache=True` parameter to your `generate_video()` call:
|
||||
|
||||
```python
|
||||
# ... previous code
|
||||
generator.generate_video(
|
||||
prompt="Your prompt here",
|
||||
sampling_param=params,
|
||||
enable_teacache=True
|
||||
)
|
||||
# more code ...
|
||||
```
|
||||
|
||||
### Complete Example
|
||||
|
||||
At the bottom is a complete example of using TeaCache for faster video generation. You can run it using the following command:
|
||||
|
||||
```bash
|
||||
python examples/inference/optimizations/teacache_example.py
|
||||
```
|
||||
|
||||
### Advanced Configuration
|
||||
|
||||
While TeaCache works well with default settings, you can fine-tune its behavior by adjusting the threshold value:
|
||||
|
||||
1. Lower threshold values (e.g., 0.1) will result in more skipped calculations and faster generation with slightly more potential for quality degradation
|
||||
2. Higher threshold values (e.g., 0.15-0.23) will skip fewer calculations but maintain quality closer to the original
|
||||
|
||||
Note that the optimal threshold depends on your specific model and content.
|
||||
|
||||
## Benchmarking different optimizations
|
||||
|
||||
To benchmark the performance improvement, try generating the same video with and without TeaCache enabled and compare the generation times:
|
||||
|
||||
```python
|
||||
# Without TeaCache
|
||||
start_time = time.time()
|
||||
generator.generate_video(prompt="Your prompt", enable_teacache=False)
|
||||
standard_time = time.time() - start_time
|
||||
|
||||
# With TeaCache
|
||||
start_time = time.time()
|
||||
generator.generate_video(prompt="Your prompt", enable_teacache=True)
|
||||
teacache_time = time.time() - start_time
|
||||
|
||||
print(f"Standard generation: {standard_time:.2f} seconds")
|
||||
print(f"TeaCache generation: {teacache_time:.2f} seconds")
|
||||
print(f"Speedup: {standard_time/teacache_time:.2f}x")
|
||||
```
|
||||
|
||||
Note: If you want to benchmark different attention backends, you'll need to reinstantiate `VideoGenerator`.
|
||||
@@ -0,0 +1,33 @@
|
||||
import os
|
||||
import time
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
def main():
|
||||
# set the attention backend
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "FLASH_ATTN"
|
||||
|
||||
start_time = time.time()
|
||||
gen = VideoGenerator.from_pretrained(
|
||||
model_path="Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
num_gpus=1,
|
||||
)
|
||||
load_time = time.time() - start_time
|
||||
print(f"Model loading time: {load_time:.2f} seconds")
|
||||
|
||||
gen_start_time = time.time()
|
||||
|
||||
gen.generate_video(
|
||||
prompt=
|
||||
"Will Smith casually eats noodles, his relaxed demeanor contrasting with the energetic background of a bustling street food market. The scene captures a mix of humor and authenticity. Mid-shot framing, vibrant lighting.",
|
||||
seed=1024,
|
||||
output_path="example_outputs/")
|
||||
|
||||
generation_time = time.time() - gen_start_time
|
||||
print(f"Video generation time: {generation_time:.2f} seconds")
|
||||
|
||||
total_time = time.time() - start_time
|
||||
print(f"Total execution time: {total_time:.2f} seconds")
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,44 @@
|
||||
import time
|
||||
|
||||
from fastvideo import VideoGenerator, SamplingParam
|
||||
|
||||
|
||||
def main():
|
||||
start_time = time.time()
|
||||
|
||||
gen = VideoGenerator.from_pretrained(
|
||||
model_path="Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
num_gpus=1,
|
||||
use_cpu_offload=False,
|
||||
)
|
||||
load_time = time.time() - start_time
|
||||
print(f"Model loading time: {load_time:.2f} seconds")
|
||||
|
||||
gen_start_time = time.time()
|
||||
|
||||
params = SamplingParam.from_pretrained(
|
||||
model_path="Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
)
|
||||
# this controls the threshold for the tea cache
|
||||
params.teacache_params.teacache_thresh = 0.08
|
||||
gen.generate_video(
|
||||
prompt=
|
||||
"Will Smith casually eats noodles, his relaxed demeanor contrasting with the energetic background of a bustling street food market. The scene captures a mix of humor and authenticity. Mid-shot framing, vibrant lighting.",
|
||||
sampling_param=params,
|
||||
height=480,
|
||||
width=832,
|
||||
num_frames=61, # 85 ,77
|
||||
num_inference_steps=50,
|
||||
enable_teacache=True,
|
||||
seed=1024,
|
||||
output_path="example_outputs/")
|
||||
|
||||
generation_time = time.time() - gen_start_time
|
||||
print(f"Video generation time: {generation_time:.2f} seconds")
|
||||
|
||||
total_time = time.time() - start_time
|
||||
print(f"Total execution time: {total_time:.2f} seconds")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,3 +1,5 @@
|
||||
from fastvideo.v1.configs.pipelines import PipelineConfig
|
||||
from fastvideo.v1.configs.sample import SamplingParam
|
||||
from fastvideo.v1.entrypoints.video_generator import VideoGenerator
|
||||
|
||||
__all__ = ["VideoGenerator"]
|
||||
__all__ = ["VideoGenerator", "PipelineConfig", "SamplingParam"]
|
||||
|
||||
@@ -56,20 +56,6 @@ class AttentionMetadata:
|
||||
# Current step of diffusion process
|
||||
current_timestep: int
|
||||
|
||||
# @property
|
||||
# @abstractmethod
|
||||
# def inference_metadata(self) -> Optional["AttentionMetadata"]:
|
||||
# """Return the attention metadata that's required to run prefill
|
||||
# attention."""
|
||||
# pass
|
||||
|
||||
# @property
|
||||
# @abstractmethod
|
||||
# def training_metadata(self) -> Optional["AttentionMetadata"]:
|
||||
# """Return the attention metadata that's required to run decode
|
||||
# attention."""
|
||||
# pass
|
||||
|
||||
def asdict_zerocopy(self,
|
||||
skip_fields: Optional[Set[str]] = None
|
||||
) -> Dict[str, Any]:
|
||||
@@ -86,55 +72,6 @@ class AttentionMetadata:
|
||||
|
||||
T = TypeVar("T", bound=AttentionMetadata)
|
||||
|
||||
# class AttentionState(ABC, Generic[T]):
|
||||
# """Holds attention backend-specific objects reused during the
|
||||
# lifetime of the model runner."""
|
||||
|
||||
# @abstractmethod
|
||||
# def __init__(self, runner: "ModelRunnerBase"):
|
||||
# ...
|
||||
|
||||
# @abstractmethod
|
||||
# @contextmanager
|
||||
# def graph_capture(self, max_batch_size: int):
|
||||
# """Context manager used when capturing CUDA graphs."""
|
||||
# yield
|
||||
|
||||
# @abstractmethod
|
||||
# def graph_clone(self, batch_size: int) -> "AttentionState[T]":
|
||||
# """Clone attention state to save in CUDA graph metadata."""
|
||||
# ...
|
||||
|
||||
# @abstractmethod
|
||||
# def graph_capture_get_metadata_for_batch(
|
||||
# self,
|
||||
# batch_size: int,
|
||||
# is_encoder_decoder_model: bool = False) -> T:
|
||||
# """Get attention metadata for CUDA graph capture of batch_size."""
|
||||
# ...
|
||||
|
||||
# @abstractmethod
|
||||
# def get_graph_input_buffers(
|
||||
# self,
|
||||
# attn_metadata: T,
|
||||
# is_encoder_decoder_model: bool = False) -> Dict[str, Any]:
|
||||
# """Get attention-specific input buffers for CUDA graph capture."""
|
||||
# ...
|
||||
|
||||
# @abstractmethod
|
||||
# def prepare_graph_input_buffers(
|
||||
# self,
|
||||
# input_buffers: Dict[str, Any],
|
||||
# attn_metadata: T,
|
||||
# is_encoder_decoder_model: bool = False) -> None:
|
||||
# """In-place modify input buffers dict for CUDA graph replay."""
|
||||
# ...
|
||||
|
||||
# @abstractmethod
|
||||
# def begin_forward(self, model_input: "ModelRunnerInputBase") -> None:
|
||||
# """Prepare state for forward pass."""
|
||||
# ...
|
||||
|
||||
|
||||
class AttentionMetadataBuilder(ABC, Generic[T]):
|
||||
"""Abstract class for attention metadata builders."""
|
||||
|
||||
@@ -0,0 +1,66 @@
|
||||
from typing import List, Optional, Type
|
||||
|
||||
import torch
|
||||
from sageattention import sageattn
|
||||
|
||||
from fastvideo.v1.attention.backends.abstract import (
|
||||
AttentionBackend) # FlashAttentionMetadata,
|
||||
from fastvideo.v1.attention.backends.abstract import (AttentionImpl,
|
||||
AttentionMetadata)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class SageAttentionBackend(AttentionBackend):
|
||||
|
||||
accept_output_buffer: bool = True
|
||||
|
||||
@staticmethod
|
||||
def get_supported_head_sizes() -> List[int]:
|
||||
return [32, 64, 96, 128, 160, 192, 224, 256]
|
||||
|
||||
@staticmethod
|
||||
def get_name() -> str:
|
||||
return "SAGE_ATTN"
|
||||
|
||||
@staticmethod
|
||||
def get_impl_cls() -> Type["SageAttentionImpl"]:
|
||||
return SageAttentionImpl
|
||||
|
||||
# @staticmethod
|
||||
# def get_metadata_cls() -> Type["AttentionMetadata"]:
|
||||
# return FlashAttentionMetadata
|
||||
|
||||
|
||||
class SageAttentionImpl(AttentionImpl):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
num_heads: int,
|
||||
head_size: int,
|
||||
causal: bool,
|
||||
softmax_scale: float,
|
||||
num_kv_heads: Optional[int] = 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:
|
||||
output = sageattn(
|
||||
query,
|
||||
key,
|
||||
value,
|
||||
# since input is (batch_size, seq_len, head_num, head_dim)
|
||||
tensor_layout="NHD",
|
||||
is_causal=self.causal)
|
||||
return output
|
||||
@@ -0,0 +1,48 @@
|
||||
{
|
||||
"embedded_cfg_scale": 6,
|
||||
"flow_shift": 17,
|
||||
"use_cpu_offload": false,
|
||||
"disable_autocast": false,
|
||||
"precision": "bf16",
|
||||
"vae_precision": "fp16",
|
||||
"vae_tiling": true,
|
||||
"vae_sp": true,
|
||||
"vae_config": {
|
||||
"load_encoder": false,
|
||||
"load_decoder": true,
|
||||
"tile_sample_min_height": 256,
|
||||
"tile_sample_min_width": 256,
|
||||
"tile_sample_min_num_frames": 16,
|
||||
"tile_sample_stride_height": 192,
|
||||
"tile_sample_stride_width": 192,
|
||||
"tile_sample_stride_num_frames": 12,
|
||||
"blend_num_frames": 4,
|
||||
"use_tiling": true,
|
||||
"use_temporal_tiling": true,
|
||||
"use_parallel_tiling": true
|
||||
},
|
||||
"dit_config": {
|
||||
"prefix": "Hunyuan",
|
||||
"quant_config": null
|
||||
},
|
||||
"text_encoder_precisions": [
|
||||
"fp16",
|
||||
"fp16"
|
||||
],
|
||||
"text_encoder_configs": [
|
||||
{
|
||||
"prefix": "llama",
|
||||
"quant_config": null,
|
||||
"lora_config": null
|
||||
},
|
||||
{
|
||||
"prefix": "clip",
|
||||
"quant_config": null,
|
||||
"lora_config": null,
|
||||
"num_hidden_layers_override": null,
|
||||
"require_post_norm": null
|
||||
}
|
||||
],
|
||||
"mask_strategy_file_path": null,
|
||||
"enable_torch_compile": false
|
||||
}
|
||||
@@ -1,7 +1,6 @@
|
||||
from fastvideo.v1.configs.models.base import ArchConfig, ModelConfig
|
||||
from fastvideo.v1.configs.models.vaes.base import VAEArchConfig, VAEConfig
|
||||
from fastvideo.v1.configs.models.base import ModelConfig
|
||||
from fastvideo.v1.configs.models.dits.base import DiTConfig
|
||||
from fastvideo.v1.configs.models.encoders.base import EncoderConfig
|
||||
from fastvideo.v1.configs.models.vaes.base import VAEConfig
|
||||
|
||||
__all__ = [
|
||||
"ArchConfig", "ModelConfig",
|
||||
"VAEArchConfig", "VAEConfig"
|
||||
]
|
||||
__all__ = ["ModelConfig", "VAEConfig", "DiTConfig", "EncoderConfig"]
|
||||
|
||||
@@ -1,47 +1,80 @@
|
||||
from dataclasses import dataclass, fields
|
||||
from typing import Dict, Any
|
||||
from dataclasses import dataclass, field, fields
|
||||
from typing import Any, Dict
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
# 1. ArchConfig contains all fields from diffuser's/transformer's config.json (i.e. all fields related to the architecture of the model)
|
||||
# 2. ArchConfig should be inherited & overriden by each model arch_config
|
||||
# 2. ArchConfig should be inherited & overridden by each model arch_config
|
||||
# 3. Any field in ArchConfig is fixed upon initialization, and should be hidden away from users
|
||||
@dataclass
|
||||
class ArchConfig:
|
||||
pass
|
||||
|
||||
|
||||
@dataclass
|
||||
class ModelConfig:
|
||||
# Every model config parameter can be categorized into either ArchConfig or everything else
|
||||
# Diffuser/Transformer parameters
|
||||
arch_config: ArchConfig = ArchConfig()
|
||||
arch_config: ArchConfig = field(default_factory=ArchConfig)
|
||||
|
||||
# FastVideo-specific parameters here
|
||||
# i.e. STA, quantization, teacache
|
||||
|
||||
def __getattr__(self, name):
|
||||
# Only called if 'name' is not found in ModelConfig directly
|
||||
if hasattr(self.arch_config, name):
|
||||
return getattr(self.arch_config, name)
|
||||
raise AttributeError(
|
||||
f"'{type(self).__name__}' object has no attribute '{name}'")
|
||||
|
||||
def __getstate__(self):
|
||||
# Return a dictionary of attributes to pickle
|
||||
# Convert to dict and exclude any problematic attributes
|
||||
state = self.__dict__.copy()
|
||||
return state
|
||||
|
||||
def __setstate__(self, state):
|
||||
# Restore instance attributes from the unpickled state
|
||||
self.__dict__.update(state)
|
||||
|
||||
# This should be used only when loading from transformers/diffusers
|
||||
def update_model_arch(
|
||||
self,
|
||||
source_model_dict: Dict[str, Any]
|
||||
) -> None:
|
||||
def update_model_arch(self, source_model_dict: Dict[str, Any]) -> None:
|
||||
# Remove all keys that start with "_"
|
||||
keys_to_remove = [
|
||||
key for key in list(source_model_dict.keys())
|
||||
if str(key).startswith("_")
|
||||
]
|
||||
for key in keys_to_remove:
|
||||
source_model_dict.pop(key)
|
||||
|
||||
arch_config = self.arch_config
|
||||
valid_fields = {f.name for f in fields(arch_config)}
|
||||
|
||||
|
||||
for key, value in source_model_dict.items():
|
||||
if key in valid_fields:
|
||||
setattr(arch_config, key, value)
|
||||
else:
|
||||
raise AttributeError(f"{type(arch_config).__name__} has no field '{key}'")
|
||||
raise AttributeError(
|
||||
f"{type(arch_config).__name__} has no field '{key}'")
|
||||
|
||||
if hasattr(arch_config, "__post_init__"):
|
||||
arch_config.__post_init__()
|
||||
|
||||
def update_model_config(self, source_model_dict: Dict[str, Any]) -> None:
|
||||
assert "arch_config" not in source_model_dict, "Source model config shouldn't contain arch_config."
|
||||
|
||||
def update_model_config(
|
||||
self,
|
||||
source_model_dict: Dict[str, Any]
|
||||
) -> None:
|
||||
assert "arch_config" not in source_model_dict.keys(), "Source model config shouldn't contain arch_config."
|
||||
|
||||
valid_fields = {f.name for f in fields(self)}
|
||||
|
||||
for key, value in source_model_dict.items():
|
||||
if key in valid_fields:
|
||||
setattr(self, key, value)
|
||||
else:
|
||||
print(f"{type(self).__name__} does not contain field '{key}'!")
|
||||
raise AttributeError(f"Invalid field: {key}")
|
||||
logger.warning("%s does not contain field '%s'!",
|
||||
type(self).__name__, key)
|
||||
raise AttributeError(f"Invalid field: {key}")
|
||||
|
||||
if hasattr(self, "__post_init__"):
|
||||
self.__post_init__()
|
||||
|
||||
@@ -0,0 +1,5 @@
|
||||
from fastvideo.v1.configs.models.dits.flux import FluxImageConfig
|
||||
from fastvideo.v1.configs.models.dits.hunyuanvideo import HunyuanVideoConfig
|
||||
from fastvideo.v1.configs.models.dits.wanvideo import WanVideoConfig
|
||||
|
||||
__all__ = ["HunyuanVideoConfig", "WanVideoConfig", "FluxImageConfig"]
|
||||
@@ -0,0 +1,35 @@
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Optional, Tuple
|
||||
|
||||
from fastvideo.v1.configs.models.base import ArchConfig, ModelConfig
|
||||
from fastvideo.v1.configs.quantization import QuantizationConfig
|
||||
from fastvideo.v1.platforms import _Backend
|
||||
|
||||
|
||||
@dataclass
|
||||
class DiTArchConfig(ArchConfig):
|
||||
_fsdp_shard_conditions: list = field(default_factory=list)
|
||||
_compile_conditions: list = field(default_factory=list)
|
||||
_param_names_mapping: dict = field(default_factory=dict)
|
||||
_supported_attention_backends: Tuple[_Backend,
|
||||
...] = (_Backend.SLIDING_TILE_ATTN,
|
||||
_Backend.SAGE_ATTN,
|
||||
_Backend.FLASH_ATTN,
|
||||
_Backend.TORCH_SDPA)
|
||||
|
||||
hidden_size: int = 0
|
||||
num_attention_heads: int = 0
|
||||
num_channels_latents: int = 0
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if not self._compile_conditions:
|
||||
self._compile_conditions = self._fsdp_shard_conditions.copy()
|
||||
|
||||
|
||||
@dataclass
|
||||
class DiTConfig(ModelConfig):
|
||||
arch_config: DiTArchConfig = field(default_factory=DiTArchConfig)
|
||||
|
||||
# FastVideoDiT-specific parameters
|
||||
prefix: str = ""
|
||||
quant_config: Optional[QuantizationConfig] = None
|
||||
@@ -0,0 +1,130 @@
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Optional, Tuple
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
|
||||
|
||||
def is_blocks(n: str, m) -> bool:
|
||||
return "blocks" in n and str.isdigit(n.split(".")[-1])
|
||||
|
||||
|
||||
@dataclass
|
||||
class FluxImageArchConfig(DiTArchConfig):
|
||||
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_blocks])
|
||||
|
||||
_param_names_mapping: dict = field(
|
||||
default_factory=lambda: {
|
||||
# 1. context_embedder to txt_in mapping:
|
||||
r"^context_embedder\.(.*)$":
|
||||
r"txt_in.\1",
|
||||
|
||||
# 2. x_embedder to img_in mapping:
|
||||
r"^x_embedder\.(.*)$":
|
||||
r"img_in.\1",
|
||||
|
||||
# 3. Top-level time_text_embed mappings:
|
||||
r"^time_text_embed\.timestep_embedder\.linear_1\.(.*)$":
|
||||
r"time_in.mlp.fc_in.\1",
|
||||
r"^time_text_embed\.timestep_embedder\.linear_2\.(.*)$":
|
||||
r"time_in.mlp.fc_out.\1",
|
||||
r"^time_text_embed\.guidance_embedder\.linear_1\.(.*)$":
|
||||
r"guidance_in.mlp.fc_in.\1",
|
||||
r"^time_text_embed\.guidance_embedder\.linear_2\.(.*)$":
|
||||
r"guidance_in.mlp.fc_out.\1",
|
||||
r"^time_text_embed\.text_embedder\.linear_1\.(.*)$":
|
||||
r"txt2_in.fc_in.\1",
|
||||
r"^time_text_embed\.text_embedder\.linear_2\.(.*)$":
|
||||
r"txt2_in.fc_out.\1",
|
||||
|
||||
# 4. transformer_blocks mapping:
|
||||
r"^transformer_blocks\.(\d+)\.norm1\.linear\.(.*)$":
|
||||
r"double_blocks.\1.img_mod.linear.\2",
|
||||
r"^transformer_blocks\.(\d+)\.norm1_context\.linear\.(.*)$":
|
||||
r"double_blocks.\1.txt_mod.linear.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn\.norm_q\.(.*)$":
|
||||
r"double_blocks.\1.img_attn_q_norm.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn\.norm_k\.(.*)$":
|
||||
r"double_blocks.\1.img_attn_k_norm.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn\.to_q\.(.*)$":
|
||||
(r"double_blocks.\1.img_attn_qkv.\2", 0, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.to_k\.(.*)$":
|
||||
(r"double_blocks.\1.img_attn_qkv.\2", 1, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.to_v\.(.*)$":
|
||||
(r"double_blocks.\1.img_attn_qkv.\2", 2, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.add_q_proj\.(.*)$":
|
||||
(r"double_blocks.\1.txt_attn_qkv.\2", 0, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.add_k_proj\.(.*)$":
|
||||
(r"double_blocks.\1.txt_attn_qkv.\2", 1, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.add_v_proj\.(.*)$":
|
||||
(r"double_blocks.\1.txt_attn_qkv.\2", 2, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.to_out\.0\.(.*)$":
|
||||
r"double_blocks.\1.img_attn_proj.\2",
|
||||
# Corrected: merge attn.to_add_out into the main projection.
|
||||
r"^transformer_blocks\.(\d+)\.attn\.to_add_out\.(.*)$":
|
||||
r"double_blocks.\1.txt_attn_proj.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn\.norm_added_q\.(.*)$":
|
||||
r"double_blocks.\1.txt_attn_q_norm.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn\.norm_added_k\.(.*)$":
|
||||
r"double_blocks.\1.txt_attn_k_norm.\2",
|
||||
r"^transformer_blocks\.(\d+)\.ff\.net\.0(?:\.proj)?\.(.*)$":
|
||||
r"double_blocks.\1.img_mlp.fc_in.\2",
|
||||
r"^transformer_blocks\.(\d+)\.ff\.net\.2(?:\.proj)?\.(.*)$":
|
||||
r"double_blocks.\1.img_mlp.fc_out.\2",
|
||||
r"^transformer_blocks\.(\d+)\.ff_context\.net\.0(?:\.proj)?\.(.*)$":
|
||||
r"double_blocks.\1.txt_mlp.fc_in.\2",
|
||||
r"^transformer_blocks\.(\d+)\.ff_context\.net\.2(?:\.proj)?\.(.*)$":
|
||||
r"double_blocks.\1.txt_mlp.fc_out.\2",
|
||||
|
||||
# 5. single_transformer_blocks mapping:
|
||||
r"^single_transformer_blocks\.(\d+)\.attn\.norm_q\.(.*)$":
|
||||
r"single_blocks.\1.q_norm.\2",
|
||||
r"^single_transformer_blocks\.(\d+)\.attn\.norm_k\.(.*)$":
|
||||
r"single_blocks.\1.k_norm.\2",
|
||||
r"^single_transformer_blocks\.(\d+)\.attn\.to_q\.(.*)$":
|
||||
(r"single_blocks.\1.linear1.\2", 0, 4),
|
||||
r"^single_transformer_blocks\.(\d+)\.attn\.to_k\.(.*)$":
|
||||
(r"single_blocks.\1.linear1.\2", 1, 4),
|
||||
r"^single_transformer_blocks\.(\d+)\.attn\.to_v\.(.*)$":
|
||||
(r"single_blocks.\1.linear1.\2", 2, 4),
|
||||
r"^single_transformer_blocks\.(\d+)\.proj_mlp\.(.*)$":
|
||||
(r"single_blocks.\1.linear1.\2", 3, 4),
|
||||
# Corrected: map proj_out to modulation.linear rather than a separate proj_out branch.
|
||||
r"^single_transformer_blocks\.(\d+)\.proj_out\.(.*)$":
|
||||
r"single_blocks.\1.linear2.\2",
|
||||
r"^single_transformer_blocks\.(\d+)\.norm\.linear\.(.*)$":
|
||||
r"single_blocks.\1.modulation.linear.\2",
|
||||
|
||||
# 6. Final layers mapping:
|
||||
r"^norm_out\.linear\.(.*)$":
|
||||
r"final_layer.adaLN_modulation.linear.\1",
|
||||
r"^proj_out\.(.*)$":
|
||||
r"final_layer.linear.\1",
|
||||
})
|
||||
|
||||
patch_size: int = 1
|
||||
in_channels: int = 64
|
||||
out_channels: Optional[int] = None
|
||||
num_layers: int = 19
|
||||
num_single_layers: int = 38
|
||||
attention_head_dim: int = 128
|
||||
num_attention_heads: int = 24
|
||||
joint_attention_dim: int = 4096
|
||||
pooled_projection_dim: int = 768
|
||||
guidance_embeds: bool = False
|
||||
axes_dims_rope: Tuple[int, ...] = (16, 56, 56)
|
||||
rope_theta: int = 10000
|
||||
dtype: Optional[torch.dtype] = torch.bfloat16
|
||||
|
||||
def __post_init__(self):
|
||||
self.out_channels = self.out_channels or self.in_channels
|
||||
self.hidden_size = self.num_attention_heads * self.attention_head_dim
|
||||
self.num_channels_latents = self.in_channels // 4
|
||||
|
||||
|
||||
@dataclass
|
||||
class FluxImageConfig(DiTConfig):
|
||||
arch_config: DiTArchConfig = field(default_factory=FluxImageArchConfig)
|
||||
|
||||
prefix: str = "Flux"
|
||||
@@ -0,0 +1,177 @@
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Optional, Tuple
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
|
||||
|
||||
def is_double_block(n: str, m) -> bool:
|
||||
return "double" in n and str.isdigit(n.split(".")[-1])
|
||||
|
||||
|
||||
def is_single_block(n: str, m) -> bool:
|
||||
return "single" in n and str.isdigit(n.split(".")[-1])
|
||||
|
||||
|
||||
def is_refiner_block(n: str, m) -> bool:
|
||||
return "refiner" in n and str.isdigit(n.split(".")[-1])
|
||||
|
||||
|
||||
def is_txt_in(n: str, m) -> bool:
|
||||
return n.split(".")[-1] == "txt_in"
|
||||
|
||||
|
||||
@dataclass
|
||||
class HunyuanVideoArchConfig(DiTArchConfig):
|
||||
_fsdp_shard_conditions: list = field(
|
||||
default_factory=lambda:
|
||||
[is_double_block, is_single_block, is_refiner_block])
|
||||
|
||||
_compile_conditions: list = field(
|
||||
default_factory=lambda: [is_double_block, is_single_block, is_txt_in])
|
||||
|
||||
_param_names_mapping: dict = field(
|
||||
default_factory=lambda: {
|
||||
# 1. context_embedder.time_text_embed submodules (specific rules, applied first):
|
||||
r"^context_embedder\.time_text_embed\.timestep_embedder\.linear_1\.(.*)$":
|
||||
r"txt_in.t_embedder.mlp.fc_in.\1",
|
||||
r"^context_embedder\.time_text_embed\.timestep_embedder\.linear_2\.(.*)$":
|
||||
r"txt_in.t_embedder.mlp.fc_out.\1",
|
||||
r"^context_embedder\.proj_in\.(.*)$":
|
||||
r"txt_in.input_embedder.\1",
|
||||
r"^context_embedder\.time_text_embed\.text_embedder\.linear_1\.(.*)$":
|
||||
r"txt_in.c_embedder.fc_in.\1",
|
||||
r"^context_embedder\.time_text_embed\.text_embedder\.linear_2\.(.*)$":
|
||||
r"txt_in.c_embedder.fc_out.\1",
|
||||
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.norm1\.(.*)$":
|
||||
r"txt_in.refiner_blocks.\1.norm1.\2",
|
||||
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.norm2\.(.*)$":
|
||||
r"txt_in.refiner_blocks.\1.norm2.\2",
|
||||
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_q\.(.*)$":
|
||||
(r"txt_in.refiner_blocks.\1.self_attn_qkv.\2", 0, 3),
|
||||
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_k\.(.*)$":
|
||||
(r"txt_in.refiner_blocks.\1.self_attn_qkv.\2", 1, 3),
|
||||
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_v\.(.*)$":
|
||||
(r"txt_in.refiner_blocks.\1.self_attn_qkv.\2", 2, 3),
|
||||
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_out\.0\.(.*)$":
|
||||
r"txt_in.refiner_blocks.\1.self_attn_proj.\2",
|
||||
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.ff\.net\.0(?:\.proj)?\.(.*)$":
|
||||
r"txt_in.refiner_blocks.\1.mlp.fc_in.\2",
|
||||
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.ff\.net\.2(?:\.proj)?\.(.*)$":
|
||||
r"txt_in.refiner_blocks.\1.mlp.fc_out.\2",
|
||||
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.norm_out\.linear\.(.*)$":
|
||||
r"txt_in.refiner_blocks.\1.adaLN_modulation.linear.\2",
|
||||
|
||||
# 3. x_embedder mapping:
|
||||
r"^x_embedder\.proj\.(.*)$":
|
||||
r"img_in.proj.\1",
|
||||
|
||||
# 4. Top-level time_text_embed mappings:
|
||||
r"^time_text_embed\.timestep_embedder\.linear_1\.(.*)$":
|
||||
r"time_in.mlp.fc_in.\1",
|
||||
r"^time_text_embed\.timestep_embedder\.linear_2\.(.*)$":
|
||||
r"time_in.mlp.fc_out.\1",
|
||||
r"^time_text_embed\.guidance_embedder\.linear_1\.(.*)$":
|
||||
r"guidance_in.mlp.fc_in.\1",
|
||||
r"^time_text_embed\.guidance_embedder\.linear_2\.(.*)$":
|
||||
r"guidance_in.mlp.fc_out.\1",
|
||||
r"^time_text_embed\.text_embedder\.linear_1\.(.*)$":
|
||||
r"vector_in.fc_in.\1",
|
||||
r"^time_text_embed\.text_embedder\.linear_2\.(.*)$":
|
||||
r"vector_in.fc_out.\1",
|
||||
|
||||
# 5. transformer_blocks mapping:
|
||||
r"^transformer_blocks\.(\d+)\.norm1\.linear\.(.*)$":
|
||||
r"double_blocks.\1.img_mod.linear.\2",
|
||||
r"^transformer_blocks\.(\d+)\.norm1_context\.linear\.(.*)$":
|
||||
r"double_blocks.\1.txt_mod.linear.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn\.norm_q\.(.*)$":
|
||||
r"double_blocks.\1.img_attn_q_norm.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn\.norm_k\.(.*)$":
|
||||
r"double_blocks.\1.img_attn_k_norm.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn\.to_q\.(.*)$":
|
||||
(r"double_blocks.\1.img_attn_qkv.\2", 0, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.to_k\.(.*)$":
|
||||
(r"double_blocks.\1.img_attn_qkv.\2", 1, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.to_v\.(.*)$":
|
||||
(r"double_blocks.\1.img_attn_qkv.\2", 2, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.add_q_proj\.(.*)$":
|
||||
(r"double_blocks.\1.txt_attn_qkv.\2", 0, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.add_k_proj\.(.*)$":
|
||||
(r"double_blocks.\1.txt_attn_qkv.\2", 1, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.add_v_proj\.(.*)$":
|
||||
(r"double_blocks.\1.txt_attn_qkv.\2", 2, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.to_out\.0\.(.*)$":
|
||||
r"double_blocks.\1.img_attn_proj.\2",
|
||||
# Corrected: merge attn.to_add_out into the main projection.
|
||||
r"^transformer_blocks\.(\d+)\.attn\.to_add_out\.(.*)$":
|
||||
r"double_blocks.\1.txt_attn_proj.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn\.norm_added_q\.(.*)$":
|
||||
r"double_blocks.\1.txt_attn_q_norm.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn\.norm_added_k\.(.*)$":
|
||||
r"double_blocks.\1.txt_attn_k_norm.\2",
|
||||
r"^transformer_blocks\.(\d+)\.ff\.net\.0(?:\.proj)?\.(.*)$":
|
||||
r"double_blocks.\1.img_mlp.fc_in.\2",
|
||||
r"^transformer_blocks\.(\d+)\.ff\.net\.2(?:\.proj)?\.(.*)$":
|
||||
r"double_blocks.\1.img_mlp.fc_out.\2",
|
||||
r"^transformer_blocks\.(\d+)\.ff_context\.net\.0(?:\.proj)?\.(.*)$":
|
||||
r"double_blocks.\1.txt_mlp.fc_in.\2",
|
||||
r"^transformer_blocks\.(\d+)\.ff_context\.net\.2(?:\.proj)?\.(.*)$":
|
||||
r"double_blocks.\1.txt_mlp.fc_out.\2",
|
||||
|
||||
# 6. single_transformer_blocks mapping:
|
||||
r"^single_transformer_blocks\.(\d+)\.attn\.norm_q\.(.*)$":
|
||||
r"single_blocks.\1.q_norm.\2",
|
||||
r"^single_transformer_blocks\.(\d+)\.attn\.norm_k\.(.*)$":
|
||||
r"single_blocks.\1.k_norm.\2",
|
||||
r"^single_transformer_blocks\.(\d+)\.attn\.to_q\.(.*)$":
|
||||
(r"single_blocks.\1.linear1.\2", 0, 4),
|
||||
r"^single_transformer_blocks\.(\d+)\.attn\.to_k\.(.*)$":
|
||||
(r"single_blocks.\1.linear1.\2", 1, 4),
|
||||
r"^single_transformer_blocks\.(\d+)\.attn\.to_v\.(.*)$":
|
||||
(r"single_blocks.\1.linear1.\2", 2, 4),
|
||||
r"^single_transformer_blocks\.(\d+)\.proj_mlp\.(.*)$":
|
||||
(r"single_blocks.\1.linear1.\2", 3, 4),
|
||||
# Corrected: map proj_out to modulation.linear rather than a separate proj_out branch.
|
||||
r"^single_transformer_blocks\.(\d+)\.proj_out\.(.*)$":
|
||||
r"single_blocks.\1.linear2.\2",
|
||||
r"^single_transformer_blocks\.(\d+)\.norm\.linear\.(.*)$":
|
||||
r"single_blocks.\1.modulation.linear.\2",
|
||||
|
||||
# 7. Final layers mapping:
|
||||
r"^norm_out\.linear\.(.*)$":
|
||||
r"final_layer.adaLN_modulation.linear.\1",
|
||||
r"^proj_out\.(.*)$":
|
||||
r"final_layer.linear.\1",
|
||||
})
|
||||
|
||||
patch_size: int = 2
|
||||
patch_size_t: int = 1
|
||||
in_channels: int = 16
|
||||
out_channels: int = 16
|
||||
num_attention_heads: int = 24
|
||||
attention_head_dim: int = 128
|
||||
mlp_ratio: float = 4.0
|
||||
num_layers: int = 20
|
||||
num_single_layers: int = 40
|
||||
num_refiner_layers: int = 2
|
||||
rope_axes_dim: Tuple[int, int, int] = (16, 56, 56)
|
||||
guidance_embeds: bool = False
|
||||
dtype: Optional[torch.dtype] = None
|
||||
text_embed_dim: int = 4096
|
||||
pooled_projection_dim: int = 768
|
||||
rope_theta: int = 256
|
||||
qk_norm: str = "rms_norm"
|
||||
|
||||
def __post_init__(self):
|
||||
super().__post_init__()
|
||||
self.hidden_size: int = self.attention_head_dim * self.num_attention_heads
|
||||
self.num_channels_latents: int = self.in_channels
|
||||
|
||||
|
||||
@dataclass
|
||||
class HunyuanVideoConfig(DiTConfig):
|
||||
arch_config: DiTArchConfig = field(default_factory=HunyuanVideoArchConfig)
|
||||
|
||||
prefix: str = "Hunyuan"
|
||||
@@ -0,0 +1,83 @@
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Optional, Tuple
|
||||
|
||||
from fastvideo.v1.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
|
||||
|
||||
def is_blocks(n: str, m) -> bool:
|
||||
return "blocks" in n and str.isdigit(n.split(".")[-1])
|
||||
|
||||
|
||||
@dataclass
|
||||
class WanVideoArchConfig(DiTArchConfig):
|
||||
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_blocks])
|
||||
|
||||
_param_names_mapping: dict = field(
|
||||
default_factory=lambda: {
|
||||
r"^patch_embedding\.(.*)$":
|
||||
r"patch_embedding.proj.\1",
|
||||
r"^condition_embedder\.text_embedder\.linear_1\.(.*)$":
|
||||
r"condition_embedder.text_embedder.fc_in.\1",
|
||||
r"^condition_embedder\.text_embedder\.linear_2\.(.*)$":
|
||||
r"condition_embedder.text_embedder.fc_out.\1",
|
||||
r"^condition_embedder\.time_embedder\.linear_1\.(.*)$":
|
||||
r"condition_embedder.time_embedder.mlp.fc_in.\1",
|
||||
r"^condition_embedder\.time_embedder\.linear_2\.(.*)$":
|
||||
r"condition_embedder.time_embedder.mlp.fc_out.\1",
|
||||
r"^condition_embedder\.time_proj\.(.*)$":
|
||||
r"condition_embedder.time_modulation.linear.\1",
|
||||
r"^condition_embedder\.image_embedder\.ff\.net\.0\.proj\.(.*)$":
|
||||
r"condition_embedder.image_embedder.ff.fc_in.\1",
|
||||
r"^condition_embedder\.image_embedder\.ff\.net\.2\.(.*)$":
|
||||
r"condition_embedder.image_embedder.ff.fc_out.\1",
|
||||
r"^blocks\.(\d+)\.attn1\.to_q\.(.*)$":
|
||||
r"blocks.\1.to_q.\2",
|
||||
r"^blocks\.(\d+)\.attn1\.to_k\.(.*)$":
|
||||
r"blocks.\1.to_k.\2",
|
||||
r"^blocks\.(\d+)\.attn1\.to_v\.(.*)$":
|
||||
r"blocks.\1.to_v.\2",
|
||||
r"^blocks\.(\d+)\.attn1\.to_out\.0\.(.*)$":
|
||||
r"blocks.\1.to_out.\2",
|
||||
r"^blocks\.(\d+)\.attn1\.norm_q\.(.*)$":
|
||||
r"blocks.\1.norm_q.\2",
|
||||
r"^blocks\.(\d+)\.attn1\.norm_k\.(.*)$":
|
||||
r"blocks.\1.norm_k.\2",
|
||||
r"^blocks\.(\d+)\.attn2\.to_out\.0\.(.*)$":
|
||||
r"blocks.\1.attn2.to_out.\2",
|
||||
r"^blocks\.(\d+)\.ffn\.net\.0\.proj\.(.*)$":
|
||||
r"blocks.\1.ffn.fc_in.\2",
|
||||
r"^blocks\.(\d+)\.ffn\.net\.2\.(.*)$":
|
||||
r"blocks.\1.ffn.fc_out.\2",
|
||||
r"blocks\.(\d+)\.norm2\.(.*)$":
|
||||
r"blocks.\1.self_attn_residual_norm.norm.\2",
|
||||
})
|
||||
|
||||
patch_size: Tuple[int, int, int] = (1, 2, 2)
|
||||
text_len = 512
|
||||
num_attention_heads: int = 40
|
||||
attention_head_dim: int = 128
|
||||
in_channels: int = 16
|
||||
out_channels: int = 16
|
||||
text_dim: int = 4096
|
||||
freq_dim: int = 256
|
||||
ffn_dim: int = 13824
|
||||
num_layers: int = 40
|
||||
cross_attn_norm: bool = True
|
||||
qk_norm: str = "rms_norm_across_heads"
|
||||
eps: float = 1e-6
|
||||
image_dim: Optional[int] = None
|
||||
added_kv_proj_dim: Optional[int] = None
|
||||
rope_max_seq_len: int = 1024
|
||||
|
||||
def __post_init__(self):
|
||||
super().__post_init__()
|
||||
self.out_channels = self.out_channels or self.in_channels
|
||||
self.hidden_size = self.num_attention_heads * self.attention_head_dim
|
||||
self.num_channels_latents = self.in_channels if self.added_kv_proj_dim is None else self.out_channels
|
||||
|
||||
|
||||
@dataclass
|
||||
class WanVideoConfig(DiTConfig):
|
||||
arch_config: DiTArchConfig = field(default_factory=WanVideoArchConfig)
|
||||
|
||||
prefix: str = "Wan"
|
||||
@@ -0,0 +1,14 @@
|
||||
from fastvideo.v1.configs.models.encoders.base import (BaseEncoderOutput,
|
||||
EncoderConfig,
|
||||
ImageEncoderConfig,
|
||||
TextEncoderConfig)
|
||||
from fastvideo.v1.configs.models.encoders.clip import (CLIPTextConfig,
|
||||
CLIPVisionConfig)
|
||||
from fastvideo.v1.configs.models.encoders.llama import LlamaConfig
|
||||
from fastvideo.v1.configs.models.encoders.t5 import T5Config
|
||||
|
||||
__all__ = [
|
||||
"EncoderConfig", "TextEncoderConfig", "ImageEncoderConfig",
|
||||
"BaseEncoderOutput", "CLIPTextConfig", "CLIPVisionConfig", "LlamaConfig",
|
||||
"T5Config"
|
||||
]
|
||||
@@ -0,0 +1,75 @@
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.configs.models.base import ArchConfig, ModelConfig
|
||||
from fastvideo.v1.configs.quantization import QuantizationConfig
|
||||
from fastvideo.v1.platforms import _Backend
|
||||
|
||||
|
||||
@dataclass
|
||||
class EncoderArchConfig(ArchConfig):
|
||||
architectures: List[str] = field(default_factory=lambda: [])
|
||||
_supported_attention_backends: Tuple[_Backend, ...] = (_Backend.FLASH_ATTN,
|
||||
_Backend.TORCH_SDPA)
|
||||
output_hidden_states: bool = False
|
||||
use_return_dict: bool = True
|
||||
|
||||
|
||||
@dataclass
|
||||
class TextEncoderArchConfig(EncoderArchConfig):
|
||||
vocab_size: int = 0
|
||||
hidden_size: int = 0
|
||||
num_hidden_layers: int = 0
|
||||
num_attention_heads: int = 0
|
||||
pad_token_id: int = 0
|
||||
eos_token_id: int = 0
|
||||
text_len: int = 0
|
||||
hidden_state_skip_layer: int = 0
|
||||
decoder_start_token_id: int = 0
|
||||
output_past: bool = True
|
||||
scalable_attention: bool = True
|
||||
tie_word_embeddings: bool = False
|
||||
|
||||
tokenizer_kwargs: Dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self.tokenizer_kwargs = {
|
||||
"truncation": True,
|
||||
"max_length": self.text_len,
|
||||
"return_tensors": "pt",
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class ImageEncoderArchConfig(EncoderArchConfig):
|
||||
pass
|
||||
|
||||
|
||||
@dataclass
|
||||
class BaseEncoderOutput:
|
||||
last_hidden_state: Optional[torch.FloatTensor] = None
|
||||
pooler_output: Optional[torch.FloatTensor] = None
|
||||
hidden_states: Optional[Tuple[torch.FloatTensor, ...]] = None
|
||||
attentions: Optional[Tuple[torch.FloatTensor, ...]] = None
|
||||
attention_mask: Optional[torch.Tensor] = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class EncoderConfig(ModelConfig):
|
||||
arch_config: ArchConfig = field(default_factory=EncoderArchConfig)
|
||||
|
||||
prefix: str = ""
|
||||
quant_config: Optional[QuantizationConfig] = None
|
||||
lora_config: Optional[Any] = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class TextEncoderConfig(EncoderConfig):
|
||||
arch_config: ArchConfig = field(default_factory=TextEncoderArchConfig)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ImageEncoderConfig(EncoderConfig):
|
||||
arch_config: ArchConfig = field(default_factory=ImageEncoderArchConfig)
|
||||
@@ -0,0 +1,66 @@
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Optional
|
||||
|
||||
from fastvideo.v1.configs.models.encoders.base import (ImageEncoderArchConfig,
|
||||
ImageEncoderConfig,
|
||||
TextEncoderArchConfig,
|
||||
TextEncoderConfig)
|
||||
|
||||
|
||||
@dataclass
|
||||
class CLIPTextArchConfig(TextEncoderArchConfig):
|
||||
vocab_size: int = 49408
|
||||
hidden_size: int = 512
|
||||
intermediate_size: int = 2048
|
||||
projection_dim: int = 512
|
||||
num_hidden_layers: int = 12
|
||||
num_attention_heads: int = 8
|
||||
max_position_embeddings: int = 77
|
||||
hidden_act: str = "quick_gelu"
|
||||
layer_norm_eps: float = 1e-5
|
||||
dropout: float = 0.0
|
||||
attention_dropout: float = 0.0
|
||||
initializer_range: float = 0.02
|
||||
initializer_factor: float = 1.0
|
||||
pad_token_id: int = 1
|
||||
bos_token_id: int = 49406
|
||||
eos_token_id: int = 49407
|
||||
text_len: int = 77
|
||||
|
||||
|
||||
@dataclass
|
||||
class CLIPVisionArchConfig(ImageEncoderArchConfig):
|
||||
hidden_size: int = 768
|
||||
intermediate_size: int = 3072
|
||||
projection_dim: int = 512
|
||||
num_hidden_layers: int = 12
|
||||
num_attention_heads: int = 12
|
||||
num_channels: int = 3
|
||||
image_size: int = 224
|
||||
patch_size: int = 32
|
||||
hidden_act: str = "quick_gelu"
|
||||
layer_norm_eps: float = 1e-5
|
||||
dropout: float = 0.0
|
||||
attention_dropout: float = 0.0
|
||||
initializer_range: float = 0.02
|
||||
initializer_factor: float = 1.0
|
||||
|
||||
|
||||
@dataclass
|
||||
class CLIPTextConfig(TextEncoderConfig):
|
||||
arch_config: TextEncoderArchConfig = field(
|
||||
default_factory=CLIPTextArchConfig)
|
||||
|
||||
num_hidden_layers_override: Optional[int] = None
|
||||
require_post_norm: Optional[bool] = None
|
||||
prefix: str = "clip"
|
||||
|
||||
|
||||
@dataclass
|
||||
class CLIPVisionConfig(ImageEncoderConfig):
|
||||
arch_config: ImageEncoderArchConfig = field(
|
||||
default_factory=CLIPVisionArchConfig)
|
||||
|
||||
num_hidden_layers_override: Optional[int] = None
|
||||
require_post_norm: Optional[bool] = None
|
||||
prefix: str = "clip"
|
||||
@@ -0,0 +1,40 @@
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Optional
|
||||
|
||||
from fastvideo.v1.configs.models.encoders.base import (TextEncoderArchConfig,
|
||||
TextEncoderConfig)
|
||||
|
||||
|
||||
@dataclass
|
||||
class LlamaArchConfig(TextEncoderArchConfig):
|
||||
vocab_size: int = 32000
|
||||
hidden_size: int = 4096
|
||||
intermediate_size: int = 11008
|
||||
num_hidden_layers: int = 32
|
||||
num_attention_heads: int = 32
|
||||
num_key_value_heads: Optional[int] = None
|
||||
hidden_act: str = "silu"
|
||||
max_position_embeddings: int = 2048
|
||||
initializer_range: float = 0.02
|
||||
rms_norm_eps: float = 1e-6
|
||||
use_cache: bool = True
|
||||
pad_token_id: int = 0
|
||||
bos_token_id: int = 1
|
||||
eos_token_id: int = 2
|
||||
pretraining_tp: int = 1
|
||||
tie_word_embeddings: bool = False
|
||||
rope_theta: float = 10000.0
|
||||
rope_scaling: Optional[float] = None
|
||||
attention_bias: bool = False
|
||||
attention_dropout: float = 0.0
|
||||
mlp_bias: bool = False
|
||||
head_dim: Optional[int] = None
|
||||
hidden_state_skip_layer: int = 2
|
||||
text_len: int = 256
|
||||
|
||||
|
||||
@dataclass
|
||||
class LlamaConfig(TextEncoderConfig):
|
||||
arch_config: TextEncoderArchConfig = field(default_factory=LlamaArchConfig)
|
||||
|
||||
prefix: str = "llama"
|
||||
@@ -0,0 +1,55 @@
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Optional
|
||||
|
||||
from fastvideo.v1.configs.models.encoders.base import (TextEncoderArchConfig,
|
||||
TextEncoderConfig)
|
||||
|
||||
|
||||
@dataclass
|
||||
class T5ArchConfig(TextEncoderArchConfig):
|
||||
vocab_size: int = 32128
|
||||
d_model: int = 512
|
||||
d_kv: int = 64
|
||||
d_ff: int = 2048
|
||||
num_layers: int = 6
|
||||
num_decoder_layers: Optional[int] = None
|
||||
num_heads: int = 8
|
||||
relative_attention_num_buckets: int = 32
|
||||
relative_attention_max_distance: int = 128
|
||||
dropout_rate: float = 0.1
|
||||
layer_norm_epsilon: float = 1e-6
|
||||
initializer_factor: float = 1.0
|
||||
feed_forward_proj: str = "relu"
|
||||
dense_act_fn: str = ""
|
||||
is_gated_act: bool = False
|
||||
is_encoder_decoder: bool = True
|
||||
use_cache: bool = True
|
||||
pad_token_id: int = 0
|
||||
eos_token_id: int = 1
|
||||
classifier_dropout: float = 0.0
|
||||
text_len: int = 512
|
||||
|
||||
# Referenced from https://github.com/huggingface/transformers/blob/main/src/transformers/models/t5/configuration_t5.py
|
||||
def __post_init__(self):
|
||||
super().__post_init__()
|
||||
act_info = self.feed_forward_proj.split("-")
|
||||
self.dense_act_fn: str = act_info[-1]
|
||||
self.is_gated_act: bool = act_info[0] == "gated"
|
||||
if self.feed_forward_proj == "gated-gelu":
|
||||
self.dense_act_fn = "gelu_new"
|
||||
|
||||
self.tokenizer_kwargs = {
|
||||
"padding": "max_length",
|
||||
"truncation": True,
|
||||
"max_length": self.text_len,
|
||||
"add_special_tokens": True,
|
||||
"return_attention_mask": True,
|
||||
"return_tensors": "pt",
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class T5Config(TextEncoderConfig):
|
||||
arch_config: TextEncoderArchConfig = field(default_factory=T5ArchConfig)
|
||||
|
||||
prefix: str = "t5"
|
||||
@@ -1,7 +1,9 @@
|
||||
from fastvideo.v1.configs.models.vaes.hunyuanvae import HunyuanVAEConfig, HunyuanVAEArchConfig
|
||||
from fastvideo.v1.configs.models.vaes.wanvae import WanVAEConfig, WanVAEArchConfig
|
||||
from fastvideo.v1.configs.models.vaes.hunyuanvae import HunyuanVAEConfig
|
||||
from fastvideo.v1.configs.models.vaes.image_vae import ImageVAEConfig
|
||||
from fastvideo.v1.configs.models.vaes.wanvae import WanVAEConfig
|
||||
|
||||
__all__ = [
|
||||
"HunyuanVAEConfig", "HunyuanVAEArchConfig",
|
||||
"WanVAEConfig", "WanVAEArchConfig"
|
||||
]
|
||||
"HunyuanVAEConfig",
|
||||
"WanVAEConfig",
|
||||
"ImageVAEConfig",
|
||||
]
|
||||
|
||||
@@ -1,9 +1,10 @@
|
||||
from dataclasses import dataclass
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Union
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.configs.models import ArchConfig, ModelConfig
|
||||
from fastvideo.v1.configs.models.base import ArchConfig, ModelConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class VAEArchConfig(ArchConfig):
|
||||
@@ -12,9 +13,10 @@ class VAEArchConfig(ArchConfig):
|
||||
temporal_compression_ratio: int = 4
|
||||
spatial_compression_ratio: int = 8
|
||||
|
||||
|
||||
@dataclass
|
||||
class VAEConfig(ModelConfig):
|
||||
arch_config: VAEArchConfig = VAEArchConfig()
|
||||
arch_config: VAEArchConfig = field(default_factory=VAEArchConfig)
|
||||
|
||||
# FastVideoVAE-specific parameters
|
||||
load_encoder: bool = True
|
||||
|
||||
@@ -1,7 +1,8 @@
|
||||
from dataclasses import dataclass
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Tuple
|
||||
|
||||
from fastvideo.v1.configs.models.vaes.base import VAEConfig, VAEArchConfig
|
||||
from fastvideo.v1.configs.models.vaes.base import VAEArchConfig, VAEConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class HunyuanVAEArchConfig(VAEArchConfig):
|
||||
@@ -30,8 +31,10 @@ class HunyuanVAEArchConfig(VAEArchConfig):
|
||||
mid_block_add_attention: bool = True
|
||||
|
||||
def __post_init__(self):
|
||||
self.spatial_compression_ratio: int = 2**(len(self.block_out_channels)-1)
|
||||
self.spatial_compression_ratio: int = 2**(len(self.block_out_channels) -
|
||||
1)
|
||||
|
||||
|
||||
@dataclass
|
||||
class HunyuanVAEConfig(VAEConfig):
|
||||
arch_config: VAEArchConfig = HunyuanVAEArchConfig()
|
||||
arch_config: VAEArchConfig = field(default_factory=HunyuanVAEArchConfig)
|
||||
|
||||
@@ -0,0 +1,41 @@
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Optional, Tuple
|
||||
|
||||
from fastvideo.v1.configs.models.vaes.base import VAEArchConfig, VAEConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class ImageVAEArchConfig(VAEArchConfig):
|
||||
in_channels: int = 3
|
||||
out_channels: int = 3
|
||||
down_block_types: Tuple[str] = ("DownEncoderBlock2D", )
|
||||
up_block_types: Tuple[str] = ("UpDecoderBlock2D", )
|
||||
block_out_channels: Tuple[int] = (64, )
|
||||
layers_per_block: int = 1
|
||||
act_fn: str = "silu"
|
||||
latent_channels: int = 4
|
||||
norm_num_groups: int = 32
|
||||
sample_size: int = 32
|
||||
scaling_factor: float = 0.18215
|
||||
shift_factor: Optional[float] = None
|
||||
latents_mean: Optional[Tuple[float]] = None
|
||||
latents_std: Optional[Tuple[float]] = None
|
||||
force_upcast: float = True
|
||||
use_quant_conv: bool = True
|
||||
use_post_quant_conv: bool = True
|
||||
mid_block_add_attention: bool = True
|
||||
|
||||
def __post_init__(self):
|
||||
self.spatial_compression_ratio: int = 2**(len(self.block_out_channels) -
|
||||
1)
|
||||
self.temporal_compression_ratio = 1
|
||||
|
||||
|
||||
@dataclass
|
||||
class ImageVAEConfig(VAEConfig):
|
||||
arch_config: VAEArchConfig = field(default_factory=ImageVAEArchConfig)
|
||||
|
||||
# overrides VAEConfig
|
||||
use_tiling: bool = False
|
||||
use_temporal_tiling: bool = False
|
||||
use_parallel_tiling: bool = False
|
||||
@@ -1,9 +1,10 @@
|
||||
from dataclasses import dataclass
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Tuple
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.configs.models.vaes.base import VAEConfig, VAEArchConfig
|
||||
from fastvideo.v1.configs.models.vaes.base import VAEArchConfig, VAEConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class WanVAEArchConfig(VAEArchConfig):
|
||||
@@ -52,16 +53,17 @@ class WanVAEArchConfig(VAEArchConfig):
|
||||
)
|
||||
temporal_compression_ratio = 4
|
||||
spatial_compression_ratio = 8
|
||||
|
||||
|
||||
def __post_init__(self):
|
||||
self.scaling_factor: torch.tensor = 1.0 / torch.tensor(self.latents_std).view(
|
||||
1, self.z_dim, 1, 1, 1)
|
||||
self.scaling_factor: torch.tensor = 1.0 / torch.tensor(
|
||||
self.latents_std).view(1, self.z_dim, 1, 1, 1)
|
||||
self.shift_factor: torch.tensor = torch.tensor(self.latents_mean).view(
|
||||
1, self.z_dim, 1, 1, 1)
|
||||
|
||||
|
||||
@dataclass
|
||||
class WanVAEConfig(VAEConfig):
|
||||
arch_config: VAEArchConfig = WanVAEArchConfig()
|
||||
arch_config: VAEArchConfig = field(default_factory=WanVAEArchConfig)
|
||||
use_feature_cache: bool = True
|
||||
|
||||
use_tiling: bool = False
|
||||
@@ -69,4 +71,5 @@ class WanVAEConfig(VAEConfig):
|
||||
use_parallel_tiling: bool = False
|
||||
|
||||
def __post_init__(self):
|
||||
self.blend_num_frames = (self.tile_sample_min_num_frames - self.tile_sample_stride_num_frames) * 2
|
||||
self.blend_num_frames = (self.tile_sample_min_num_frames -
|
||||
self.tile_sample_stride_num_frames) * 2
|
||||
|
||||
@@ -1,9 +1,16 @@
|
||||
from fastvideo.v1.configs.pipelines.hunyuan import HunyuanConfig, FastHunyuanConfig
|
||||
from fastvideo.v1.configs.pipelines.wan import WanT2V480PConfig, WanI2V480PConfig
|
||||
from fastvideo.v1.configs.pipelines.base import BaseConfig, SlidingTileAttnConfig
|
||||
from fastvideo.v1.configs.pipelines.registry import get_pipeline_config_cls_for_name
|
||||
from fastvideo.v1.configs.pipelines.base import (PipelineConfig,
|
||||
SlidingTileAttnConfig)
|
||||
from fastvideo.v1.configs.pipelines.hunyuan import (FastHunyuanConfig,
|
||||
HunyuanConfig)
|
||||
from fastvideo.v1.configs.pipelines.registry import (
|
||||
get_pipeline_config_cls_for_name)
|
||||
from fastvideo.v1.configs.pipelines.wan import (WanI2V480PConfig,
|
||||
WanI2V720PConfig,
|
||||
WanT2V480PConfig,
|
||||
WanT2V720PConfig)
|
||||
|
||||
__all__ = [
|
||||
"HunyuanConfig", "FastHunyuanConfig", "BaseConfig", "SlidingTileAttnConfig",
|
||||
"WanT2V480PConfig", "WanI2V480PConfig", "get_pipeline_config_cls_for_name"
|
||||
]
|
||||
"HunyuanConfig", "FastHunyuanConfig", "PipelineConfig",
|
||||
"SlidingTileAttnConfig", "WanT2V480PConfig", "WanI2V480PConfig",
|
||||
"WanT2V720PConfig", "WanI2V720PConfig", "get_pipeline_config_cls_for_name"
|
||||
]
|
||||
|
||||
@@ -1,25 +1,30 @@
|
||||
from dataclasses import dataclass, asdict, fields
|
||||
from typing import Optional, Dict, Any
|
||||
import json
|
||||
from dataclasses import asdict, dataclass, field, fields
|
||||
from typing import Any, Callable, Dict, Optional, Tuple, cast
|
||||
|
||||
from fastvideo.v1.configs.models import ModelConfig, VAEConfig
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.configs.models import (DiTConfig, EncoderConfig, ModelConfig,
|
||||
VAEConfig)
|
||||
from fastvideo.v1.configs.models.encoders import BaseEncoderOutput
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.utils import shallow_asdict
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def preprocess_text(prompt: str) -> str:
|
||||
return prompt
|
||||
|
||||
|
||||
def postprocess_text(output: BaseEncoderOutput) -> torch.tensor:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
@dataclass
|
||||
class BaseConfig:
|
||||
class PipelineConfig:
|
||||
"""Base configuration for all pipeline architectures."""
|
||||
|
||||
# Video parameters
|
||||
height: int = 720
|
||||
width: int = 1280
|
||||
num_frames: int = 125
|
||||
fps: int = 24
|
||||
|
||||
# Video generation parameters
|
||||
num_inference_steps: int = 50
|
||||
guidance_scale: float = 1.0
|
||||
seed: int = 1024
|
||||
guidance_rescale: float = 0.0
|
||||
embedded_cfg_scale: float = 6.0
|
||||
flow_shift: Optional[float] = None
|
||||
use_cpu_offload: bool = False
|
||||
@@ -32,47 +37,78 @@ class BaseConfig:
|
||||
vae_precision: str = "fp16"
|
||||
vae_tiling: bool = True
|
||||
vae_sp: bool = True
|
||||
# vae_scale_factor: Optional[int] = None # Deprecated
|
||||
vae_config: VAEConfig = VAEConfig()
|
||||
vae_config: VAEConfig = field(default_factory=VAEConfig)
|
||||
|
||||
# DiT configuration
|
||||
num_channels_latents: Optional[int] = None # Deprecated
|
||||
|
||||
# Image encoder configuration
|
||||
image_encoder_precision: str = "fp32"
|
||||
dit_config: DiTConfig = field(default_factory=DiTConfig)
|
||||
|
||||
# Text encoder configuration
|
||||
text_encoder_precision: str = "fp16"
|
||||
text_len: int = -1 # Deprecated
|
||||
hidden_state_skip_layer: int = 0 # Deprecated
|
||||
text_encoder_precisions: Tuple[str, ...] = field(
|
||||
default_factory=lambda: ("fp16", ))
|
||||
text_encoder_configs: Tuple[EncoderConfig, ...] = field(
|
||||
default_factory=lambda: (EncoderConfig(), ))
|
||||
preprocess_text_funcs: Tuple[Callable[[str], str], ...] = field(
|
||||
default_factory=lambda: (preprocess_text, ))
|
||||
postprocess_text_funcs: Tuple[Callable[[BaseEncoderOutput], torch.tensor],
|
||||
...] = field(default_factory=lambda:
|
||||
(postprocess_text, ))
|
||||
|
||||
# STA (Spatial-Temporal Attention) parameters
|
||||
mask_strategy_file_path: Optional[str] = None
|
||||
|
||||
# Compilation
|
||||
enable_torch_compile: bool = False
|
||||
|
||||
neg_prompt: Optional[str] = None
|
||||
@classmethod
|
||||
def from_pretrained(cls, model_path: str) -> "PipelineConfig":
|
||||
from fastvideo.v1.configs.pipelines.registry import (
|
||||
get_pipeline_config_cls_for_name)
|
||||
pipeline_config_cls = get_pipeline_config_cls_for_name(model_path)
|
||||
if pipeline_config_cls is not None:
|
||||
pipeline_config = pipeline_config_cls()
|
||||
else:
|
||||
logger.warning(
|
||||
"Couldn't find an optimal sampling param for %s. Using the default sampling param.",
|
||||
model_path)
|
||||
pipeline_config = cls()
|
||||
|
||||
return cast(PipelineConfig, pipeline_config)
|
||||
|
||||
def dump_to_json(self, file_path: str):
|
||||
output_dict = shallow_asdict(self)
|
||||
del_keys = []
|
||||
for key, value in output_dict.items():
|
||||
if isinstance(value, ModelConfig):
|
||||
model_dict = asdict(value)
|
||||
# Model Arch Config should be hidden away from the users
|
||||
model_dict.pop("arch_config")
|
||||
output_dict[key] = model_dict
|
||||
|
||||
elif isinstance(value, tuple) and all(
|
||||
isinstance(v, ModelConfig) for v in value):
|
||||
model_dicts = []
|
||||
for v in value:
|
||||
model_dict = asdict(v)
|
||||
# Model Arch Config should be hidden away from the users
|
||||
model_dict.pop("arch_config")
|
||||
model_dicts.append(model_dict)
|
||||
output_dict[key] = model_dicts
|
||||
elif isinstance(value, tuple) and all(callable(f) for f in value):
|
||||
# Skip dumping functions
|
||||
del_keys.append(key)
|
||||
|
||||
for key in del_keys:
|
||||
output_dict.pop(key, None)
|
||||
|
||||
with open(file_path, "w") as f:
|
||||
json.dump(output_dict, f, indent=2)
|
||||
|
||||
def load_from_json(self, file_path: str):
|
||||
with open(file_path, "r") as f:
|
||||
with open(file_path) as f:
|
||||
input_pipeline_dict = json.load(f)
|
||||
self.update_pipeline_config(input_pipeline_dict)
|
||||
|
||||
def update_pipeline_config(
|
||||
self,
|
||||
source_pipeline_dict: Dict[str, Any]
|
||||
) -> None:
|
||||
def update_pipeline_config(self, source_pipeline_dict: Dict[str,
|
||||
Any]) -> None:
|
||||
for f in fields(self):
|
||||
key = f.name
|
||||
if key in source_pipeline_dict:
|
||||
@@ -82,11 +118,23 @@ class BaseConfig:
|
||||
# If it's a nested ModelConfig, update it recursively
|
||||
if isinstance(current_value, ModelConfig):
|
||||
current_value.update_model_config(new_value)
|
||||
elif isinstance(current_value, tuple) and all(
|
||||
isinstance(v, ModelConfig) for v in current_value):
|
||||
assert len(current_value) == len(
|
||||
new_value
|
||||
), "Users shouldn't delete or add text encoder config objects in your json"
|
||||
for target_config, source_config in zip(
|
||||
current_value, new_value):
|
||||
target_config.update_model_config(source_config)
|
||||
else:
|
||||
setattr(self, key, new_value)
|
||||
|
||||
if hasattr(self, "__post_init__"):
|
||||
self.__post_init__()
|
||||
|
||||
|
||||
@dataclass
|
||||
class SlidingTileAttnConfig(BaseConfig):
|
||||
class SlidingTileAttnConfig(PipelineConfig):
|
||||
"""Configuration for sliding tile attention."""
|
||||
|
||||
# Override any BaseConfig defaults as needed
|
||||
@@ -100,4 +148,4 @@ class SlidingTileAttnConfig(BaseConfig):
|
||||
|
||||
# Additional configuration specific to sliding tile attention
|
||||
pad_to_square: bool = False
|
||||
use_overlap_optimization: bool = True
|
||||
use_overlap_optimization: bool = True
|
||||
|
||||
@@ -0,0 +1,68 @@
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Callable, Tuple
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.configs.models import DiTConfig, EncoderConfig, VAEConfig
|
||||
from fastvideo.v1.configs.models.dits import FluxImageConfig
|
||||
from fastvideo.v1.configs.models.encoders import (BaseEncoderOutput,
|
||||
CLIPTextConfig, T5Config)
|
||||
from fastvideo.v1.configs.models.vaes import ImageVAEConfig
|
||||
from fastvideo.v1.configs.pipelines.base import PipelineConfig
|
||||
|
||||
|
||||
def t5_preprocess_text(prompt: str) -> str:
|
||||
return prompt
|
||||
|
||||
|
||||
def t5_postprocess_text(outputs: BaseEncoderOutput) -> torch.tensor:
|
||||
hidden_state: torch.tensor = outputs.last_hidden_state
|
||||
assert torch.isnan(hidden_state).sum() == 0
|
||||
prompt_embeds_tensor: torch.tensor = torch.stack([
|
||||
torch.cat([u, u.new_zeros(512 - u.size(0), u.size(1))])
|
||||
for u in hidden_state
|
||||
],
|
||||
dim=0)
|
||||
return prompt_embeds_tensor
|
||||
|
||||
|
||||
def clip_preprocess_text(prompt: str) -> str:
|
||||
return prompt
|
||||
|
||||
|
||||
def clip_postprocess_text(outputs: BaseEncoderOutput) -> torch.tensor:
|
||||
pooler_output: torch.tensor = outputs.pooler_output
|
||||
return pooler_output
|
||||
|
||||
|
||||
@dataclass
|
||||
class FluxConfig(PipelineConfig):
|
||||
"""Base configuration for Flux pipeline architecture."""
|
||||
|
||||
# FluxConfig-specific parameters with defaults
|
||||
# DiT
|
||||
dit_config: DiTConfig = field(default_factory=FluxImageConfig)
|
||||
# VAE
|
||||
vae_config: VAEConfig = field(default_factory=ImageVAEConfig)
|
||||
# Denoising stage
|
||||
embedded_cfg_scale: float = 3.5
|
||||
|
||||
# Text encoding stage
|
||||
text_encoder_configs: Tuple[EncoderConfig, ...] = field(
|
||||
default_factory=lambda: (CLIPTextConfig(), T5Config()))
|
||||
preprocess_text_funcs: Tuple[Callable[[str], str], ...] = field(
|
||||
default_factory=lambda: (clip_preprocess_text, t5_preprocess_text))
|
||||
postprocess_text_funcs: Tuple[
|
||||
Callable[[BaseEncoderOutput], torch.tensor],
|
||||
...] = field(default_factory=lambda:
|
||||
(clip_postprocess_text, t5_postprocess_text))
|
||||
|
||||
# Precision for each component
|
||||
precision: str = "bf16"
|
||||
vae_precision: str = "fp16"
|
||||
text_encoder_precisions: Tuple[str, ...] = field(
|
||||
default_factory=lambda: ("bf16", "bf16"))
|
||||
|
||||
def __post_init__(self):
|
||||
self.vae_config.load_encoder = False
|
||||
self.vae_config.load_decoder = True
|
||||
@@ -1,34 +1,88 @@
|
||||
from dataclasses import dataclass
|
||||
from fastvideo.v1.configs.pipelines.base import BaseConfig
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Callable, Tuple, TypedDict
|
||||
|
||||
from fastvideo.v1.configs.models import VAEConfig
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.configs.models import DiTConfig, EncoderConfig, VAEConfig
|
||||
from fastvideo.v1.configs.models.dits import HunyuanVideoConfig
|
||||
from fastvideo.v1.configs.models.encoders import (BaseEncoderOutput,
|
||||
CLIPTextConfig, LlamaConfig)
|
||||
from fastvideo.v1.configs.models.vaes import HunyuanVAEConfig
|
||||
from fastvideo.v1.configs.pipelines.base import PipelineConfig
|
||||
|
||||
PROMPT_TEMPLATE_ENCODE_VIDEO = (
|
||||
"<|start_header_id|>system<|end_header_id|>\n\nDescribe the video by detailing the following aspects: "
|
||||
"1. The main content and theme of the video."
|
||||
"2. The color, shape, size, texture, quantity, text, and spatial relationships of the objects."
|
||||
"3. Actions, events, behaviors temporal relationships, physical movement changes of the objects."
|
||||
"4. background environment, light, style and atmosphere."
|
||||
"5. camera angles, movements, and transitions used in the video:<|eot_id|>"
|
||||
"<|start_header_id|>user<|end_header_id|>\n\n{}<|eot_id|>")
|
||||
|
||||
|
||||
class PromptTemplate(TypedDict):
|
||||
template: str
|
||||
crop_start: int
|
||||
|
||||
|
||||
prompt_template_video: PromptTemplate = {
|
||||
"template": PROMPT_TEMPLATE_ENCODE_VIDEO,
|
||||
"crop_start": 95,
|
||||
}
|
||||
|
||||
|
||||
def llama_preprocess_text(prompt: str) -> str:
|
||||
return prompt_template_video["template"].format(prompt)
|
||||
|
||||
|
||||
def llama_postprocess_text(outputs: BaseEncoderOutput) -> torch.tensor:
|
||||
hidden_state_skip_layer = 2
|
||||
assert outputs.hidden_states is not None
|
||||
hidden_states: Tuple[torch.Tensor, ...] = outputs.hidden_states
|
||||
last_hidden_state: torch.tensor = hidden_states[-(hidden_state_skip_layer +
|
||||
1)]
|
||||
crop_start = prompt_template_video.get("crop_start", -1)
|
||||
last_hidden_state = last_hidden_state[:, crop_start:]
|
||||
return last_hidden_state
|
||||
|
||||
|
||||
def clip_preprocess_text(prompt: str) -> str:
|
||||
return prompt
|
||||
|
||||
|
||||
def clip_postprocess_text(outputs: BaseEncoderOutput) -> torch.tensor:
|
||||
pooler_output: torch.tensor = outputs.pooler_output
|
||||
return pooler_output
|
||||
|
||||
|
||||
@dataclass
|
||||
class HunyuanConfig(BaseConfig):
|
||||
class HunyuanConfig(PipelineConfig):
|
||||
"""Base configuration for HunYuan pipeline architecture."""
|
||||
|
||||
# HunyuanConfig-specific parameters with defaults
|
||||
# DiT
|
||||
dit_config: DiTConfig = field(default_factory=HunyuanVideoConfig)
|
||||
# VAE
|
||||
vae_config: VAEConfig = HunyuanVAEConfig()
|
||||
vae_config: VAEConfig = field(default_factory=HunyuanVAEConfig)
|
||||
# Denoising stage
|
||||
embedded_cfg_scale: int = 6
|
||||
flow_shift: int = 7
|
||||
num_inference_steps: int = 50
|
||||
|
||||
# Text encoding stage
|
||||
hidden_state_skip_layer: int = 2
|
||||
text_len: int = 256
|
||||
text_encoder_configs: Tuple[EncoderConfig, ...] = field(
|
||||
default_factory=lambda: (LlamaConfig(), CLIPTextConfig()))
|
||||
preprocess_text_funcs: Tuple[Callable[[str], str], ...] = field(
|
||||
default_factory=lambda: (llama_preprocess_text, clip_preprocess_text))
|
||||
postprocess_text_funcs: Tuple[
|
||||
Callable[[BaseEncoderOutput], torch.tensor],
|
||||
...] = field(default_factory=lambda:
|
||||
(llama_postprocess_text, clip_postprocess_text))
|
||||
|
||||
# Precision for each component
|
||||
precision: str = "bf16"
|
||||
vae_precision: str = "fp16"
|
||||
text_encoder_precision: str = "fp16"
|
||||
|
||||
# HunyuanConfig-specific added parameters
|
||||
# Secondary text encoder
|
||||
text_encoder_precision_2: str = "fp16"
|
||||
text_len_2: int = 77
|
||||
text_encoder_precisions: Tuple[str, ...] = field(
|
||||
default_factory=lambda: ("fp16", "fp16"))
|
||||
|
||||
def __post_init__(self):
|
||||
self.vae_config.load_encoder = False
|
||||
@@ -40,7 +94,6 @@ class FastHunyuanConfig(HunyuanConfig):
|
||||
"""Configuration specifically optimized for FastHunyuan weights."""
|
||||
|
||||
# Override HunyuanConfig defaults
|
||||
num_inference_steps: int = 6
|
||||
flow_shift: int = 17
|
||||
|
||||
# No need to re-specify guidance_scale or embedded_cfg_scale as they
|
||||
|
||||
@@ -3,10 +3,14 @@
|
||||
import os
|
||||
from typing import Callable, Dict, Optional, Type
|
||||
|
||||
from fastvideo.v1.configs.pipelines.base import BaseConfig
|
||||
from fastvideo.v1.configs.pipelines.hunyuan import HunyuanConfig, FastHunyuanConfig
|
||||
from fastvideo.v1.configs.pipelines.wan import WanT2V480PConfig, WanI2V480PConfig
|
||||
|
||||
from fastvideo.v1.configs.pipelines.base import PipelineConfig
|
||||
from fastvideo.v1.configs.pipelines.flux import FluxConfig
|
||||
from fastvideo.v1.configs.pipelines.hunyuan import (FastHunyuanConfig,
|
||||
HunyuanConfig)
|
||||
from fastvideo.v1.configs.pipelines.wan import (WanI2V480PConfig,
|
||||
WanI2V720PConfig,
|
||||
WanT2V480PConfig,
|
||||
WanT2V720PConfig)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.utils import (maybe_download_model_index,
|
||||
verify_model_config_and_directory)
|
||||
@@ -14,11 +18,15 @@ from fastvideo.v1.utils import (maybe_download_model_index,
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# Registry maps specific model weights to their config classes
|
||||
WEIGHT_CONFIG_REGISTRY: Dict[str, Type[BaseConfig]] = {
|
||||
"FastVideo/FastHunyuan-Diffusers": FastHunyuanConfig,
|
||||
WEIGHT_CONFIG_REGISTRY: Dict[str, Type[PipelineConfig]] = {
|
||||
"FastVideo/FastHunyuan-diffusers": FastHunyuanConfig,
|
||||
"hunyuanvideo-community/HunyuanVideo": HunyuanConfig,
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WanT2V480PConfig,
|
||||
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers": WanI2V480PConfig
|
||||
"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,
|
||||
"black-forest-labs/FLUX.1-dev": FluxConfig,
|
||||
"black-forest-labs/FLUX.1-schnell": FluxConfig,
|
||||
# Add other specific weight variants
|
||||
}
|
||||
|
||||
@@ -27,22 +35,24 @@ PIPELINE_DETECTOR: Dict[str, Callable[[str], bool]] = {
|
||||
"hunyuan": lambda id: "hunyuan" in id.lower(),
|
||||
"wanpipeline": lambda id: "wanpipeline" in id.lower(),
|
||||
"wanimagetovideo": lambda id: "wanimagetovideo" in id.lower(),
|
||||
"flux": lambda id: "flux" in id.lower(),
|
||||
# Add other pipeline architecture detectors
|
||||
}
|
||||
|
||||
# Fallback configs when exact match isn't found but architecture is detected
|
||||
PIPELINE_FALLBACK_CONFIG: Dict[str, Type[BaseConfig]] = {
|
||||
PIPELINE_FALLBACK_CONFIG: Dict[str, Type[PipelineConfig]] = {
|
||||
"hunyuan":
|
||||
HunyuanConfig, # Base Hunyuan config as fallback for any Hunyuan variant
|
||||
"wanpipeline":
|
||||
WanT2V480PConfig, # Base Wan config as fallback for any Wan variant
|
||||
"wanimagetovideo": WanI2V480PConfig,
|
||||
"flux": FluxConfig,
|
||||
# Other fallbacks by architecture
|
||||
}
|
||||
|
||||
|
||||
def get_pipeline_config_cls_for_name(
|
||||
pipeline_name_or_path: str) -> Optional[type[BaseConfig]]:
|
||||
pipeline_name_or_path: str) -> Optional[type[PipelineConfig]]:
|
||||
"""Get the appropriate config class for specific pretrained weights."""
|
||||
|
||||
if os.path.exists(pipeline_name_or_path):
|
||||
@@ -66,7 +76,6 @@ def get_pipeline_config_cls_for_name(
|
||||
|
||||
# If no match, try to use the fallback config
|
||||
fallback_config = None
|
||||
print(pipeline_name)
|
||||
# Try to determine pipeline architecture for fallback
|
||||
for pipeline_type, detector in PIPELINE_DETECTOR.items():
|
||||
if detector(pipeline_name.lower()):
|
||||
|
||||
@@ -1,39 +1,60 @@
|
||||
from dataclasses import dataclass
|
||||
from fastvideo.v1.configs.pipelines.base import BaseConfig
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Callable, Tuple
|
||||
|
||||
from fastvideo.v1.configs.models import VAEConfig
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.configs.models import DiTConfig, EncoderConfig, VAEConfig
|
||||
from fastvideo.v1.configs.models.dits import WanVideoConfig
|
||||
from fastvideo.v1.configs.models.encoders import (BaseEncoderOutput,
|
||||
CLIPVisionConfig, T5Config)
|
||||
from fastvideo.v1.configs.models.vaes import WanVAEConfig
|
||||
from fastvideo.v1.configs.pipelines.base import PipelineConfig
|
||||
|
||||
|
||||
def t5_postprocess_text(outputs: BaseEncoderOutput) -> torch.tensor:
|
||||
mask: torch.tensor = outputs.attention_mask
|
||||
hidden_state: torch.tensor = outputs.last_hidden_state
|
||||
seq_lens = mask.gt(0).sum(dim=1).long()
|
||||
assert torch.isnan(hidden_state).sum() == 0
|
||||
prompt_embeds = [u[:v] for u, v in zip(hidden_state, seq_lens)]
|
||||
prompt_embeds_tensor: torch.tensor = torch.stack([
|
||||
torch.cat([u, u.new_zeros(512 - u.size(0), u.size(1))])
|
||||
for u in prompt_embeds
|
||||
],
|
||||
dim=0)
|
||||
return prompt_embeds_tensor
|
||||
|
||||
|
||||
@dataclass
|
||||
class WanT2V480PConfig(BaseConfig):
|
||||
class WanT2V480PConfig(PipelineConfig):
|
||||
"""Base configuration for Wan T2V 1.3B pipeline architecture."""
|
||||
|
||||
# WanConfig-specific parameters with defaults
|
||||
# DiT
|
||||
dit_config: DiTConfig = field(default_factory=WanVideoConfig)
|
||||
# VAE
|
||||
vae_config: VAEConfig = WanVAEConfig()
|
||||
vae_config: VAEConfig = field(default_factory=WanVAEConfig)
|
||||
vae_tiling: bool = False
|
||||
vae_sp: bool = False
|
||||
|
||||
# Video parameters
|
||||
height: int = 480
|
||||
width: int = 832
|
||||
num_frames: int = 81
|
||||
fps: int = 16
|
||||
use_cpu_offload: bool = True
|
||||
|
||||
# Denoising stage
|
||||
guidance_scale: float = 3.0
|
||||
neg_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"
|
||||
flow_shift: int = 3
|
||||
num_inference_steps: int = 50
|
||||
|
||||
# Text encoding stage
|
||||
text_len: int = 512
|
||||
text_encoder_configs: Tuple[EncoderConfig, ...] = field(
|
||||
default_factory=lambda: (T5Config(), ))
|
||||
postprocess_text_funcs: Tuple[Callable[[BaseEncoderOutput], torch.tensor],
|
||||
...] = field(default_factory=lambda:
|
||||
(t5_postprocess_text, ))
|
||||
|
||||
# Precision for each component
|
||||
precision: str = "bf16"
|
||||
vae_precision: str = "fp16"
|
||||
text_encoder_precision: str = "fp32"
|
||||
text_encoder_precisions: Tuple[str, ...] = field(
|
||||
default_factory=lambda: ("fp32", ))
|
||||
|
||||
# WanConfig-specific added parameters
|
||||
|
||||
@@ -41,18 +62,38 @@ class WanT2V480PConfig(BaseConfig):
|
||||
self.vae_config.load_encoder = False
|
||||
self.vae_config.load_decoder = True
|
||||
|
||||
|
||||
@dataclass
|
||||
class WanT2V720PConfig(WanT2V480PConfig):
|
||||
"""Base configuration for Wan T2V 14B 720P pipeline architecture."""
|
||||
|
||||
# WanConfig-specific parameters with defaults
|
||||
|
||||
# Denoising stage
|
||||
flow_shift: int = 5
|
||||
|
||||
|
||||
@dataclass
|
||||
class WanI2V480PConfig(WanT2V480PConfig):
|
||||
"""Base configuration for Wan I2V 14B 480P pipeline architecture."""
|
||||
|
||||
# WanConfig-specific parameters with defaults
|
||||
# Denoising stage
|
||||
guidance_scale: float = 5.0
|
||||
num_inference_steps: int = 40
|
||||
|
||||
# Precision for each component
|
||||
image_encoder_config: EncoderConfig = field(
|
||||
default_factory=CLIPVisionConfig)
|
||||
image_encoder_precision: str = "fp32"
|
||||
|
||||
def __post_init__(self):
|
||||
self.vae_config.load_encoder = True
|
||||
self.vae_config.load_decoder = True
|
||||
|
||||
|
||||
@dataclass
|
||||
class WanI2V720PConfig(WanI2V480PConfig):
|
||||
"""Base configuration for Wan I2V 14B 720P pipeline architecture."""
|
||||
|
||||
# WanConfig-specific parameters with defaults
|
||||
|
||||
# Denoising stage
|
||||
flow_shift: int = 5
|
||||
|
||||
@@ -0,0 +1,3 @@
|
||||
from fastvideo.v1.configs.quantization.base import QuantizationConfig
|
||||
|
||||
__all__ = ["QuantizationConfig"]
|
||||
@@ -0,0 +1,6 @@
|
||||
from dataclasses import dataclass
|
||||
|
||||
|
||||
@dataclass
|
||||
class QuantizationConfig:
|
||||
pass
|
||||
@@ -0,0 +1,3 @@
|
||||
from fastvideo.v1.configs.sample.base import SamplingParam
|
||||
|
||||
__all__ = ["SamplingParam"]
|
||||
@@ -0,0 +1,84 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict, List, Optional, Union
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class SamplingParam:
|
||||
"""
|
||||
Sampling parameters for video generation.
|
||||
"""
|
||||
# All fields below are copied from ForwardBatch
|
||||
data_type: str = "video"
|
||||
|
||||
# Image inputs
|
||||
image_path: Optional[str] = None
|
||||
|
||||
# Text inputs
|
||||
prompt: Optional[Union[str, List[str]]] = None
|
||||
negative_prompt: Optional[str] = None
|
||||
prompt_path: Optional[str] = None
|
||||
output_path: str = "outputs/"
|
||||
|
||||
# Batch info
|
||||
num_videos_per_prompt: int = 1
|
||||
seed: int = 1024
|
||||
|
||||
# Original dimensions (before VAE scaling)
|
||||
num_frames: int = 125
|
||||
num_frames_round_down: bool = False # Whether to round down num_frames if it's not divisible by num_gpus
|
||||
height: int = 720
|
||||
width: int = 1280
|
||||
fps: int = 24
|
||||
|
||||
# Denoising parameters
|
||||
num_inference_steps: int = 50
|
||||
guidance_scale: float = 1.0
|
||||
guidance_rescale: float = 0.0
|
||||
|
||||
# TeaCache parameters
|
||||
enable_teacache: bool = False
|
||||
|
||||
# Misc
|
||||
save_video: bool = True
|
||||
return_frames: bool = False
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self.data_type = "video" if self.num_frames > 1 else "image"
|
||||
|
||||
def check_sampling_param(self):
|
||||
if self.prompt_path and not self.prompt_path.endswith(".txt"):
|
||||
raise ValueError("prompt_path must be a txt file")
|
||||
|
||||
def update(self, source_dict: Dict[str, Any]) -> None:
|
||||
for key, value in source_dict.items():
|
||||
if hasattr(self, key):
|
||||
setattr(self, key, value)
|
||||
else:
|
||||
logger.exception("%s has no attribute %s",
|
||||
type(self).__name__, key)
|
||||
|
||||
self.__post_init__()
|
||||
|
||||
@classmethod
|
||||
def from_pretrained(cls, model_path: str) -> "SamplingParam":
|
||||
from fastvideo.v1.configs.sample.registry import (
|
||||
get_sampling_param_cls_for_name)
|
||||
sampling_cls = get_sampling_param_cls_for_name(model_path)
|
||||
if sampling_cls is not None:
|
||||
sampling_param: SamplingParam = sampling_cls()
|
||||
else:
|
||||
logger.warning(
|
||||
"Couldn't find an optimal sampling param for %s. Using the default sampling param.",
|
||||
model_path)
|
||||
sampling_param = cls()
|
||||
|
||||
return sampling_param
|
||||
|
||||
|
||||
@dataclass
|
||||
class CacheParams:
|
||||
cache_type: str = "none"
|
||||
@@ -0,0 +1,29 @@
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.v1.configs.sample.base import SamplingParam
|
||||
from fastvideo.v1.configs.sample.teacache import TeaCacheParams
|
||||
|
||||
|
||||
@dataclass
|
||||
class HunyuanSamplingParam(SamplingParam):
|
||||
num_inference_steps: int = 50
|
||||
|
||||
num_frames: int = 125
|
||||
height: int = 720
|
||||
width: int = 1280
|
||||
fps: int = 24
|
||||
|
||||
guidance_scale: float = 1.0
|
||||
|
||||
teacache_params: TeaCacheParams = field(
|
||||
default_factory=lambda: TeaCacheParams(
|
||||
teacache_thresh=0.15,
|
||||
coefficients=[
|
||||
7.33226126e+02, -4.01131952e+02, 6.75869174e+01,
|
||||
-3.14987800e+00, 9.61237896e-02
|
||||
]))
|
||||
|
||||
|
||||
@dataclass
|
||||
class FastHunyuanSamplingParam(HunyuanSamplingParam):
|
||||
num_inference_steps: int = 6
|
||||
@@ -0,0 +1,79 @@
|
||||
import os
|
||||
from typing import Any, Callable, Dict, Optional
|
||||
|
||||
from fastvideo.v1.configs.sample.hunyuan import (FastHunyuanSamplingParam,
|
||||
HunyuanSamplingParam)
|
||||
from fastvideo.v1.configs.sample.wan import (WanI2V_14B_480P_SamplingParam,
|
||||
WanI2V_14B_720P_SamplingParam,
|
||||
WanT2V_1_3B_SamplingParam,
|
||||
WanT2V_14B_SamplingParam)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.utils import (maybe_download_model_index,
|
||||
verify_model_config_and_directory)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
# Registry maps specific model weights to their config classes
|
||||
SAMPLING_PARAM_REGISTRY: Dict[str, Any] = {
|
||||
"FastVideo/FastHunyuan-diffusers": FastHunyuanSamplingParam,
|
||||
"hunyuanvideo-community/HunyuanVideo": HunyuanSamplingParam,
|
||||
"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,
|
||||
# Add other specific weight variants
|
||||
}
|
||||
|
||||
# For determining pipeline type from model ID
|
||||
SAMPLING_PARAM_DETECTOR: Dict[str, Callable[[str], bool]] = {
|
||||
"hunyuan": lambda id: "hunyuan" in id.lower(),
|
||||
"wanpipeline": lambda id: "wanpipeline" in id.lower(),
|
||||
"wanimagetovideo": lambda id: "wanimagetovideo" in id.lower(),
|
||||
# Add other pipeline architecture detectors
|
||||
}
|
||||
|
||||
# Fallback configs when exact match isn't found but architecture is detected
|
||||
SAMPLING_FALLBACK_PARAM: Dict[str, Any] = {
|
||||
"hunyuan":
|
||||
HunyuanSamplingParam, # Base Hunyuan config as fallback for any Hunyuan variant
|
||||
"wanpipeline":
|
||||
WanT2V_1_3B_SamplingParam, # Base Wan config as fallback for any Wan variant
|
||||
"wanimagetovideo": WanI2V_14B_480P_SamplingParam,
|
||||
# Other fallbacks by architecture
|
||||
}
|
||||
|
||||
|
||||
def get_sampling_param_cls_for_name(
|
||||
pipeline_name_or_path: str) -> Optional[Any]:
|
||||
"""Get the appropriate sampling param for specific pretrained weights."""
|
||||
|
||||
if os.path.exists(pipeline_name_or_path):
|
||||
config = verify_model_config_and_directory(pipeline_name_or_path)
|
||||
logger.warning(
|
||||
"FastVideo may not correctly identify the optimal sampling param for this model, as the local directory may have been renamed."
|
||||
)
|
||||
else:
|
||||
config = maybe_download_model_index(pipeline_name_or_path)
|
||||
|
||||
pipeline_name = config["_class_name"]
|
||||
|
||||
# First try exact match for specific weights
|
||||
if pipeline_name_or_path in SAMPLING_PARAM_REGISTRY:
|
||||
return SAMPLING_PARAM_REGISTRY[pipeline_name_or_path]
|
||||
|
||||
# Try partial matches (for local paths that might include the weight ID)
|
||||
for registered_id, config_class in SAMPLING_PARAM_REGISTRY.items():
|
||||
if registered_id in pipeline_name_or_path:
|
||||
return config_class
|
||||
|
||||
# If no match, try to use the fallback config
|
||||
fallback_config = None
|
||||
# Try to determine pipeline architecture for fallback
|
||||
for pipeline_type, detector in SAMPLING_PARAM_DETECTOR.items():
|
||||
if detector(pipeline_name.lower()):
|
||||
fallback_config = SAMPLING_FALLBACK_PARAM.get(pipeline_type)
|
||||
break
|
||||
|
||||
logger.warning(
|
||||
"No match found for pipeline %s, using fallback sampling param %s.",
|
||||
pipeline_name_or_path, fallback_config)
|
||||
return fallback_config
|
||||
@@ -0,0 +1,40 @@
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.v1.configs.sample.base import CacheParams
|
||||
|
||||
|
||||
@dataclass
|
||||
class TeaCacheParams(CacheParams):
|
||||
cache_type: str = "teacache"
|
||||
teacache_thresh: float = 0.0
|
||||
coefficients: list[float] = field(default_factory=list)
|
||||
|
||||
|
||||
@dataclass
|
||||
class WanTeaCacheParams(CacheParams):
|
||||
# Unfortunately, TeaCache is very different for Wan than other models
|
||||
cache_type: str = "teacache"
|
||||
teacache_thresh: float = 0.0
|
||||
use_ret_steps: bool = True
|
||||
ret_steps_coeffs: list[float] = field(default_factory=list)
|
||||
non_ret_steps_coeffs: list[float] = field(default_factory=list)
|
||||
|
||||
@property
|
||||
def coefficients(self) -> list[float]:
|
||||
if self.use_ret_steps:
|
||||
return self.ret_steps_coeffs
|
||||
else:
|
||||
return self.non_ret_steps_coeffs
|
||||
|
||||
@property
|
||||
def ret_steps(self) -> int:
|
||||
if self.use_ret_steps:
|
||||
return 5 * 2
|
||||
else:
|
||||
return 1 * 2
|
||||
|
||||
def get_cutoff_steps(self, num_inference_steps: int) -> int:
|
||||
if self.use_ret_steps:
|
||||
return num_inference_steps * 2
|
||||
else:
|
||||
return num_inference_steps * 2 - 2
|
||||
@@ -0,0 +1,95 @@
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.v1.configs.sample.base import SamplingParam
|
||||
from fastvideo.v1.configs.sample.teacache import WanTeaCacheParams
|
||||
|
||||
|
||||
@dataclass
|
||||
class WanT2V_1_3B_SamplingParam(SamplingParam):
|
||||
# Video parameters
|
||||
height: int = 480
|
||||
width: int = 832
|
||||
num_frames: int = 81
|
||||
fps: int = 16
|
||||
|
||||
# Denoising stage
|
||||
guidance_scale: float = 3.0
|
||||
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"
|
||||
num_inference_steps: int = 50
|
||||
|
||||
teacache_params: WanTeaCacheParams = field(
|
||||
default_factory=lambda: WanTeaCacheParams(
|
||||
teacache_thresh=0.08,
|
||||
ret_steps_coeffs=[
|
||||
-5.21862437e+04, 9.23041404e+03, -5.28275948e+02,
|
||||
1.36987616e+01, -4.99875664e-02
|
||||
],
|
||||
non_ret_steps_coeffs=[
|
||||
2.39676752e+03, -1.31110545e+03, 2.01331979e+02,
|
||||
-8.29855975e+00, 1.37887774e-01
|
||||
]))
|
||||
|
||||
|
||||
@dataclass
|
||||
class WanT2V_14B_SamplingParam(SamplingParam):
|
||||
# Video parameters
|
||||
height: int = 720
|
||||
width: int = 1280
|
||||
num_frames: int = 81
|
||||
fps: int = 16
|
||||
|
||||
# Denoising stage
|
||||
guidance_scale: float = 5.0
|
||||
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"
|
||||
num_inference_steps: int = 50
|
||||
|
||||
teacache_params: WanTeaCacheParams = field(
|
||||
default_factory=lambda: WanTeaCacheParams(
|
||||
teacache_thresh=0.20,
|
||||
use_ret_steps=False,
|
||||
ret_steps_coeffs=[
|
||||
-3.03318725e+05, 4.90537029e+04, -2.65530556e+03,
|
||||
5.87365115e+01, -3.15583525e-01
|
||||
],
|
||||
non_ret_steps_coeffs=[
|
||||
-5784.54975374, 5449.50911966, -1811.16591783, 256.27178429,
|
||||
-13.02252404
|
||||
]))
|
||||
|
||||
|
||||
@dataclass
|
||||
class WanI2V_14B_480P_SamplingParam(WanT2V_1_3B_SamplingParam):
|
||||
# Denoising stage
|
||||
guidance_scale: float = 5.0
|
||||
num_inference_steps: int = 40
|
||||
|
||||
teacache_params: WanTeaCacheParams = field(
|
||||
default_factory=lambda: WanTeaCacheParams(
|
||||
teacache_thresh=0.26,
|
||||
ret_steps_coeffs=[
|
||||
-3.03318725e+05, 4.90537029e+04, -2.65530556e+03,
|
||||
5.87365115e+01, -3.15583525e-01
|
||||
],
|
||||
non_ret_steps_coeffs=[
|
||||
-5784.54975374, 5449.50911966, -1811.16591783, 256.27178429,
|
||||
-13.02252404
|
||||
]))
|
||||
|
||||
|
||||
@dataclass
|
||||
class WanI2V_14B_720P_SamplingParam(WanT2V_14B_SamplingParam):
|
||||
# Denoising stage
|
||||
guidance_scale: float = 5.0
|
||||
num_inference_steps: int = 40
|
||||
|
||||
teacache_params: WanTeaCacheParams = field(
|
||||
default_factory=lambda: WanTeaCacheParams(
|
||||
teacache_thresh=0.3,
|
||||
ret_steps_coeffs=[
|
||||
-3.03318725e+05, 4.90537029e+04, -2.65530556e+03,
|
||||
5.87365115e+01, -3.15583525e-01
|
||||
],
|
||||
non_ret_steps_coeffs=[
|
||||
-5784.54975374, 5449.50911966, -1811.16591783, 256.27178429,
|
||||
-13.02252404
|
||||
]))
|
||||
@@ -1,12 +1,4 @@
|
||||
{
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 81,
|
||||
"fps": 16,
|
||||
"num_inference_steps": 50,
|
||||
"guidance_scale": 3.0,
|
||||
"seed": 1024,
|
||||
"guidance_rescale": 0.0,
|
||||
"embedded_cfg_scale": 6.0,
|
||||
"flow_shift": 3,
|
||||
"use_cpu_offload": true,
|
||||
@@ -30,12 +22,20 @@
|
||||
"use_parallel_tiling": false,
|
||||
"use_feature_cache": true
|
||||
},
|
||||
"num_channels_latents": null,
|
||||
"image_encoder_precision": "fp32",
|
||||
"text_encoder_precision": "fp32",
|
||||
"text_len": 512,
|
||||
"hidden_state_skip_layer": 0,
|
||||
"dit_config": {
|
||||
"prefix": "Wan",
|
||||
"quant_config": null
|
||||
},
|
||||
"text_encoder_precisions": [
|
||||
"fp32"
|
||||
],
|
||||
"text_encoder_configs": [
|
||||
{
|
||||
"prefix": "t5",
|
||||
"quant_config": null,
|
||||
"lora_config": null
|
||||
}
|
||||
],
|
||||
"mask_strategy_file_path": null,
|
||||
"enable_torch_compile": false,
|
||||
"neg_prompt": "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"
|
||||
"enable_torch_compile": false
|
||||
}
|
||||
@@ -1,12 +1,4 @@
|
||||
{
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 81,
|
||||
"fps": 16,
|
||||
"num_inference_steps": 40,
|
||||
"guidance_scale": 5.0,
|
||||
"seed": 1024,
|
||||
"guidance_rescale": 0.0,
|
||||
"embedded_cfg_scale": 6.0,
|
||||
"flow_shift": 3,
|
||||
"use_cpu_offload": true,
|
||||
@@ -30,12 +22,28 @@
|
||||
"use_parallel_tiling": false,
|
||||
"use_feature_cache": true
|
||||
},
|
||||
"num_channels_latents": null,
|
||||
"image_encoder_precision": "fp32",
|
||||
"text_encoder_precision": "fp32",
|
||||
"text_len": 512,
|
||||
"hidden_state_skip_layer": 0,
|
||||
"dit_config": {
|
||||
"prefix": "Wan",
|
||||
"quant_config": null
|
||||
},
|
||||
"text_encoder_precisions": [
|
||||
"fp32"
|
||||
],
|
||||
"text_encoder_configs": [
|
||||
{
|
||||
"prefix": "t5",
|
||||
"quant_config": null,
|
||||
"lora_config": null
|
||||
}
|
||||
],
|
||||
"mask_strategy_file_path": null,
|
||||
"enable_torch_compile": false,
|
||||
"neg_prompt": "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"
|
||||
"image_encoder_config": {
|
||||
"prefix": "clip",
|
||||
"quant_config": null,
|
||||
"lora_config": null,
|
||||
"num_hidden_layers_override": null,
|
||||
"require_post_norm": null
|
||||
},
|
||||
"image_encoder_precision": "fp32"
|
||||
}
|
||||
@@ -6,10 +6,10 @@ This module provides a consolidated interface for generating videos using
|
||||
diffusion models.
|
||||
"""
|
||||
|
||||
import gc
|
||||
import os
|
||||
import time
|
||||
from dataclasses import asdict
|
||||
from typing import Any, Callable, Dict, List, Optional, Union
|
||||
from typing import Any, Dict, List, Optional, Union
|
||||
|
||||
import imageio
|
||||
import numpy as np
|
||||
@@ -17,7 +17,9 @@ import torch
|
||||
import torchvision
|
||||
from einops import rearrange
|
||||
|
||||
from fastvideo.v1.configs.pipelines import get_pipeline_config_cls_for_name, BaseConfig
|
||||
from fastvideo.v1.configs.pipelines import (PipelineConfig,
|
||||
get_pipeline_config_cls_for_name)
|
||||
from fastvideo.v1.configs.sample import SamplingParam
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines import ForwardBatch
|
||||
@@ -52,7 +54,9 @@ class VideoGenerator:
|
||||
model_path: str,
|
||||
device: Optional[str] = None,
|
||||
torch_dtype: Optional[torch.dtype] = None,
|
||||
pipeline_config: Optional[Union[str | BaseConfig]] = None,
|
||||
pipeline_config: Optional[
|
||||
Union[str
|
||||
| PipelineConfig]] = None,
|
||||
**kwargs) -> "VideoGenerator":
|
||||
"""
|
||||
Create a video generator from a pretrained model.
|
||||
@@ -71,7 +75,7 @@ class VideoGenerator:
|
||||
|
||||
config = None
|
||||
# 1. If users provide a pipeline config, it will override the default pipeline config
|
||||
if isinstance(pipeline_config, BaseConfig):
|
||||
if isinstance(pipeline_config, PipelineConfig):
|
||||
config = pipeline_config
|
||||
else:
|
||||
config_cls = get_pipeline_config_cls_for_name(model_path)
|
||||
@@ -124,20 +128,8 @@ class VideoGenerator:
|
||||
def generate_video(
|
||||
self,
|
||||
prompt: str,
|
||||
image_path: Optional[str] = None,
|
||||
negative_prompt: Optional[str] = None,
|
||||
output_path: Optional[str] = None,
|
||||
save_video: bool = True,
|
||||
return_frames: bool = False,
|
||||
num_inference_steps: Optional[int] = None,
|
||||
guidance_scale: Optional[float] = None,
|
||||
num_frames: Optional[int] = None,
|
||||
height: Optional[int] = None,
|
||||
width: Optional[int] = None,
|
||||
fps: Optional[int] = None,
|
||||
seed: Optional[int] = None,
|
||||
callback: Optional[Callable[[int, int, torch.Tensor], None]] = None,
|
||||
callback_steps: int = 1,
|
||||
sampling_param: Optional[SamplingParam] = None,
|
||||
**kwargs,
|
||||
) -> Union[Dict[str, Any], List[np.ndarray]]:
|
||||
"""
|
||||
Generate a video based on the given prompt.
|
||||
@@ -164,90 +156,71 @@ class VideoGenerator:
|
||||
# Create a copy of inference args to avoid modifying the original
|
||||
fastvideo_args = self.fastvideo_args
|
||||
|
||||
# Override parameters if provided
|
||||
if image_path is not None:
|
||||
fastvideo_args.image_path = image_path
|
||||
if negative_prompt is not None:
|
||||
fastvideo_args.neg_prompt = negative_prompt
|
||||
if num_inference_steps is not None:
|
||||
fastvideo_args.num_inference_steps = num_inference_steps
|
||||
if guidance_scale is not None:
|
||||
fastvideo_args.guidance_scale = guidance_scale
|
||||
if num_frames is not None:
|
||||
fastvideo_args.num_frames = num_frames
|
||||
if height is not None:
|
||||
fastvideo_args.height = height
|
||||
if width is not None:
|
||||
fastvideo_args.width = width
|
||||
if fps is not None:
|
||||
fastvideo_args.fps = fps
|
||||
if seed is not None:
|
||||
fastvideo_args.seed = seed
|
||||
|
||||
# Validate inputs
|
||||
if not isinstance(prompt, str):
|
||||
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)
|
||||
kwargs["prompt"] = prompt
|
||||
sampling_param.update(kwargs)
|
||||
|
||||
# Process negative prompt
|
||||
if fastvideo_args.neg_prompt is not None:
|
||||
fastvideo_args.neg_prompt = fastvideo_args.neg_prompt.strip()
|
||||
if sampling_param.negative_prompt is not None:
|
||||
sampling_param.negative_prompt = sampling_param.negative_prompt.strip(
|
||||
)
|
||||
|
||||
# Validate dimensions
|
||||
if (fastvideo_args.height <= 0 or fastvideo_args.width <= 0
|
||||
or fastvideo_args.num_frames <= 0):
|
||||
if (sampling_param.height <= 0 or sampling_param.width <= 0
|
||||
or sampling_param.num_frames <= 0):
|
||||
raise ValueError(
|
||||
f"Height, width, and num_frames must be positive integers, got "
|
||||
f"height={fastvideo_args.height}, width={fastvideo_args.width}, "
|
||||
f"num_frames={fastvideo_args.num_frames}")
|
||||
f"height={sampling_param.height}, width={sampling_param.width}, "
|
||||
f"num_frames={sampling_param.num_frames}")
|
||||
|
||||
if (fastvideo_args.num_frames - 1) % 4 != 0:
|
||||
if (
|
||||
sampling_param.num_frames - 1
|
||||
) % fastvideo_args.vae_config.arch_config.temporal_compression_ratio != 0:
|
||||
raise ValueError(
|
||||
f"num_frames-1 must be a multiple of 4, got {fastvideo_args.num_frames}"
|
||||
f"num_frames-1 must be a multiple of {fastvideo_args.vae_config.arch_config.temporal_compression_ratio}, got {sampling_param.num_frames}"
|
||||
)
|
||||
|
||||
# Calculate sizes
|
||||
target_height = align_to(fastvideo_args.height, 16)
|
||||
target_width = align_to(fastvideo_args.width, 16)
|
||||
target_height = align_to(sampling_param.height, 16)
|
||||
target_width = align_to(sampling_param.width, 16)
|
||||
|
||||
# Calculate latent sizes
|
||||
latents_size = [(fastvideo_args.num_frames - 1) // 4 + 1,
|
||||
fastvideo_args.height // 8, fastvideo_args.width // 8]
|
||||
latents_size = [(sampling_param.num_frames - 1) // 4 + 1,
|
||||
sampling_param.height // 8, sampling_param.width // 8]
|
||||
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
|
||||
|
||||
# Log parameters
|
||||
debug_str = f"""
|
||||
height: {target_height}
|
||||
width: {target_width}
|
||||
video_length: {fastvideo_args.num_frames}
|
||||
video_length: {sampling_param.num_frames}
|
||||
prompt: {prompt}
|
||||
neg_prompt: {fastvideo_args.neg_prompt}
|
||||
seed: {fastvideo_args.seed}
|
||||
infer_steps: {fastvideo_args.num_inference_steps}
|
||||
num_videos_per_prompt: {fastvideo_args.num_videos}
|
||||
guidance_scale: {fastvideo_args.guidance_scale}
|
||||
neg_prompt: {sampling_param.negative_prompt}
|
||||
seed: {sampling_param.seed}
|
||||
infer_steps: {sampling_param.num_inference_steps}
|
||||
num_videos_per_prompt: {sampling_param.num_videos_per_prompt}
|
||||
guidance_scale: {sampling_param.guidance_scale}
|
||||
n_tokens: {n_tokens}
|
||||
flow_shift: {fastvideo_args.flow_shift}
|
||||
embedded_guidance_scale: {fastvideo_args.embedded_cfg_scale}"""
|
||||
embedded_guidance_scale: {fastvideo_args.embedded_cfg_scale}
|
||||
save_video: {sampling_param.save_video}
|
||||
output_path: {sampling_param.output_path}
|
||||
""" # type: ignore[attr-defined]
|
||||
logger.info(debug_str)
|
||||
|
||||
# Prepare batch
|
||||
device = torch.device(fastvideo_args.device_str)
|
||||
batch = ForwardBatch(
|
||||
prompt=prompt,
|
||||
image_path=fastvideo_args.image_path,
|
||||
negative_prompt=fastvideo_args.neg_prompt,
|
||||
num_videos_per_prompt=fastvideo_args.num_videos,
|
||||
height=fastvideo_args.height,
|
||||
width=fastvideo_args.width,
|
||||
num_frames=fastvideo_args.num_frames,
|
||||
num_inference_steps=fastvideo_args.num_inference_steps,
|
||||
guidance_scale=fastvideo_args.guidance_scale,
|
||||
**shallow_asdict(sampling_param),
|
||||
eta=0.0,
|
||||
n_tokens=n_tokens,
|
||||
data_type="video" if fastvideo_args.num_frames > 1 else "image",
|
||||
device=device,
|
||||
extra={},
|
||||
)
|
||||
|
||||
@@ -268,23 +241,31 @@ class VideoGenerator:
|
||||
frames.append((x * 255).numpy().astype(np.uint8))
|
||||
|
||||
# Save video if requested
|
||||
if save_video:
|
||||
save_path = output_path or fastvideo_args.output_path
|
||||
if batch.save_video:
|
||||
save_path = batch.output_path
|
||||
if save_path:
|
||||
os.makedirs(os.path.dirname(save_path), exist_ok=True)
|
||||
video_path = os.path.join(save_path, f"{prompt[:100]}.mp4")
|
||||
imageio.mimsave(video_path, frames, fps=fastvideo_args.fps)
|
||||
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")
|
||||
|
||||
if return_frames:
|
||||
if batch.return_frames:
|
||||
return frames
|
||||
else:
|
||||
return {
|
||||
"samples": samples,
|
||||
"prompts": prompt,
|
||||
"size":
|
||||
(target_height, target_width, fastvideo_args.num_frames),
|
||||
"size": (target_height, target_width, batch.num_frames),
|
||||
"generation_time": gen_time
|
||||
}
|
||||
|
||||
def shutdown(self):
|
||||
"""
|
||||
Shutdown the video generator.
|
||||
"""
|
||||
self.executor.shutdown()
|
||||
del self.executor
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
@@ -170,7 +170,8 @@ environment_variables: Dict[str, Callable[[], Any]] = {
|
||||
# Available options:
|
||||
# - "TORCH_SDPA": use torch.nn.MultiheadAttention
|
||||
# - "FLASH_ATTN": use FlashAttention
|
||||
# - "STA" : use sliding tile attention
|
||||
# - "SLIDING_TILE_ATTN" : use Sliding Tile Attention
|
||||
# - "SAGE_ATTN": use Sage Attention
|
||||
"FASTVIDEO_ATTENTION_BACKEND":
|
||||
lambda: os.getenv("FASTVIDEO_ATTENTION_BACKEND", None),
|
||||
|
||||
|
||||
@@ -1,28 +0,0 @@
|
||||
# Basic Video Generation Tutorial
|
||||
The `VideoGenerator` class provides the primary Python interface for doing offline video generation, which is interacting with a diffusion pipeline without using a separate inference api server.
|
||||
|
||||
|
||||
## Usage
|
||||
The first script in this example shows the most basic usage of FastVideo. If you are new to Python and FastVideo, you should start here.
|
||||
|
||||
```bash
|
||||
python fastvideo/v1/examples/inference/basic/basic.py
|
||||
```
|
||||
|
||||
# Basic Walkthrough
|
||||
|
||||
All you need to generate videos using multi-gpus from state-of-the-art diffusion pipelines is the following few lines!
|
||||
|
||||
```python
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/FastHunyuan-Diffusers",
|
||||
num_gpus=2,
|
||||
)
|
||||
|
||||
prompt = "A beautiful woman in a red dress walking down a street"
|
||||
video = generator.generate_video(prompt)
|
||||
```
|
||||
|
||||
More to come! These examples and APIs are still under construction!
|
||||
@@ -1,4 +1,5 @@
|
||||
from fastvideo import VideoGenerator
|
||||
# from fastvideo.v1.configs.sample import SamplingParam
|
||||
|
||||
def main():
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
@@ -6,14 +7,18 @@ def main():
|
||||
# If a local path is provided, FastVideo will make a best effort
|
||||
# attempt to identify the optimal arguments.
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/FastHunyuan-Diffusers",
|
||||
"FastVideo/FastHunyuan-diffusers",
|
||||
# if num_gpus > 1, FastVideo will automatically handle distributed setup
|
||||
num_gpus=4,
|
||||
)
|
||||
|
||||
# sampling_param = SamplingParam.from_pretrained("/workspace/data/Wan-AI/Wan2.1-I2V-14B-480P-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 beautiful woman in a red dress walking down a street"
|
||||
video = generator.generate_video(prompt)
|
||||
# video = generator.generate_video(prompt, sampling_param=sampling_param, output_path="wan_t2v_videos/")
|
||||
|
||||
# Generate another video with a different prompt, without reloading the
|
||||
# model!
|
||||
@@ -22,4 +27,4 @@ def main():
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
main()
|
||||
@@ -1,141 +0,0 @@
|
||||
import os
|
||||
import gradio as gr
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
args = FastVideoArgs(model_path="FastVideo/FastHunyuan-Diffusers", num_gpus=2)
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
model_path=args.model_path,
|
||||
num_gpus=args.num_gpus
|
||||
)
|
||||
|
||||
def generate_video(
|
||||
prompt,
|
||||
negative_prompt,
|
||||
use_negative_prompt,
|
||||
seed,
|
||||
guidance_scale,
|
||||
num_frames,
|
||||
height,
|
||||
width,
|
||||
num_inference_steps,
|
||||
randomize_seed=False,
|
||||
):
|
||||
if randomize_seed:
|
||||
seed = torch.randint(0, 1000000, (1, )).item()
|
||||
|
||||
if not use_negative_prompt:
|
||||
negative_prompt = None
|
||||
|
||||
generator.generate_video(
|
||||
prompt=prompt,
|
||||
negative_prompt=negative_prompt,
|
||||
num_inference_steps=num_inference_steps,
|
||||
num_frames=num_frames,
|
||||
height=height,
|
||||
width=width,
|
||||
guidance_scale=guidance_scale,
|
||||
seed=seed
|
||||
)
|
||||
|
||||
output_path = os.path.join(args.output_path, f"{prompt[:100]}.mp4")
|
||||
|
||||
return output_path, seed
|
||||
|
||||
examples = [
|
||||
"A hand enters the frame, pulling a sheet of plastic wrap over three balls of dough placed on a wooden surface. The plastic wrap is stretched to cover the dough more securely. The hand adjusts the wrap, ensuring that it is tight and smooth over the dough. The scene focuses on the hand’s movements as it secures the edges of the plastic wrap. No new objects appear, and the camera remains stationary, focusing on the action of covering the dough.",
|
||||
"A vintage train snakes through the mountains, its plume of white steam rising dramatically against the jagged peaks. The cars glint in the late afternoon sun, their deep crimson and gold accents lending a touch of elegance. The tracks carve a precarious path along the cliffside, revealing glimpses of a roaring river far below. Inside, passengers peer out the large windows, their faces lit with awe as the landscape unfolds.",
|
||||
"A crowded rooftop bar buzzes with energy, the city skyline twinkling like a field of stars in the background. Strings of fairy lights hang above, casting a warm, golden glow over the scene. Groups of people gather around high tables, their laughter blending with the soft rhythm of live jazz. The aroma of freshly mixed cocktails and charred appetizers wafts through the air, mingling with the cool night breeze.",
|
||||
]
|
||||
|
||||
with gr.Blocks() as demo:
|
||||
gr.Markdown("# FastVideo Inference Demo")
|
||||
|
||||
with gr.Group():
|
||||
with gr.Row():
|
||||
prompt = gr.Text(
|
||||
label="Prompt",
|
||||
show_label=False,
|
||||
max_lines=1,
|
||||
placeholder="Enter your prompt",
|
||||
container=False,
|
||||
)
|
||||
run_button = gr.Button("Run", scale=0)
|
||||
result = gr.Video(label="Result", show_label=False)
|
||||
|
||||
with gr.Accordion("Advanced options", open=False):
|
||||
with gr.Group():
|
||||
with gr.Row():
|
||||
height = gr.Slider(
|
||||
label="Height",
|
||||
minimum=256,
|
||||
maximum=1024,
|
||||
step=32,
|
||||
value=args.height,
|
||||
)
|
||||
width = gr.Slider(label="Width", minimum=256, maximum=1024, step=32, value=args.width)
|
||||
|
||||
with gr.Row():
|
||||
num_frames = gr.Slider(
|
||||
label="Number of Frames",
|
||||
minimum=21,
|
||||
maximum=163,
|
||||
value=45,
|
||||
)
|
||||
guidance_scale = gr.Slider(
|
||||
label="Guidance Scale",
|
||||
minimum=1,
|
||||
maximum=12,
|
||||
value=args.guidance_scale,
|
||||
)
|
||||
num_inference_steps = gr.Slider(
|
||||
label="Inference Steps",
|
||||
minimum=4,
|
||||
maximum=100,
|
||||
value=6,
|
||||
)
|
||||
|
||||
with gr.Row():
|
||||
use_negative_prompt = gr.Checkbox(label="Use negative prompt", value=False)
|
||||
negative_prompt = gr.Text(
|
||||
label="Negative prompt",
|
||||
max_lines=1,
|
||||
placeholder="Enter a negative prompt",
|
||||
visible=False,
|
||||
)
|
||||
|
||||
seed = gr.Slider(label="Seed", minimum=0, maximum=1000000, step=1, value=args.seed)
|
||||
randomize_seed = gr.Checkbox(label="Randomize seed", value=True)
|
||||
seed_output = gr.Number(label="Used Seed")
|
||||
|
||||
gr.Examples(examples=examples, inputs=prompt)
|
||||
|
||||
use_negative_prompt.change(
|
||||
fn=lambda x: gr.update(visible=x),
|
||||
inputs=use_negative_prompt,
|
||||
outputs=negative_prompt,
|
||||
)
|
||||
|
||||
run_button.click(
|
||||
fn=generate_video,
|
||||
inputs=[
|
||||
prompt,
|
||||
negative_prompt,
|
||||
use_negative_prompt,
|
||||
seed,
|
||||
guidance_scale,
|
||||
num_frames,
|
||||
height,
|
||||
width,
|
||||
num_inference_steps,
|
||||
randomize_seed,
|
||||
],
|
||||
outputs=[result, seed_output],
|
||||
)
|
||||
|
||||
demo.queue(max_size=20).launch(server_name="0.0.0.0", server_port=7860)
|
||||
+52
-172
@@ -5,21 +5,32 @@
|
||||
import argparse
|
||||
import dataclasses
|
||||
from contextlib import contextmanager
|
||||
from typing import List, Optional
|
||||
from dataclasses import field
|
||||
from typing import Any, Callable, List, Optional, Tuple
|
||||
|
||||
from fastvideo.v1.configs.models import DiTConfig, EncoderConfig, VAEConfig
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.utils import FlexibleArgumentParser
|
||||
|
||||
from fastvideo.v1.configs.models import VAEConfig
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def preprocess_text(prompt: str) -> str:
|
||||
return prompt
|
||||
|
||||
|
||||
def postprocess_text(output: Any) -> Any:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
class FastVideoArgs:
|
||||
# Model and path configuration
|
||||
model_path: str
|
||||
|
||||
# Cache strategy
|
||||
cache_strategy: str = "none"
|
||||
|
||||
# Distributed executor backend
|
||||
distributed_executor_backend: str = "mp"
|
||||
|
||||
@@ -36,56 +47,40 @@ class FastVideoArgs:
|
||||
dist_timeout: Optional[int] = None # timeout for torch.distributed
|
||||
|
||||
# Video generation parameters
|
||||
height: int = 720
|
||||
width: int = 1280
|
||||
num_frames: int = 117
|
||||
num_inference_steps: int = 50
|
||||
guidance_scale: float = 1.0
|
||||
guidance_rescale: float = 0.0
|
||||
embedded_cfg_scale: float = 6.0
|
||||
flow_shift: Optional[float] = None
|
||||
|
||||
output_type: str = "pil"
|
||||
|
||||
# Model configuration
|
||||
# DiT configuration
|
||||
dit_config: DiTConfig = field(default_factory=DiTConfig)
|
||||
precision: str = "bf16"
|
||||
|
||||
# VAE configuration
|
||||
vae_precision: str = "fp16"
|
||||
vae_tiling: bool = True # Might change in between forward passes
|
||||
vae_sp: bool = False # Might change in between forward passes
|
||||
vae_tiling: bool = True # Might change in between forward passes
|
||||
vae_sp: bool = False # Might change in between forward passes
|
||||
# vae_scale_factor: Optional[int] = None # Deprecated
|
||||
vae_config: VAEConfig = VAEConfig()
|
||||
|
||||
# DiT configuration
|
||||
num_channels_latents: Optional[int] = None
|
||||
vae_config: VAEConfig = field(default_factory=VAEConfig)
|
||||
|
||||
# Image encoder configuration
|
||||
image_encoder_precision: str = "fp32"
|
||||
image_encoder_config: EncoderConfig = field(default_factory=EncoderConfig)
|
||||
|
||||
# Text encoder configuration
|
||||
text_encoder_precision: str = "fp16"
|
||||
text_len: int = 256
|
||||
hidden_state_skip_layer: int = 2
|
||||
|
||||
# Secondary text encoder
|
||||
text_encoder_precision_2: str = "fp16"
|
||||
text_len_2: int = 77
|
||||
|
||||
# Flow Matching parameters
|
||||
flow_solver: str = "euler"
|
||||
denoise_type: str = "flow" # Deprecated. Will use scheduler_config.json
|
||||
text_encoder_precisions: Tuple[str, ...] = field(
|
||||
default_factory=lambda: ("fp16", ))
|
||||
text_encoder_configs: Tuple[EncoderConfig, ...] = field(
|
||||
default_factory=lambda: (EncoderConfig(), ))
|
||||
preprocess_text_funcs: Tuple[Callable[[str], str], ...] = field(
|
||||
default_factory=lambda: (preprocess_text, ))
|
||||
postprocess_text_funcs: Tuple[Callable[[Any], Any], ...] = field(
|
||||
default_factory=lambda: (postprocess_text, ))
|
||||
|
||||
# STA (Spatial-Temporal Attention) parameters
|
||||
mask_strategy_file_path: Optional[str] = None
|
||||
enable_torch_compile: bool = False
|
||||
|
||||
# Scheduler options
|
||||
scheduler_type: str = "euler" # Deprecated. Will use the param in scheduler_config.json
|
||||
|
||||
neg_prompt: Optional[str] = None
|
||||
num_videos: int = 1
|
||||
fps: int = 24
|
||||
use_cpu_offload: bool = False
|
||||
disable_autocast: bool = False
|
||||
|
||||
@@ -93,11 +88,6 @@ class FastVideoArgs:
|
||||
log_level: str = "info"
|
||||
|
||||
# Inference parameters
|
||||
image_path: Optional[str] = None
|
||||
prompt: Optional[str] = None
|
||||
prompt_path: Optional[str] = None
|
||||
output_path: str = "outputs/"
|
||||
seed: int = 1024
|
||||
device_str: Optional[str] = None
|
||||
device = None
|
||||
|
||||
@@ -177,43 +167,6 @@ class FastVideoArgs:
|
||||
help="Set timeout for torch.distributed initialization.",
|
||||
)
|
||||
|
||||
# Video generation parameters
|
||||
parser.add_argument(
|
||||
"--height",
|
||||
type=int,
|
||||
default=FastVideoArgs.height,
|
||||
help="Height of generated video",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--width",
|
||||
type=int,
|
||||
default=FastVideoArgs.width,
|
||||
help="Width of generated video",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--num-frames",
|
||||
type=int,
|
||||
default=FastVideoArgs.num_frames,
|
||||
help="Number of frames to generate",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--num-inference-steps",
|
||||
type=int,
|
||||
default=FastVideoArgs.num_inference_steps,
|
||||
help="Number of inference steps",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--guidance-scale",
|
||||
type=float,
|
||||
default=FastVideoArgs.guidance_scale,
|
||||
help="Guidance scale for classifier-free guidance",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--guidance-rescale",
|
||||
type=float,
|
||||
default=FastVideoArgs.guidance_rescale,
|
||||
help="Guidance rescale for classifier-free guidance",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--embedded-cfg-scale",
|
||||
type=float,
|
||||
@@ -265,16 +218,11 @@ class FastVideoArgs:
|
||||
|
||||
parser.add_argument(
|
||||
"--text-encoder-precision",
|
||||
nargs="+",
|
||||
type=str,
|
||||
default=FastVideoArgs.text_encoder_precision,
|
||||
default=FastVideoArgs.text_encoder_precisions,
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
help="Precision for text encoder",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--text-len",
|
||||
type=int,
|
||||
default=FastVideoArgs.text_len,
|
||||
help="Maximum text length",
|
||||
help="Precision for each text encoder",
|
||||
)
|
||||
|
||||
# Image encoder config
|
||||
@@ -286,36 +234,6 @@ class FastVideoArgs:
|
||||
help="Precision for image encoder",
|
||||
)
|
||||
|
||||
# Secondary text encoder
|
||||
|
||||
parser.add_argument(
|
||||
"--text-encoder-precision-2",
|
||||
type=str,
|
||||
default=FastVideoArgs.text_encoder_precision_2,
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
help="Precision for secondary text encoder",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--text-len-2",
|
||||
type=int,
|
||||
default=FastVideoArgs.text_len_2,
|
||||
help="Maximum secondary text length",
|
||||
)
|
||||
|
||||
# Flow Matching parameters
|
||||
parser.add_argument(
|
||||
"--flow-solver",
|
||||
type=str,
|
||||
default=FastVideoArgs.flow_solver,
|
||||
help="Solver for flow matching",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--denoise-type",
|
||||
type=str,
|
||||
default=FastVideoArgs.denoise_type,
|
||||
help="Denoise type for noised inputs",
|
||||
)
|
||||
|
||||
# STA (Spatial-Temporal Attention) parameters
|
||||
parser.add_argument(
|
||||
"--mask-strategy-file-path",
|
||||
@@ -329,33 +247,6 @@ class FastVideoArgs:
|
||||
"Use torch.compile for speeding up STA inference without teacache",
|
||||
)
|
||||
|
||||
# Scheduler options
|
||||
parser.add_argument(
|
||||
"--scheduler-type",
|
||||
type=str,
|
||||
default=FastVideoArgs.scheduler_type,
|
||||
help="Type of scheduler to use",
|
||||
)
|
||||
|
||||
# HunYuan specific parameters
|
||||
parser.add_argument(
|
||||
"--neg-prompt",
|
||||
type=str,
|
||||
default=FastVideoArgs.neg_prompt,
|
||||
help="Negative prompt for sampling",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--num-videos",
|
||||
type=int,
|
||||
default=FastVideoArgs.num_videos,
|
||||
help="Number of videos to generate per prompt",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--fps",
|
||||
type=int,
|
||||
default=FastVideoArgs.fps,
|
||||
help="Frames per second for output video",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--use-cpu-offload",
|
||||
action="store_true",
|
||||
@@ -376,36 +267,6 @@ class FastVideoArgs:
|
||||
help="The logging level of all loggers.",
|
||||
)
|
||||
|
||||
# Inference parameters
|
||||
prompt_group = parser.add_mutually_exclusive_group(required=True)
|
||||
prompt_group.add_argument(
|
||||
"--prompt",
|
||||
type=str,
|
||||
help="Text prompt for video generation",
|
||||
)
|
||||
prompt_group.add_argument(
|
||||
"--prompt-path",
|
||||
type=str,
|
||||
help="Path to a text file containing the prompt",
|
||||
)
|
||||
|
||||
parser.add_argument("--image-path",
|
||||
type=str,
|
||||
help="Path to the image for I2V generation")
|
||||
|
||||
parser.add_argument(
|
||||
"--output-path",
|
||||
type=str,
|
||||
default=FastVideoArgs.output_path,
|
||||
help="Directory to save generated videos",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--seed",
|
||||
type=int,
|
||||
default=FastVideoArgs.seed,
|
||||
help="Random seed for reproducibility",
|
||||
)
|
||||
|
||||
return parser
|
||||
|
||||
@classmethod
|
||||
@@ -454,8 +315,27 @@ class FastVideoArgs:
|
||||
raise ValueError(
|
||||
"Currently enabling vae_sp requires enabling vae_tiling, please set --vae-tiling to True."
|
||||
)
|
||||
if self.prompt_path and not self.prompt_path.endswith(".txt"):
|
||||
raise ValueError("prompt_path must be a text file")
|
||||
|
||||
if len(self.text_encoder_configs) != len(self.text_encoder_precisions):
|
||||
raise ValueError(
|
||||
f"Length of text encoder configs ({len(self.text_encoder_configs)}) must be equal to length of text encoder precisions ({len(self.text_encoder_precisions)})"
|
||||
)
|
||||
|
||||
if len(self.text_encoder_configs) != len(self.preprocess_text_funcs):
|
||||
raise ValueError(
|
||||
f"Length of text encoder configs ({len(self.text_encoder_configs)}) must be equal to length of text preprocessing functions ({len(self.preprocess_text_funcs)})"
|
||||
)
|
||||
|
||||
if len(self.preprocess_text_funcs) != len(self.postprocess_text_funcs):
|
||||
raise ValueError(
|
||||
f"Length of text postprocess functions ({len(self.postprocess_text_funcs)}) must be equal to length of text preprocessing functions ({len(self.preprocess_text_funcs)})"
|
||||
)
|
||||
|
||||
if self.enable_torch_compile and self.num_gpus > 1:
|
||||
logger.warning(
|
||||
"Currently torch compile does not work with multi-gpu. Setting enable_torch_compile to False"
|
||||
)
|
||||
self.enable_torch_compile = False
|
||||
|
||||
|
||||
_current_fastvideo_args = None
|
||||
|
||||
@@ -11,6 +11,7 @@ import torch
|
||||
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from fastvideo.v1.attention import AttentionMetadata
|
||||
@@ -30,11 +31,13 @@ batchsize_forward_time: defaultdict = defaultdict(list)
|
||||
#
|
||||
@dataclass
|
||||
class ForwardContext:
|
||||
current_timestep: int
|
||||
# TODO(will): check this arg
|
||||
# copy from vllm_config.compilation_config.static_forward_context
|
||||
# attn_layers: Dict[str, Any]
|
||||
# TODO: extend to support per-layer dynamic forward context
|
||||
attn_metadata: "AttentionMetadata" # set dynamically for each forward pass
|
||||
forward_batch: Optional[ForwardBatch] = None
|
||||
|
||||
|
||||
_forward_context: Optional[ForwardContext] = None
|
||||
@@ -52,6 +55,7 @@ def get_forward_context() -> ForwardContext:
|
||||
@contextmanager
|
||||
def set_forward_context(current_timestep,
|
||||
attn_metadata,
|
||||
forward_batch: Optional[ForwardBatch] = None,
|
||||
fastvideo_args: Optional[FastVideoArgs] = None):
|
||||
"""A context manager that stores the current forward context,
|
||||
can be attention metadata, etc.
|
||||
@@ -63,7 +67,9 @@ def set_forward_context(current_timestep,
|
||||
forward_start_time = time.perf_counter()
|
||||
global _forward_context
|
||||
prev_context = _forward_context
|
||||
_forward_context = ForwardContext(attn_metadata=attn_metadata)
|
||||
_forward_context = ForwardContext(current_timestep=current_timestep,
|
||||
attn_metadata=attn_metadata,
|
||||
forward_batch=forward_batch)
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
@@ -76,10 +82,6 @@ def set_forward_context(current_timestep,
|
||||
else:
|
||||
# for v1 attention backends
|
||||
batchsize = attn_metadata.num_input_tokens
|
||||
# we use synchronous scheduling right now,
|
||||
# adding a sync point here should not affect
|
||||
# scheduling of the next batch
|
||||
torch.cuda.synchronize()
|
||||
now = time.perf_counter()
|
||||
# time measurement is in milliseconds
|
||||
batchsize_forward_time[batchsize].append(
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
# type: ignore
|
||||
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Inference module for diffusion models.
|
||||
|
||||
@@ -1,27 +1,30 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import List, Optional, Tuple, Union
|
||||
from typing import Any, List, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from fastvideo.v1.configs.models import DiTConfig
|
||||
from fastvideo.v1.platforms import _Backend
|
||||
|
||||
|
||||
# TODO
|
||||
class BaseDiT(nn.Module, ABC):
|
||||
_fsdp_shard_conditions: list = []
|
||||
attention_head_dim: int | None = None
|
||||
_compile_conditions: list = []
|
||||
_param_names_mapping: dict
|
||||
hidden_size: int
|
||||
num_attention_heads: int
|
||||
num_channels_latents: int
|
||||
# always supports torch_sdpa
|
||||
_supported_attention_backends: Tuple[_Backend,
|
||||
...] = (_Backend.TORCH_SDPA, )
|
||||
_supported_attention_backends: Tuple[
|
||||
_Backend, ...] = DiTConfig()._supported_attention_backends
|
||||
|
||||
def __init_subclass__(cls) -> None:
|
||||
required_class_attrs = [
|
||||
"_fsdp_shard_conditions", "_param_names_mapping"
|
||||
"_fsdp_shard_conditions", "_param_names_mapping",
|
||||
"_compile_conditions"
|
||||
]
|
||||
super().__init_subclass__()
|
||||
for attr in required_class_attrs:
|
||||
@@ -30,8 +33,9 @@ class BaseDiT(nn.Module, ABC):
|
||||
f"Subclasses of BaseDiT must define '{attr}' class variable"
|
||||
)
|
||||
|
||||
def __init__(self, *args, **kwargs) -> None:
|
||||
def __init__(self, config: DiTConfig, **kwargs) -> None:
|
||||
super().__init__()
|
||||
self.config = config
|
||||
if not self.supported_attention_backends:
|
||||
raise ValueError(
|
||||
f"Subclass {self.__class__.__name__} must define _supported_attention_backends"
|
||||
@@ -49,7 +53,9 @@ class BaseDiT(nn.Module, ABC):
|
||||
pass
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
required_attrs = ["hidden_size", "num_attention_heads"]
|
||||
required_attrs = [
|
||||
"hidden_size", "num_attention_heads", "num_channels_latents"
|
||||
]
|
||||
for attr in required_attrs:
|
||||
if not hasattr(self, attr):
|
||||
raise AttributeError(
|
||||
@@ -59,3 +65,58 @@ class BaseDiT(nn.Module, ABC):
|
||||
@property
|
||||
def supported_attention_backends(self) -> Tuple[_Backend, ...]:
|
||||
return self._supported_attention_backends
|
||||
|
||||
|
||||
class CachableDiT(BaseDiT):
|
||||
"""
|
||||
An intermediate base class that adds TeaCache optimization functionality to DiT models.
|
||||
TeaCache accelerates inference by selectively skipping redundant computation when consecutive
|
||||
diffusion steps are similar enough.
|
||||
"""
|
||||
# These are required class attributes that should be overridden by concrete implementations
|
||||
_fsdp_shard_conditions = []
|
||||
_param_names_mapping = {}
|
||||
# Ensure these instance attributes are properly defined in subclasses
|
||||
hidden_size: int
|
||||
num_attention_heads: int
|
||||
num_channels_latents: int
|
||||
# always supports torch_sdpa
|
||||
_supported_attention_backends: Tuple[
|
||||
_Backend, ...] = DiTConfig()._supported_attention_backends
|
||||
|
||||
def __init__(self, config: DiTConfig, **kwargs) -> None:
|
||||
super().__init__(config, **kwargs)
|
||||
|
||||
self.cnt = 0
|
||||
self.teacache_thresh = 0
|
||||
self.coefficients: list[float] = []
|
||||
|
||||
# NOTE(will): Only wan2.1 needs these, so we are hardcoding it here
|
||||
if self.config.prefix == "wan":
|
||||
self.use_ret_steps = self.config.cache_config.use_ret_steps
|
||||
self.is_even = False
|
||||
self.previous_e0_even: torch.Tensor | None = None
|
||||
self.previous_e0_odd: torch.Tensor | None = None
|
||||
self.previous_residual_even: torch.Tensor | None = None
|
||||
self.previous_residual_odd: torch.Tensor | None = None
|
||||
self.accumulated_rel_l1_distance_even = 0
|
||||
self.accumulated_rel_l1_distance_odd = 0
|
||||
self.should_calc_even = True
|
||||
self.should_calc_odd = True
|
||||
else:
|
||||
self.accumulated_rel_l1_distance = 0
|
||||
self.previous_modulated_input = None
|
||||
self.previous_residual = None
|
||||
|
||||
def maybe_cache_states(self, hidden_states: torch.Tensor,
|
||||
original_hidden_states: torch.Tensor) -> None:
|
||||
pass
|
||||
|
||||
def should_skip_forward_for_cached_states(self,
|
||||
**kwargs: dict[str, Any]) -> bool:
|
||||
return False
|
||||
|
||||
def retrieve_cached_states(self,
|
||||
hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
raise NotImplementedError(
|
||||
"maybe_retrieve_cached_states is not implemented")
|
||||
|
||||
@@ -0,0 +1,560 @@
|
||||
from typing import List, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from diffusers.models.normalization import RMSNorm
|
||||
|
||||
from fastvideo.v1.attention import DistributedAttention
|
||||
from fastvideo.v1.configs.models.dits import FluxImageConfig
|
||||
from fastvideo.v1.layers.layernorm import (LayerNormScaleShift, ScaleResidual,
|
||||
ScaleResidualLayerNormScaleShift)
|
||||
from fastvideo.v1.layers.linear import ReplicatedLinear
|
||||
from fastvideo.v1.layers.mlp import MLP
|
||||
from fastvideo.v1.layers.rotary_embedding import (_apply_rotary_emb,
|
||||
get_rotary_pos_embed)
|
||||
from fastvideo.v1.layers.visual_embedding import (ModulateProjection,
|
||||
TimestepEmbedder)
|
||||
from fastvideo.v1.models.dits.base import BaseDiT
|
||||
from fastvideo.v1.platforms import _Backend
|
||||
|
||||
|
||||
class MMDoubleStreamBlock(nn.Module):
|
||||
"""
|
||||
A multimodal DiT block with separate modulation for text and image/video,
|
||||
using distributed attention and linear layers.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int,
|
||||
num_attention_heads: int,
|
||||
mlp_ratio: float = 4.0,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.deterministic = False
|
||||
self.num_attention_heads = num_attention_heads
|
||||
head_dim = hidden_size // num_attention_heads
|
||||
mlp_hidden_dim = int(hidden_size * mlp_ratio)
|
||||
|
||||
# Image modulation components
|
||||
self.img_mod = ModulateProjection(
|
||||
hidden_size,
|
||||
factor=6,
|
||||
act_layer="silu",
|
||||
dtype=dtype,
|
||||
prefix=f"{prefix}.img_mod",
|
||||
)
|
||||
|
||||
# Fused operations for image stream
|
||||
self.img_attn_norm = LayerNormScaleShift(hidden_size,
|
||||
norm_type="layer",
|
||||
elementwise_affine=False,
|
||||
dtype=dtype)
|
||||
self.img_attn_residual_mlp_norm = ScaleResidualLayerNormScaleShift(
|
||||
hidden_size,
|
||||
norm_type="layer",
|
||||
elementwise_affine=False,
|
||||
dtype=dtype)
|
||||
self.img_mlp_residual = ScaleResidual()
|
||||
|
||||
# Image attention components
|
||||
self.img_attn_qkv = ReplicatedLinear(hidden_size,
|
||||
hidden_size * 3,
|
||||
bias=True,
|
||||
params_dtype=dtype,
|
||||
prefix=f"{prefix}.img_attn_qkv")
|
||||
|
||||
self.img_attn_q_norm = RMSNorm(head_dim, eps=1e-6)
|
||||
self.img_attn_k_norm = RMSNorm(head_dim, eps=1e-6)
|
||||
|
||||
self.img_attn_proj = ReplicatedLinear(hidden_size,
|
||||
hidden_size,
|
||||
bias=True,
|
||||
params_dtype=dtype,
|
||||
prefix=f"{prefix}.img_attn_proj")
|
||||
|
||||
self.img_mlp = MLP(hidden_size,
|
||||
mlp_hidden_dim,
|
||||
bias=True,
|
||||
dtype=dtype,
|
||||
prefix=f"{prefix}.img_mlp")
|
||||
|
||||
# Text modulation components
|
||||
self.txt_mod = ModulateProjection(
|
||||
hidden_size,
|
||||
factor=6,
|
||||
act_layer="silu",
|
||||
dtype=dtype,
|
||||
prefix=f"{prefix}.txt_mod",
|
||||
)
|
||||
|
||||
# Fused operations for text stream
|
||||
self.txt_attn_norm = LayerNormScaleShift(hidden_size,
|
||||
norm_type="layer",
|
||||
elementwise_affine=False,
|
||||
dtype=dtype)
|
||||
self.txt_attn_residual_mlp_norm = ScaleResidualLayerNormScaleShift(
|
||||
hidden_size,
|
||||
norm_type="layer",
|
||||
elementwise_affine=False,
|
||||
dtype=dtype)
|
||||
self.txt_mlp_residual = ScaleResidual()
|
||||
|
||||
# Text attention components
|
||||
self.txt_attn_qkv = ReplicatedLinear(hidden_size,
|
||||
hidden_size * 3,
|
||||
bias=True,
|
||||
params_dtype=dtype)
|
||||
|
||||
# QK norm layers for text
|
||||
self.txt_attn_q_norm = RMSNorm(head_dim, eps=1e-6)
|
||||
self.txt_attn_k_norm = RMSNorm(head_dim, eps=1e-6)
|
||||
|
||||
self.txt_attn_proj = ReplicatedLinear(hidden_size,
|
||||
hidden_size,
|
||||
bias=True,
|
||||
params_dtype=dtype)
|
||||
|
||||
self.txt_mlp = MLP(hidden_size, mlp_hidden_dim, bias=True, dtype=dtype)
|
||||
|
||||
# Distributed attention
|
||||
self.attn = DistributedAttention(
|
||||
num_heads=num_attention_heads,
|
||||
head_size=head_dim,
|
||||
causal=False,
|
||||
supported_attention_backends=supported_attention_backends,
|
||||
prefix=f"{prefix}.attn")
|
||||
|
||||
def forward(
|
||||
self,
|
||||
img: torch.Tensor,
|
||||
txt: torch.Tensor,
|
||||
vec: torch.Tensor,
|
||||
freqs_cis_img: Tuple[torch.Tensor, torch.Tensor],
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
# Process modulation vectors
|
||||
img_mod_outputs = self.img_mod(vec)
|
||||
(
|
||||
img_attn_shift,
|
||||
img_attn_scale,
|
||||
img_attn_gate,
|
||||
img_mlp_shift,
|
||||
img_mlp_scale,
|
||||
img_mlp_gate,
|
||||
) = torch.chunk(img_mod_outputs, 6, dim=-1)
|
||||
|
||||
txt_mod_outputs = self.txt_mod(vec)
|
||||
(
|
||||
txt_attn_shift,
|
||||
txt_attn_scale,
|
||||
txt_attn_gate,
|
||||
txt_mlp_shift,
|
||||
txt_mlp_scale,
|
||||
txt_mlp_gate,
|
||||
) = torch.chunk(txt_mod_outputs, 6, dim=-1)
|
||||
|
||||
# Prepare image for attention using fused operation
|
||||
img_attn_input = self.img_attn_norm(img, img_attn_shift, img_attn_scale)
|
||||
# Get QKV for image
|
||||
img_qkv, _ = self.img_attn_qkv(img_attn_input)
|
||||
batch_size, image_seq_len = img_qkv.shape[0], img_qkv.shape[1]
|
||||
|
||||
# Split QKV
|
||||
img_qkv = img_qkv.view(batch_size, image_seq_len, 3,
|
||||
self.num_attention_heads, -1)
|
||||
img_q, img_k, img_v = img_qkv[:, :, 0], img_qkv[:, :, 1], img_qkv[:, :,
|
||||
2]
|
||||
|
||||
# Apply QK-Norm if needed
|
||||
|
||||
img_q = self.img_attn_q_norm(img_q).to(img_v)
|
||||
img_k = self.img_attn_k_norm(img_k).to(img_v)
|
||||
# Apply rotary embeddings for image
|
||||
cos, sin = freqs_cis_img
|
||||
img_q, img_k = _apply_rotary_emb(
|
||||
img_q, cos, sin,
|
||||
is_neox_style=False), _apply_rotary_emb(img_k,
|
||||
cos,
|
||||
sin,
|
||||
is_neox_style=False)
|
||||
# Prepare text for attention using fused operation
|
||||
txt_attn_input = self.txt_attn_norm(txt, txt_attn_shift, txt_attn_scale)
|
||||
|
||||
# Get QKV for text
|
||||
txt_qkv, _ = self.txt_attn_qkv(txt_attn_input)
|
||||
batch_size, text_seq_len = txt_qkv.shape[0], txt_qkv.shape[1]
|
||||
|
||||
# Split QKV
|
||||
txt_qkv = txt_qkv.view(batch_size, text_seq_len, 3,
|
||||
self.num_attention_heads, -1)
|
||||
txt_q, txt_k, txt_v = txt_qkv[:, :, 0], txt_qkv[:, :, 1], txt_qkv[:, :,
|
||||
2]
|
||||
|
||||
# Apply QK-Norm if needed
|
||||
txt_q = self.txt_attn_q_norm(txt_q).to(txt_q.dtype)
|
||||
txt_k = self.txt_attn_k_norm(txt_k).to(txt_k.dtype)
|
||||
|
||||
# Run distributed attention
|
||||
img_attn, txt_attn = self.attn(img_q, img_k, img_v, txt_q, txt_k, txt_v)
|
||||
img_attn_out, _ = self.img_attn_proj(
|
||||
img_attn.view(batch_size, image_seq_len, -1))
|
||||
# Use fused operation for residual connection, normalization, and modulation
|
||||
img_mlp_input, img_residual = self.img_attn_residual_mlp_norm(
|
||||
img, img_attn_out, img_attn_gate, img_mlp_shift, img_mlp_scale)
|
||||
|
||||
# Process image MLP
|
||||
img_mlp_out = self.img_mlp(img_mlp_input)
|
||||
img = self.img_mlp_residual(img_residual, img_mlp_out, img_mlp_gate)
|
||||
|
||||
# Process text attention output
|
||||
txt_attn_out, _ = self.txt_attn_proj(
|
||||
txt_attn.reshape(batch_size, text_seq_len, -1))
|
||||
|
||||
# Use fused operation for residual connection, normalization, and modulation
|
||||
txt_mlp_input, txt_residual = self.txt_attn_residual_mlp_norm(
|
||||
txt, txt_attn_out, txt_attn_gate, txt_mlp_shift, txt_mlp_scale)
|
||||
|
||||
# Process text MLP
|
||||
txt_mlp_out = self.txt_mlp(txt_mlp_input)
|
||||
txt = self.txt_mlp_residual(txt_residual, txt_mlp_out, txt_mlp_gate)
|
||||
|
||||
return img, txt
|
||||
|
||||
|
||||
class MMSingleStreamBlock(nn.Module):
|
||||
"""
|
||||
A DiT block with parallel linear layers using distributed attention
|
||||
and tensor parallelism.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int,
|
||||
num_attention_heads: int,
|
||||
mlp_ratio: float = 4.0,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.deterministic = False
|
||||
self.hidden_size = hidden_size
|
||||
self.num_attention_heads = num_attention_heads
|
||||
head_dim = hidden_size // num_attention_heads
|
||||
mlp_hidden_dim = int(hidden_size * mlp_ratio)
|
||||
self.mlp_hidden_dim = mlp_hidden_dim
|
||||
|
||||
# Combined QKV and MLP input projection
|
||||
self.linear1 = ReplicatedLinear(hidden_size,
|
||||
hidden_size * 3 + mlp_hidden_dim,
|
||||
bias=True,
|
||||
params_dtype=dtype,
|
||||
prefix=f"{prefix}.linear1")
|
||||
|
||||
# Combined projection and MLP output
|
||||
self.linear2 = ReplicatedLinear(hidden_size + mlp_hidden_dim,
|
||||
hidden_size,
|
||||
bias=True,
|
||||
params_dtype=dtype,
|
||||
prefix=f"{prefix}.linear2")
|
||||
|
||||
# QK norm layers
|
||||
self.q_norm = RMSNorm(head_dim, eps=1e-6)
|
||||
self.k_norm = RMSNorm(head_dim, eps=1e-6)
|
||||
|
||||
# Fused operations with better naming
|
||||
self.input_norm_scale_shift = LayerNormScaleShift(
|
||||
hidden_size,
|
||||
norm_type="layer",
|
||||
eps=1e-6,
|
||||
elementwise_affine=False,
|
||||
dtype=dtype)
|
||||
self.output_residual = ScaleResidual()
|
||||
|
||||
# Activation function
|
||||
self.mlp_act = nn.GELU(approximate="tanh")
|
||||
|
||||
# Modulation
|
||||
self.modulation = ModulateProjection(hidden_size,
|
||||
factor=3,
|
||||
act_layer="silu",
|
||||
dtype=dtype,
|
||||
prefix=f"{prefix}.modulation")
|
||||
|
||||
# Distributed attention
|
||||
self.attn = DistributedAttention(
|
||||
num_heads=num_attention_heads,
|
||||
head_size=head_dim,
|
||||
causal=False,
|
||||
supported_attention_backends=supported_attention_backends,
|
||||
prefix=f"{prefix}.attn")
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
vec: torch.Tensor,
|
||||
txt_len: int,
|
||||
freqs_cis_img: Tuple[torch.Tensor, torch.Tensor],
|
||||
) -> torch.Tensor:
|
||||
# Process modulation
|
||||
mod_shift, mod_scale, mod_gate = self.modulation(vec).chunk(3, dim=-1)
|
||||
|
||||
# Apply pre-norm and modulation using fused operation
|
||||
x_mod = self.input_norm_scale_shift(x, mod_shift, mod_scale)
|
||||
|
||||
# Get combined projections
|
||||
linear1_out, _ = self.linear1(x_mod)
|
||||
|
||||
# Split into QKV and MLP parts
|
||||
qkv, mlp = torch.split(linear1_out,
|
||||
[3 * self.hidden_size, self.mlp_hidden_dim],
|
||||
dim=-1)
|
||||
|
||||
# Process QKV
|
||||
batch_size, seq_len = qkv.shape[0], qkv.shape[1]
|
||||
qkv = qkv.view(batch_size, seq_len, 3, self.num_attention_heads, -1)
|
||||
q, k, v = qkv[:, :, 0], qkv[:, :, 1], qkv[:, :, 2]
|
||||
|
||||
# Apply QK-Norm
|
||||
q = self.q_norm(q).to(v.dtype)
|
||||
k = self.k_norm(k).to(v.dtype)
|
||||
|
||||
# Split into image and text parts
|
||||
img_q, txt_q = q[:, :-txt_len], q[:, -txt_len:]
|
||||
img_k, txt_k = k[:, :-txt_len], k[:, -txt_len:]
|
||||
img_v, txt_v = v[:, :-txt_len], v[:, -txt_len:]
|
||||
# Apply rotary embeddings to image parts
|
||||
cos, sin = freqs_cis_img
|
||||
img_q, img_k = _apply_rotary_emb(
|
||||
img_q, cos, sin,
|
||||
is_neox_style=False), _apply_rotary_emb(img_k,
|
||||
cos,
|
||||
sin,
|
||||
is_neox_style=False)
|
||||
|
||||
# Run distributed attention
|
||||
img_attn_output, txt_attn_output = self.attn(img_q, img_k, img_v, txt_q,
|
||||
txt_k, txt_v)
|
||||
attn_output = torch.cat((img_attn_output, txt_attn_output),
|
||||
dim=1).view(batch_size, seq_len, -1)
|
||||
# Process MLP activation
|
||||
mlp_output = self.mlp_act(mlp)
|
||||
|
||||
# Combine attention and MLP outputs
|
||||
combined = torch.cat((attn_output, mlp_output), dim=-1)
|
||||
|
||||
# Final projection
|
||||
output, _ = self.linear2(combined)
|
||||
|
||||
# Apply residual connection with gating using fused operation
|
||||
return self.output_residual(x, output, mod_gate)
|
||||
|
||||
|
||||
class FinalLayer(nn.Module):
|
||||
"""
|
||||
The final layer of DiT that projects features to pixel space.
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
hidden_size,
|
||||
patch_size,
|
||||
out_channels,
|
||||
dtype=None,
|
||||
prefix: str = "") -> None:
|
||||
super().__init__()
|
||||
|
||||
# Normalization
|
||||
self.norm_final = nn.LayerNorm(hidden_size,
|
||||
eps=1e-6,
|
||||
elementwise_affine=False,
|
||||
dtype=dtype)
|
||||
|
||||
output_dim = patch_size**3 * out_channels
|
||||
|
||||
self.linear = ReplicatedLinear(hidden_size,
|
||||
output_dim,
|
||||
bias=True,
|
||||
params_dtype=dtype,
|
||||
prefix=f"{prefix}.linear")
|
||||
|
||||
# Modulation
|
||||
self.adaLN_modulation = ModulateProjection(
|
||||
hidden_size,
|
||||
factor=2,
|
||||
act_layer="silu",
|
||||
dtype=dtype,
|
||||
prefix=f"{prefix}.adaLN_modulation")
|
||||
|
||||
def forward(self, img, vec):
|
||||
scale, shift = self.adaLN_modulation(vec).chunk(2, dim=-1)
|
||||
img = self.norm_final(img) * (1.0 +
|
||||
scale.unsqueeze(1)) + shift.unsqueeze(1)
|
||||
img, _ = self.linear(img)
|
||||
return img
|
||||
|
||||
|
||||
class FluxTransformer2DModel(BaseDiT):
|
||||
_fsdp_shard_conditions = FluxImageConfig()._fsdp_shard_conditions
|
||||
_supported_attention_backends = FluxImageConfig(
|
||||
)._supported_attention_backends
|
||||
_param_names_mapping = FluxImageConfig()._param_names_mapping
|
||||
|
||||
def __init__(self, config: FluxImageConfig) -> None:
|
||||
super().__init__(config=config)
|
||||
|
||||
self.hidden_size = config.hidden_size
|
||||
self.num_attention_heads = config.num_attention_heads
|
||||
self.num_channels_latents = config.num_channels_latents
|
||||
self.text_states_dim = config.joint_attention_dim
|
||||
self.text_states_dim_2 = config.pooled_projection_dim
|
||||
self.rope_dim_list = list(config.axes_dims_rope)
|
||||
self.rope_theta = config.rope_theta
|
||||
self.out_channels = config.out_channels
|
||||
self.patch_size = config.patch_size
|
||||
|
||||
self.img_in = ReplicatedLinear(config.in_channels,
|
||||
self.hidden_size,
|
||||
params_dtype=config.dtype,
|
||||
prefix=f"{config.prefix}.img_in")
|
||||
self.txt_in = ReplicatedLinear(self.text_states_dim,
|
||||
self.hidden_size,
|
||||
params_dtype=config.dtype,
|
||||
prefix=f"{config.prefix}.txt_in")
|
||||
self.time_in = TimestepEmbedder(self.hidden_size,
|
||||
act_layer="silu",
|
||||
dtype=config.dtype,
|
||||
prefix=f"{config.prefix}.time_in")
|
||||
self.txt2_in = MLP(self.text_states_dim_2,
|
||||
self.hidden_size,
|
||||
self.hidden_size,
|
||||
act_type="silu",
|
||||
dtype=config.dtype,
|
||||
prefix=f"{config.prefix}.txt2_in")
|
||||
self.guidance_in = (TimestepEmbedder(
|
||||
self.hidden_size,
|
||||
act_layer="silu",
|
||||
dtype=config.dtype,
|
||||
prefix=f"{config.prefix}.guidance_in")
|
||||
if config.guidance_embeds else None)
|
||||
|
||||
# Double blocks
|
||||
self.double_blocks = nn.ModuleList([
|
||||
MMDoubleStreamBlock(
|
||||
self.hidden_size,
|
||||
self.num_attention_heads,
|
||||
dtype=config.dtype,
|
||||
supported_attention_backends=self._supported_attention_backends,
|
||||
prefix=f"{config.prefix}.double_blocks.{i}")
|
||||
for i in range(config.num_layers)
|
||||
])
|
||||
|
||||
# Single blocks
|
||||
self.single_blocks = nn.ModuleList([
|
||||
MMSingleStreamBlock(
|
||||
self.hidden_size,
|
||||
self.num_attention_heads,
|
||||
dtype=config.dtype,
|
||||
supported_attention_backends=self._supported_attention_backends,
|
||||
prefix=f"{config.prefix}.single_blocks.{i+config.num_layers}")
|
||||
for i in range(config.num_single_layers)
|
||||
])
|
||||
|
||||
self.final_layer = FinalLayer(config.hidden_size,
|
||||
self.patch_size,
|
||||
self.out_channels,
|
||||
dtype=config.dtype,
|
||||
prefix=f"{config.prefix}.final_layer")
|
||||
|
||||
self.__post_init__()
|
||||
|
||||
def forward(self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: Union[torch.Tensor, List[torch.Tensor]],
|
||||
timestep: torch.LongTensor,
|
||||
encoder_hidden_states_image: Optional[Union[
|
||||
torch.Tensor, List[torch.Tensor]]] = None,
|
||||
guidance=None,
|
||||
**kwargs):
|
||||
"""
|
||||
Forward pass of the FluxTransformer2DModel.
|
||||
|
||||
Args:
|
||||
hidden_states: Input image latents [B, N, C]
|
||||
encoder_hidden_states: Text embeddings [B, L, D]
|
||||
timestep: Diffusion timestep
|
||||
guidance: Guidance scale for CFG
|
||||
|
||||
Returns:
|
||||
Tuple of (output)
|
||||
"""
|
||||
h = kwargs.pop("height_latents") or None
|
||||
w = kwargs.pop("width_latents") or None
|
||||
assert h is not None and w is not None
|
||||
|
||||
img = x = hidden_states
|
||||
|
||||
# Match diffusers implementation by multiplying timestep by 1000
|
||||
t = timestep.to(img.dtype)
|
||||
|
||||
# Split text embeddings - first token is global, rest are per-token
|
||||
txt = encoder_hidden_states[1]
|
||||
text_states_2 = encoder_hidden_states[0]
|
||||
|
||||
# Get spatial dimensions
|
||||
# _, _, oh, ow = img.shape
|
||||
th, tw = (h // self.patch_size // 2, w // self.patch_size // 2)
|
||||
|
||||
# Get rotary embeddings
|
||||
freqs_cos_img, freqs_sin_img = get_rotary_pos_embed(
|
||||
(1, th, tw),
|
||||
self.hidden_size,
|
||||
self.num_attention_heads,
|
||||
self.rope_dim_list,
|
||||
self.rope_theta,
|
||||
shard_dim=1,
|
||||
)
|
||||
freqs_cos_img = freqs_cos_img.to(img.device)
|
||||
freqs_sin_img = freqs_sin_img.to(img.device)
|
||||
freqs_cis_img = (freqs_cos_img, freqs_sin_img)
|
||||
|
||||
# Prepare modulation vectors
|
||||
vec = self.time_in(t)
|
||||
|
||||
# Add text modulation
|
||||
vec = vec + self.txt2_in(text_states_2)
|
||||
|
||||
# Add guidance modulation
|
||||
if self.guidance_in is not None and guidance is not None:
|
||||
vec = vec + self.guidance_in(guidance)
|
||||
|
||||
# embed text and image
|
||||
img, _ = self.img_in(img)
|
||||
txt, _ = self.txt_in(txt)
|
||||
img_seq_len = img.shape[1]
|
||||
txt_seq_len = txt.shape[1]
|
||||
|
||||
# Process through double stream blocks
|
||||
for index, block in enumerate(self.double_blocks):
|
||||
double_block_args = [img, txt, vec, freqs_cis_img]
|
||||
img, txt = block(*double_block_args)
|
||||
|
||||
# Merge txt and img to pass through single stream blocks
|
||||
x = torch.cat((img, txt), 1)
|
||||
|
||||
# Process through single stream blocks
|
||||
for index, block in enumerate(self.single_blocks):
|
||||
single_block_args = [x, vec, txt_seq_len, freqs_cis_img]
|
||||
x = block(*single_block_args)
|
||||
|
||||
# Extract image features
|
||||
img = x[:, :img_seq_len, ...]
|
||||
|
||||
# Final layer
|
||||
img = self.final_layer(img, vec)
|
||||
|
||||
return img
|
||||
@@ -2,12 +2,16 @@
|
||||
|
||||
from typing import List, Optional, Tuple, Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from fastvideo.v1.attention import DistributedAttention, LocalAttention
|
||||
from fastvideo.v1.configs.models.dits import HunyuanVideoConfig
|
||||
from fastvideo.v1.configs.sample.teacache import TeaCacheParams
|
||||
from fastvideo.v1.distributed.parallel_state import (
|
||||
get_sequence_model_parallel_world_size)
|
||||
from fastvideo.v1.forward_context import get_forward_context
|
||||
from fastvideo.v1.layers.layernorm import (LayerNormScaleShift, ScaleResidual,
|
||||
ScaleResidualLayerNormScaleShift)
|
||||
from fastvideo.v1.layers.linear import ReplicatedLinear
|
||||
@@ -18,7 +22,8 @@ from fastvideo.v1.layers.rotary_embedding import (_apply_rotary_emb,
|
||||
from fastvideo.v1.layers.visual_embedding import (ModulateProjection,
|
||||
PatchEmbed, TimestepEmbedder,
|
||||
unpatchify)
|
||||
from fastvideo.v1.models.dits.base import BaseDiT
|
||||
from fastvideo.v1.models.dits.base import CachableDiT
|
||||
from fastvideo.v1.models.utils import modulate
|
||||
from fastvideo.v1.platforms import _Backend
|
||||
|
||||
|
||||
@@ -417,7 +422,7 @@ class MMSingleStreamBlock(nn.Module):
|
||||
return self.output_residual(x, output, mod_gate)
|
||||
|
||||
|
||||
class HunyuanVideoTransformer3DModel(BaseDiT):
|
||||
class HunyuanVideoTransformer3DModel(CachableDiT):
|
||||
"""
|
||||
HunyuanVideo Transformer backbone adapted for distributed training.
|
||||
|
||||
@@ -431,238 +436,105 @@ class HunyuanVideoTransformer3DModel(BaseDiT):
|
||||
# PY: we make the input args the same as HF config
|
||||
|
||||
# shard single stream, double stream blocks, and refiner_blocks
|
||||
_fsdp_shard_conditions = [
|
||||
lambda n, m: "double" in n and str.isdigit(n.split(".")[-1]),
|
||||
lambda n, m: "single" in n and str.isdigit(n.split(".")[-1]),
|
||||
lambda n, m: "refiner" in n and str.isdigit(n.split(".")[-1]),
|
||||
]
|
||||
_supported_attention_backends = (_Backend.SLIDING_TILE_ATTN,
|
||||
_Backend.FLASH_ATTN, _Backend.TORCH_SDPA)
|
||||
_param_names_mapping = {
|
||||
# 1. context_embedder.time_text_embed submodules (specific rules, applied first):
|
||||
r"^context_embedder\.time_text_embed\.timestep_embedder\.linear_1\.(.*)$":
|
||||
r"txt_in.t_embedder.mlp.fc_in.\1",
|
||||
r"^context_embedder\.time_text_embed\.timestep_embedder\.linear_2\.(.*)$":
|
||||
r"txt_in.t_embedder.mlp.fc_out.\1",
|
||||
r"^context_embedder\.proj_in\.(.*)$":
|
||||
r"txt_in.input_embedder.\1",
|
||||
r"^context_embedder\.time_text_embed\.text_embedder\.linear_1\.(.*)$":
|
||||
r"txt_in.c_embedder.fc_in.\1",
|
||||
r"^context_embedder\.time_text_embed\.text_embedder\.linear_2\.(.*)$":
|
||||
r"txt_in.c_embedder.fc_out.\1",
|
||||
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.norm1\.(.*)$":
|
||||
r"txt_in.refiner_blocks.\1.norm1.\2",
|
||||
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.norm2\.(.*)$":
|
||||
r"txt_in.refiner_blocks.\1.norm2.\2",
|
||||
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_q\.(.*)$":
|
||||
(r"txt_in.refiner_blocks.\1.self_attn_qkv.\2", 0, 3),
|
||||
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_k\.(.*)$":
|
||||
(r"txt_in.refiner_blocks.\1.self_attn_qkv.\2", 1, 3),
|
||||
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_v\.(.*)$":
|
||||
(r"txt_in.refiner_blocks.\1.self_attn_qkv.\2", 2, 3),
|
||||
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_out\.0\.(.*)$":
|
||||
r"txt_in.refiner_blocks.\1.self_attn_proj.\2",
|
||||
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.ff\.net\.0(?:\.proj)?\.(.*)$":
|
||||
r"txt_in.refiner_blocks.\1.mlp.fc_in.\2",
|
||||
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.ff\.net\.2(?:\.proj)?\.(.*)$":
|
||||
r"txt_in.refiner_blocks.\1.mlp.fc_out.\2",
|
||||
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.norm_out\.linear\.(.*)$":
|
||||
r"txt_in.refiner_blocks.\1.adaLN_modulation.linear.\2",
|
||||
_fsdp_shard_conditions = HunyuanVideoConfig()._fsdp_shard_conditions
|
||||
_compile_conditions = HunyuanVideoConfig()._compile_conditions
|
||||
_supported_attention_backends = HunyuanVideoConfig(
|
||||
)._supported_attention_backends
|
||||
_param_names_mapping = HunyuanVideoConfig()._param_names_mapping
|
||||
|
||||
# 3. x_embedder mapping:
|
||||
r"^x_embedder\.proj\.(.*)$":
|
||||
r"img_in.proj.\1",
|
||||
def __init__(self, config: HunyuanVideoConfig):
|
||||
super().__init__(config=config)
|
||||
|
||||
# 4. Top-level time_text_embed mappings:
|
||||
r"^time_text_embed\.timestep_embedder\.linear_1\.(.*)$":
|
||||
r"time_in.mlp.fc_in.\1",
|
||||
r"^time_text_embed\.timestep_embedder\.linear_2\.(.*)$":
|
||||
r"time_in.mlp.fc_out.\1",
|
||||
r"^time_text_embed\.guidance_embedder\.linear_1\.(.*)$":
|
||||
r"guidance_in.mlp.fc_in.\1",
|
||||
r"^time_text_embed\.guidance_embedder\.linear_2\.(.*)$":
|
||||
r"guidance_in.mlp.fc_out.\1",
|
||||
r"^time_text_embed\.text_embedder\.linear_1\.(.*)$":
|
||||
r"vector_in.fc_in.\1",
|
||||
r"^time_text_embed\.text_embedder\.linear_2\.(.*)$":
|
||||
r"vector_in.fc_out.\1",
|
||||
|
||||
# 5. transformer_blocks mapping:
|
||||
r"^transformer_blocks\.(\d+)\.norm1\.linear\.(.*)$":
|
||||
r"double_blocks.\1.img_mod.linear.\2",
|
||||
r"^transformer_blocks\.(\d+)\.norm1_context\.linear\.(.*)$":
|
||||
r"double_blocks.\1.txt_mod.linear.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn\.norm_q\.(.*)$":
|
||||
r"double_blocks.\1.img_attn_q_norm.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn\.norm_k\.(.*)$":
|
||||
r"double_blocks.\1.img_attn_k_norm.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn\.to_q\.(.*)$":
|
||||
(r"double_blocks.\1.img_attn_qkv.\2", 0, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.to_k\.(.*)$":
|
||||
(r"double_blocks.\1.img_attn_qkv.\2", 1, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.to_v\.(.*)$":
|
||||
(r"double_blocks.\1.img_attn_qkv.\2", 2, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.add_q_proj\.(.*)$":
|
||||
(r"double_blocks.\1.txt_attn_qkv.\2", 0, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.add_k_proj\.(.*)$":
|
||||
(r"double_blocks.\1.txt_attn_qkv.\2", 1, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.add_v_proj\.(.*)$":
|
||||
(r"double_blocks.\1.txt_attn_qkv.\2", 2, 3),
|
||||
r"^transformer_blocks\.(\d+)\.attn\.to_out\.0\.(.*)$":
|
||||
r"double_blocks.\1.img_attn_proj.\2",
|
||||
# Corrected: merge attn.to_add_out into the main projection.
|
||||
r"^transformer_blocks\.(\d+)\.attn\.to_add_out\.(.*)$":
|
||||
r"double_blocks.\1.txt_attn_proj.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn\.norm_added_q\.(.*)$":
|
||||
r"double_blocks.\1.txt_attn_q_norm.\2",
|
||||
r"^transformer_blocks\.(\d+)\.attn\.norm_added_k\.(.*)$":
|
||||
r"double_blocks.\1.txt_attn_k_norm.\2",
|
||||
r"^transformer_blocks\.(\d+)\.ff\.net\.0(?:\.proj)?\.(.*)$":
|
||||
r"double_blocks.\1.img_mlp.fc_in.\2",
|
||||
r"^transformer_blocks\.(\d+)\.ff\.net\.2(?:\.proj)?\.(.*)$":
|
||||
r"double_blocks.\1.img_mlp.fc_out.\2",
|
||||
r"^transformer_blocks\.(\d+)\.ff_context\.net\.0(?:\.proj)?\.(.*)$":
|
||||
r"double_blocks.\1.txt_mlp.fc_in.\2",
|
||||
r"^transformer_blocks\.(\d+)\.ff_context\.net\.2(?:\.proj)?\.(.*)$":
|
||||
r"double_blocks.\1.txt_mlp.fc_out.\2",
|
||||
|
||||
# 6. single_transformer_blocks mapping:
|
||||
r"^single_transformer_blocks\.(\d+)\.attn\.norm_q\.(.*)$":
|
||||
r"single_blocks.\1.q_norm.\2",
|
||||
r"^single_transformer_blocks\.(\d+)\.attn\.norm_k\.(.*)$":
|
||||
r"single_blocks.\1.k_norm.\2",
|
||||
r"^single_transformer_blocks\.(\d+)\.attn\.to_q\.(.*)$":
|
||||
(r"single_blocks.\1.linear1.\2", 0, 4),
|
||||
r"^single_transformer_blocks\.(\d+)\.attn\.to_k\.(.*)$":
|
||||
(r"single_blocks.\1.linear1.\2", 1, 4),
|
||||
r"^single_transformer_blocks\.(\d+)\.attn\.to_v\.(.*)$":
|
||||
(r"single_blocks.\1.linear1.\2", 2, 4),
|
||||
r"^single_transformer_blocks\.(\d+)\.proj_mlp\.(.*)$":
|
||||
(r"single_blocks.\1.linear1.\2", 3, 4),
|
||||
# Corrected: map proj_out to modulation.linear rather than a separate proj_out branch.
|
||||
r"^single_transformer_blocks\.(\d+)\.proj_out\.(.*)$":
|
||||
r"single_blocks.\1.linear2.\2",
|
||||
r"^single_transformer_blocks\.(\d+)\.norm\.linear\.(.*)$":
|
||||
r"single_blocks.\1.modulation.linear.\2",
|
||||
|
||||
# 7. Final layers mapping:
|
||||
r"^norm_out\.linear\.(.*)$":
|
||||
r"final_layer.adaLN_modulation.linear.\1",
|
||||
r"^proj_out\.(.*)$":
|
||||
r"final_layer.linear.\1",
|
||||
}
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
patch_size: int = 2,
|
||||
patch_size_t: int = 1,
|
||||
in_channels: int = 16,
|
||||
out_channels: int = 16,
|
||||
num_attention_heads: int = 24,
|
||||
attention_head_dim: int = 128,
|
||||
mlp_ratio: float = 4.0,
|
||||
num_layers: int = 20,
|
||||
num_single_layers: int = 40,
|
||||
num_refiner_layers: int = 2,
|
||||
rope_axes_dim: Tuple[int, int, int] = (16, 56, 56),
|
||||
guidance_embeds: bool = False,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
text_embed_dim: int = 4096,
|
||||
pooled_projection_dim: int = 768,
|
||||
rope_theta: int = 256,
|
||||
qk_norm: str = "rms_norm", #TODO(PY)
|
||||
prefix="Hunyuan",
|
||||
):
|
||||
super().__init__()
|
||||
hidden_size = attention_head_dim * num_attention_heads
|
||||
self.patch_size = [patch_size_t, patch_size, patch_size]
|
||||
self.in_channels = in_channels
|
||||
self.out_channels = in_channels if out_channels is None else out_channels
|
||||
self.patch_size = [
|
||||
config.patch_size_t, config.patch_size, config.patch_size
|
||||
]
|
||||
self.in_channels = config.in_channels
|
||||
self.num_channels_latents = config.num_channels_latents
|
||||
self.out_channels = config.in_channels if config.out_channels is None else config.out_channels
|
||||
self.unpatchify_channels = self.out_channels
|
||||
self.guidance_embeds = guidance_embeds
|
||||
self.rope_dim_list = list(rope_axes_dim)
|
||||
self.rope_theta = rope_theta
|
||||
self.text_states_dim = text_embed_dim
|
||||
self.text_states_dim_2 = pooled_projection_dim
|
||||
self.guidance_embeds = config.guidance_embeds
|
||||
self.rope_dim_list = list(config.rope_axes_dim)
|
||||
self.rope_theta = config.rope_theta
|
||||
self.text_states_dim = config.text_embed_dim
|
||||
self.text_states_dim_2 = config.pooled_projection_dim
|
||||
# TODO(will): hack?
|
||||
self.dtype = dtype
|
||||
self.dtype = config.dtype
|
||||
|
||||
if hidden_size % num_attention_heads != 0:
|
||||
pe_dim = config.hidden_size // config.num_attention_heads
|
||||
if sum(config.rope_axes_dim) != pe_dim:
|
||||
raise ValueError(
|
||||
f"Hidden size {hidden_size} must be divisible by num_attention_heads {num_attention_heads}"
|
||||
f"Got {config.rope_axes_dim} but expected positional dim {pe_dim}"
|
||||
)
|
||||
|
||||
pe_dim = hidden_size // num_attention_heads
|
||||
if sum(rope_axes_dim) != pe_dim:
|
||||
raise ValueError(
|
||||
f"Got {rope_axes_dim} but expected positional dim {pe_dim}")
|
||||
|
||||
self.hidden_size = hidden_size
|
||||
self.num_attention_heads = num_attention_heads
|
||||
self.hidden_size = config.hidden_size
|
||||
self.num_attention_heads = config.num_attention_heads
|
||||
self.num_channels_latents = config.num_channels_latents
|
||||
|
||||
# Image projection
|
||||
self.img_in = PatchEmbed(self.patch_size,
|
||||
self.in_channels,
|
||||
self.hidden_size,
|
||||
dtype=dtype,
|
||||
prefix=f"{prefix}.img_in")
|
||||
dtype=config.dtype,
|
||||
prefix=f"{config.prefix}.img_in")
|
||||
|
||||
self.txt_in = SingleTokenRefiner(self.text_states_dim,
|
||||
hidden_size,
|
||||
num_attention_heads,
|
||||
depth=num_refiner_layers,
|
||||
dtype=dtype,
|
||||
prefix=f"{prefix}.txt_in")
|
||||
config.hidden_size,
|
||||
config.num_attention_heads,
|
||||
depth=config.num_refiner_layers,
|
||||
dtype=config.dtype,
|
||||
prefix=f"{config.prefix}.txt_in")
|
||||
|
||||
# Time modulation
|
||||
self.time_in = TimestepEmbedder(self.hidden_size,
|
||||
act_layer="silu",
|
||||
dtype=dtype,
|
||||
prefix=f"{prefix}.time_in")
|
||||
dtype=config.dtype,
|
||||
prefix=f"{config.prefix}.time_in")
|
||||
|
||||
# Text modulation
|
||||
self.vector_in = MLP(self.text_states_dim_2,
|
||||
self.hidden_size,
|
||||
self.hidden_size,
|
||||
act_type="silu",
|
||||
dtype=dtype,
|
||||
prefix=f"{prefix}.vector_in")
|
||||
dtype=config.dtype,
|
||||
prefix=f"{config.prefix}.vector_in")
|
||||
|
||||
# Guidance modulation
|
||||
self.guidance_in = (TimestepEmbedder(self.hidden_size,
|
||||
act_layer="silu",
|
||||
dtype=dtype,
|
||||
prefix=f"{prefix}.guidance_in")
|
||||
self.guidance_in = (TimestepEmbedder(
|
||||
self.hidden_size,
|
||||
act_layer="silu",
|
||||
dtype=config.dtype,
|
||||
prefix=f"{config.prefix}.guidance_in")
|
||||
if self.guidance_embeds else None)
|
||||
|
||||
# Double blocks
|
||||
self.double_blocks = nn.ModuleList([
|
||||
MMDoubleStreamBlock(
|
||||
hidden_size,
|
||||
num_attention_heads,
|
||||
mlp_ratio=mlp_ratio,
|
||||
dtype=dtype,
|
||||
config.hidden_size,
|
||||
config.num_attention_heads,
|
||||
mlp_ratio=config.mlp_ratio,
|
||||
dtype=config.dtype,
|
||||
supported_attention_backends=self._supported_attention_backends,
|
||||
prefix=f"{prefix}.double_blocks.{i}") for i in range(num_layers)
|
||||
prefix=f"{config.prefix}.double_blocks.{i}")
|
||||
for i in range(config.num_layers)
|
||||
])
|
||||
|
||||
# Single blocks
|
||||
self.single_blocks = nn.ModuleList([
|
||||
MMSingleStreamBlock(
|
||||
hidden_size,
|
||||
num_attention_heads,
|
||||
mlp_ratio=mlp_ratio,
|
||||
dtype=dtype,
|
||||
config.hidden_size,
|
||||
config.num_attention_heads,
|
||||
mlp_ratio=config.mlp_ratio,
|
||||
dtype=config.dtype,
|
||||
supported_attention_backends=self._supported_attention_backends,
|
||||
prefix=f"{prefix}.single_blocks.{i+num_layers}")
|
||||
for i in range(num_single_layers)
|
||||
prefix=f"{config.prefix}.single_blocks.{i+config.num_layers}")
|
||||
for i in range(config.num_single_layers)
|
||||
])
|
||||
|
||||
self.final_layer = FinalLayer(hidden_size,
|
||||
self.final_layer = FinalLayer(config.hidden_size,
|
||||
self.patch_size,
|
||||
self.out_channels,
|
||||
dtype=dtype,
|
||||
prefix=f"{prefix}.final_layer")
|
||||
dtype=config.dtype,
|
||||
prefix=f"{config.prefix}.final_layer")
|
||||
|
||||
self.__post_init__()
|
||||
|
||||
@@ -688,6 +560,11 @@ class HunyuanVideoTransformer3DModel(BaseDiT):
|
||||
Returns:
|
||||
Tuple of (output)
|
||||
"""
|
||||
forward_context = get_forward_context()
|
||||
forward_batch = forward_context.forward_batch
|
||||
assert forward_batch is not None
|
||||
enable_teacache = forward_batch.enable_teacache
|
||||
|
||||
if guidance is None:
|
||||
guidance = torch.tensor([6016.0],
|
||||
device=hidden_states.device,
|
||||
@@ -735,26 +612,40 @@ class HunyuanVideoTransformer3DModel(BaseDiT):
|
||||
img_seq_len = img.shape[1]
|
||||
|
||||
freqs_cis = (freqs_cos, freqs_sin) if freqs_cos is not None else None
|
||||
# Process through double stream blocks
|
||||
for index, block in enumerate(self.double_blocks):
|
||||
double_block_args = [img, txt, vec, freqs_cis]
|
||||
img, txt = block(*double_block_args)
|
||||
# Merge txt and img to pass through single stream blocks
|
||||
x = torch.cat((img, txt), 1)
|
||||
|
||||
# Process through single stream blocks
|
||||
if len(self.single_blocks) > 0:
|
||||
for index, block in enumerate(self.single_blocks):
|
||||
single_block_args = [
|
||||
x,
|
||||
vec,
|
||||
txt_seq_len,
|
||||
freqs_cis,
|
||||
]
|
||||
x = block(*single_block_args)
|
||||
should_skip_forward = self.should_skip_forward_for_cached_states(
|
||||
img=img, vec=vec)
|
||||
|
||||
if should_skip_forward:
|
||||
img = self.retrieve_cached_states(img)
|
||||
else:
|
||||
if enable_teacache:
|
||||
original_img = img.clone()
|
||||
|
||||
# Process through double stream blocks
|
||||
for index, block in enumerate(self.double_blocks):
|
||||
double_block_args = [img, txt, vec, freqs_cis]
|
||||
img, txt = block(*double_block_args)
|
||||
# Merge txt and img to pass through single stream blocks
|
||||
x = torch.cat((img, txt), 1)
|
||||
|
||||
# Process through single stream blocks
|
||||
if len(self.single_blocks) > 0:
|
||||
for index, block in enumerate(self.single_blocks):
|
||||
single_block_args = [
|
||||
x,
|
||||
vec,
|
||||
txt_seq_len,
|
||||
freqs_cis,
|
||||
]
|
||||
x = block(*single_block_args)
|
||||
|
||||
# Extract image features
|
||||
img = x[:, :img_seq_len, ...]
|
||||
|
||||
if enable_teacache:
|
||||
self.maybe_cache_states(img, original_img)
|
||||
|
||||
# Extract image features
|
||||
img = x[:, :img_seq_len, ...]
|
||||
# Final layer processing
|
||||
img = self.final_layer(img, vec)
|
||||
# Unpatchify to get original shape
|
||||
@@ -762,6 +653,99 @@ class HunyuanVideoTransformer3DModel(BaseDiT):
|
||||
|
||||
return img
|
||||
|
||||
def maybe_cache_states(self, hidden_states: torch.Tensor,
|
||||
original_hidden_states: torch.Tensor) -> None:
|
||||
self.previous_residual = hidden_states - original_hidden_states
|
||||
|
||||
def should_skip_forward_for_cached_states(self, **kwargs) -> bool:
|
||||
|
||||
forward_context = get_forward_context()
|
||||
forward_batch = forward_context.forward_batch
|
||||
assert forward_batch is not None
|
||||
current_timestep = forward_context.current_timestep
|
||||
enable_teacache = forward_batch.enable_teacache
|
||||
|
||||
if not enable_teacache:
|
||||
return False
|
||||
raise NotImplementedError(
|
||||
"teacache is not supported yet for HunyuanVideo")
|
||||
|
||||
teacache_params = forward_batch.teacache_params
|
||||
assert teacache_params is not None, "teacache_params is not initialized"
|
||||
assert isinstance(
|
||||
teacache_params,
|
||||
TeaCacheParams), "teacache_params is not a TeaCacheParams"
|
||||
num_inference_steps = forward_batch.num_inference_steps
|
||||
teache_thresh = teacache_params.teacache_thresh
|
||||
|
||||
coefficients = teacache_params.coefficients
|
||||
|
||||
if current_timestep == 0:
|
||||
self.cnt = 0
|
||||
|
||||
inp = kwargs["img"].clone()
|
||||
vec_ = kwargs["vec"].clone()
|
||||
# convert to DTensor
|
||||
vec_ = torch.distributed.tensor.DTensor.from_local(
|
||||
vec_,
|
||||
torch.distributed.DeviceMesh(
|
||||
"cuda",
|
||||
list(range(get_sequence_model_parallel_world_size())),
|
||||
mesh_dim_names=("dp", )),
|
||||
[torch.distributed.tensor.Replicate()])
|
||||
|
||||
inp = torch.distributed.tensor.DTensor.from_local(
|
||||
inp,
|
||||
torch.distributed.DeviceMesh(
|
||||
"cuda",
|
||||
list(range(get_sequence_model_parallel_world_size())),
|
||||
mesh_dim_names=("dp", )),
|
||||
[torch.distributed.tensor.Replicate()])
|
||||
|
||||
# txt_ = kwargs["txt"].clone()
|
||||
|
||||
# inp = img.clone()
|
||||
# vec_ = vec.clone()
|
||||
# txt_ = txt.clone()
|
||||
(
|
||||
img_mod1_shift,
|
||||
img_mod1_scale,
|
||||
img_mod1_gate,
|
||||
img_mod2_shift,
|
||||
img_mod2_scale,
|
||||
img_mod2_gate,
|
||||
) = self.double_blocks[0].img_mod(vec_).chunk(6, dim=-1)
|
||||
normed_inp = self.double_blocks[0].img_attn_norm.norm(inp)
|
||||
modulated_inp = modulate(normed_inp,
|
||||
shift=img_mod1_shift,
|
||||
scale=img_mod1_scale)
|
||||
if self.cnt == 0 or self.cnt == num_inference_steps - 1:
|
||||
should_calc = True
|
||||
self.accumulated_rel_l1_distance = 0
|
||||
else:
|
||||
coefficients = [
|
||||
7.33226126e+02, -4.01131952e+02, 6.75869174e+01,
|
||||
-3.14987800e+00, 9.61237896e-02
|
||||
]
|
||||
rescale_func = np.poly1d(coefficients)
|
||||
assert self.previous_modulated_input is not None, "previous_modulated_input is not initialized"
|
||||
self.accumulated_rel_l1_distance += rescale_func(
|
||||
((modulated_inp - self.previous_modulated_input).abs().mean() /
|
||||
self.previous_modulated_input.abs().mean()).cpu().item())
|
||||
if self.accumulated_rel_l1_distance < teache_thresh:
|
||||
should_calc = False
|
||||
else:
|
||||
should_calc = True
|
||||
self.accumulated_rel_l1_distance = 0
|
||||
self.previous_modulated_input = modulated_inp
|
||||
self.cnt += 1
|
||||
|
||||
return not should_calc
|
||||
|
||||
def retrieve_cached_states(self,
|
||||
hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
return hidden_states + self.previous_residual
|
||||
|
||||
|
||||
class SingleTokenRefiner(nn.Module):
|
||||
"""
|
||||
|
||||
@@ -3,12 +3,16 @@
|
||||
import math
|
||||
from typing import List, Optional, Tuple, Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from fastvideo.v1.attention import DistributedAttention, LocalAttention
|
||||
from fastvideo.v1.configs.models.dits import WanVideoConfig
|
||||
from fastvideo.v1.configs.sample.wan import WanTeaCacheParams
|
||||
from fastvideo.v1.distributed.parallel_state import (
|
||||
get_sequence_model_parallel_world_size)
|
||||
from fastvideo.v1.forward_context import get_forward_context
|
||||
from fastvideo.v1.layers.layernorm import (LayerNormScaleShift, RMSNorm,
|
||||
ScaleResidual,
|
||||
ScaleResidualLayerNormScaleShift)
|
||||
@@ -20,7 +24,7 @@ from fastvideo.v1.layers.rotary_embedding import (_apply_rotary_emb,
|
||||
get_rotary_pos_embed)
|
||||
from fastvideo.v1.layers.visual_embedding import (ModulateProjection,
|
||||
PatchEmbed, TimestepEmbedder)
|
||||
from fastvideo.v1.models.dits.base import BaseDiT
|
||||
from fastvideo.v1.models.dits.base import CachableDiT
|
||||
from fastvideo.v1.platforms import _Backend
|
||||
|
||||
|
||||
@@ -349,115 +353,61 @@ class WanTransformerBlock(nn.Module):
|
||||
return hidden_states
|
||||
|
||||
|
||||
class WanTransformer3DModel(BaseDiT):
|
||||
_fsdp_shard_conditions = [
|
||||
lambda n, m: "blocks" in n and str.isdigit(n.split(".")[-1]),
|
||||
]
|
||||
_supported_attention_backends = (_Backend.SLIDING_TILE_ATTN,
|
||||
_Backend.FLASH_ATTN, _Backend.TORCH_SDPA)
|
||||
_param_names_mapping = {
|
||||
r"^patch_embedding\.(.*)$":
|
||||
r"patch_embedding.proj.\1",
|
||||
r"^condition_embedder\.text_embedder\.linear_1\.(.*)$":
|
||||
r"condition_embedder.text_embedder.fc_in.\1",
|
||||
r"^condition_embedder\.text_embedder\.linear_2\.(.*)$":
|
||||
r"condition_embedder.text_embedder.fc_out.\1",
|
||||
r"^condition_embedder\.time_embedder\.linear_1\.(.*)$":
|
||||
r"condition_embedder.time_embedder.mlp.fc_in.\1",
|
||||
r"^condition_embedder\.time_embedder\.linear_2\.(.*)$":
|
||||
r"condition_embedder.time_embedder.mlp.fc_out.\1",
|
||||
r"^condition_embedder\.time_proj\.(.*)$":
|
||||
r"condition_embedder.time_modulation.linear.\1",
|
||||
r"^condition_embedder\.image_embedder\.ff\.net\.0\.proj\.(.*)$":
|
||||
r"condition_embedder.image_embedder.ff.fc_in.\1",
|
||||
r"^condition_embedder\.image_embedder\.ff\.net\.2\.(.*)$":
|
||||
r"condition_embedder.image_embedder.ff.fc_out.\1",
|
||||
r"^blocks\.(\d+)\.attn1\.to_q\.(.*)$":
|
||||
r"blocks.\1.to_q.\2",
|
||||
r"^blocks\.(\d+)\.attn1\.to_k\.(.*)$":
|
||||
r"blocks.\1.to_k.\2",
|
||||
r"^blocks\.(\d+)\.attn1\.to_v\.(.*)$":
|
||||
r"blocks.\1.to_v.\2",
|
||||
r"^blocks\.(\d+)\.attn1\.to_out\.0\.(.*)$":
|
||||
r"blocks.\1.to_out.\2",
|
||||
r"^blocks\.(\d+)\.attn1\.norm_q\.(.*)$":
|
||||
r"blocks.\1.norm_q.\2",
|
||||
r"^blocks\.(\d+)\.attn1\.norm_k\.(.*)$":
|
||||
r"blocks.\1.norm_k.\2",
|
||||
r"^blocks\.(\d+)\.attn2\.to_out\.0\.(.*)$":
|
||||
r"blocks.\1.attn2.to_out.\2",
|
||||
r"^blocks\.(\d+)\.ffn\.net\.0\.proj\.(.*)$":
|
||||
r"blocks.\1.ffn.fc_in.\2",
|
||||
r"^blocks\.(\d+)\.ffn\.net\.2\.(.*)$":
|
||||
r"blocks.\1.ffn.fc_out.\2",
|
||||
r"blocks\.(\d+)\.norm2\.(.*)$":
|
||||
r"blocks.\1.self_attn_residual_norm.norm.\2",
|
||||
}
|
||||
class WanTransformer3DModel(CachableDiT):
|
||||
_fsdp_shard_conditions = WanVideoConfig()._fsdp_shard_conditions
|
||||
_compile_conditions = WanVideoConfig()._compile_conditions
|
||||
_supported_attention_backends = WanVideoConfig(
|
||||
)._supported_attention_backends
|
||||
_param_names_mapping = WanVideoConfig()._param_names_mapping
|
||||
|
||||
def __init__(self,
|
||||
patch_size: Tuple[int, int, int] = (1, 2, 2),
|
||||
text_len=512,
|
||||
num_attention_heads: int = 40,
|
||||
attention_head_dim: int = 128,
|
||||
in_channels: int = 16,
|
||||
out_channels: int = 16,
|
||||
text_dim: int = 4096,
|
||||
freq_dim: int = 256,
|
||||
ffn_dim: int = 13824,
|
||||
num_layers: int = 40,
|
||||
cross_attn_norm: bool = True,
|
||||
qk_norm: str = "rms_norm_across_heads",
|
||||
eps: float = 1e-6,
|
||||
image_dim: Optional[int] = None,
|
||||
added_kv_proj_dim: Optional[int] = None,
|
||||
rope_max_seq_len: int = 1024,
|
||||
prefix="Wan") -> None:
|
||||
super().__init__()
|
||||
def __init__(self, config: WanVideoConfig) -> None:
|
||||
super().__init__(config=config)
|
||||
|
||||
inner_dim = num_attention_heads * attention_head_dim
|
||||
self.hidden_size = inner_dim
|
||||
self.num_attention_heads = num_attention_heads
|
||||
self.in_channels = in_channels
|
||||
self.out_channels = out_channels or in_channels
|
||||
self.patch_size = patch_size
|
||||
self.text_len = text_len
|
||||
inner_dim = config.num_attention_heads * config.attention_head_dim
|
||||
self.hidden_size = config.hidden_size
|
||||
self.num_attention_heads = config.num_attention_heads
|
||||
self.in_channels = config.in_channels
|
||||
self.out_channels = config.out_channels
|
||||
self.num_channels_latents = config.num_channels_latents
|
||||
self.patch_size = config.patch_size
|
||||
self.text_len = config.text_len
|
||||
|
||||
# 1. Patch & position embedding
|
||||
self.patch_embedding = PatchEmbed(in_chans=in_channels,
|
||||
self.patch_embedding = PatchEmbed(in_chans=config.in_channels,
|
||||
embed_dim=inner_dim,
|
||||
patch_size=patch_size,
|
||||
patch_size=config.patch_size,
|
||||
flatten=False)
|
||||
|
||||
# 2. Condition embeddings
|
||||
self.condition_embedder = WanTimeTextImageEmbedding(
|
||||
dim=inner_dim,
|
||||
time_freq_dim=freq_dim,
|
||||
text_embed_dim=text_dim,
|
||||
image_embed_dim=image_dim,
|
||||
time_freq_dim=config.freq_dim,
|
||||
text_embed_dim=config.text_dim,
|
||||
image_embed_dim=config.image_dim,
|
||||
)
|
||||
|
||||
# 3. Transformer blocks
|
||||
self.blocks = nn.ModuleList([
|
||||
WanTransformerBlock(inner_dim,
|
||||
ffn_dim,
|
||||
num_attention_heads,
|
||||
qk_norm,
|
||||
cross_attn_norm,
|
||||
eps,
|
||||
added_kv_proj_dim,
|
||||
config.ffn_dim,
|
||||
config.num_attention_heads,
|
||||
config.qk_norm,
|
||||
config.cross_attn_norm,
|
||||
config.eps,
|
||||
config.added_kv_proj_dim,
|
||||
self._supported_attention_backends,
|
||||
prefix=f"{prefix}.blocks.{i}")
|
||||
for i in range(num_layers)
|
||||
prefix=f"{config.prefix}.blocks.{i}")
|
||||
for i in range(config.num_layers)
|
||||
])
|
||||
|
||||
# 4. Output norm & projection
|
||||
self.norm_out = LayerNormScaleShift(inner_dim,
|
||||
norm_type="layer",
|
||||
eps=eps,
|
||||
eps=config.eps,
|
||||
elementwise_affine=False,
|
||||
dtype=torch.float32)
|
||||
self.proj_out = nn.Linear(inner_dim,
|
||||
out_channels * math.prod(patch_size))
|
||||
self.proj_out = nn.Linear(
|
||||
inner_dim, config.out_channels * math.prod(config.patch_size))
|
||||
self.scale_shift_table = nn.Parameter(
|
||||
torch.randn(1, 2, inner_dim) / inner_dim**0.5)
|
||||
|
||||
@@ -473,6 +423,10 @@ class WanTransformer3DModel(BaseDiT):
|
||||
torch.Tensor, List[torch.Tensor]]] = None,
|
||||
guidance=None,
|
||||
**kwargs) -> torch.Tensor:
|
||||
forward_batch = get_forward_context().forward_batch
|
||||
assert forward_batch is not None
|
||||
enable_teacache = forward_batch.enable_teacache
|
||||
|
||||
orig_dtype = hidden_states.dtype
|
||||
if not isinstance(encoder_hidden_states, torch.Tensor):
|
||||
encoder_hidden_states = encoder_hidden_states[0]
|
||||
@@ -516,16 +470,32 @@ class WanTransformer3DModel(BaseDiT):
|
||||
[encoder_hidden_states_image, encoder_hidden_states], dim=1)
|
||||
|
||||
assert encoder_hidden_states.dtype == orig_dtype
|
||||
|
||||
# 4. Transformer blocks
|
||||
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
||||
for block in self.blocks:
|
||||
hidden_states = self._gradient_checkpointing_func(
|
||||
block, hidden_states, encoder_hidden_states, timestep_proj,
|
||||
freqs_cis)
|
||||
# if caching is enabled, we might be able to skip the forward pass
|
||||
should_skip_forward = self.should_skip_forward_for_cached_states(
|
||||
timestep_proj=timestep_proj, temb=temb)
|
||||
|
||||
if should_skip_forward:
|
||||
hidden_states = self.retrieve_cached_states(hidden_states)
|
||||
else:
|
||||
for block in self.blocks:
|
||||
hidden_states = block(hidden_states, encoder_hidden_states,
|
||||
timestep_proj, freqs_cis)
|
||||
# if teacache is enabled, we need to cache the original hidden states
|
||||
if enable_teacache:
|
||||
original_hidden_states = hidden_states.clone()
|
||||
|
||||
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
||||
for block in self.blocks:
|
||||
hidden_states = self._gradient_checkpointing_func(
|
||||
block, hidden_states, encoder_hidden_states,
|
||||
timestep_proj, freqs_cis)
|
||||
else:
|
||||
for block in self.blocks:
|
||||
hidden_states = block(hidden_states, encoder_hidden_states,
|
||||
timestep_proj, freqs_cis)
|
||||
|
||||
# if teacache is enabled, we need to cache the original hidden states
|
||||
if enable_teacache:
|
||||
self.maybe_cache_states(hidden_states, original_hidden_states)
|
||||
|
||||
# 5. Output norm, projection & unpatchify
|
||||
shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2,
|
||||
@@ -541,3 +511,96 @@ class WanTransformer3DModel(BaseDiT):
|
||||
output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
|
||||
|
||||
return output
|
||||
|
||||
def maybe_cache_states(self, hidden_states: torch.Tensor,
|
||||
original_hidden_states: torch.Tensor) -> None:
|
||||
if self.is_even:
|
||||
self.previous_residual_even = hidden_states.squeeze(
|
||||
0) - original_hidden_states
|
||||
else:
|
||||
self.previous_residual_odd = hidden_states.squeeze(
|
||||
0) - original_hidden_states
|
||||
|
||||
def should_skip_forward_for_cached_states(self, **kwargs) -> bool:
|
||||
|
||||
forward_context = get_forward_context()
|
||||
forward_batch = forward_context.forward_batch
|
||||
assert forward_batch is not None
|
||||
if not forward_batch.enable_teacache:
|
||||
return False
|
||||
teacache_params = forward_batch.teacache_params
|
||||
assert teacache_params is not None, "teacache_params is not initialized"
|
||||
assert isinstance(
|
||||
teacache_params,
|
||||
WanTeaCacheParams), "teacache_params is not a WanTeaCacheParams"
|
||||
current_timestep = forward_context.current_timestep
|
||||
num_inference_steps = forward_batch.num_inference_steps
|
||||
|
||||
# initialize the coefficients, cutoff_steps, and ret_steps
|
||||
coefficients = teacache_params.coefficients
|
||||
use_ret_steps = teacache_params.use_ret_steps
|
||||
cutoff_steps = teacache_params.get_cutoff_steps(num_inference_steps)
|
||||
ret_steps = teacache_params.ret_steps
|
||||
teacache_thresh = teacache_params.teacache_thresh
|
||||
|
||||
if current_timestep == 0:
|
||||
self.cnt = 0
|
||||
|
||||
timestep_proj = kwargs["timestep_proj"]
|
||||
temb = kwargs["temb"]
|
||||
modulated_inp = timestep_proj if use_ret_steps else temb
|
||||
|
||||
if self.cnt % 2 == 0: # even -> condition
|
||||
self.is_even = True
|
||||
if self.cnt < ret_steps or self.cnt >= cutoff_steps:
|
||||
self.should_calc_even = True
|
||||
self.accumulated_rel_l1_distance_even = 0
|
||||
else:
|
||||
assert self.previous_e0_even is not None, "previous_e0_even is not initialized"
|
||||
assert self.accumulated_rel_l1_distance_even is not None, "accumulated_rel_l1_distance_even is not initialized"
|
||||
rescale_func = np.poly1d(coefficients)
|
||||
self.accumulated_rel_l1_distance_even += rescale_func(
|
||||
((modulated_inp - self.previous_e0_even).abs().mean() /
|
||||
self.previous_e0_even.abs().mean()).cpu().item())
|
||||
if self.accumulated_rel_l1_distance_even < teacache_thresh:
|
||||
self.should_calc_even = False
|
||||
else:
|
||||
self.should_calc_even = True
|
||||
self.accumulated_rel_l1_distance_even = 0
|
||||
self.previous_e0_even = modulated_inp.clone()
|
||||
|
||||
else: # odd -> unconditon
|
||||
self.is_even = False
|
||||
if self.cnt < ret_steps or self.cnt >= cutoff_steps:
|
||||
self.should_calc_odd = True
|
||||
self.accumulated_rel_l1_distance_odd = 0
|
||||
else:
|
||||
assert self.previous_e0_odd is not None, "previous_e0_odd is not initialized"
|
||||
assert self.accumulated_rel_l1_distance_odd is not None, "accumulated_rel_l1_distance_odd is not initialized"
|
||||
rescale_func = np.poly1d(coefficients)
|
||||
self.accumulated_rel_l1_distance_odd += rescale_func(
|
||||
((modulated_inp - self.previous_e0_odd).abs().mean() /
|
||||
self.previous_e0_odd.abs().mean()).cpu().item())
|
||||
if self.accumulated_rel_l1_distance_odd < teacache_thresh:
|
||||
self.should_calc_odd = False
|
||||
else:
|
||||
self.should_calc_odd = True
|
||||
self.accumulated_rel_l1_distance_odd = 0
|
||||
self.previous_e0_odd = modulated_inp.clone()
|
||||
self.cnt += 1
|
||||
should_skip_forward = False
|
||||
if self.is_even:
|
||||
if not self.should_calc_even:
|
||||
should_skip_forward = True
|
||||
else:
|
||||
if not self.should_calc_odd:
|
||||
should_skip_forward = True
|
||||
|
||||
return should_skip_forward
|
||||
|
||||
def retrieve_cached_states(self,
|
||||
hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
if self.is_even:
|
||||
return hidden_states + self.previous_residual_even
|
||||
else:
|
||||
return hidden_states + self.previous_residual_odd
|
||||
|
||||
@@ -1,22 +1,57 @@
|
||||
from typing import Tuple
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Optional, Tuple
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from fastvideo.v1.configs.models.encoders import (BaseEncoderOutput,
|
||||
ImageEncoderConfig,
|
||||
TextEncoderConfig)
|
||||
from fastvideo.v1.platforms import _Backend
|
||||
|
||||
|
||||
class BaseEncoder(nn.Module):
|
||||
_supported_attention_backends: Tuple[_Backend,
|
||||
...] = (_Backend.TORCH_SDPA, )
|
||||
class TextEncoder(nn.Module, ABC):
|
||||
_supported_attention_backends: Tuple[
|
||||
_Backend, ...] = TextEncoderConfig()._supported_attention_backends
|
||||
|
||||
def __init__(self, *args, **kwargs) -> None:
|
||||
def __init__(self, config: TextEncoderConfig) -> None:
|
||||
super().__init__()
|
||||
self.config = config
|
||||
if not self.supported_attention_backends:
|
||||
raise ValueError(
|
||||
f"Subclass {self.__class__.__name__} must define _supported_attention_backends"
|
||||
)
|
||||
|
||||
def forward(self, *args, **kwargs):
|
||||
@abstractmethod
|
||||
def forward(self,
|
||||
input_ids: Optional[torch.Tensor],
|
||||
position_ids: Optional[torch.Tensor] = None,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
inputs_embeds: Optional[torch.Tensor] = None,
|
||||
output_hidden_states: Optional[bool] = None,
|
||||
**kwargs) -> BaseEncoderOutput:
|
||||
pass
|
||||
|
||||
@property
|
||||
def supported_attention_backends(self) -> Tuple[_Backend, ...]:
|
||||
return self._supported_attention_backends
|
||||
|
||||
|
||||
class ImageEncoder(nn.Module, ABC):
|
||||
_supported_attention_backends: Tuple[
|
||||
_Backend, ...] = ImageEncoderConfig()._supported_attention_backends
|
||||
|
||||
def __init__(self, config: ImageEncoderConfig) -> None:
|
||||
super().__init__()
|
||||
self.config = config
|
||||
if not self.supported_attention_backends:
|
||||
raise ValueError(
|
||||
f"Subclass {self.__class__.__name__} must define _supported_attention_backends"
|
||||
)
|
||||
|
||||
@abstractmethod
|
||||
def forward(self, pixel_values: torch.Tensor,
|
||||
**kwargs) -> BaseEncoderOutput:
|
||||
pass
|
||||
|
||||
@property
|
||||
|
||||
@@ -3,62 +3,32 @@
|
||||
# Adapted from transformers: https://github.com/huggingface/transformers/blob/v4.39.0/src/transformers/models/clip/modeling_clip.py
|
||||
"""Minimal implementation of CLIPVisionModel intended to be only used
|
||||
within a vision language model."""
|
||||
from typing import Iterable, Optional, Set, Tuple, Union, cast
|
||||
from typing import Iterable, Optional, Set, Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from transformers import CLIPTextConfig, CLIPVisionConfig
|
||||
from transformers.modeling_outputs import BaseModelOutputWithPooling
|
||||
from vllm.model_executor.models.interfaces import SupportsQuant
|
||||
|
||||
# from transformers.modeling_attn_mask_utils import _create_4d_causal_attention_mask, _prepare_4d_attention_mask
|
||||
from fastvideo.v1.attention import LocalAttention
|
||||
from fastvideo.v1.configs.models.encoders import (BaseEncoderOutput,
|
||||
CLIPTextConfig,
|
||||
CLIPVisionConfig)
|
||||
from fastvideo.v1.configs.quantization import QuantizationConfig
|
||||
from fastvideo.v1.distributed import (divide,
|
||||
get_tensor_model_parallel_world_size)
|
||||
from fastvideo.v1.layers.activation import get_act_fn
|
||||
from fastvideo.v1.layers.linear import (ColumnParallelLinear, QKVParallelLinear,
|
||||
RowParallelLinear)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.encoders.base import BaseEncoder
|
||||
from fastvideo.v1.models.encoders.vision import (VisionEncoderInfo,
|
||||
resolve_visual_encoder_outputs)
|
||||
from fastvideo.v1.models.encoders.base import ImageEncoder, TextEncoder
|
||||
from fastvideo.v1.models.encoders.vision import resolve_visual_encoder_outputs
|
||||
# TODO: support quantization
|
||||
# from vllm.model_executor.layers.quantization import QuantizationConfig
|
||||
from fastvideo.v1.models.loader.weight_utils import default_weight_loader
|
||||
from fastvideo.v1.platforms import _Backend
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class QuantizationConfig:
|
||||
pass
|
||||
|
||||
|
||||
class CLIPEncoderInfo(VisionEncoderInfo[CLIPVisionConfig]):
|
||||
|
||||
def get_num_image_tokens(
|
||||
self,
|
||||
*,
|
||||
image_width: int,
|
||||
image_height: int,
|
||||
) -> int:
|
||||
return self.get_patch_grid_length()**2 + 1
|
||||
|
||||
def get_max_image_tokens(self) -> int:
|
||||
return self.get_patch_grid_length()**2 + 1
|
||||
|
||||
def get_image_size(self) -> int:
|
||||
return cast(int, self.vision_config.image_size)
|
||||
|
||||
def get_patch_size(self) -> int:
|
||||
return cast(int, self.vision_config.patch_size)
|
||||
|
||||
def get_patch_grid_length(self) -> int:
|
||||
image_size, patch_size = self.get_image_size(), self.get_patch_size()
|
||||
assert image_size % patch_size == 0
|
||||
return image_size // patch_size
|
||||
|
||||
|
||||
# Adapted from https://github.com/huggingface/transformers/blob/v4.39.0/src/transformers/models/clip/modeling_clip.py#L164 # noqa
|
||||
class CLIPVisionEmbeddings(nn.Module):
|
||||
|
||||
@@ -158,7 +128,7 @@ class CLIPAttention(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: CLIPVisionConfig,
|
||||
config: Union[CLIPVisionConfig, CLIPTextConfig],
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
@@ -193,13 +163,13 @@ class CLIPAttention(nn.Module):
|
||||
self.tp_size = get_tensor_model_parallel_world_size()
|
||||
self.num_heads_per_partition = divide(self.num_heads, self.tp_size)
|
||||
|
||||
self.attn = LocalAttention(self.num_heads_per_partition,
|
||||
self.head_dim,
|
||||
self.num_heads_per_partition,
|
||||
softmax_scale=self.scale,
|
||||
causal=True,
|
||||
supported_attention_backends=self.config.
|
||||
supported_attention_backends)
|
||||
self.attn = LocalAttention(
|
||||
self.num_heads_per_partition,
|
||||
self.head_dim,
|
||||
self.num_heads_per_partition,
|
||||
softmax_scale=self.scale,
|
||||
causal=True,
|
||||
supported_attention_backends=config._supported_attention_backends)
|
||||
|
||||
def _shape(self, tensor: torch.Tensor, seq_len: int, bsz: int):
|
||||
return tensor.view(bsz, seq_len, self.num_heads,
|
||||
@@ -239,7 +209,7 @@ class CLIPMLP(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: CLIPVisionConfig,
|
||||
config: Union[CLIPVisionConfig, CLIPTextConfig],
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = "",
|
||||
) -> None:
|
||||
@@ -269,7 +239,7 @@ class CLIPEncoderLayer(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: CLIPTextConfig,
|
||||
config: Union[CLIPTextConfig, CLIPVisionConfig],
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = "",
|
||||
) -> None:
|
||||
@@ -314,7 +284,7 @@ class CLIPEncoder(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: CLIPVisionConfig,
|
||||
config: Union[CLIPVisionConfig, CLIPTextConfig],
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
num_hidden_layers_override: Optional[int] = None,
|
||||
prefix: str = "",
|
||||
@@ -356,7 +326,6 @@ class CLIPTextTransformer(nn.Module):
|
||||
def __init__(self,
|
||||
config: CLIPTextConfig,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
*,
|
||||
num_hidden_layers_override: Optional[int] = None,
|
||||
prefix: str = ""):
|
||||
super().__init__()
|
||||
@@ -377,27 +346,21 @@ class CLIPTextTransformer(nn.Module):
|
||||
# For `pooled_output` computation
|
||||
self.eos_token_id = config.eos_token_id
|
||||
|
||||
# For attention mask, it differs between `flash_attention_2` and other attention implementations
|
||||
self._use_flash_attention_2 = config._attn_implementation == "flash_attention_2"
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_ids: Optional[torch.Tensor] = None,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
input_ids: Optional[torch.Tensor],
|
||||
position_ids: Optional[torch.Tensor] = None,
|
||||
output_attentions: Optional[bool] = None,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
inputs_embeds: Optional[torch.Tensor] = None,
|
||||
output_hidden_states: Optional[bool] = None,
|
||||
return_dict: Optional[bool] = None,
|
||||
) -> Union[Tuple, BaseModelOutputWithPooling]:
|
||||
) -> BaseEncoderOutput:
|
||||
r"""
|
||||
Returns:
|
||||
|
||||
"""
|
||||
output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
|
||||
output_hidden_states = (output_hidden_states
|
||||
if output_hidden_states is not None else
|
||||
self.config.output_hidden_states)
|
||||
return_dict = return_dict if return_dict is not None else self.config.use_return_dict
|
||||
|
||||
if input_ids is None:
|
||||
raise ValueError("You have to specify input_ids")
|
||||
@@ -456,11 +419,7 @@ class CLIPTextTransformer(nn.Module):
|
||||
) == self.eos_token_id).int().argmax(dim=-1),
|
||||
]
|
||||
|
||||
if not return_dict:
|
||||
return (last_hidden_state, pooled_output) + encoder_outputs[1:]
|
||||
|
||||
# return last_hidden_state
|
||||
return BaseModelOutputWithPooling(
|
||||
return BaseEncoderOutput(
|
||||
last_hidden_state=last_hidden_state,
|
||||
pooler_output=pooled_output,
|
||||
hidden_states=encoder_outputs,
|
||||
@@ -468,42 +427,34 @@ class CLIPTextTransformer(nn.Module):
|
||||
)
|
||||
|
||||
|
||||
class CLIPTextModel(BaseEncoder):
|
||||
_supported_attention_backends = (_Backend.FLASH_ATTN, _Backend.TORCH_SDPA)
|
||||
class CLIPTextModel(TextEncoder):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: CLIPTextConfig,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = "",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
self.config = config
|
||||
self.config.supported_attention_backends = self._supported_attention_backends
|
||||
super().__init__(config)
|
||||
self.text_model = CLIPTextTransformer(config=config,
|
||||
quant_config=quant_config,
|
||||
prefix=prefix)
|
||||
quant_config=config.quant_config,
|
||||
prefix=config.prefix)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_ids: Optional[torch.Tensor] = None,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
input_ids: Optional[torch.Tensor],
|
||||
position_ids: Optional[torch.Tensor] = None,
|
||||
output_attentions: Optional[bool] = None,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
inputs_embeds: Optional[torch.Tensor] = None,
|
||||
output_hidden_states: Optional[bool] = None,
|
||||
return_dict: Optional[bool] = None,
|
||||
) -> Union[Tuple, BaseModelOutputWithPooling]:
|
||||
return_dict = return_dict if return_dict is not None else self.config.use_return_dict
|
||||
**kwargs,
|
||||
) -> BaseEncoderOutput:
|
||||
|
||||
return self.text_model(
|
||||
outputs: BaseEncoderOutput = self.text_model(
|
||||
input_ids=input_ids,
|
||||
attention_mask=attention_mask,
|
||||
position_ids=position_ids,
|
||||
output_attentions=output_attentions,
|
||||
output_hidden_states=output_hidden_states,
|
||||
return_dict=None,
|
||||
)
|
||||
return outputs
|
||||
|
||||
def load_weights(self, weights: Iterable[Tuple[str,
|
||||
torch.Tensor]]) -> Set[str]:
|
||||
@@ -548,7 +499,6 @@ class CLIPVisionTransformer(nn.Module):
|
||||
self,
|
||||
config: CLIPVisionConfig,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
*,
|
||||
num_hidden_layers_override: Optional[int] = None,
|
||||
require_post_norm: Optional[bool] = None,
|
||||
prefix: str = "",
|
||||
@@ -615,37 +565,29 @@ class CLIPVisionTransformer(nn.Module):
|
||||
return encoder_outputs
|
||||
|
||||
|
||||
class CLIPVisionModel(BaseEncoder, SupportsQuant):
|
||||
class CLIPVisionModel(ImageEncoder):
|
||||
config_class = CLIPVisionConfig
|
||||
main_input_name = "pixel_values"
|
||||
packed_modules_mapping = {"qkv_proj": ["q_proj", "k_proj", "v_proj"]}
|
||||
_supported_attention_backends = (_Backend.FLASH_ATTN, _Backend.TORCH_SDPA)
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: CLIPVisionConfig,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
*,
|
||||
num_hidden_layers_override: Optional[int] = None,
|
||||
require_post_norm: Optional[bool] = None,
|
||||
prefix: str = "",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.config = config
|
||||
self.config.supported_attention_backends = self._supported_attention_backends
|
||||
def __init__(self, config: CLIPVisionConfig) -> None:
|
||||
super().__init__(config)
|
||||
self.vision_model = CLIPVisionTransformer(
|
||||
config=config,
|
||||
quant_config=quant_config,
|
||||
num_hidden_layers_override=num_hidden_layers_override,
|
||||
require_post_norm=require_post_norm,
|
||||
prefix=f"{prefix}.vision_model")
|
||||
quant_config=config.quant_config,
|
||||
num_hidden_layers_override=config.num_hidden_layers_override,
|
||||
require_post_norm=config.require_post_norm,
|
||||
prefix=f"{config.prefix}.vision_model")
|
||||
|
||||
def forward(
|
||||
self,
|
||||
pixel_values: torch.Tensor,
|
||||
feature_sample_layers: Optional[list[int]] = None,
|
||||
) -> torch.Tensor:
|
||||
return self.vision_model(pixel_values, feature_sample_layers)
|
||||
**kwargs,
|
||||
) -> BaseEncoderOutput:
|
||||
last_hidden_state = self.vision_model(pixel_values,
|
||||
feature_sample_layers)
|
||||
return BaseEncoderOutput(last_hidden_state=last_hidden_state)
|
||||
|
||||
@property
|
||||
def device(self):
|
||||
|
||||
@@ -23,15 +23,16 @@
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
"""Inference-only LLaMA model compatible with HuggingFace weights."""
|
||||
from typing import Any, Dict, Iterable, Optional, Set, Tuple, Type
|
||||
from typing import Any, Dict, Iterable, Optional, Set, Tuple
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
from transformers import LlamaConfig
|
||||
from transformers.modeling_outputs import BaseModelOutputWithPast
|
||||
|
||||
# from vllm.model_executor.layers.quantization import QuantizationConfig
|
||||
from fastvideo.v1.attention import LocalAttention
|
||||
# from ..utils import (extract_layer_index)
|
||||
from fastvideo.v1.configs.models.encoders import BaseEncoderOutput, LlamaConfig
|
||||
from fastvideo.v1.configs.quantization import QuantizationConfig
|
||||
from fastvideo.v1.distributed import get_tensor_model_parallel_world_size
|
||||
from fastvideo.v1.layers.activation import SiluAndMul
|
||||
from fastvideo.v1.layers.layernorm import RMSNorm
|
||||
@@ -39,15 +40,9 @@ from fastvideo.v1.layers.linear import (MergedColumnParallelLinear,
|
||||
QKVParallelLinear, RowParallelLinear)
|
||||
from fastvideo.v1.layers.rotary_embedding import get_rope
|
||||
from fastvideo.v1.layers.vocab_parallel_embedding import VocabParallelEmbedding
|
||||
from fastvideo.v1.models.encoders.base import BaseEncoder
|
||||
from fastvideo.v1.models.encoders.base import TextEncoder
|
||||
from fastvideo.v1.models.loader.weight_utils import (default_weight_loader,
|
||||
maybe_remap_kv_scale_name)
|
||||
# from ..utils import (extract_layer_index)
|
||||
from fastvideo.v1.platforms import _Backend
|
||||
|
||||
|
||||
class QuantizationConfig:
|
||||
pass
|
||||
|
||||
|
||||
class LlamaMLP(nn.Module):
|
||||
@@ -171,7 +166,7 @@ class LlamaAttention(nn.Module):
|
||||
self.num_kv_heads,
|
||||
softmax_scale=self.scaling,
|
||||
causal=True,
|
||||
supported_attention_backends=config.supported_attention_backends)
|
||||
supported_attention_backends=config._supported_attention_backends)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
@@ -279,28 +274,23 @@ class LlamaDecoderLayer(nn.Module):
|
||||
return hidden_states, residual
|
||||
|
||||
|
||||
class LlamaModel(BaseEncoder):
|
||||
_supported_attention_backends = (_Backend.FLASH_ATTN, _Backend.TORCH_SDPA)
|
||||
class LlamaModel(TextEncoder):
|
||||
|
||||
def __init__(self,
|
||||
config: LlamaConfig,
|
||||
prefix: str = "",
|
||||
layer_type: Type[LlamaDecoderLayer] = LlamaDecoderLayer):
|
||||
super().__init__()
|
||||
|
||||
quant_config = None
|
||||
lora_config = None
|
||||
def __init__(
|
||||
self,
|
||||
config: LlamaConfig,
|
||||
):
|
||||
super().__init__(config)
|
||||
|
||||
self.config = config
|
||||
self.config.supported_attention_backends = self._supported_attention_backends
|
||||
self.quant_config = quant_config
|
||||
if lora_config is not None:
|
||||
self.quant_config = self.config.quant_config
|
||||
if config.lora_config is not None:
|
||||
max_loras = 1
|
||||
lora_vocab_size = 1
|
||||
if hasattr(lora_config, "max_loras"):
|
||||
max_loras = lora_config.max_loras
|
||||
if hasattr(lora_config, "lora_extra_vocab_size"):
|
||||
lora_vocab_size = lora_config.lora_extra_vocab_size
|
||||
if hasattr(config.lora_config, "max_loras"):
|
||||
max_loras = config.lora_config.max_loras
|
||||
if hasattr(config.lora_config, "lora_extra_vocab_size"):
|
||||
lora_vocab_size = config.lora_config.lora_extra_vocab_size
|
||||
lora_vocab = lora_vocab_size * max_loras
|
||||
else:
|
||||
lora_vocab = 0
|
||||
@@ -311,13 +301,13 @@ class LlamaModel(BaseEncoder):
|
||||
self.vocab_size,
|
||||
config.hidden_size,
|
||||
org_num_embeddings=config.vocab_size,
|
||||
quant_config=quant_config,
|
||||
quant_config=config.quant_config,
|
||||
)
|
||||
|
||||
self.layers = nn.ModuleList([
|
||||
layer_type(config=config,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.layers.{i}")
|
||||
LlamaDecoderLayer(config=config,
|
||||
quant_config=config.quant_config,
|
||||
prefix=f"{config.prefix}.layers.{i}")
|
||||
for i in range(config.num_hidden_layers)
|
||||
])
|
||||
|
||||
@@ -329,11 +319,12 @@ class LlamaModel(BaseEncoder):
|
||||
def forward(
|
||||
self,
|
||||
input_ids: Optional[torch.Tensor],
|
||||
positions: Optional[torch.Tensor] = None,
|
||||
position_ids: Optional[torch.Tensor] = None,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
inputs_embeds: Optional[torch.Tensor] = None,
|
||||
output_hidden_states: Optional[bool] = None,
|
||||
) -> torch.Tensor:
|
||||
**kwargs,
|
||||
) -> BaseEncoderOutput:
|
||||
output_hidden_states = (output_hidden_states
|
||||
if output_hidden_states is not None else
|
||||
self.config.output_hidden_states)
|
||||
@@ -343,10 +334,10 @@ class LlamaModel(BaseEncoder):
|
||||
hidden_states = self.get_input_embeddings(input_ids)
|
||||
residual = None
|
||||
|
||||
if positions is None:
|
||||
positions = torch.arange(0,
|
||||
hidden_states.shape[1],
|
||||
device=hidden_states.device).unsqueeze(0)
|
||||
if position_ids is None:
|
||||
position_ids = torch.arange(
|
||||
0, hidden_states.shape[1],
|
||||
device=hidden_states.device).unsqueeze(0)
|
||||
|
||||
all_hidden_states: Optional[Tuple[Any, ...]] = (
|
||||
) if output_hidden_states else None
|
||||
@@ -356,7 +347,8 @@ class LlamaModel(BaseEncoder):
|
||||
all_hidden_states += (
|
||||
hidden_states, ) if residual is None else (hidden_states +
|
||||
residual, )
|
||||
hidden_states, residual = layer(positions, hidden_states, residual)
|
||||
hidden_states, residual = layer(position_ids, hidden_states,
|
||||
residual)
|
||||
|
||||
hidden_states, _ = self.norm(hidden_states, residual)
|
||||
|
||||
@@ -366,7 +358,7 @@ class LlamaModel(BaseEncoder):
|
||||
|
||||
# TODO(will): maybe unify the output format with other models and use
|
||||
# our own class
|
||||
output = BaseModelOutputWithPast(
|
||||
output = BaseEncoderOutput(
|
||||
last_hidden_state=hidden_states,
|
||||
# past_key_values=past_key_values if use_cache else None,
|
||||
hidden_states=all_hidden_states,
|
||||
@@ -395,17 +387,17 @@ class LlamaModel(BaseEncoder):
|
||||
# Models trained using ColossalAI may include these tensors in
|
||||
# the checkpoint. Skip them.
|
||||
continue
|
||||
if (self.quant_config is not None and
|
||||
(scale_name := self.quant_config.get_cache_scale(name))):
|
||||
# Loading kv cache quantization scales
|
||||
param = params_dict[scale_name]
|
||||
weight_loader = getattr(param, "weight_loader",
|
||||
default_weight_loader)
|
||||
loaded_weight = (loaded_weight if loaded_weight.dim() == 0 else
|
||||
loaded_weight[0])
|
||||
weight_loader(param, loaded_weight)
|
||||
loaded_params.add(scale_name)
|
||||
continue
|
||||
# if (self.quant_config is not None and
|
||||
# (scale_name := self.quant_config.get_cache_scale(name))):
|
||||
# # Loading kv cache quantization scales
|
||||
# param = params_dict[scale_name]
|
||||
# weight_loader = getattr(param, "weight_loader",
|
||||
# default_weight_loader)
|
||||
# loaded_weight = (loaded_weight if loaded_weight.dim() == 0 else
|
||||
# loaded_weight[0])
|
||||
# weight_loader(param, loaded_weight)
|
||||
# loaded_params.add(scale_name)
|
||||
# continue
|
||||
if "scale" in name:
|
||||
# Remapping the name of FP8 kv-scale.
|
||||
kv_scale_name: Optional[str] = maybe_remap_kv_scale_name(
|
||||
|
||||
@@ -26,8 +26,9 @@ from typing import Iterable, Optional, Set, Tuple
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch import nn
|
||||
from transformers import T5Config
|
||||
|
||||
from fastvideo.v1.configs.models.encoders import BaseEncoderOutput, T5Config
|
||||
from fastvideo.v1.configs.quantization import QuantizationConfig
|
||||
from fastvideo.v1.distributed import (get_tensor_model_parallel_rank,
|
||||
get_tensor_model_parallel_world_size)
|
||||
from fastvideo.v1.layers.activation import get_act_fn
|
||||
@@ -35,13 +36,10 @@ from fastvideo.v1.layers.layernorm import RMSNorm
|
||||
from fastvideo.v1.layers.linear import (MergedColumnParallelLinear,
|
||||
QKVParallelLinear, RowParallelLinear)
|
||||
from fastvideo.v1.layers.vocab_parallel_embedding import VocabParallelEmbedding
|
||||
from fastvideo.v1.models.encoders.base import TextEncoder
|
||||
from fastvideo.v1.models.loader.weight_utils import default_weight_loader
|
||||
|
||||
|
||||
class QuantizationConfig:
|
||||
pass
|
||||
|
||||
|
||||
class AttentionType:
|
||||
"""
|
||||
Attention type.
|
||||
@@ -501,10 +499,10 @@ class T5Stack(nn.Module):
|
||||
return hidden_states
|
||||
|
||||
|
||||
class T5EncoderModel(nn.Module):
|
||||
class T5EncoderModel(TextEncoder):
|
||||
|
||||
def __init__(self, config: T5Config, prefix: str = ""):
|
||||
super().__init__()
|
||||
super().__init__(config)
|
||||
|
||||
quant_config = None
|
||||
|
||||
@@ -526,22 +524,21 @@ class T5EncoderModel(nn.Module):
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_ids: Optional[torch.LongTensor] = None,
|
||||
attention_mask: Optional[torch.FloatTensor] = None,
|
||||
head_mask: Optional[torch.FloatTensor] = None,
|
||||
inputs_embeds: Optional[torch.FloatTensor] = None,
|
||||
output_attentions: Optional[bool] = None,
|
||||
input_ids: Optional[torch.Tensor],
|
||||
position_ids: Optional[torch.Tensor] = None,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
inputs_embeds: Optional[torch.Tensor] = None,
|
||||
output_hidden_states: Optional[bool] = None,
|
||||
return_dict: Optional[bool] = None,
|
||||
) -> torch.Tensor:
|
||||
**kwargs,
|
||||
) -> BaseEncoderOutput:
|
||||
attn_metadata = AttentionMetadata(None)
|
||||
encoder_outputs = self.encoder(
|
||||
hidden_states = self.encoder(
|
||||
input_ids=input_ids,
|
||||
attention_mask=attention_mask,
|
||||
attn_metadata=attn_metadata,
|
||||
)
|
||||
|
||||
return encoder_outputs
|
||||
return BaseEncoderOutput(last_hidden_state=hidden_states)
|
||||
|
||||
def load_weights(self, weights: Iterable[Tuple[str,
|
||||
torch.Tensor]]) -> Set[str]:
|
||||
@@ -589,10 +586,10 @@ class T5EncoderModel(nn.Module):
|
||||
return loaded_params
|
||||
|
||||
|
||||
class UMT5EncoderModel(nn.Module):
|
||||
class UMT5EncoderModel(TextEncoder):
|
||||
|
||||
def __init__(self, config: T5Config, prefix: str = ""):
|
||||
super().__init__()
|
||||
super().__init__(config)
|
||||
|
||||
quant_config = None
|
||||
|
||||
@@ -614,22 +611,24 @@ class UMT5EncoderModel(nn.Module):
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_ids: Optional[torch.LongTensor] = None,
|
||||
attention_mask: Optional[torch.FloatTensor] = None,
|
||||
head_mask: Optional[torch.FloatTensor] = None,
|
||||
inputs_embeds: Optional[torch.FloatTensor] = None,
|
||||
output_attentions: Optional[bool] = None,
|
||||
input_ids: Optional[torch.Tensor],
|
||||
position_ids: Optional[torch.Tensor] = None,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
inputs_embeds: Optional[torch.Tensor] = None,
|
||||
output_hidden_states: Optional[bool] = None,
|
||||
return_dict: Optional[bool] = None,
|
||||
) -> torch.Tensor:
|
||||
**kwargs,
|
||||
) -> BaseEncoderOutput:
|
||||
attn_metadata = AttentionMetadata(None)
|
||||
encoder_outputs = self.encoder(
|
||||
hidden_states = self.encoder(
|
||||
input_ids=input_ids,
|
||||
attention_mask=attention_mask,
|
||||
attn_metadata=attn_metadata,
|
||||
)
|
||||
|
||||
return encoder_outputs
|
||||
return BaseEncoderOutput(
|
||||
last_hidden_state=hidden_states,
|
||||
attention_mask=attention_mask,
|
||||
)
|
||||
|
||||
def load_weights(self, weights: Iterable[Tuple[str,
|
||||
torch.Tensor]]) -> Set[str]:
|
||||
|
||||
@@ -1,8 +1,8 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import dataclasses
|
||||
from dataclasses import asdict
|
||||
import glob
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
from abc import ABC, abstractmethod
|
||||
@@ -11,13 +11,12 @@ from typing import Any, Generator, Iterable, List, Optional, Tuple, cast
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from safetensors.torch import load_file as safetensors_load_file
|
||||
from transformers import AutoImageProcessor, AutoTokenizer, PretrainedConfig
|
||||
from transformers import AutoImageProcessor, AutoTokenizer
|
||||
from transformers.utils import SAFE_WEIGHTS_INDEX_NAME
|
||||
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.hf_transformer_utils import (get_diffusers_config,
|
||||
get_hf_config)
|
||||
from fastvideo.v1.models.hf_transformer_utils import get_diffusers_config
|
||||
from fastvideo.v1.models.loader.fsdp_load import load_fsdp_model
|
||||
from fastvideo.v1.models.loader.utils import set_default_torch_dtype
|
||||
from fastvideo.v1.models.loader.weight_utils import (
|
||||
@@ -202,18 +201,35 @@ class TextEncoderLoader(ComponentLoader):
|
||||
def load(self, model_path: str, architecture: str,
|
||||
fastvideo_args: FastVideoArgs):
|
||||
"""Load the text encoders based on the model path, architecture, and inference args."""
|
||||
model_config: PretrainedConfig = get_hf_config(
|
||||
model=model_path,
|
||||
trust_remote_code=fastvideo_args.trust_remote_code,
|
||||
revision=fastvideo_args.revision,
|
||||
model_override_args=None,
|
||||
)
|
||||
# model_config: PretrainedConfig = get_hf_config(
|
||||
# model=model_path,
|
||||
# trust_remote_code=fastvideo_args.trust_remote_code,
|
||||
# revision=fastvideo_args.revision,
|
||||
# model_override_args=None,
|
||||
# )
|
||||
with open(os.path.join(model_path, "config.json")) as f:
|
||||
model_config = json.load(f)
|
||||
model_config.pop("_name_or_path", None)
|
||||
model_config.pop("transformers_version", None)
|
||||
model_config.pop("model_type", None)
|
||||
model_config.pop("tokenizer_class", None)
|
||||
model_config.pop("torch_dtype", None)
|
||||
logger.info("HF Model config: %s", model_config)
|
||||
|
||||
# @TODO(Wei): Better way to handle this?
|
||||
try:
|
||||
encoder_config = fastvideo_args.text_encoder_configs[0]
|
||||
encoder_config.update_model_arch(model_config)
|
||||
encoder_precision = fastvideo_args.text_encoder_precisions[0]
|
||||
except Exception:
|
||||
encoder_config = fastvideo_args.text_encoder_configs[1]
|
||||
encoder_config.update_model_arch(model_config)
|
||||
encoder_precision = fastvideo_args.text_encoder_precisions[1]
|
||||
|
||||
target_device = torch.device(fastvideo_args.device_str)
|
||||
# TODO(will): add support for other dtypes
|
||||
return self.load_model(model_path, model_config, target_device,
|
||||
fastvideo_args.text_encoder_precision)
|
||||
return self.load_model(model_path, encoder_config, target_device,
|
||||
encoder_precision)
|
||||
|
||||
def load_model(self,
|
||||
model_path: str,
|
||||
@@ -252,17 +268,26 @@ class ImageEncoderLoader(TextEncoderLoader):
|
||||
def load(self, model_path: str, architecture: str,
|
||||
fastvideo_args: FastVideoArgs):
|
||||
"""Load the text encoders based on the model path, architecture, and inference args."""
|
||||
model_config: PretrainedConfig = get_hf_config(
|
||||
model=model_path,
|
||||
trust_remote_code=fastvideo_args.trust_remote_code,
|
||||
revision=fastvideo_args.revision,
|
||||
model_override_args=None,
|
||||
)
|
||||
# model_config: PretrainedConfig = get_hf_config(
|
||||
# model=model_path,
|
||||
# trust_remote_code=fastvideo_args.trust_remote_code,
|
||||
# revision=fastvideo_args.revision,
|
||||
# model_override_args=None,
|
||||
# )
|
||||
with open(os.path.join(model_path, "config.json")) as f:
|
||||
model_config = json.load(f)
|
||||
model_config.pop("_name_or_path", None)
|
||||
model_config.pop("transformers_version", None)
|
||||
model_config.pop("torch_dtype", None)
|
||||
model_config.pop("model_type", None)
|
||||
logger.info("HF Model config: %s", model_config)
|
||||
|
||||
encoder_config = fastvideo_args.image_encoder_config
|
||||
encoder_config.update_model_arch(model_config)
|
||||
|
||||
target_device = torch.device(fastvideo_args.device_str)
|
||||
# TODO(will): add support for other dtypes
|
||||
return self.load_model(model_path, model_config, target_device,
|
||||
return self.load_model(model_path, encoder_config, target_device,
|
||||
fastvideo_args.image_encoder_precision)
|
||||
|
||||
|
||||
@@ -289,7 +314,8 @@ class TokenizerLoader(ComponentLoader):
|
||||
logger.info("Loading tokenizer from %s", model_path)
|
||||
|
||||
tokenizer = AutoTokenizer.from_pretrained(
|
||||
model_path,
|
||||
model_path, # "<path to model>/tokenizer"
|
||||
# in v0, this was same string as encoder_name "ClipTextModel"
|
||||
# TODO(will): pass these tokenizer kwargs from inference args? Maybe
|
||||
# other method of config?
|
||||
padding_size='right',
|
||||
@@ -325,7 +351,8 @@ class VAELoader(ComponentLoader):
|
||||
safetensors_list
|
||||
) == 1, f"Found {len(safetensors_list)} safetensors files in {model_path}"
|
||||
loaded = safetensors_load_file(safetensors_list[0])
|
||||
vae.load_state_dict(loaded, strict=False) # We might only load encoder or decoder
|
||||
vae.load_state_dict(
|
||||
loaded, strict=False) # We might only load encoder or decoder
|
||||
dtype = PRECISION_TO_TYPE[fastvideo_args.vae_precision]
|
||||
vae = vae.eval().to(dtype)
|
||||
|
||||
@@ -338,17 +365,16 @@ class TransformerLoader(ComponentLoader):
|
||||
def load(self, model_path: str, architecture: str,
|
||||
fastvideo_args: FastVideoArgs):
|
||||
"""Load the transformer based on the model path, architecture, and inference args."""
|
||||
model_config = get_diffusers_config(model=model_path)
|
||||
cls_name = model_config.pop("_class_name")
|
||||
config = get_diffusers_config(model=model_path)
|
||||
cls_name = config.pop("_class_name")
|
||||
if cls_name is None:
|
||||
raise ValueError(
|
||||
"Model config does not contain a _class_name attribute. "
|
||||
"Only diffusers format is supported.")
|
||||
model_config.pop("_diffusers_version")
|
||||
|
||||
# Config from Diffusers supercedes fastvideo's model config
|
||||
# dit_config = fastvideo_args.dit_config
|
||||
# model_config.update(dit_config)
|
||||
# Config from Diffusers supersedes fastvideo's model config
|
||||
dit_config = fastvideo_args.dit_config
|
||||
dit_config.update_model_arch(config)
|
||||
|
||||
model_cls, _ = ModelRegistry.resolve_model_cls(cls_name)
|
||||
|
||||
@@ -367,11 +393,30 @@ class TransformerLoader(ComponentLoader):
|
||||
# Load the model using FSDP loader
|
||||
logger.info("Loading model from %s", cls_name)
|
||||
model = load_fsdp_model(model_cls=model_cls,
|
||||
init_params=model_config,
|
||||
init_params={"config": dit_config},
|
||||
weight_dir_list=safetensors_list,
|
||||
device=fastvideo_args.device,
|
||||
cpu_offload=fastvideo_args.use_cpu_offload,
|
||||
default_dtype=default_dtype)
|
||||
if fastvideo_args.enable_torch_compile:
|
||||
logger.info("Torch Compile enabled for DiT")
|
||||
for n, m in reversed(list(model.named_modules())):
|
||||
if any([
|
||||
compile_condition(n, m)
|
||||
for compile_condition in model._compile_conditions
|
||||
]):
|
||||
parts = n.split(".")
|
||||
parent = model
|
||||
attr = parts[-1]
|
||||
for part in parts[:-1]:
|
||||
if part.isdigit():
|
||||
parent = parent[int(part)]
|
||||
else:
|
||||
parent = getattr(parent, part)
|
||||
if attr.isdigit():
|
||||
parent[int(attr)] = torch.compile(m)
|
||||
else:
|
||||
setattr(parent, attr, torch.compile(m))
|
||||
|
||||
total_params = sum(p.numel() for p in model.parameters())
|
||||
logger.info("Loaded model with %.2fB parameters", total_params / 1e9)
|
||||
|
||||
@@ -23,6 +23,7 @@ _TEXT_TO_VIDEO_DIT_MODELS = {
|
||||
"HunyuanVideoTransformer3DModel":
|
||||
("dits", "hunyuanvideo", "HunyuanVideoTransformer3DModel"),
|
||||
"WanTransformer3DModel": ("dits", "wanvideo", "WanTransformer3DModel"),
|
||||
"FluxTransformer2DModel": ("dits", "flux", "FluxTransformer2DModel"),
|
||||
}
|
||||
|
||||
_IMAGE_TO_VIDEO_DIT_MODELS = {
|
||||
@@ -34,6 +35,7 @@ _TEXT_ENCODER_MODELS = {
|
||||
"CLIPTextModel": ("encoders", "clip", "CLIPTextModel"),
|
||||
"LlamaModel": ("encoders", "llama", "LlamaModel"),
|
||||
"UMT5EncoderModel": ("encoders", "t5", "UMT5EncoderModel"),
|
||||
"T5EncoderModel": ("encoders", "t5", "T5EncoderModel"),
|
||||
}
|
||||
|
||||
_IMAGE_ENCODER_MODELS: dict[str, tuple] = {
|
||||
@@ -45,12 +47,13 @@ _VAE_MODELS = {
|
||||
"AutoencoderKLHunyuanVideo":
|
||||
("vaes", "hunyuanvae", "AutoencoderKLHunyuanVideo"),
|
||||
"AutoencoderKLWan": ("vaes", "wanvae", "AutoencoderKLWan"),
|
||||
"AutoencoderKL": ("vaes", "image_vae", "AutoencoderKL"),
|
||||
}
|
||||
|
||||
_SCHEDULERS = {
|
||||
"FlowMatchEulerDiscreteScheduler":
|
||||
("schedulers", "scheduling_flow_match_euler_discrete",
|
||||
"FlowMatchDiscreteScheduler"),
|
||||
"FlowMatchEulerDiscreteScheduler"),
|
||||
"UniPCMultistepScheduler":
|
||||
("schedulers", "scheduling_unipc_multistep", "UniPCMultistepScheduler"),
|
||||
}
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
# type: ignore
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
# Copyright 2024 Stability AI, Katherine Crowson and The HuggingFace Team. All rights reserved.
|
||||
@@ -19,13 +20,16 @@
|
||||
#
|
||||
# ==============================================================================
|
||||
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Optional, Tuple, Union
|
||||
from typing import List, Optional, Tuple, Union
|
||||
|
||||
import numpy as np
|
||||
import scipy
|
||||
import torch
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.schedulers.scheduling_utils import SchedulerMixin
|
||||
from diffusers.utils import BaseOutput, logging
|
||||
from diffusers.utils import BaseOutput, is_scipy_available, logging
|
||||
|
||||
from fastvideo.v1.models.schedulers.base import BaseScheduler
|
||||
|
||||
@@ -33,7 +37,7 @@ logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
|
||||
@dataclass
|
||||
class FlowMatchDiscreteSchedulerOutput(BaseOutput):
|
||||
class FlowMatchEulerDiscreteSchedulerOutput(BaseOutput):
|
||||
"""
|
||||
Output class for the scheduler's `step` function output.
|
||||
|
||||
@@ -46,7 +50,8 @@ class FlowMatchDiscreteSchedulerOutput(BaseOutput):
|
||||
prev_sample: torch.FloatTensor
|
||||
|
||||
|
||||
class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
|
||||
class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
|
||||
BaseScheduler):
|
||||
"""
|
||||
Euler scheduler.
|
||||
|
||||
@@ -56,16 +61,37 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
|
||||
Args:
|
||||
num_train_timesteps (`int`, defaults to 1000):
|
||||
The number of diffusion steps to train the model.
|
||||
timestep_spacing (`str`, defaults to `"linspace"`):
|
||||
The way the timesteps should be scaled. Refer to Table 2 of the [Common Diffusion Noise Schedules and
|
||||
Sample Steps are Flawed](https://huggingface.co/papers/2305.08891) for more information.
|
||||
shift (`float`, defaults to 1.0):
|
||||
The shift value for the timestep schedule.
|
||||
reverse (`bool`, defaults to `True`):
|
||||
Whether to reverse the timestep schedule.
|
||||
use_dynamic_shifting (`bool`, defaults to False):
|
||||
Whether to apply timestep shifting on-the-fly based on the image resolution.
|
||||
base_shift (`float`, defaults to 0.5):
|
||||
Value to stabilize image generation. Increasing `base_shift` reduces variation and image is more consistent
|
||||
with desired output.
|
||||
max_shift (`float`, defaults to 1.15):
|
||||
Value change allowed to latent vectors. Increasing `max_shift` encourages more variation and image may be
|
||||
more exaggerated or stylized.
|
||||
base_image_seq_len (`int`, defaults to 256):
|
||||
The base image sequence length.
|
||||
max_image_seq_len (`int`, defaults to 4096):
|
||||
The maximum image sequence length.
|
||||
invert_sigmas (`bool`, defaults to False):
|
||||
Whether to invert the sigmas.
|
||||
shift_terminal (`float`, defaults to None):
|
||||
The end value of the shifted timestep schedule.
|
||||
use_karras_sigmas (`bool`, defaults to False):
|
||||
Whether to use Karras sigmas for step sizes in the noise schedule during sampling.
|
||||
use_exponential_sigmas (`bool`, defaults to False):
|
||||
Whether to use exponential sigmas for step sizes in the noise schedule during sampling.
|
||||
use_beta_sigmas (`bool`, defaults to False):
|
||||
Whether to use beta sigmas for step sizes in the noise schedule during sampling.
|
||||
time_shift_type (`str`, defaults to "exponential"):
|
||||
The type of dynamic resolution-dependent timestep shifting to apply. Either "exponential" or "linear".
|
||||
stochastic_sampling (`bool`, defaults to False):
|
||||
Whether to use stochastic sampling.
|
||||
"""
|
||||
|
||||
_compatibles: list[Any] = []
|
||||
_compatibles = []
|
||||
order = 1
|
||||
|
||||
@register_to_config
|
||||
@@ -73,31 +99,62 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
|
||||
self,
|
||||
num_train_timesteps: int = 1000,
|
||||
shift: float = 1.0,
|
||||
reverse: bool = True,
|
||||
solver: str = "euler",
|
||||
n_tokens: Optional[int] = None,
|
||||
**kwargs,
|
||||
use_dynamic_shifting: bool = False,
|
||||
base_shift: Optional[float] = 0.5,
|
||||
max_shift: Optional[float] = 1.15,
|
||||
base_image_seq_len: Optional[int] = 256,
|
||||
max_image_seq_len: Optional[int] = 4096,
|
||||
invert_sigmas: bool = False,
|
||||
shift_terminal: Optional[float] = None,
|
||||
use_karras_sigmas: Optional[bool] = False,
|
||||
use_exponential_sigmas: Optional[bool] = False,
|
||||
use_beta_sigmas: Optional[bool] = False,
|
||||
time_shift_type: str = "exponential",
|
||||
stochastic_sampling: bool = False,
|
||||
):
|
||||
sigmas = torch.linspace(1, 0, num_train_timesteps + 1)
|
||||
|
||||
if not reverse:
|
||||
sigmas = sigmas.flip(0)
|
||||
|
||||
self.sigmas = sigmas
|
||||
# the value fed to model
|
||||
self.timesteps = (sigmas[:-1] *
|
||||
num_train_timesteps).to(dtype=torch.float32)
|
||||
|
||||
self._step_index: int | None = None
|
||||
self._begin_index = 0
|
||||
|
||||
self.supported_solver = ["euler"]
|
||||
if solver not in self.supported_solver:
|
||||
if self.config.use_beta_sigmas and not is_scipy_available():
|
||||
raise ImportError(
|
||||
"Make sure to install scipy if you want to use beta sigmas.")
|
||||
if sum([
|
||||
self.config.use_beta_sigmas, self.config.use_exponential_sigmas,
|
||||
self.config.use_karras_sigmas
|
||||
]) > 1:
|
||||
raise ValueError(
|
||||
f"Solver {solver} not supported. Supported solvers: {self.supported_solver}"
|
||||
"Only one of `config.use_beta_sigmas`, `config.use_exponential_sigmas`, `config.use_karras_sigmas` can be used."
|
||||
)
|
||||
if time_shift_type not in {"exponential", "linear"}:
|
||||
raise ValueError(
|
||||
"`time_shift_type` must either be 'exponential' or 'linear'.")
|
||||
|
||||
BaseScheduler.__init__(self)
|
||||
timesteps = np.linspace(1,
|
||||
num_train_timesteps,
|
||||
num_train_timesteps,
|
||||
dtype=np.float32)[::-1].copy()
|
||||
timesteps = torch.from_numpy(timesteps).to(dtype=torch.float32)
|
||||
|
||||
sigmas = timesteps / num_train_timesteps
|
||||
if not use_dynamic_shifting:
|
||||
# when use_dynamic_shifting is True, we apply the timestep shifting on the fly based on the image resolution
|
||||
sigmas = shift * sigmas / (1 + (shift - 1) * sigmas)
|
||||
|
||||
self.timesteps = sigmas * num_train_timesteps
|
||||
|
||||
self._step_index = None
|
||||
self._begin_index = None
|
||||
|
||||
self._shift = shift
|
||||
|
||||
self.sigmas = sigmas.to(
|
||||
"cpu") # to avoid too much CPU/GPU communication
|
||||
self.sigma_min = self.sigmas[-1].item()
|
||||
self.sigma_max = self.sigmas[0].item()
|
||||
|
||||
@property
|
||||
def shift(self):
|
||||
"""
|
||||
The value used for shifting.
|
||||
"""
|
||||
return self._shift
|
||||
|
||||
@property
|
||||
def step_index(self):
|
||||
@@ -124,42 +181,207 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
|
||||
"""
|
||||
self._begin_index = begin_index
|
||||
|
||||
def set_shift(self, shift: float):
|
||||
self._shift = shift
|
||||
|
||||
def scale_noise(
|
||||
self,
|
||||
sample: torch.FloatTensor,
|
||||
timestep: Union[float, torch.FloatTensor],
|
||||
noise: Optional[torch.FloatTensor] = None,
|
||||
) -> torch.FloatTensor:
|
||||
"""
|
||||
Forward process in flow-matching
|
||||
|
||||
Args:
|
||||
sample (`torch.FloatTensor`):
|
||||
The input sample.
|
||||
timestep (`int`, *optional*):
|
||||
The current timestep in the diffusion chain.
|
||||
|
||||
Returns:
|
||||
`torch.FloatTensor`:
|
||||
A scaled input sample.
|
||||
"""
|
||||
# Make sure sigmas and timesteps have the same device and dtype as original_samples
|
||||
sigmas = self.sigmas.to(device=sample.device, dtype=sample.dtype)
|
||||
|
||||
if sample.device.type == "mps" and torch.is_floating_point(timestep):
|
||||
# mps does not support float64
|
||||
schedule_timesteps = self.timesteps.to(sample.device,
|
||||
dtype=torch.float32)
|
||||
timestep = timestep.to(sample.device, dtype=torch.float32)
|
||||
else:
|
||||
schedule_timesteps = self.timesteps.to(sample.device)
|
||||
timestep = timestep.to(sample.device)
|
||||
|
||||
# self.begin_index is None when scheduler is used for training, or pipeline does not implement set_begin_index
|
||||
if self.begin_index is None:
|
||||
step_indices = [
|
||||
self.index_for_timestep(t, schedule_timesteps) for t in timestep
|
||||
]
|
||||
elif self.step_index is not None:
|
||||
# add_noise is called after first denoising step (for inpainting)
|
||||
step_indices = [self.step_index] * timestep.shape[0]
|
||||
else:
|
||||
# add noise is called before first denoising step to create initial latent(img2img)
|
||||
step_indices = [self.begin_index] * timestep.shape[0]
|
||||
|
||||
sigma = sigmas[step_indices].flatten()
|
||||
while len(sigma.shape) < len(sample.shape):
|
||||
sigma = sigma.unsqueeze(-1)
|
||||
|
||||
sample = sigma * noise + (1.0 - sigma) * sample
|
||||
|
||||
return sample
|
||||
|
||||
def _sigma_to_t(self, sigma):
|
||||
return sigma * self.config.num_train_timesteps
|
||||
|
||||
def time_shift(self, mu: float, sigma: float, t: torch.Tensor):
|
||||
if self.config.time_shift_type == "exponential":
|
||||
return self._time_shift_exponential(mu, sigma, t)
|
||||
elif self.config.time_shift_type == "linear":
|
||||
return self._time_shift_linear(mu, sigma, t)
|
||||
|
||||
def stretch_shift_to_terminal(self, t: torch.Tensor) -> torch.Tensor:
|
||||
r"""
|
||||
Stretches and shifts the timestep schedule to ensure it terminates at the configured `shift_terminal` config
|
||||
value.
|
||||
|
||||
Reference:
|
||||
https://github.com/Lightricks/LTX-Video/blob/a01a171f8fe3d99dce2728d60a73fecf4d4238ae/ltx_video/schedulers/rf.py#L51
|
||||
|
||||
Args:
|
||||
t (`torch.Tensor`):
|
||||
A tensor of timesteps to be stretched and shifted.
|
||||
|
||||
Returns:
|
||||
`torch.Tensor`:
|
||||
A tensor of adjusted timesteps such that the final value equals `self.config.shift_terminal`.
|
||||
"""
|
||||
one_minus_z = 1 - t
|
||||
scale_factor = one_minus_z[-1] / (1 - self.config.shift_terminal)
|
||||
stretched_t = 1 - (one_minus_z / scale_factor)
|
||||
return stretched_t
|
||||
|
||||
def set_timesteps(
|
||||
self,
|
||||
num_inference_steps: int,
|
||||
num_inference_steps: Optional[int] = None,
|
||||
device: Union[str, torch.device] = None,
|
||||
n_tokens: int = 0,
|
||||
sigmas: Optional[List[float]] = None,
|
||||
mu: Optional[float] = None,
|
||||
timesteps: Optional[List[float]] = None,
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
Sets the discrete timesteps used for the diffusion chain (to be run before inference).
|
||||
|
||||
Args:
|
||||
num_inference_steps (`int`):
|
||||
num_inference_steps (`int`, *optional*):
|
||||
The number of diffusion steps used when generating samples with a pre-trained model.
|
||||
device (`str` or `torch.device`, *optional*):
|
||||
The device to which the timesteps should be moved to. If `None`, the timesteps are not moved.
|
||||
n_tokens (`int`, *optional*):
|
||||
Number of tokens in the input sequence.
|
||||
sigmas (`List[float]`, *optional*):
|
||||
Custom values for sigmas to be used for each diffusion step. If `None`, the sigmas are computed
|
||||
automatically.
|
||||
mu (`float`, *optional*):
|
||||
Determines the amount of shifting applied to sigmas when performing resolution-dependent timestep
|
||||
shifting.
|
||||
timesteps (`List[float]`, *optional*):
|
||||
Custom values for timesteps to be used for each diffusion step. If `None`, the timesteps are computed
|
||||
automatically.
|
||||
"""
|
||||
if self.config.use_dynamic_shifting and mu is None:
|
||||
raise ValueError(
|
||||
"`mu` must be passed when `use_dynamic_shifting` is set to be `True`"
|
||||
)
|
||||
|
||||
if sigmas is not None and timesteps is not None and len(sigmas) != len(
|
||||
timesteps):
|
||||
raise ValueError(
|
||||
"`sigmas` and `timesteps` should have the same length")
|
||||
|
||||
if num_inference_steps is not None:
|
||||
if (sigmas is not None and len(sigmas) != num_inference_steps) or (
|
||||
timesteps is not None
|
||||
and len(timesteps) != num_inference_steps):
|
||||
raise ValueError(
|
||||
"`sigmas` and `timesteps` should have the same length as num_inference_steps, if `num_inference_steps` is provided"
|
||||
)
|
||||
else:
|
||||
num_inference_steps = len(sigmas) if sigmas is not None else len(
|
||||
timesteps)
|
||||
|
||||
self.num_inference_steps = num_inference_steps
|
||||
|
||||
sigmas = torch.linspace(1, 0, num_inference_steps + 1)
|
||||
sigmas = self.sd3_time_shift(sigmas)
|
||||
# 1. Prepare default sigmas
|
||||
is_timesteps_provided = timesteps is not None
|
||||
|
||||
if not self.config.reverse:
|
||||
sigmas = 1 - sigmas
|
||||
if is_timesteps_provided:
|
||||
timesteps = np.array(timesteps).astype(np.float32)
|
||||
|
||||
if sigmas is None:
|
||||
if timesteps is None:
|
||||
timesteps = np.linspace(self._sigma_to_t(self.sigma_max),
|
||||
self._sigma_to_t(self.sigma_min),
|
||||
num_inference_steps)
|
||||
sigmas = timesteps / self.config.num_train_timesteps
|
||||
else:
|
||||
sigmas = np.array(sigmas).astype(np.float32)
|
||||
num_inference_steps = len(sigmas)
|
||||
|
||||
# 2. Perform timestep shifting. Either no shifting is applied, or resolution-dependent shifting of
|
||||
# "exponential" or "linear" type is applied
|
||||
if self.config.use_dynamic_shifting:
|
||||
sigmas = self.time_shift(mu, 1.0, sigmas)
|
||||
else:
|
||||
sigmas = self.shift * sigmas / (1 + (self.shift - 1) * sigmas)
|
||||
|
||||
# 3. If required, stretch the sigmas schedule to terminate at the configured `shift_terminal` value
|
||||
if self.config.shift_terminal:
|
||||
sigmas = self.stretch_shift_to_terminal(sigmas)
|
||||
|
||||
# 4. If required, convert sigmas to one of karras, exponential, or beta sigma schedules
|
||||
if self.config.use_karras_sigmas:
|
||||
sigmas = self._convert_to_karras(
|
||||
in_sigmas=sigmas, num_inference_steps=num_inference_steps)
|
||||
elif self.config.use_exponential_sigmas:
|
||||
sigmas = self._convert_to_exponential(
|
||||
in_sigmas=sigmas, num_inference_steps=num_inference_steps)
|
||||
elif self.config.use_beta_sigmas:
|
||||
sigmas = self._convert_to_beta(
|
||||
in_sigmas=sigmas, num_inference_steps=num_inference_steps)
|
||||
|
||||
# 5. Convert sigmas and timesteps to tensors and move to specified device
|
||||
sigmas = torch.from_numpy(sigmas).to(dtype=torch.float32, device=device)
|
||||
if not is_timesteps_provided:
|
||||
timesteps = sigmas * self.config.num_train_timesteps
|
||||
else:
|
||||
timesteps = torch.from_numpy(timesteps).to(dtype=torch.float32,
|
||||
device=device)
|
||||
|
||||
# 6. Append the terminal sigma value.
|
||||
# If a model requires inverted sigma schedule for denoising but timesteps without inversion, the
|
||||
# `invert_sigmas` flag can be set to `True`. This case is only required in Mochi
|
||||
if self.config.invert_sigmas:
|
||||
sigmas = 1.0 - sigmas
|
||||
timesteps = sigmas * self.config.num_train_timesteps
|
||||
sigmas = torch.cat([sigmas, torch.ones(1, device=sigmas.device)])
|
||||
else:
|
||||
sigmas = torch.cat([sigmas, torch.zeros(1, device=sigmas.device)])
|
||||
|
||||
self.timesteps = timesteps
|
||||
self.sigmas = sigmas
|
||||
self.timesteps = (sigmas[:-1] * self.config.num_train_timesteps).to(
|
||||
dtype=torch.float32, device=device)
|
||||
|
||||
# Reset step index
|
||||
self._step_index = None
|
||||
self._begin_index = None
|
||||
|
||||
def index_for_timestep(self, timestep, schedule_timesteps=None) -> int:
|
||||
def scale_model_input(self,
|
||||
sample: torch.Tensor,
|
||||
timestep: Optional[int] = None) -> torch.Tensor:
|
||||
return sample
|
||||
|
||||
def index_for_timestep(self, timestep, schedule_timesteps=None):
|
||||
if schedule_timesteps is None:
|
||||
schedule_timesteps = self.timesteps
|
||||
|
||||
@@ -171,14 +393,9 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
|
||||
# case we start in the middle of the denoising schedule (e.g. for image-to-image)
|
||||
pos = 1 if len(indices) > 1 else 0
|
||||
|
||||
idx: int = indices[pos].item()
|
||||
return indices[pos].item()
|
||||
|
||||
return idx
|
||||
|
||||
def set_shift(self, shift: float) -> None:
|
||||
self.config.shift = shift
|
||||
|
||||
def _init_step_index(self, timestep) -> None:
|
||||
def _init_step_index(self, timestep):
|
||||
if self.begin_index is None:
|
||||
if isinstance(timestep, torch.Tensor):
|
||||
timestep = timestep.to(self.timesteps.device)
|
||||
@@ -186,22 +403,19 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
|
||||
else:
|
||||
self._step_index = self._begin_index
|
||||
|
||||
def scale_model_input(self,
|
||||
sample: torch.Tensor,
|
||||
timestep: Optional[int] = None) -> torch.Tensor:
|
||||
return sample
|
||||
|
||||
def sd3_time_shift(self, t: torch.Tensor):
|
||||
return (self.config.shift * t) / (1 + (self.config.shift - 1) * t)
|
||||
|
||||
def step(
|
||||
self,
|
||||
model_output: torch.FloatTensor,
|
||||
timestep: Union[float, torch.FloatTensor],
|
||||
sample: torch.FloatTensor,
|
||||
s_churn: float = 0.0,
|
||||
s_tmin: float = 0.0,
|
||||
s_tmax: float = float("inf"),
|
||||
s_noise: float = 1.0,
|
||||
generator: Optional[torch.Generator] = None,
|
||||
per_token_timesteps: Optional[torch.Tensor] = None,
|
||||
return_dict: bool = True,
|
||||
**kwargs,
|
||||
) -> Union[FlowMatchDiscreteSchedulerOutput, Tuple]:
|
||||
) -> Union[FlowMatchEulerDiscreteSchedulerOutput, Tuple]:
|
||||
"""
|
||||
Predict the sample from the previous timestep by reversing the SDE. This function propagates the diffusion
|
||||
process from the learned model outputs (most often the predicted noise).
|
||||
@@ -213,24 +427,30 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
|
||||
The current discrete timestep in the diffusion chain.
|
||||
sample (`torch.FloatTensor`):
|
||||
A current instance of a sample created by the diffusion process.
|
||||
s_churn (`float`):
|
||||
s_tmin (`float`):
|
||||
s_tmax (`float`):
|
||||
s_noise (`float`, defaults to 1.0):
|
||||
Scaling factor for noise added to the sample.
|
||||
generator (`torch.Generator`, *optional*):
|
||||
A random number generator.
|
||||
n_tokens (`int`, *optional*):
|
||||
Number of tokens in the input sequence.
|
||||
per_token_timesteps (`torch.Tensor`, *optional*):
|
||||
The timesteps for each token in the sample.
|
||||
return_dict (`bool`):
|
||||
Whether or not to return a [`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] or
|
||||
tuple.
|
||||
Whether or not to return a
|
||||
[`~schedulers.scheduling_flow_match_euler_discrete.FlowMatchEulerDiscreteSchedulerOutput`] or tuple.
|
||||
|
||||
Returns:
|
||||
[`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] or `tuple`:
|
||||
If return_dict is `True`, [`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] is
|
||||
returned, otherwise a tuple is returned where the first element is the sample tensor.
|
||||
[`~schedulers.scheduling_flow_match_euler_discrete.FlowMatchEulerDiscreteSchedulerOutput`] or `tuple`:
|
||||
If return_dict is `True`,
|
||||
[`~schedulers.scheduling_flow_match_euler_discrete.FlowMatchEulerDiscreteSchedulerOutput`] is returned,
|
||||
otherwise a tuple is returned where the first element is the sample tensor.
|
||||
"""
|
||||
|
||||
if isinstance(timestep, (int, torch.IntTensor, torch.LongTensor)):
|
||||
raise ValueError((
|
||||
"Passing integer indices (e.g. from `enumerate(timesteps)`) as timesteps to"
|
||||
" `EulerDiscreteScheduler.step()` is not supported. Make sure to pass"
|
||||
" `FlowMatchEulerDiscreteScheduler.step()` is not supported. Make sure to pass"
|
||||
" one of the `scheduler.timesteps` as a timestep."), )
|
||||
|
||||
if self.step_index is None:
|
||||
@@ -239,24 +459,132 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
|
||||
# Upcast to avoid precision issues when computing prev_sample
|
||||
sample = sample.to(torch.float32)
|
||||
|
||||
assert self.step_index is not None
|
||||
dt = self.sigmas[self.step_index + 1] - self.sigmas[self.step_index]
|
||||
if per_token_timesteps is not None:
|
||||
per_token_sigmas = per_token_timesteps / self.config.num_train_timesteps
|
||||
|
||||
if self.config.solver == "euler":
|
||||
prev_sample = sample + model_output.to(torch.float32) * dt
|
||||
sigmas = self.sigmas[:, None, None]
|
||||
lower_mask = sigmas < per_token_sigmas[None] - 1e-6
|
||||
lower_sigmas = lower_mask * sigmas
|
||||
lower_sigmas, _ = lower_sigmas.max(dim=0)
|
||||
|
||||
current_sigma = per_token_sigmas[..., None]
|
||||
next_sigma = lower_sigmas[..., None]
|
||||
dt = current_sigma - next_sigma
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Solver {self.config.solver} not supported. Supported solvers: {self.supported_solver}"
|
||||
)
|
||||
sigma_idx = self.step_index
|
||||
sigma = self.sigmas[sigma_idx]
|
||||
sigma_next = self.sigmas[sigma_idx + 1]
|
||||
|
||||
current_sigma = sigma
|
||||
next_sigma = sigma_next
|
||||
dt = sigma_next - sigma
|
||||
|
||||
if self.config.stochastic_sampling:
|
||||
x0 = sample - current_sigma * model_output
|
||||
noise = torch.randn_like(sample)
|
||||
prev_sample = (1.0 - next_sigma) * x0 + next_sigma * noise
|
||||
else:
|
||||
prev_sample = sample + dt * model_output
|
||||
|
||||
# upon completion increase step index by one
|
||||
assert self._step_index is not None
|
||||
self._step_index += 1
|
||||
if per_token_timesteps is None:
|
||||
# Cast sample back to model compatible dtype
|
||||
prev_sample = prev_sample.to(model_output.dtype)
|
||||
|
||||
if not return_dict:
|
||||
return (prev_sample, )
|
||||
|
||||
return FlowMatchDiscreteSchedulerOutput(prev_sample=prev_sample)
|
||||
return FlowMatchEulerDiscreteSchedulerOutput(prev_sample=prev_sample)
|
||||
|
||||
# Copied from diffusers.schedulers.scheduling_euler_discrete.EulerDiscreteScheduler._convert_to_karras
|
||||
def _convert_to_karras(self, in_sigmas: torch.Tensor,
|
||||
num_inference_steps) -> torch.Tensor:
|
||||
"""Constructs the noise schedule of Karras et al. (2022)."""
|
||||
|
||||
# Hack to make sure that other schedulers which copy this function don't break
|
||||
# TODO: Add this logic to the other schedulers
|
||||
if hasattr(self.config, "sigma_min"):
|
||||
sigma_min = self.config.sigma_min
|
||||
else:
|
||||
sigma_min = None
|
||||
|
||||
if hasattr(self.config, "sigma_max"):
|
||||
sigma_max = self.config.sigma_max
|
||||
else:
|
||||
sigma_max = None
|
||||
|
||||
sigma_min = sigma_min if sigma_min is not None else in_sigmas[-1].item()
|
||||
sigma_max = sigma_max if sigma_max is not None else in_sigmas[0].item()
|
||||
|
||||
rho = 7.0 # 7.0 is the value used in the paper
|
||||
ramp = np.linspace(0, 1, num_inference_steps)
|
||||
min_inv_rho = sigma_min**(1 / rho)
|
||||
max_inv_rho = sigma_max**(1 / rho)
|
||||
sigmas = (max_inv_rho + ramp * (min_inv_rho - max_inv_rho))**rho
|
||||
return sigmas
|
||||
|
||||
# Copied from diffusers.schedulers.scheduling_euler_discrete.EulerDiscreteScheduler._convert_to_exponential
|
||||
def _convert_to_exponential(self, in_sigmas: torch.Tensor,
|
||||
num_inference_steps: int) -> torch.Tensor:
|
||||
"""Constructs an exponential noise schedule."""
|
||||
|
||||
# Hack to make sure that other schedulers which copy this function don't break
|
||||
# TODO: Add this logic to the other schedulers
|
||||
if hasattr(self.config, "sigma_min"):
|
||||
sigma_min = self.config.sigma_min
|
||||
else:
|
||||
sigma_min = None
|
||||
|
||||
if hasattr(self.config, "sigma_max"):
|
||||
sigma_max = self.config.sigma_max
|
||||
else:
|
||||
sigma_max = None
|
||||
|
||||
sigma_min = sigma_min if sigma_min is not None else in_sigmas[-1].item()
|
||||
sigma_max = sigma_max if sigma_max is not None else in_sigmas[0].item()
|
||||
|
||||
sigmas = np.exp(
|
||||
np.linspace(math.log(sigma_max), math.log(sigma_min),
|
||||
num_inference_steps))
|
||||
return sigmas
|
||||
|
||||
# Copied from diffusers.schedulers.scheduling_euler_discrete.EulerDiscreteScheduler._convert_to_beta
|
||||
def _convert_to_beta(self,
|
||||
in_sigmas: torch.Tensor,
|
||||
num_inference_steps: int,
|
||||
alpha: float = 0.6,
|
||||
beta: float = 0.6) -> torch.Tensor:
|
||||
"""From "Beta Sampling is All You Need" [arXiv:2407.12173] (Lee et. al, 2024)"""
|
||||
|
||||
# Hack to make sure that other schedulers which copy this function don't break
|
||||
# TODO: Add this logic to the other schedulers
|
||||
if hasattr(self.config, "sigma_min"):
|
||||
sigma_min = self.config.sigma_min
|
||||
else:
|
||||
sigma_min = None
|
||||
|
||||
if hasattr(self.config, "sigma_max"):
|
||||
sigma_max = self.config.sigma_max
|
||||
else:
|
||||
sigma_max = None
|
||||
|
||||
sigma_min = sigma_min if sigma_min is not None else in_sigmas[-1].item()
|
||||
sigma_max = sigma_max if sigma_max is not None else in_sigmas[0].item()
|
||||
|
||||
sigmas = np.array([
|
||||
sigma_min + (ppf * (sigma_max - sigma_min)) for ppf in [
|
||||
scipy.stats.beta.ppf(timestep, alpha, beta)
|
||||
for timestep in 1 - np.linspace(0, 1, num_inference_steps)
|
||||
]
|
||||
])
|
||||
return sigmas
|
||||
|
||||
def _time_shift_exponential(self, mu, sigma, t):
|
||||
return math.exp(mu) / (math.exp(mu) + (1 / t - 1)**sigma)
|
||||
|
||||
def _time_shift_linear(self, mu, sigma, t):
|
||||
return mu / (mu + (1 / t - 1)**sigma)
|
||||
|
||||
def __len__(self):
|
||||
return self.config.num_train_timesteps
|
||||
|
||||
@@ -118,3 +118,27 @@ def extract_layer_index(layer_name: str) -> int:
|
||||
assert len(int_vals) == 1, (f"layer name {layer_name} should"
|
||||
" only contain one integer")
|
||||
return int_vals[0]
|
||||
|
||||
|
||||
def modulate(x: torch.Tensor,
|
||||
shift: Optional[torch.Tensor] = None,
|
||||
scale: Optional[torch.Tensor] = None) -> torch.Tensor:
|
||||
"""modulate by shift and scale
|
||||
|
||||
Args:
|
||||
x (torch.Tensor): input tensor.
|
||||
shift (torch.Tensor, optional): shift tensor. Defaults to None.
|
||||
scale (torch.Tensor, optional): scale tensor. Defaults to None.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: the output tensor after modulate.
|
||||
"""
|
||||
if scale is None and shift is None:
|
||||
return x
|
||||
elif shift is None:
|
||||
return x * (1 + scale.unsqueeze(1)) # type: ignore[union-attr]
|
||||
elif scale is None:
|
||||
return x + shift.unsqueeze(1) # type: ignore[union-attr]
|
||||
else:
|
||||
return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze(
|
||||
1) # type: ignore[union-attr]
|
||||
|
||||
@@ -2,16 +2,16 @@
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from math import prod
|
||||
from typing import Iterator, Optional, Tuple, Union
|
||||
from typing import Iterator, Optional, Tuple, Union, cast
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
|
||||
from fastvideo.v1.configs.models import VAEConfig
|
||||
from fastvideo.v1.distributed import (get_sequence_model_parallel_rank,
|
||||
get_sequence_model_parallel_world_size)
|
||||
from fastvideo.v1.configs.models import VAEConfig
|
||||
|
||||
|
||||
class ParallelTiledVAE(ABC):
|
||||
@@ -28,7 +28,6 @@ class ParallelTiledVAE(ABC):
|
||||
|
||||
def __init__(self, config: VAEConfig, **kwargs) -> None:
|
||||
self.config = config
|
||||
self.arch_config = config.arch_config
|
||||
self.tile_sample_min_height = config.tile_sample_min_height
|
||||
self.tile_sample_min_width = config.tile_sample_min_width
|
||||
self.tile_sample_min_num_frames = config.tile_sample_min_num_frames
|
||||
@@ -42,15 +41,15 @@ class ParallelTiledVAE(ABC):
|
||||
|
||||
@property
|
||||
def temporal_compression_ratio(self) -> int:
|
||||
return self.arch_config.temporal_compression_ratio
|
||||
|
||||
return cast(int, self.config.temporal_compression_ratio)
|
||||
|
||||
@property
|
||||
def spatial_compression_ratio(self) -> int:
|
||||
return self.arch_config.spatial_compression_ratio
|
||||
|
||||
return cast(int, self.config.spatial_compression_ratio)
|
||||
|
||||
@property
|
||||
def scaling_factor(self) -> Union[float, torch.tensor]:
|
||||
return self.arch_config.scaling_factor
|
||||
return cast(Union[float, torch.tensor], self.config.scaling_factor)
|
||||
|
||||
@abstractmethod
|
||||
def _encode(self, *args, **kwargs) -> torch.Tensor:
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user