Compare commits
14
Commits
wei/api
...
model_config
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
ab55216810 | ||
|
|
477555384e | ||
|
|
24c3ac6ea2 | ||
|
|
d9689cef95 | ||
|
|
60f01567d1 | ||
|
|
e16222a5cc | ||
|
|
4e297a47e5 | ||
|
|
81bc6c9943 | ||
|
|
80450962ef | ||
|
|
7e6236a863 | ||
|
|
684f7feee1 | ||
|
|
c614f154e2 | ||
|
|
f1098c77dc | ||
|
|
0405b618f8 |
+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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -22,3 +22,4 @@ help:
|
||||
clean:
|
||||
@$(SPHINXBUILD) -M clean "$(SOURCEDIR)" "$(BUILDDIR)" $(SPHINXOPTS) $(O)
|
||||
rm -rf "$(SOURCEDIR)/getting_started/examples"
|
||||
rm -rf "$(SOURCEDIR)/inference/examples"
|
||||
|
||||
@@ -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,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,31 @@
|
||||
# 🐳 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.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/
|
||||
```
|
||||
@@ -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/
|
||||
```
|
||||
@@ -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())
|
||||
|
||||
@@ -15,7 +15,7 @@ FastVideo has been tested on the following GPUs, but it should work on any GPUs
|
||||
|
||||
- OS: Linux
|
||||
- Python: 3.10-3.12
|
||||
- CUDA 12.4+ (Untested on CUDA < 12.4)
|
||||
- CUDA 12.4+
|
||||
|
||||
## Installation Options
|
||||
|
||||
@@ -73,7 +73,7 @@ To try Sliding Tile Attention (optional), please follow the instructions in [csr
|
||||
## 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
|
||||
|
||||
|
||||
+27
-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,31 @@ 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/examples/examples_index -->
|
||||
:::
|
||||
|
||||
:::{toctree}
|
||||
:caption: Inference
|
||||
:maxdepth: 1
|
||||
|
||||
inference/examples/examples_inference_index
|
||||
inference/v0_inference
|
||||
:::
|
||||
|
||||
:::{toctree}
|
||||
:caption: Training
|
||||
:maxdepth: 1
|
||||
|
||||
training/data_preprocess
|
||||
training/distillation
|
||||
training/finetune
|
||||
:::
|
||||
|
||||
% What is STA Kernel?
|
||||
@@ -60,27 +78,22 @@ 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
|
||||
:::
|
||||
|
||||
## 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,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"]
|
||||
|
||||
@@ -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
|
||||
@@ -1,9 +0,0 @@
|
||||
from fastvideo.v1.configs.base import BaseConfig, SlidingTileAttnConfig
|
||||
from fastvideo.v1.configs.hunyuan import FastHunyuanConfig, HunyuanConfig
|
||||
from fastvideo.v1.configs.registry import get_pipeline_config_cls_for_name
|
||||
from fastvideo.v1.configs.wan import WanI2V480PConfig, WanT2V480PConfig
|
||||
|
||||
__all__ = [
|
||||
"HunyuanConfig", "FastHunyuanConfig", "BaseConfig", "SlidingTileAttnConfig",
|
||||
"WanT2V480PConfig", "WanI2V480PConfig", "get_pipeline_config_cls_for_name"
|
||||
]
|
||||
|
||||
@@ -1,67 +0,0 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
|
||||
|
||||
@dataclass
|
||||
class BaseConfig:
|
||||
"""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
|
||||
disable_autocast: bool = False
|
||||
|
||||
# Model configuration
|
||||
precision: str = "bf16"
|
||||
|
||||
# VAE configuration
|
||||
vae_precision: str = "fp16"
|
||||
vae_tiling: bool = True
|
||||
vae_sp: bool = True
|
||||
vae_scale_factor: Optional[int] = None
|
||||
|
||||
# DiT configuration
|
||||
num_channels_latents: Optional[int] = None
|
||||
|
||||
# Image encoder configuration
|
||||
image_encoder_precision: str = "fp32"
|
||||
|
||||
# Text encoder configuration
|
||||
text_encoder_precision: str = "fp16"
|
||||
text_len: int = -1
|
||||
hidden_state_skip_layer: int = 0
|
||||
|
||||
# STA (Spatial-Temporal Attention) parameters
|
||||
mask_strategy_file_path: Optional[str] = None
|
||||
enable_torch_compile: bool = False
|
||||
|
||||
neg_prompt: Optional[str] = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class SlidingTileAttnConfig(BaseConfig):
|
||||
"""Configuration for sliding tile attention."""
|
||||
|
||||
# Override any BaseConfig defaults as needed
|
||||
# Add sliding tile specific parameters
|
||||
window_size: int = 16
|
||||
stride: int = 8
|
||||
|
||||
# You can provide custom defaults for inherited fields
|
||||
height: int = 576
|
||||
width: int = 1024
|
||||
|
||||
# Additional configuration specific to sliding tile attention
|
||||
pad_to_square: bool = False
|
||||
use_overlap_optimization: bool = True
|
||||
@@ -0,0 +1,6 @@
|
||||
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__ = ["ModelConfig", "VAEConfig", "DiTConfig", "EncoderConfig"]
|
||||
@@ -0,0 +1,62 @@
|
||||
from dataclasses import dataclass, 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 & 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()
|
||||
|
||||
# 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}'")
|
||||
|
||||
# This should be used only when loading from transformers/diffusers
|
||||
def update_model_arch(self, source_model_dict: Dict[str, Any]) -> None:
|
||||
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}'")
|
||||
|
||||
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."
|
||||
|
||||
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:
|
||||
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,4 @@
|
||||
from fastvideo.v1.configs.models.dits.hunyuanvideo import HunyuanVideoConfig
|
||||
from fastvideo.v1.configs.models.dits.wanvideo import WanVideoConfig
|
||||
|
||||
__all__ = ["HunyuanVideoConfig", "WanVideoConfig"]
|
||||
@@ -0,0 +1,30 @@
|
||||
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)
|
||||
_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
|
||||
|
||||
|
||||
@dataclass
|
||||
class DiTConfig(ModelConfig):
|
||||
arch_config: DiTArchConfig = DiTArchConfig()
|
||||
|
||||
# FastVideoDiT-specific parameters
|
||||
prefix: str = ""
|
||||
quant_config: Optional[QuantizationConfig] = None
|
||||
@@ -0,0 +1,169 @@
|
||||
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])
|
||||
|
||||
|
||||
@dataclass
|
||||
class HunyuanVideoArchConfig(DiTArchConfig):
|
||||
_fsdp_shard_conditions: list = field(
|
||||
default_factory=lambda:
|
||||
[is_double_block, is_single_block, is_refiner_block])
|
||||
|
||||
_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):
|
||||
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 = HunyuanVideoArchConfig()
|
||||
|
||||
prefix: str = "Hunyuan"
|
||||
@@ -0,0 +1,82 @@
|
||||
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):
|
||||
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 = WanVideoArchConfig()
|
||||
|
||||
prefix: str = "Wan"
|
||||
@@ -0,0 +1,12 @@
|
||||
from fastvideo.v1.configs.models.encoders.base import (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",
|
||||
"CLIPTextConfig", "CLIPVisionConfig", "LlamaConfig", "T5Config"
|
||||
]
|
||||
@@ -0,0 +1,55 @@
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, List, 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 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
|
||||
|
||||
|
||||
@dataclass
|
||||
class ImageEncoderArchConfig(EncoderArchConfig):
|
||||
pass
|
||||
|
||||
|
||||
@dataclass
|
||||
class EncoderConfig(ModelConfig):
|
||||
arch_config: ArchConfig = EncoderArchConfig()
|
||||
|
||||
prefix: str = ""
|
||||
quant_config: Optional[QuantizationConfig] = None
|
||||
lora_config: Optional[Any] = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class TextEncoderConfig(EncoderConfig):
|
||||
arch_config: ArchConfig = TextEncoderArchConfig()
|
||||
|
||||
|
||||
@dataclass
|
||||
class ImageEncoderConfig(EncoderConfig):
|
||||
arch_config: ArchConfig = ImageEncoderArchConfig()
|
||||
@@ -0,0 +1,64 @@
|
||||
from dataclasses import dataclass
|
||||
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 = CLIPTextArchConfig()
|
||||
|
||||
num_hidden_layers_override: Optional[int] = None
|
||||
require_post_norm: Optional[bool] = None
|
||||
prefix: str = "clip"
|
||||
|
||||
|
||||
@dataclass
|
||||
class CLIPVisionConfig(ImageEncoderConfig):
|
||||
arch_config: ImageEncoderArchConfig = 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
|
||||
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 = LlamaArchConfig()
|
||||
|
||||
prefix: str = "llama"
|
||||
@@ -0,0 +1,45 @@
|
||||
from dataclasses import dataclass
|
||||
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):
|
||||
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"
|
||||
|
||||
|
||||
@dataclass
|
||||
class T5Config(TextEncoderConfig):
|
||||
arch_config: TextEncoderArchConfig = T5ArchConfig()
|
||||
|
||||
prefix: str = "t5"
|
||||
@@ -0,0 +1,7 @@
|
||||
from fastvideo.v1.configs.models.vaes.hunyuanvae import HunyuanVAEConfig
|
||||
from fastvideo.v1.configs.models.vaes.wanvae import WanVAEConfig
|
||||
|
||||
__all__ = [
|
||||
"HunyuanVAEConfig",
|
||||
"WanVAEConfig",
|
||||
]
|
||||
@@ -0,0 +1,38 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import Union
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.configs.models.base import ArchConfig, ModelConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class VAEArchConfig(ArchConfig):
|
||||
scaling_factor: Union[float, torch.tensor] = 0
|
||||
|
||||
temporal_compression_ratio: int = 4
|
||||
spatial_compression_ratio: int = 8
|
||||
|
||||
|
||||
@dataclass
|
||||
class VAEConfig(ModelConfig):
|
||||
arch_config: VAEArchConfig = VAEArchConfig()
|
||||
|
||||
# FastVideoVAE-specific parameters
|
||||
load_encoder: bool = True
|
||||
load_decoder: bool = True
|
||||
|
||||
tile_sample_min_height: int = 256
|
||||
tile_sample_min_width: int = 256
|
||||
tile_sample_min_num_frames: int = 16
|
||||
tile_sample_stride_height: int = 192
|
||||
tile_sample_stride_width: int = 192
|
||||
tile_sample_stride_num_frames: int = 12
|
||||
blend_num_frames: int = 0
|
||||
|
||||
use_tiling: bool = True
|
||||
use_temporal_tiling: bool = True
|
||||
use_parallel_tiling: bool = True
|
||||
|
||||
def __post_init__(self):
|
||||
self.blend_num_frames = self.tile_sample_min_num_frames - self.tile_sample_stride_num_frames
|
||||
@@ -0,0 +1,40 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import Tuple
|
||||
|
||||
from fastvideo.v1.configs.models.vaes.base import VAEArchConfig, VAEConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class HunyuanVAEArchConfig(VAEArchConfig):
|
||||
in_channels: int = 3
|
||||
out_channels: int = 3
|
||||
latent_channels: int = 16
|
||||
down_block_types: Tuple[str, ...] = (
|
||||
"HunyuanVideoDownBlock3D",
|
||||
"HunyuanVideoDownBlock3D",
|
||||
"HunyuanVideoDownBlock3D",
|
||||
"HunyuanVideoDownBlock3D",
|
||||
)
|
||||
up_block_types: Tuple[str, ...] = (
|
||||
"HunyuanVideoUpBlock3D",
|
||||
"HunyuanVideoUpBlock3D",
|
||||
"HunyuanVideoUpBlock3D",
|
||||
"HunyuanVideoUpBlock3D",
|
||||
)
|
||||
block_out_channels: Tuple[int, ...] = (128, 256, 512, 512)
|
||||
layers_per_block: int = 2
|
||||
act_fn: str = "silu"
|
||||
norm_num_groups: int = 32
|
||||
scaling_factor: float = 0.476986
|
||||
spatial_compression_ratio: int = 8
|
||||
temporal_compression_ratio: int = 4
|
||||
mid_block_add_attention: bool = True
|
||||
|
||||
def __post_init__(self):
|
||||
self.spatial_compression_ratio: int = 2**(len(self.block_out_channels) -
|
||||
1)
|
||||
|
||||
|
||||
@dataclass
|
||||
class HunyuanVAEConfig(VAEConfig):
|
||||
arch_config: VAEArchConfig = HunyuanVAEArchConfig()
|
||||
@@ -0,0 +1,75 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import Tuple
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.configs.models.vaes.base import VAEArchConfig, VAEConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class WanVAEArchConfig(VAEArchConfig):
|
||||
base_dim: int = 96
|
||||
z_dim: int = 16
|
||||
dim_mult: Tuple[int, ...] = (1, 2, 4, 4)
|
||||
num_res_blocks: int = 2
|
||||
attn_scales: Tuple[float, ...] = ()
|
||||
temperal_downsample: Tuple[bool, ...] = (False, True, True)
|
||||
dropout: float = 0.0
|
||||
latents_mean: Tuple[float, ...] = (
|
||||
-0.7571,
|
||||
-0.7089,
|
||||
-0.9113,
|
||||
0.1075,
|
||||
-0.1745,
|
||||
0.9653,
|
||||
-0.1517,
|
||||
1.5508,
|
||||
0.4134,
|
||||
-0.0715,
|
||||
0.5517,
|
||||
-0.3632,
|
||||
-0.1922,
|
||||
-0.9497,
|
||||
0.2503,
|
||||
-0.2921,
|
||||
)
|
||||
latents_std: Tuple[float, ...] = (
|
||||
2.8184,
|
||||
1.4541,
|
||||
2.3275,
|
||||
2.6558,
|
||||
1.2196,
|
||||
1.7708,
|
||||
2.6052,
|
||||
2.0743,
|
||||
3.2687,
|
||||
2.1526,
|
||||
2.8652,
|
||||
1.5579,
|
||||
1.6382,
|
||||
1.1253,
|
||||
2.8251,
|
||||
1.9160,
|
||||
)
|
||||
temporal_compression_ratio = 4
|
||||
spatial_compression_ratio = 8
|
||||
|
||||
def __post_init__(self):
|
||||
self.scaling_factor: torch.tensor = 1.0 / torch.tensor(
|
||||
self.latents_std).view(1, self.z_dim, 1, 1, 1)
|
||||
self.shift_factor: torch.tensor = torch.tensor(self.latents_mean).view(
|
||||
1, self.z_dim, 1, 1, 1)
|
||||
|
||||
|
||||
@dataclass
|
||||
class WanVAEConfig(VAEConfig):
|
||||
arch_config: VAEArchConfig = WanVAEArchConfig()
|
||||
use_feature_cache: bool = True
|
||||
|
||||
use_tiling: bool = False
|
||||
use_temporal_tiling: bool = False
|
||||
use_parallel_tiling: bool = False
|
||||
|
||||
def __post_init__(self):
|
||||
self.blend_num_frames = (self.tile_sample_min_num_frames -
|
||||
self.tile_sample_stride_num_frames) * 2
|
||||
@@ -0,0 +1,14 @@
|
||||
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,
|
||||
WanT2V480PConfig)
|
||||
|
||||
__all__ = [
|
||||
"HunyuanConfig", "FastHunyuanConfig", "PipelineConfig",
|
||||
"SlidingTileAttnConfig", "WanT2V480PConfig", "WanI2V480PConfig",
|
||||
"get_pipeline_config_cls_for_name"
|
||||
]
|
||||
@@ -0,0 +1,107 @@
|
||||
import json
|
||||
from dataclasses import asdict, dataclass, fields
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
from fastvideo.v1.configs.models import (DiTConfig, EncoderConfig, ModelConfig,
|
||||
VAEConfig)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.utils import shallow_asdict
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class PipelineConfig:
|
||||
"""Base configuration for all pipeline architectures."""
|
||||
# Video generation parameters
|
||||
embedded_cfg_scale: float = 6.0
|
||||
flow_shift: Optional[float] = None
|
||||
use_cpu_offload: bool = False
|
||||
disable_autocast: bool = False
|
||||
|
||||
# Model configuration
|
||||
precision: str = "bf16"
|
||||
|
||||
# VAE configuration
|
||||
vae_precision: str = "fp16"
|
||||
vae_tiling: bool = True
|
||||
vae_sp: bool = True
|
||||
vae_config: VAEConfig = VAEConfig()
|
||||
|
||||
# DiT configuration
|
||||
dit_config: DiTConfig = DiTConfig()
|
||||
|
||||
# Text encoder configuration
|
||||
text_encoder_precision: str = "fp16"
|
||||
text_encoder_config: EncoderConfig = EncoderConfig()
|
||||
|
||||
# STA (Spatial-Temporal Attention) parameters
|
||||
mask_strategy_file_path: Optional[str] = None
|
||||
enable_torch_compile: bool = False
|
||||
|
||||
@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 pipeline_config
|
||||
|
||||
def dump_to_json(self, file_path: str):
|
||||
output_dict = shallow_asdict(self)
|
||||
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
|
||||
|
||||
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) 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:
|
||||
for f in fields(self):
|
||||
key = f.name
|
||||
if key in source_pipeline_dict:
|
||||
current_value = getattr(self, key)
|
||||
new_value = source_pipeline_dict[key]
|
||||
|
||||
# If it's a nested ModelConfig, update it recursively
|
||||
if isinstance(current_value, ModelConfig):
|
||||
current_value.update_model_config(new_value)
|
||||
else:
|
||||
setattr(self, key, new_value)
|
||||
|
||||
if hasattr(self, "__post_init__"):
|
||||
self.__post_init__()
|
||||
|
||||
|
||||
@dataclass
|
||||
class SlidingTileAttnConfig(PipelineConfig):
|
||||
"""Configuration for sliding tile attention."""
|
||||
|
||||
# Override any BaseConfig defaults as needed
|
||||
# Add sliding tile specific parameters
|
||||
window_size: int = 16
|
||||
stride: int = 8
|
||||
|
||||
# You can provide custom defaults for inherited fields
|
||||
height: int = 576
|
||||
width: int = 1024
|
||||
|
||||
# Additional configuration specific to sliding tile attention
|
||||
pad_to_square: bool = False
|
||||
use_overlap_optimization: bool = True
|
||||
@@ -1,21 +1,27 @@
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastvideo.v1.configs.base import BaseConfig
|
||||
from fastvideo.v1.configs.models import DiTConfig, EncoderConfig, VAEConfig
|
||||
from fastvideo.v1.configs.models.dits import HunyuanVideoConfig
|
||||
from fastvideo.v1.configs.models.encoders import CLIPTextConfig, LlamaConfig
|
||||
from fastvideo.v1.configs.models.vaes import HunyuanVAEConfig
|
||||
from fastvideo.v1.configs.pipelines.base import PipelineConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class HunyuanConfig(BaseConfig):
|
||||
class HunyuanConfig(PipelineConfig):
|
||||
"""Base configuration for HunYuan pipeline architecture."""
|
||||
|
||||
# HunyuanConfig-specific parameters with defaults
|
||||
# DiT
|
||||
dit_config: DiTConfig = HunyuanVideoConfig()
|
||||
# VAE
|
||||
vae_config: VAEConfig = 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_config: EncoderConfig = LlamaConfig()
|
||||
|
||||
# Precision for each component
|
||||
precision: str = "bf16"
|
||||
@@ -24,8 +30,12 @@ class HunyuanConfig(BaseConfig):
|
||||
|
||||
# HunyuanConfig-specific added parameters
|
||||
# Secondary text encoder
|
||||
text_encoder_config_2: EncoderConfig = CLIPTextConfig()
|
||||
text_encoder_precision_2: str = "fp16"
|
||||
text_len_2: int = 77
|
||||
|
||||
def __post_init__(self):
|
||||
self.vae_config.load_encoder = False
|
||||
self.vae_config.load_decoder = True
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -33,7 +43,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,9 +3,11 @@
|
||||
import os
|
||||
from typing import Callable, Dict, Optional, Type
|
||||
|
||||
from fastvideo.v1.configs.base import BaseConfig
|
||||
from fastvideo.v1.configs.hunyuan import FastHunyuanConfig, HunyuanConfig
|
||||
from fastvideo.v1.configs.wan import WanI2V480PConfig, WanT2V480PConfig
|
||||
from fastvideo.v1.configs.pipelines.base import PipelineConfig
|
||||
from fastvideo.v1.configs.pipelines.hunyuan import (FastHunyuanConfig,
|
||||
HunyuanConfig)
|
||||
from fastvideo.v1.configs.pipelines.wan import (WanI2V480PConfig,
|
||||
WanT2V480PConfig)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.utils import (maybe_download_model_index,
|
||||
verify_model_config_and_directory)
|
||||
@@ -13,8 +15,8 @@ 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
|
||||
@@ -30,7 +32,7 @@ PIPELINE_DETECTOR: Dict[str, Callable[[str], bool]] = {
|
||||
}
|
||||
|
||||
# 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":
|
||||
@@ -41,7 +43,7 @@ PIPELINE_FALLBACK_CONFIG: Dict[str, Type[BaseConfig]] = {
|
||||
|
||||
|
||||
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):
|
||||
@@ -65,7 +67,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()):
|
||||
@@ -0,0 +1,55 @@
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastvideo.v1.configs.models import DiTConfig, EncoderConfig, VAEConfig
|
||||
from fastvideo.v1.configs.models.dits import WanVideoConfig
|
||||
from fastvideo.v1.configs.models.encoders import CLIPVisionConfig, T5Config
|
||||
from fastvideo.v1.configs.models.vaes import WanVAEConfig
|
||||
from fastvideo.v1.configs.pipelines.base import PipelineConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class WanT2V480PConfig(PipelineConfig):
|
||||
"""Base configuration for Wan T2V 1.3B pipeline architecture."""
|
||||
|
||||
# WanConfig-specific parameters with defaults
|
||||
# DiT
|
||||
dit_config: DiTConfig = WanVideoConfig()
|
||||
# VAE
|
||||
vae_config: VAEConfig = WanVAEConfig()
|
||||
vae_tiling: bool = False
|
||||
vae_sp: bool = False
|
||||
|
||||
# Video parameters
|
||||
use_cpu_offload: bool = True
|
||||
|
||||
# Denoising stage
|
||||
flow_shift: int = 3
|
||||
|
||||
# Text encoding stage
|
||||
text_encoder_config: EncoderConfig = T5Config()
|
||||
|
||||
# Precision for each component
|
||||
precision: str = "bf16"
|
||||
vae_precision: str = "fp16"
|
||||
text_encoder_precision: str = "fp32"
|
||||
|
||||
# WanConfig-specific added parameters
|
||||
|
||||
def __post_init__(self):
|
||||
self.vae_config.load_encoder = False
|
||||
self.vae_config.load_decoder = True
|
||||
|
||||
|
||||
@dataclass
|
||||
class WanI2V480PConfig(WanT2V480PConfig):
|
||||
"""Base configuration for Wan I2V 14B 480P pipeline architecture."""
|
||||
|
||||
# WanConfig-specific parameters with defaults
|
||||
|
||||
# Precision for each component
|
||||
image_encoder_config: EncoderConfig = CLIPVisionConfig()
|
||||
image_encoder_precision: str = "fp32"
|
||||
|
||||
def __post_init__(self):
|
||||
self.vae_config.load_encoder = True
|
||||
self.vae_config.load_decoder = True
|
||||
@@ -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,72 @@
|
||||
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:
|
||||
# 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
|
||||
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
|
||||
|
||||
# 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
|
||||
@@ -0,0 +1,20 @@
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastvideo.v1.configs.sample.base import SamplingParam
|
||||
|
||||
|
||||
@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
|
||||
|
||||
|
||||
@dataclass
|
||||
class FastHunyuanSamplingParam(HunyuanSamplingParam):
|
||||
num_inference_steps: int = 6
|
||||
@@ -0,0 +1,75 @@
|
||||
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 (WanI2V480PSamplingParam,
|
||||
WanT2V480PSamplingParam)
|
||||
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": WanT2V480PSamplingParam,
|
||||
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers": WanI2V480PSamplingParam
|
||||
# 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":
|
||||
WanT2V480PSamplingParam, # Base Wan config as fallback for any Wan variant
|
||||
"wanimagetovideo": WanI2V480PSamplingParam,
|
||||
# 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,24 @@
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastvideo.v1.configs.sample.base import SamplingParam
|
||||
|
||||
|
||||
@dataclass
|
||||
class WanT2V480PSamplingParam(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
|
||||
|
||||
|
||||
@dataclass
|
||||
class WanI2V480PSamplingParam(WanT2V480PSamplingParam):
|
||||
# Denoising stage
|
||||
guidance_scale: float = 5.0
|
||||
num_inference_steps: int = 40
|
||||
@@ -1,45 +0,0 @@
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastvideo.v1.configs.base import BaseConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class WanT2V480PConfig(BaseConfig):
|
||||
"""Base configuration for Wan T2V 1.3B pipeline architecture."""
|
||||
|
||||
# WanConfig-specific parameters with defaults
|
||||
# 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
|
||||
|
||||
# Precision for each component
|
||||
precision: str = "bf16"
|
||||
vae_precision: str = "fp16"
|
||||
text_encoder_precision: str = "fp32"
|
||||
|
||||
# WanConfig-specific added parameters
|
||||
|
||||
|
||||
@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_precision: str = "fp32"
|
||||
@@ -0,0 +1,37 @@
|
||||
{
|
||||
"embedded_cfg_scale": 6.0,
|
||||
"flow_shift": 3,
|
||||
"use_cpu_offload": true,
|
||||
"disable_autocast": false,
|
||||
"precision": "bf16",
|
||||
"vae_precision": "fp16",
|
||||
"vae_tiling": false,
|
||||
"vae_sp": false,
|
||||
"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": 8,
|
||||
"use_tiling": false,
|
||||
"use_temporal_tiling": false,
|
||||
"use_parallel_tiling": false,
|
||||
"use_feature_cache": true
|
||||
},
|
||||
"dit_config": {
|
||||
"prefix": "Wan",
|
||||
"quant_config": null
|
||||
},
|
||||
"text_encoder_precision": "fp32",
|
||||
"text_encoder_config": {
|
||||
"prefix": "t5",
|
||||
"quant_config": null,
|
||||
"lora_config": null
|
||||
},
|
||||
"mask_strategy_file_path": null,
|
||||
"enable_torch_compile": false
|
||||
}
|
||||
@@ -0,0 +1,45 @@
|
||||
{
|
||||
"embedded_cfg_scale": 6.0,
|
||||
"flow_shift": 3,
|
||||
"use_cpu_offload": true,
|
||||
"disable_autocast": false,
|
||||
"precision": "bf16",
|
||||
"vae_precision": "fp16",
|
||||
"vae_tiling": false,
|
||||
"vae_sp": false,
|
||||
"vae_config": {
|
||||
"load_encoder": true,
|
||||
"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": 8,
|
||||
"use_tiling": false,
|
||||
"use_temporal_tiling": false,
|
||||
"use_parallel_tiling": false,
|
||||
"use_feature_cache": true
|
||||
},
|
||||
"dit_config": {
|
||||
"prefix": "Wan",
|
||||
"quant_config": null
|
||||
},
|
||||
"text_encoder_precision": "fp32",
|
||||
"text_encoder_config": {
|
||||
"prefix": "t5",
|
||||
"quant_config": null,
|
||||
"lora_config": null
|
||||
},
|
||||
"mask_strategy_file_path": null,
|
||||
"enable_torch_compile": false,
|
||||
"image_encoder_config": {
|
||||
"prefix": "clip",
|
||||
"quant_config": null,
|
||||
"lora_config": null,
|
||||
"num_hidden_layers_override": null,
|
||||
"require_post_norm": null
|
||||
},
|
||||
"image_encoder_precision": "fp32"
|
||||
}
|
||||
@@ -9,7 +9,7 @@ diffusion models.
|
||||
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,11 +17,13 @@ import torch
|
||||
import torchvision
|
||||
from einops import rearrange
|
||||
|
||||
from fastvideo.v1.configs import get_pipeline_config_cls_for_name
|
||||
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
|
||||
from fastvideo.v1.utils import align_to
|
||||
from fastvideo.v1.utils import align_to, shallow_asdict
|
||||
from fastvideo.v1.worker.executor import Executor
|
||||
|
||||
logger = init_logger(__name__)
|
||||
@@ -52,6 +54,9 @@ class VideoGenerator:
|
||||
model_path: str,
|
||||
device: Optional[str] = None,
|
||||
torch_dtype: Optional[torch.dtype] = None,
|
||||
pipeline_config: Optional[
|
||||
Union[str
|
||||
| PipelineConfig]] = None,
|
||||
**kwargs) -> "VideoGenerator":
|
||||
"""
|
||||
Create a video generator from a pretrained model.
|
||||
@@ -64,22 +69,30 @@ class VideoGenerator:
|
||||
|
||||
Returns:
|
||||
The created video generator
|
||||
|
||||
Priority level: Default pipeline config < User's pipeline config < User's kwargs
|
||||
"""
|
||||
|
||||
config = None
|
||||
config_cls = get_pipeline_config_cls_for_name(model_path)
|
||||
if config_cls is not None:
|
||||
config = config_cls()
|
||||
# 1. If users provide a pipeline config, it will override the default pipeline config
|
||||
if isinstance(pipeline_config, PipelineConfig):
|
||||
config = pipeline_config
|
||||
else:
|
||||
config_cls = get_pipeline_config_cls_for_name(model_path)
|
||||
if config_cls is not None:
|
||||
config = config_cls()
|
||||
if isinstance(pipeline_config, str):
|
||||
config.load_from_json(pipeline_config)
|
||||
|
||||
# 2. If users also provide some kwargs, it will override the pipeline config.
|
||||
# The user kwargs shouldn't contain model config parameters!
|
||||
if config is None:
|
||||
logger.warning("No config found for model %s, using default config",
|
||||
model_path)
|
||||
config_args = {}
|
||||
config_args = kwargs
|
||||
else:
|
||||
config_args = asdict(config)
|
||||
|
||||
# override config_args with kwargs
|
||||
config_args.update(kwargs)
|
||||
config_args = shallow_asdict(config)
|
||||
config_args.update(kwargs)
|
||||
|
||||
fastvideo_args = FastVideoArgs(
|
||||
model_path=model_path,
|
||||
@@ -115,19 +128,8 @@ class VideoGenerator:
|
||||
def generate_video(
|
||||
self,
|
||||
prompt: str,
|
||||
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.
|
||||
@@ -154,87 +156,68 @@ class VideoGenerator:
|
||||
# Create a copy of inference args to avoid modifying the original
|
||||
fastvideo_args = self.fastvideo_args
|
||||
|
||||
# Override parameters if provided
|
||||
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}"""
|
||||
logger.info(debug_str)
|
||||
|
||||
# Prepare batch
|
||||
device = torch.device(fastvideo_args.device_str)
|
||||
batch = ForwardBatch(
|
||||
prompt=prompt,
|
||||
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,
|
||||
**asdict(sampling_param),
|
||||
eta=0.0,
|
||||
n_tokens=n_tokens,
|
||||
data_type="video" if fastvideo_args.num_frames > 1 else "image",
|
||||
device=device,
|
||||
extra={},
|
||||
)
|
||||
|
||||
@@ -255,23 +238,22 @@ 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)
|
||||
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
|
||||
}
|
||||
|
||||
@@ -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),
|
||||
|
||||
|
||||
@@ -9,7 +9,7 @@ The first script in this example shows the most basic usage of FastVideo. If you
|
||||
python fastvideo/v1/examples/inference/basic/basic.py
|
||||
```
|
||||
|
||||
# Basic Walkthrough
|
||||
## Basic Walkthrough
|
||||
|
||||
All you need to generate videos using multi-gpus from state-of-the-art diffusion pipelines is the following few lines!
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -11,9 +12,13 @@ def main():
|
||||
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!
|
||||
|
||||
+10
-153
@@ -7,6 +7,7 @@ import dataclasses
|
||||
from contextlib import contextmanager
|
||||
from typing import List, Optional
|
||||
|
||||
from fastvideo.v1.configs.models import DiTConfig, EncoderConfig, VAEConfig
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.utils import FlexibleArgumentParser
|
||||
|
||||
@@ -34,55 +35,38 @@ 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 = DiTConfig()
|
||||
precision: str = "bf16"
|
||||
|
||||
# VAE configuration
|
||||
vae_precision: str = "fp16"
|
||||
vae_tiling: bool = True
|
||||
vae_sp: bool = False
|
||||
vae_scale_factor: Optional[int] = None
|
||||
|
||||
# DiT configuration
|
||||
num_channels_latents: Optional[int] = None
|
||||
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()
|
||||
|
||||
# Image encoder configuration
|
||||
image_encoder_precision: str = "fp32"
|
||||
image_encoder_config: EncoderConfig = EncoderConfig()
|
||||
|
||||
# Text encoder configuration
|
||||
text_encoder_precision: str = "fp16"
|
||||
text_len: int = 256
|
||||
hidden_state_skip_layer: int = 2
|
||||
text_encoder_config: EncoderConfig = EncoderConfig()
|
||||
|
||||
# Secondary text encoder
|
||||
text_encoder_config_2: EncoderConfig = EncoderConfig()
|
||||
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
|
||||
|
||||
# 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
|
||||
|
||||
@@ -90,11 +74,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
|
||||
|
||||
@@ -174,43 +153,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,
|
||||
@@ -267,12 +209,6 @@ class FastVideoArgs:
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
help="Precision for text encoder",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--text-len",
|
||||
type=int,
|
||||
default=FastVideoArgs.text_len,
|
||||
help="Maximum text length",
|
||||
)
|
||||
|
||||
# Image encoder config
|
||||
parser.add_argument(
|
||||
@@ -292,26 +228,6 @@ class FastVideoArgs:
|
||||
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(
|
||||
@@ -326,33 +242,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",
|
||||
@@ -373,36 +262,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
|
||||
@@ -451,8 +310,6 @@ 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")
|
||||
|
||||
|
||||
_current_fastvideo_args = None
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
# type: ignore
|
||||
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Inference module for diffusion models.
|
||||
|
||||
@@ -5,19 +5,20 @@ from typing import 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
|
||||
_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 = [
|
||||
@@ -30,8 +31,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 +51,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(
|
||||
|
||||
@@ -6,6 +6,7 @@ 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.distributed.parallel_state import (
|
||||
get_sequence_model_parallel_world_size)
|
||||
from fastvideo.v1.layers.layernorm import (LayerNormScaleShift, ScaleResidual,
|
||||
@@ -431,238 +432,104 @@ 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
|
||||
_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__()
|
||||
|
||||
|
||||
@@ -7,6 +7,7 @@ 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.distributed.parallel_state import (
|
||||
get_sequence_model_parallel_world_size)
|
||||
from fastvideo.v1.layers.layernorm import (LayerNormScaleShift, RMSNorm,
|
||||
@@ -350,114 +351,59 @@ class WanTransformerBlock(nn.Module):
|
||||
|
||||
|
||||
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",
|
||||
}
|
||||
_fsdp_shard_conditions = WanVideoConfig()._fsdp_shard_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)
|
||||
|
||||
|
||||
@@ -2,15 +2,17 @@ from typing import Tuple
|
||||
|
||||
from torch import nn
|
||||
|
||||
from fastvideo.v1.configs.models.encoders import EncoderConfig
|
||||
from fastvideo.v1.platforms import _Backend
|
||||
|
||||
|
||||
class BaseEncoder(nn.Module):
|
||||
_supported_attention_backends: Tuple[_Backend,
|
||||
...] = (_Backend.TORCH_SDPA, )
|
||||
_supported_attention_backends: Tuple[
|
||||
_Backend, ...] = EncoderConfig()._supported_attention_backends
|
||||
|
||||
def __init__(self, *args, **kwargs) -> None:
|
||||
def __init__(self, config: EncoderConfig) -> None:
|
||||
super().__init__()
|
||||
self.config = config
|
||||
if not self.supported_attention_backends:
|
||||
raise ValueError(
|
||||
f"Subclass {self.__class__.__name__} must define _supported_attention_backends"
|
||||
|
||||
@@ -3,16 +3,17 @@
|
||||
# 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 (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
|
||||
@@ -20,45 +21,14 @@ 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.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,9 +346,6 @@ 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,
|
||||
@@ -393,7 +359,6 @@ class CLIPTextTransformer(nn.Module):
|
||||
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)
|
||||
@@ -469,21 +434,15 @@ class CLIPTextTransformer(nn.Module):
|
||||
|
||||
|
||||
class CLIPTextModel(BaseEncoder):
|
||||
_supported_attention_backends = (_Backend.FLASH_ATTN, _Backend.TORCH_SDPA)
|
||||
|
||||
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,
|
||||
@@ -492,9 +451,7 @@ class CLIPTextModel(BaseEncoder):
|
||||
position_ids: Optional[torch.Tensor] = None,
|
||||
output_attentions: Optional[bool] = 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
|
||||
|
||||
return self.text_model(
|
||||
input_ids=input_ids,
|
||||
@@ -548,7 +505,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,30 +571,19 @@ class CLIPVisionTransformer(nn.Module):
|
||||
return encoder_outputs
|
||||
|
||||
|
||||
class CLIPVisionModel(BaseEncoder, SupportsQuant):
|
||||
class CLIPVisionModel(BaseEncoder):
|
||||
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,
|
||||
|
||||
@@ -23,15 +23,17 @@
|
||||
# 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 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
|
||||
@@ -42,12 +44,6 @@ from fastvideo.v1.layers.vocab_parallel_embedding import VocabParallelEmbedding
|
||||
from fastvideo.v1.models.encoders.base import BaseEncoder
|
||||
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 +167,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,
|
||||
@@ -280,27 +276,22 @@ class LlamaDecoderLayer(nn.Module):
|
||||
|
||||
|
||||
class LlamaModel(BaseEncoder):
|
||||
_supported_attention_backends = (_Backend.FLASH_ATTN, _Backend.TORCH_SDPA)
|
||||
|
||||
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 +302,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)
|
||||
])
|
||||
|
||||
@@ -395,17 +386,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 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 BaseEncoder
|
||||
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(BaseEncoder):
|
||||
|
||||
def __init__(self, config: T5Config, prefix: str = ""):
|
||||
super().__init__()
|
||||
super().__init__(config)
|
||||
|
||||
quant_config = None
|
||||
|
||||
@@ -589,10 +587,10 @@ class T5EncoderModel(nn.Module):
|
||||
return loaded_params
|
||||
|
||||
|
||||
class UMT5EncoderModel(nn.Module):
|
||||
class UMT5EncoderModel(BaseEncoder):
|
||||
|
||||
def __init__(self, config: T5Config, prefix: str = ""):
|
||||
super().__init__()
|
||||
super().__init__(config)
|
||||
|
||||
quant_config = None
|
||||
|
||||
|
||||
@@ -2,6 +2,7 @@
|
||||
|
||||
import dataclasses
|
||||
import glob
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
from abc import ABC, abstractmethod
|
||||
@@ -10,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 (
|
||||
@@ -201,18 +201,34 @@ 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)
|
||||
|
||||
try:
|
||||
encoder_config = fastvideo_args.text_encoder_config
|
||||
encoder_config.update_model_arch(model_config)
|
||||
encoder_precision = fastvideo_args.text_encoder_precision
|
||||
except Exception:
|
||||
encoder_config = fastvideo_args.text_encoder_config_2
|
||||
encoder_config.update_model_arch(model_config)
|
||||
encoder_precision = fastvideo_args.text_encoder_precision_2
|
||||
|
||||
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,
|
||||
@@ -251,17 +267,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)
|
||||
|
||||
|
||||
@@ -288,7 +313,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',
|
||||
@@ -310,9 +336,11 @@ class VAELoader(ComponentLoader):
|
||||
assert class_name is not None, "Model config does not contain a _class_name attribute. Only diffusers format is supported."
|
||||
config.pop("_diffusers_version")
|
||||
|
||||
vae_config = fastvideo_args.vae_config
|
||||
vae_config.update_model_arch(config)
|
||||
vae_cls, _ = ModelRegistry.resolve_model_cls(class_name)
|
||||
|
||||
vae = vae_cls(**config).to(fastvideo_args.device)
|
||||
vae = vae_cls(vae_config).to(fastvideo_args.device)
|
||||
|
||||
# Find all safetensors files
|
||||
safetensors_list = glob.glob(
|
||||
@@ -322,7 +350,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)
|
||||
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)
|
||||
|
||||
@@ -335,13 +364,17 @@ 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.pop("_diffusers_version")
|
||||
|
||||
# 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)
|
||||
|
||||
@@ -360,7 +393,7 @@ 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,
|
||||
|
||||
@@ -2,13 +2,14 @@
|
||||
|
||||
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)
|
||||
|
||||
@@ -20,29 +21,35 @@ class ParallelTiledVAE(ABC):
|
||||
tile_sample_stride_height: int
|
||||
tile_sample_stride_width: int
|
||||
tile_sample_stride_num_frames: int
|
||||
blend_num_frames: int
|
||||
use_tiling: bool
|
||||
use_temporal_tiling: bool
|
||||
use_parallel_tiling: bool
|
||||
temporal_compression_ratio: int
|
||||
spatial_compression_ratio: int
|
||||
scaling_factor: Union[float, torch.tensor]
|
||||
|
||||
def __init__(self, *args, **kwargs) -> None:
|
||||
# Check if subclass has defined all required properties
|
||||
required_attributes = [
|
||||
'tile_sample_min_height', 'tile_sample_min_width',
|
||||
'tile_sample_min_num_frames', 'tile_sample_stride_height',
|
||||
'tile_sample_stride_width', 'tile_sample_stride_num_frames',
|
||||
'spatial_compression_ratio', 'temporal_compression_ratio',
|
||||
'use_tiling', 'use_temporal_tiling', 'use_parallel_tiling',
|
||||
'scaling_factor'
|
||||
]
|
||||
def __init__(self, config: VAEConfig, **kwargs) -> None:
|
||||
self.config = 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
|
||||
self.tile_sample_stride_height = config.tile_sample_stride_height
|
||||
self.tile_sample_stride_width = config.tile_sample_stride_width
|
||||
self.tile_sample_stride_num_frames = config.tile_sample_stride_num_frames
|
||||
self.blend_num_frames = config.blend_num_frames
|
||||
self.use_tiling = config.use_tiling
|
||||
self.use_temporal_tiling = config.use_temporal_tiling
|
||||
self.use_parallel_tiling = config.use_parallel_tiling
|
||||
|
||||
for attr in required_attributes:
|
||||
if not hasattr(self, attr):
|
||||
raise AttributeError(
|
||||
f"Subclasses of ParallelVAE must define '{attr}' property")
|
||||
self.blend_num_frames = self.tile_sample_min_num_frames - self.tile_sample_stride_num_frames
|
||||
@property
|
||||
def temporal_compression_ratio(self) -> int:
|
||||
return cast(int, self.config.temporal_compression_ratio)
|
||||
|
||||
@property
|
||||
def spatial_compression_ratio(self) -> int:
|
||||
return cast(int, self.config.spatial_compression_ratio)
|
||||
|
||||
@property
|
||||
def scaling_factor(self) -> Union[float, torch.tensor]:
|
||||
return cast(Union[float, torch.tensor], self.config.scaling_factor)
|
||||
|
||||
@abstractmethod
|
||||
def _encode(self, *args, **kwargs) -> torch.Tensor:
|
||||
@@ -408,6 +415,10 @@ class ParallelTiledVAE(ABC):
|
||||
tile_sample_stride_height: Optional[int] = None,
|
||||
tile_sample_stride_width: Optional[int] = None,
|
||||
tile_sample_stride_num_frames: Optional[int] = None,
|
||||
blend_num_frames: Optional[int] = None,
|
||||
use_tiling: Optional[bool] = None,
|
||||
use_temporal_tiling: Optional[bool] = None,
|
||||
use_parallel_tiling: Optional[bool] = None,
|
||||
) -> None:
|
||||
r"""
|
||||
Enable tiled VAE decoding. When this option is enabled, the VAE will split the input tensor into tiles to
|
||||
@@ -439,7 +450,13 @@ class ParallelTiledVAE(ABC):
|
||||
self.tile_sample_stride_height = tile_sample_stride_height or self.tile_sample_stride_height
|
||||
self.tile_sample_stride_width = tile_sample_stride_width or self.tile_sample_stride_width
|
||||
self.tile_sample_stride_num_frames = tile_sample_stride_num_frames or self.tile_sample_stride_num_frames
|
||||
self.blend_num_frames = self.tile_sample_min_num_frames - self.tile_sample_stride_num_frames
|
||||
if blend_num_frames is not None:
|
||||
self.blend_num_frames = blend_num_frames
|
||||
else:
|
||||
self.blend_num_frames = self.tile_sample_min_num_frames - self.tile_sample_stride_num_frames
|
||||
self.use_tiling = use_tiling or self.use_tiling
|
||||
self.use_temporal_tiling = use_temporal_tiling or self.use_temporal_tiling
|
||||
self.use_parallel_tiling = use_parallel_tiling or self.use_parallel_tiling
|
||||
|
||||
def disable_tiling(self) -> None:
|
||||
r"""
|
||||
|
||||
@@ -21,10 +21,9 @@ import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import torch.utils.checkpoint
|
||||
|
||||
from fastvideo.v1.configs.models.vaes import HunyuanVAEConfig
|
||||
from fastvideo.v1.layers.activation import get_act_fn
|
||||
from fastvideo.v1.models.utils import auto_attributes
|
||||
from fastvideo.v1.models.vaes.common import ParallelTiledVAE
|
||||
|
||||
|
||||
@@ -773,95 +772,52 @@ class AutoencoderKLHunyuanVideo(nn.Module, ParallelTiledVAE):
|
||||
|
||||
_supports_gradient_checkpointing = True
|
||||
|
||||
@auto_attributes
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int = 3,
|
||||
out_channels: int = 3,
|
||||
latent_channels: int = 16,
|
||||
down_block_types: Tuple[str, ...] = (
|
||||
"HunyuanVideoDownBlock3D",
|
||||
"HunyuanVideoDownBlock3D",
|
||||
"HunyuanVideoDownBlock3D",
|
||||
"HunyuanVideoDownBlock3D",
|
||||
),
|
||||
up_block_types: Tuple[str, ...] = (
|
||||
"HunyuanVideoUpBlock3D",
|
||||
"HunyuanVideoUpBlock3D",
|
||||
"HunyuanVideoUpBlock3D",
|
||||
"HunyuanVideoUpBlock3D",
|
||||
),
|
||||
block_out_channels: Tuple[int, ...] = (128, 256, 512, 512),
|
||||
layers_per_block: int = 2,
|
||||
act_fn: str = "silu",
|
||||
norm_num_groups: int = 32,
|
||||
scaling_factor: float = 0.476986,
|
||||
spatial_compression_ratio: int = 8,
|
||||
temporal_compression_ratio: int = 4,
|
||||
mid_block_add_attention: bool = True,
|
||||
load_encoder: bool = True,
|
||||
load_decoder: bool = True,
|
||||
config: HunyuanVAEConfig,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
nn.Module.__init__(self)
|
||||
ParallelTiledVAE.__init__(self, config)
|
||||
|
||||
# TODO(will): only pass in config. We do this by manually defining a
|
||||
# config for hunyuan vae
|
||||
self.block_out_channels = block_out_channels
|
||||
self.block_out_channels = config.block_out_channels
|
||||
|
||||
if load_encoder:
|
||||
if config.load_encoder:
|
||||
self.encoder = HunyuanVideoEncoder3D(
|
||||
in_channels=in_channels,
|
||||
out_channels=latent_channels,
|
||||
down_block_types=down_block_types,
|
||||
block_out_channels=block_out_channels,
|
||||
layers_per_block=layers_per_block,
|
||||
norm_num_groups=norm_num_groups,
|
||||
act_fn=act_fn,
|
||||
in_channels=config.in_channels,
|
||||
out_channels=config.latent_channels,
|
||||
down_block_types=config.down_block_types,
|
||||
block_out_channels=config.block_out_channels,
|
||||
layers_per_block=config.layers_per_block,
|
||||
norm_num_groups=config.norm_num_groups,
|
||||
act_fn=config.act_fn,
|
||||
double_z=True,
|
||||
mid_block_add_attention=mid_block_add_attention,
|
||||
temporal_compression_ratio=temporal_compression_ratio,
|
||||
spatial_compression_ratio=spatial_compression_ratio,
|
||||
mid_block_add_attention=config.mid_block_add_attention,
|
||||
temporal_compression_ratio=config.temporal_compression_ratio,
|
||||
spatial_compression_ratio=config.spatial_compression_ratio,
|
||||
)
|
||||
self.quant_conv = nn.Conv3d(2 * latent_channels,
|
||||
2 * latent_channels,
|
||||
self.quant_conv = nn.Conv3d(2 * config.latent_channels,
|
||||
2 * config.latent_channels,
|
||||
kernel_size=1)
|
||||
|
||||
if load_decoder:
|
||||
if config.load_decoder:
|
||||
self.decoder = HunyuanVideoDecoder3D(
|
||||
in_channels=latent_channels,
|
||||
out_channels=out_channels,
|
||||
up_block_types=up_block_types,
|
||||
block_out_channels=block_out_channels,
|
||||
layers_per_block=layers_per_block,
|
||||
norm_num_groups=norm_num_groups,
|
||||
act_fn=act_fn,
|
||||
time_compression_ratio=temporal_compression_ratio,
|
||||
spatial_compression_ratio=spatial_compression_ratio,
|
||||
mid_block_add_attention=mid_block_add_attention,
|
||||
in_channels=config.latent_channels,
|
||||
out_channels=config.out_channels,
|
||||
up_block_types=config.up_block_types,
|
||||
block_out_channels=config.block_out_channels,
|
||||
layers_per_block=config.layers_per_block,
|
||||
norm_num_groups=config.norm_num_groups,
|
||||
act_fn=config.act_fn,
|
||||
time_compression_ratio=config.temporal_compression_ratio,
|
||||
spatial_compression_ratio=config.spatial_compression_ratio,
|
||||
mid_block_add_attention=config.mid_block_add_attention,
|
||||
)
|
||||
self.post_quant_conv = nn.Conv3d(latent_channels,
|
||||
latent_channels,
|
||||
self.post_quant_conv = nn.Conv3d(config.latent_channels,
|
||||
config.latent_channels,
|
||||
kernel_size=1)
|
||||
|
||||
# When decoding spatially large video latents, the memory requirement is very high. By breaking the video latent
|
||||
# frames spatially into smaller tiles and performing multiple forward passes for decoding, and then blending the
|
||||
# intermediate tiles together, the memory requirement can be lowered.
|
||||
self.use_tiling = True
|
||||
self.use_temporal_tiling = True
|
||||
self.use_parallel_tiling = True
|
||||
self.scaling_factor = scaling_factor
|
||||
|
||||
# The minimal tile height and width for spatial tiling to be used
|
||||
self.tile_sample_min_height = 256
|
||||
self.tile_sample_min_width = 256
|
||||
self.tile_sample_min_num_frames = 16
|
||||
|
||||
# The minimal distance between two spatial tiles
|
||||
self.tile_sample_stride_height = 192
|
||||
self.tile_sample_stride_width = 192
|
||||
self.tile_sample_stride_num_frames = 12
|
||||
ParallelTiledVAE.__init__(self)
|
||||
|
||||
def _encode(self, x: torch.Tensor) -> torch.Tensor:
|
||||
x = self.encoder(x)
|
||||
enc = self.quant_conv(x)
|
||||
|
||||
@@ -22,8 +22,8 @@ import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from fastvideo.v1.configs.models.vaes import WanVAEConfig
|
||||
from fastvideo.v1.layers.activation import get_act_fn
|
||||
from fastvideo.v1.models.utils import auto_attributes
|
||||
from fastvideo.v1.models.vaes.common import (DiagonalGaussianDistribution,
|
||||
ParallelTiledVAE)
|
||||
|
||||
@@ -781,96 +781,36 @@ class AutoencoderKLWan(nn.Module, ParallelTiledVAE):
|
||||
|
||||
_supports_gradient_checkpointing = False
|
||||
|
||||
@auto_attributes
|
||||
def __init__(self,
|
||||
base_dim: int = 96,
|
||||
z_dim: int = 16,
|
||||
dim_mult: Tuple[int, ...] = (1, 2, 4, 4),
|
||||
num_res_blocks: int = 2,
|
||||
attn_scales: Tuple[float, ...] = (),
|
||||
temperal_downsample: Tuple[bool, ...] = (False, True, True),
|
||||
dropout: float = 0.0,
|
||||
latents_mean: Tuple[float, ...] = (
|
||||
-0.7571,
|
||||
-0.7089,
|
||||
-0.9113,
|
||||
0.1075,
|
||||
-0.1745,
|
||||
0.9653,
|
||||
-0.1517,
|
||||
1.5508,
|
||||
0.4134,
|
||||
-0.0715,
|
||||
0.5517,
|
||||
-0.3632,
|
||||
-0.1922,
|
||||
-0.9497,
|
||||
0.2503,
|
||||
-0.2921,
|
||||
),
|
||||
latents_std: Tuple[float, ...] = (
|
||||
2.8184,
|
||||
1.4541,
|
||||
2.3275,
|
||||
2.6558,
|
||||
1.2196,
|
||||
1.7708,
|
||||
2.6052,
|
||||
2.0743,
|
||||
3.2687,
|
||||
2.1526,
|
||||
2.8652,
|
||||
1.5579,
|
||||
1.6382,
|
||||
1.1253,
|
||||
2.8251,
|
||||
1.9160,
|
||||
),
|
||||
load_encoder: bool = True,
|
||||
load_decoder: bool = True) -> None:
|
||||
super().__init__()
|
||||
def __init__(
|
||||
self,
|
||||
config: WanVAEConfig,
|
||||
) -> None:
|
||||
nn.Module.__init__(self)
|
||||
ParallelTiledVAE.__init__(self, config)
|
||||
|
||||
self.z_dim = z_dim
|
||||
self.temperal_downsample = list(temperal_downsample)
|
||||
self.temperal_upsample = list(temperal_downsample)[::-1]
|
||||
self.latents_mean = list(latents_mean)
|
||||
self.latents_std = list(latents_std)
|
||||
self.scaling_factor = 1.0 / torch.tensor(self.config.latents_std).view(
|
||||
1, self.config.z_dim, 1, 1, 1)
|
||||
self.shift_factor = torch.tensor(self.config.latents_mean).view(
|
||||
1, self.config.z_dim, 1, 1, 1)
|
||||
self.z_dim = config.z_dim
|
||||
self.temperal_downsample = list(config.temperal_downsample)
|
||||
self.temperal_upsample = list(config.temperal_downsample)[::-1]
|
||||
self.latents_mean = list(config.latents_mean)
|
||||
self.latents_std = list(config.latents_std)
|
||||
self.shift_factor = config.shift_factor
|
||||
|
||||
if load_encoder:
|
||||
self.encoder = WanEncoder3d(base_dim, z_dim * 2, dim_mult,
|
||||
num_res_blocks, attn_scales,
|
||||
self.temperal_downsample, dropout)
|
||||
self.quant_conv = WanCausalConv3d(z_dim * 2, z_dim * 2, 1)
|
||||
self.post_quant_conv = WanCausalConv3d(z_dim, z_dim, 1)
|
||||
if config.load_encoder:
|
||||
self.encoder = WanEncoder3d(config.base_dim, self.z_dim * 2,
|
||||
config.dim_mult, config.num_res_blocks,
|
||||
config.attn_scales,
|
||||
self.temperal_downsample,
|
||||
config.dropout)
|
||||
self.quant_conv = WanCausalConv3d(self.z_dim * 2, self.z_dim * 2, 1)
|
||||
self.post_quant_conv = WanCausalConv3d(self.z_dim, self.z_dim, 1)
|
||||
|
||||
if load_decoder:
|
||||
self.decoder = WanDecoder3d(base_dim, z_dim, dim_mult,
|
||||
num_res_blocks, attn_scales,
|
||||
self.temperal_upsample, dropout)
|
||||
if config.load_decoder:
|
||||
self.decoder = WanDecoder3d(config.base_dim, self.z_dim,
|
||||
config.dim_mult, config.num_res_blocks,
|
||||
config.attn_scales,
|
||||
self.temperal_upsample, config.dropout)
|
||||
|
||||
self.use_tiling = True
|
||||
self.use_temporal_tiling = False
|
||||
self.use_parallel_tiling = False
|
||||
self.spatial_compression_ratio = 8
|
||||
self.temporal_compression_ratio = 4
|
||||
|
||||
# The minimal tile height and width for spatial tiling to be used
|
||||
self.tile_sample_min_height = 256
|
||||
self.tile_sample_min_width = 256
|
||||
self.tile_sample_min_num_frames = 16
|
||||
|
||||
# The minimal distance between two spatial tiles
|
||||
self.tile_sample_stride_height = 192
|
||||
self.tile_sample_stride_width = 192
|
||||
self.tile_sample_stride_num_frames = 12
|
||||
|
||||
# Whether to use the feature cache algorithm used by diffusers and Wan2.1
|
||||
self.use_feature_cache = True # default to True for best performance
|
||||
ParallelTiledVAE.__init__(self)
|
||||
self.use_feature_cache = config.use_feature_cache
|
||||
|
||||
def clear_cache(self) -> None:
|
||||
|
||||
@@ -881,13 +821,15 @@ class AutoencoderKLWan(nn.Module, ParallelTiledVAE):
|
||||
count += 1
|
||||
return count
|
||||
|
||||
self._conv_num = _count_conv3d(self.decoder)
|
||||
self._conv_idx = 0
|
||||
self._feat_map = [None] * self._conv_num
|
||||
if self.config.load_decoder:
|
||||
self._conv_num = _count_conv3d(self.decoder)
|
||||
self._conv_idx = 0
|
||||
self._feat_map = [None] * self._conv_num
|
||||
# cache encode
|
||||
self._enc_conv_num = _count_conv3d(self.encoder)
|
||||
self._enc_conv_idx = 0
|
||||
self._enc_feat_map = [None] * self._enc_conv_num
|
||||
if self.config.load_encoder:
|
||||
self._enc_conv_num = _count_conv3d(self.encoder)
|
||||
self._enc_conv_idx = 0
|
||||
self._enc_feat_map = [None] * self._enc_conv_num
|
||||
|
||||
def encode(self, x: torch.Tensor) -> torch.Tensor:
|
||||
if self.use_feature_cache:
|
||||
|
||||
@@ -1,136 +1,3 @@
|
||||
# Adding a New Custom Pipeline
|
||||
|
||||
This guide explains how to add a new custom pipeline to the FastVideo framework. The pipeline system is designed to be modular and extensible, allowing you to implement custom video generation pipelines while reusing common components.
|
||||
|
||||
## Directory Structure
|
||||
|
||||
Create a new directory for your pipeline under `fastvideo/v1/pipelines/`:
|
||||
|
||||
```
|
||||
fastvideo/v1/pipelines/
|
||||
├── your_pipeline/
|
||||
│ ├── __init__.py
|
||||
│ └── your_pipeline.py
|
||||
```
|
||||
|
||||
## Implementation Steps
|
||||
|
||||
1. **Create Pipeline Class**
|
||||
- Your pipeline class should inherit from `ComposedPipelineBase`
|
||||
- Implement required methods and define pipeline stages
|
||||
|
||||
2. **Define EntryClass**
|
||||
- At the end of your pipeline file, define `EntryClass` to expose your pipeline
|
||||
- This is how the pipeline registry detects and loads your implementation
|
||||
|
||||
3. **Required Methods**
|
||||
- `required_config_modules()`: List required model components
|
||||
- `create_pipeline_stages()`: Define and configure pipeline stages
|
||||
- `initialize_pipeline()`: Set up any pipeline-specific initialization
|
||||
- `forward()`: Implement the main pipeline execution flow
|
||||
|
||||
## Example Implementation
|
||||
|
||||
Here's a basic template for implementing a new pipeline:
|
||||
|
||||
```python
|
||||
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
from fastvideo.v1.pipelines.stages import (
|
||||
InputValidationStage,
|
||||
ConditioningStage,
|
||||
# Import other required stages
|
||||
)
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
|
||||
class YourCustomPipeline(ComposedPipelineBase):
|
||||
def required_config_modules(self):
|
||||
return [
|
||||
"text_encoder",
|
||||
"vae",
|
||||
"transformer",
|
||||
"scheduler"
|
||||
# Add other required modules
|
||||
]
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
# Add and configure pipeline stages
|
||||
self.add_stage(
|
||||
stage_name="input_validation_stage",
|
||||
stage=InputValidationStage()
|
||||
)
|
||||
# Add more stages as needed
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
# Initialize pipeline-specific components
|
||||
pass
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
# Implement your pipeline's forward pass
|
||||
batch = self.input_validation_stage(batch, fastvideo_args)
|
||||
# Add more stage executions
|
||||
return batch
|
||||
|
||||
# This is required for pipeline registry detection
|
||||
EntryClass = YourCustomPipeline
|
||||
```
|
||||
|
||||
## Pipeline Registry
|
||||
|
||||
The pipeline registry automatically detects and loads your pipeline through the following mechanism:
|
||||
|
||||
1. It scans all packages under `fastvideo/v1/pipelines/`
|
||||
2. For each package, it looks for an `EntryClass` variable
|
||||
3. The `EntryClass` can be either:
|
||||
- A single pipeline class
|
||||
- A list of pipeline classes (for multiple implementations in one module)
|
||||
4. The registry uses the class name as the pipeline architecture identifier
|
||||
|
||||
## Available Stages
|
||||
|
||||
You can use these pre-built stages in your pipeline:
|
||||
|
||||
- `InputValidationStage`: Validates input parameters
|
||||
- `TimestepPreparationStage`: Prepares timesteps for diffusion
|
||||
- `LatentPreparationStage`: Prepares latent space
|
||||
- `ConditioningStage`: Handles conditioning inputs
|
||||
- `DenoisingStage`: Performs the denoising process
|
||||
- `DecodingStage`: Decodes the final output
|
||||
- `LlamaEncodingStage`: Text encoding with LLaMA
|
||||
- `CLIPTextEncodingStage`: Text encoding with CLIP
|
||||
|
||||
## Best Practices
|
||||
|
||||
1. **Stage Organization**
|
||||
- Organize stages in a logical order
|
||||
- Use clear, descriptive stage names
|
||||
- Document any custom stage logic
|
||||
|
||||
2. **Error Handling**
|
||||
- Implement proper error handling in each stage
|
||||
- Use the logger for debugging and monitoring
|
||||
|
||||
3. **Configuration**
|
||||
- Clearly specify required modules in `required_config_modules()`
|
||||
- Document any pipeline-specific configuration parameters
|
||||
|
||||
4. **Testing**
|
||||
- Add unit tests for your pipeline
|
||||
- Test with different input configurations
|
||||
- Verify pipeline outputs
|
||||
|
||||
## Example Usage
|
||||
|
||||
After implementing your pipeline, you can use it like this:
|
||||
|
||||
```python
|
||||
from fastvideo.v1.pipelines import PipelineRegistry
|
||||
|
||||
# Get your pipeline class
|
||||
pipeline_cls, _ = PipelineRegistry.resolve_pipeline_cls("YourCustomPipeline")
|
||||
|
||||
# Initialize and use the pipeline
|
||||
pipeline = pipeline_cls(...)
|
||||
result = pipeline(...)
|
||||
```
|
||||
Please see
|
||||
|
||||
@@ -114,12 +114,11 @@ class ComposedPipelineBase(ABC):
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
"""
|
||||
Initialize the pipeline.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
return
|
||||
|
||||
def load_modules(self, fastvideo_args: FastVideoArgs) -> Dict[str, Any]:
|
||||
"""
|
||||
|
||||
@@ -6,8 +6,6 @@ This module contains an implementation of the Hunyuan video diffusion pipeline
|
||||
using the modular pipeline architecture.
|
||||
"""
|
||||
|
||||
from diffusers.image_processor import VaeImageProcessor
|
||||
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
@@ -57,7 +55,8 @@ class HunyuanVideoPipeline(ComposedPipelineBase):
|
||||
|
||||
self.add_stage(stage_name="latent_preparation_stage",
|
||||
stage=LatentPreparationStage(
|
||||
scheduler=self.get_module("scheduler")))
|
||||
scheduler=self.get_module("scheduler"),
|
||||
transformer=self.get_module("transformer")))
|
||||
|
||||
self.add_stage(stage_name="denoising_stage",
|
||||
stage=DenoisingStage(
|
||||
@@ -67,20 +66,5 @@ class HunyuanVideoPipeline(ComposedPipelineBase):
|
||||
self.add_stage(stage_name="decoding_stage",
|
||||
stage=DecodingStage(vae=self.get_module("vae")))
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
"""
|
||||
Initialize the pipeline.
|
||||
"""
|
||||
vae_scale_factor = 2**(len(self.get_module("vae").block_out_channels) -
|
||||
1)
|
||||
fastvideo_args.vae_scale_factor = vae_scale_factor
|
||||
|
||||
self.image_processor = VaeImageProcessor(
|
||||
vae_scale_factor=vae_scale_factor)
|
||||
self.add_module("image_processor", self.image_processor)
|
||||
|
||||
num_channels_latents = self.get_module("transformer").in_channels
|
||||
fastvideo_args.num_channels_latents = num_channels_latents
|
||||
|
||||
|
||||
EntryClass = HunyuanVideoPipeline
|
||||
|
||||
@@ -36,6 +36,8 @@ class ForwardBatch:
|
||||
# Text inputs
|
||||
prompt: Optional[Union[str, List[str]]] = None
|
||||
negative_prompt: Optional[Union[str, List[str]]] = None
|
||||
prompt_path: Optional[str] = None
|
||||
output_path: str = "outputs/"
|
||||
|
||||
# Primary encoder embeddings
|
||||
prompt_embeds: List[torch.Tensor] = field(default_factory=list)
|
||||
@@ -49,6 +51,7 @@ class ForwardBatch:
|
||||
# Batch info
|
||||
batch_size: Optional[int] = None
|
||||
num_videos_per_prompt: int = 1
|
||||
seed: Optional[int] = None
|
||||
seeds: Optional[List[int]] = None
|
||||
|
||||
# Tracking if embeddings are already processed
|
||||
@@ -60,7 +63,6 @@ class ForwardBatch:
|
||||
image_latent: Optional[torch.Tensor] = None
|
||||
|
||||
# Latent dimensions
|
||||
num_channels_latents: Optional[int] = None
|
||||
height_latents: Optional[int] = None
|
||||
width_latents: Optional[int] = None
|
||||
num_frames: int = 1 # Default for image models
|
||||
@@ -68,6 +70,7 @@ class ForwardBatch:
|
||||
# Original dimensions (before VAE scaling)
|
||||
height: Optional[int] = None
|
||||
width: Optional[int] = None
|
||||
fps: Optional[int] = None
|
||||
|
||||
# Timesteps
|
||||
timesteps: Optional[torch.Tensor] = None
|
||||
@@ -95,7 +98,9 @@ class ForwardBatch:
|
||||
# Extra parameters that might be needed by specific pipeline implementations
|
||||
extra: Dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
device: torch.device = field(default_factory=lambda: torch.device("cuda"))
|
||||
# Misc
|
||||
save_video: bool = True
|
||||
return_frames: bool = False
|
||||
|
||||
def __post_init__(self):
|
||||
"""Initialize dependent fields after dataclass initialization."""
|
||||
|
||||
@@ -53,12 +53,12 @@ class CLIPImageEncodingStage(PipelineStage):
|
||||
The batch with encoded prompt embeddings.
|
||||
"""
|
||||
if fastvideo_args.use_cpu_offload:
|
||||
self.image_encoder = self.image_encoder.to(batch.device)
|
||||
self.image_encoder = self.image_encoder.to(fastvideo_args.device)
|
||||
|
||||
image = load_image(batch.image_path)
|
||||
|
||||
image_inputs = self.image_processor(
|
||||
images=image, return_tensors="pt").to(batch.device)
|
||||
images=image, return_tensors="pt").to(fastvideo_args.device)
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
image_embeds = self.image_encoder(**image_inputs)
|
||||
|
||||
|
||||
@@ -52,7 +52,7 @@ class CLIPTextEncodingStage(PipelineStage):
|
||||
The batch with encoded prompt embeddings.
|
||||
"""
|
||||
if fastvideo_args.use_cpu_offload:
|
||||
self.text_encoder = self.text_encoder.to(batch.device)
|
||||
self.text_encoder = self.text_encoder.to(fastvideo_args.device)
|
||||
|
||||
text_inputs = self.tokenizer(
|
||||
batch.prompt,
|
||||
@@ -63,7 +63,7 @@ class CLIPTextEncodingStage(PipelineStage):
|
||||
)
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
outputs = self.text_encoder(input_ids=text_inputs["input_ids"].to(
|
||||
batch.device), )
|
||||
fastvideo_args.device), )
|
||||
prompt_embeds = outputs["pooler_output"]
|
||||
|
||||
batch.prompt_embeds.append(prompt_embeds)
|
||||
@@ -79,7 +79,7 @@ class CLIPTextEncodingStage(PipelineStage):
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
negative_outputs = self.text_encoder(
|
||||
input_ids=negative_text_inputs["input_ids"].to(
|
||||
batch.device), )
|
||||
fastvideo_args.device), )
|
||||
negative_prompt_embeds = negative_outputs["pooler_output"]
|
||||
|
||||
assert batch.negative_prompt_embeds is not None
|
||||
|
||||
@@ -7,6 +7,7 @@ import torch
|
||||
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.vaes.common import ParallelTiledVAE
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.v1.utils import PRECISION_TO_TYPE
|
||||
@@ -22,8 +23,8 @@ class DecodingStage(PipelineStage):
|
||||
output format (e.g., pixel values).
|
||||
"""
|
||||
|
||||
def __init__(self, vae) -> None:
|
||||
self.vae = vae
|
||||
def __init__(self, vae: ParallelTiledVAE) -> None:
|
||||
self.vae: ParallelTiledVAE = vae
|
||||
|
||||
def forward(
|
||||
self,
|
||||
|
||||
@@ -66,7 +66,7 @@ class DenoisingStage(PipelineStage):
|
||||
"""
|
||||
# If use cpu offload, need to load the model back into gpu again
|
||||
if fastvideo_args.use_cpu_offload:
|
||||
self.transformer = self.transformer.to(batch.device)
|
||||
self.transformer = self.transformer.to(fastvideo_args.device)
|
||||
# Prepare extra step kwargs for scheduler
|
||||
extra_step_kwargs = self.prepare_extra_func_kwargs(
|
||||
self.scheduler.step,
|
||||
@@ -165,7 +165,7 @@ class DenoisingStage(PipelineStage):
|
||||
[fastvideo_args.embedded_cfg_scale] *
|
||||
latent_model_input.shape[0],
|
||||
dtype=torch.float32,
|
||||
device=batch.device,
|
||||
device=fastvideo_args.device,
|
||||
).to(target_dtype) * 1000.0 if fastvideo_args.embedded_cfg_scale
|
||||
is not None else None)
|
||||
|
||||
|
||||
@@ -9,6 +9,7 @@ import torch
|
||||
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.vaes.common import ParallelTiledVAE
|
||||
from fastvideo.v1.models.vision_utils import (get_default_height_width,
|
||||
load_image, normalize,
|
||||
numpy_to_pt, pil_to_numpy, resize)
|
||||
@@ -27,8 +28,8 @@ class EncodingStage(PipelineStage):
|
||||
input format (e.g., latents).
|
||||
"""
|
||||
|
||||
def __init__(self, vae) -> None:
|
||||
self.vae = vae
|
||||
def __init__(self, vae: ParallelTiledVAE) -> None:
|
||||
self.vae: ParallelTiledVAE = vae
|
||||
|
||||
def forward(
|
||||
self,
|
||||
@@ -49,6 +50,8 @@ class EncodingStage(PipelineStage):
|
||||
# TODO(will): remove this once we add input/output validation for stages
|
||||
if image_path is None:
|
||||
raise ValueError("Image Path must be provided")
|
||||
assert batch.height is not None
|
||||
assert batch.width is not None
|
||||
latent_height = batch.height // self.vae.spatial_compression_ratio
|
||||
latent_width = batch.width // self.vae.spatial_compression_ratio
|
||||
|
||||
@@ -57,16 +60,15 @@ class EncodingStage(PipelineStage):
|
||||
image,
|
||||
vae_scale_factor=self.vae.spatial_compression_ratio,
|
||||
height=batch.height,
|
||||
width=batch.width).to(batch.device, dtype=torch.float32)
|
||||
width=batch.width).to(fastvideo_args.device, dtype=torch.float32)
|
||||
image = image.unsqueeze(2)
|
||||
video_condition = torch.cat([
|
||||
image,
|
||||
image.new_zeros(image.shape[0], image.shape[1],
|
||||
fastvideo_args.num_frames - 1, batch.height,
|
||||
batch.width)
|
||||
batch.num_frames - 1, batch.height, batch.width)
|
||||
],
|
||||
dim=2)
|
||||
video_condition = video_condition.to(device=batch.device,
|
||||
video_condition = video_condition.to(device=fastvideo_args.device,
|
||||
dtype=torch.float32)
|
||||
|
||||
# Setup VAE precision
|
||||
@@ -106,9 +108,9 @@ class EncodingStage(PipelineStage):
|
||||
else:
|
||||
latent_condition = latent_condition * self.vae.scaling_factor
|
||||
|
||||
mask_lat_size = torch.ones(1, 1, fastvideo_args.num_frames,
|
||||
latent_height, latent_width)
|
||||
mask_lat_size[:, :, list(range(1, fastvideo_args.num_frames))] = 0
|
||||
mask_lat_size = torch.ones(1, 1, batch.num_frames, latent_height,
|
||||
latent_width)
|
||||
mask_lat_size[:, :, list(range(1, batch.num_frames))] = 0
|
||||
first_frame_mask = mask_lat_size[:, :, 0:1]
|
||||
first_frame_mask = torch.repeat_interleave(
|
||||
first_frame_mask,
|
||||
|
||||
@@ -24,9 +24,10 @@ class InputValidationStage(PipelineStage):
|
||||
def _generate_seeds(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs):
|
||||
"""Generate seeds for the inference"""
|
||||
seed = fastvideo_args.seed
|
||||
num_videos_per_prompt = fastvideo_args.num_videos
|
||||
seed = batch.seed
|
||||
num_videos_per_prompt = batch.num_videos_per_prompt
|
||||
|
||||
assert seed is not None
|
||||
seeds = [seed + i for i in range(num_videos_per_prompt)]
|
||||
batch.seeds = seeds
|
||||
# Peiyuan: using GPU seed will cause A100 and H100 to generate different results...
|
||||
@@ -85,12 +86,4 @@ class InputValidationStage(PipelineStage):
|
||||
f"Guidance scale must be positive, but got {batch.guidance_scale}"
|
||||
)
|
||||
|
||||
# Set device if not already set
|
||||
if batch.device is None:
|
||||
batch.device = self.device
|
||||
|
||||
# Set data type if not already set
|
||||
if batch.data_type is None:
|
||||
batch.data_type = fastvideo_args.precision
|
||||
|
||||
return batch
|
||||
|
||||
@@ -6,7 +6,6 @@ from diffusers.utils.torch_utils import randn_tensor
|
||||
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.vaes.common import ParallelTiledVAE
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
|
||||
@@ -21,10 +20,10 @@ class LatentPreparationStage(PipelineStage):
|
||||
denoised during the diffusion process.
|
||||
"""
|
||||
|
||||
def __init__(self, scheduler, vae=None) -> None:
|
||||
def __init__(self, scheduler, transformer) -> None:
|
||||
super().__init__()
|
||||
self.scheduler = scheduler
|
||||
self.vae = vae
|
||||
self.transformer = transformer
|
||||
|
||||
def forward(
|
||||
self,
|
||||
@@ -42,9 +41,10 @@ class LatentPreparationStage(PipelineStage):
|
||||
The batch with prepared latent variables.
|
||||
"""
|
||||
|
||||
latent_num_frames = None
|
||||
# Adjust video length based on VAE version if needed
|
||||
if hasattr(self, 'adjust_video_length'):
|
||||
batch = self.adjust_video_length(self.vae, batch, fastvideo_args)
|
||||
latent_num_frames = self.adjust_video_length(batch, fastvideo_args)
|
||||
# Determine batch size
|
||||
if isinstance(batch.prompt, list):
|
||||
batch_size = len(batch.prompt)
|
||||
@@ -58,10 +58,10 @@ class LatentPreparationStage(PipelineStage):
|
||||
|
||||
# Get required parameters
|
||||
dtype = batch.prompt_embeds[0].dtype
|
||||
device = batch.device
|
||||
device = fastvideo_args.device
|
||||
generator = batch.generator
|
||||
latents = batch.latents
|
||||
num_frames = batch.num_frames
|
||||
num_frames = latent_num_frames if latent_num_frames is not None else batch.num_frames
|
||||
height = batch.height
|
||||
width = batch.width
|
||||
|
||||
@@ -69,16 +69,15 @@ class LatentPreparationStage(PipelineStage):
|
||||
if height is None or width is None:
|
||||
raise ValueError("Height and width must be provided")
|
||||
|
||||
assert fastvideo_args.num_channels_latents is not None
|
||||
assert fastvideo_args.vae_scale_factor is not None
|
||||
|
||||
# Calculate latent shape
|
||||
shape = (
|
||||
batch_size,
|
||||
fastvideo_args.num_channels_latents,
|
||||
self.transformer.num_channels_latents,
|
||||
num_frames,
|
||||
height // fastvideo_args.vae_scale_factor,
|
||||
width // fastvideo_args.vae_scale_factor,
|
||||
height //
|
||||
fastvideo_args.vae_config.arch_config.spatial_compression_ratio,
|
||||
width //
|
||||
fastvideo_args.vae_config.arch_config.spatial_compression_ratio,
|
||||
)
|
||||
|
||||
# Validate generator if it's a list
|
||||
@@ -106,8 +105,8 @@ class LatentPreparationStage(PipelineStage):
|
||||
|
||||
return batch
|
||||
|
||||
def adjust_video_length(self, vae: ParallelTiledVAE, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
def adjust_video_length(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> int:
|
||||
"""
|
||||
Adjust video length based on VAE version.
|
||||
|
||||
@@ -119,7 +118,7 @@ class LatentPreparationStage(PipelineStage):
|
||||
The batch with adjusted video length.
|
||||
"""
|
||||
video_length = batch.num_frames
|
||||
temporal_scale_factor = vae.temporal_compression_ratio if vae is not None else 4
|
||||
temporal_scale_factor = fastvideo_args.vae_config.arch_config.temporal_compression_ratio
|
||||
# TODO
|
||||
batch.num_frames = (video_length - 1) // temporal_scale_factor + 1
|
||||
return batch
|
||||
latent_num_frames = (video_length - 1) // temporal_scale_factor + 1
|
||||
return latent_num_frames
|
||||
|
||||
@@ -74,7 +74,7 @@ class LlamaEncodingStage(PipelineStage):
|
||||
The batch with encoded prompt embeddings.
|
||||
"""
|
||||
if fastvideo_args.use_cpu_offload:
|
||||
self.text_encoder = self.text_encoder.to(batch.device)
|
||||
self.text_encoder = self.text_encoder.to(fastvideo_args.device)
|
||||
|
||||
text = prompt_template_video["template"].format(batch.prompt)
|
||||
text_inputs = self.tokenizer(
|
||||
@@ -87,7 +87,7 @@ class LlamaEncodingStage(PipelineStage):
|
||||
hidden_state_skip_layer = 2
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
outputs = self.text_encoder(
|
||||
input_ids=text_inputs["input_ids"].to(batch.device),
|
||||
input_ids=text_inputs["input_ids"].to(fastvideo_args.device),
|
||||
output_hidden_states=hidden_state_skip_layer is not None,
|
||||
)
|
||||
|
||||
@@ -111,7 +111,7 @@ class LlamaEncodingStage(PipelineStage):
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
negative_outputs = self.text_encoder(
|
||||
input_ids=negative_text_inputs["input_ids"].to(
|
||||
batch.device),
|
||||
fastvideo_args.device),
|
||||
output_hidden_states=hidden_state_skip_layer is not None,
|
||||
)
|
||||
|
||||
|
||||
@@ -51,7 +51,7 @@ class T5EncodingStage(PipelineStage):
|
||||
The batch with encoded prompt embeddings.
|
||||
"""
|
||||
if fastvideo_args.use_cpu_offload:
|
||||
self.text_encoder = self.text_encoder.to(batch.device)
|
||||
self.text_encoder = self.text_encoder.to(fastvideo_args.device)
|
||||
|
||||
text = batch.prompt
|
||||
text_inputs = self.tokenizer(
|
||||
@@ -62,7 +62,7 @@ class T5EncodingStage(PipelineStage):
|
||||
add_special_tokens=True,
|
||||
return_attention_mask=True,
|
||||
return_tensors="pt",
|
||||
).to(batch.device)
|
||||
).to(fastvideo_args.device)
|
||||
text_input_ids, mask = text_inputs.input_ids, text_inputs.attention_mask
|
||||
seq_lens = mask.gt(0).sum(dim=1).long()
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
@@ -89,7 +89,7 @@ class T5EncodingStage(PipelineStage):
|
||||
add_special_tokens=True,
|
||||
return_attention_mask=True,
|
||||
return_tensors="pt",
|
||||
).to(batch.device)
|
||||
).to(fastvideo_args.device)
|
||||
text_input_ids, mask = negative_text_inputs.input_ids, negative_text_inputs.attention_mask
|
||||
seq_lens = mask.gt(0).sum(dim=1).long()
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
|
||||
@@ -42,7 +42,7 @@ class TimestepPreparationStage(PipelineStage):
|
||||
The batch with prepared timesteps.
|
||||
"""
|
||||
scheduler = self.scheduler
|
||||
device = batch.device
|
||||
device = fastvideo_args.device
|
||||
num_inference_steps = batch.num_inference_steps
|
||||
timesteps = batch.timesteps
|
||||
sigmas = batch.sigmas
|
||||
|
||||
@@ -55,7 +55,7 @@ class WanImageToVideoPipeline(ComposedPipelineBase):
|
||||
self.add_stage(stage_name="latent_preparation_stage",
|
||||
stage=LatentPreparationStage(
|
||||
scheduler=self.get_module("scheduler"),
|
||||
vae=self.get_module("vae")))
|
||||
transformer=self.get_module("transformer")))
|
||||
|
||||
self.add_stage(stage_name="image_latent_preparation_stage",
|
||||
stage=EncodingStage(vae=self.get_module("vae")))
|
||||
@@ -68,15 +68,5 @@ class WanImageToVideoPipeline(ComposedPipelineBase):
|
||||
self.add_stage(stage_name="decoding_stage",
|
||||
stage=DecodingStage(vae=self.get_module("vae")))
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
"""
|
||||
Initialize the pipeline.
|
||||
"""
|
||||
vae_scale_factor = self.get_module("vae").spatial_compression_ratio
|
||||
fastvideo_args.vae_scale_factor = vae_scale_factor
|
||||
|
||||
num_channels_latents = self.get_module("transformer").out_channels
|
||||
fastvideo_args.num_channels_latents = num_channels_latents
|
||||
|
||||
|
||||
EntryClass = WanImageToVideoPipeline
|
||||
|
||||
@@ -48,7 +48,7 @@ class WanPipeline(ComposedPipelineBase):
|
||||
self.add_stage(stage_name="latent_preparation_stage",
|
||||
stage=LatentPreparationStage(
|
||||
scheduler=self.get_module("scheduler"),
|
||||
vae=self.get_module("vae")))
|
||||
transformer=self.get_module("transformer")))
|
||||
|
||||
self.add_stage(stage_name="denoising_stage",
|
||||
stage=DenoisingStage(
|
||||
@@ -58,15 +58,5 @@ class WanPipeline(ComposedPipelineBase):
|
||||
self.add_stage(stage_name="decoding_stage",
|
||||
stage=DecodingStage(vae=self.get_module("vae")))
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
"""
|
||||
Initialize the pipeline.
|
||||
"""
|
||||
vae_scale_factor = self.get_module("vae").spatial_compression_ratio
|
||||
fastvideo_args.vae_scale_factor = vae_scale_factor
|
||||
|
||||
num_channels_latents = self.get_module("transformer").in_channels
|
||||
fastvideo_args.num_channels_latents = num_channels_latents
|
||||
|
||||
|
||||
EntryClass = WanPipeline
|
||||
|
||||
@@ -126,12 +126,25 @@ class CudaPlatformBase(Platform):
|
||||
logger.info("Using Sliding Tile Attention backend.")
|
||||
return "fastvideo.v1.attention.backends.sliding_tile_attn.SlidingTileAttentionBackend"
|
||||
except ImportError as e:
|
||||
# TODO(will): improve error message
|
||||
logger.info(e)
|
||||
logger.info(
|
||||
"Sliding Tile Attention backend is not installed. Fall back to Flash Attention."
|
||||
)
|
||||
elif selected_backend == _Backend.SAGE_ATTN:
|
||||
try:
|
||||
from sageattention import sageattn # noqa: F401
|
||||
|
||||
from fastvideo.v1.attention.backends.sage_attn import ( # noqa: F401
|
||||
SageAttentionBackend)
|
||||
logger.info("Using Sage Attention backend.")
|
||||
return "fastvideo.v1.attention.backends.sage_attn.SageAttentionBackend"
|
||||
except ImportError as e:
|
||||
logger.info(e)
|
||||
logger.info(
|
||||
"Sage Attention backend is not installed. Fall back to Flash Attention."
|
||||
)
|
||||
elif selected_backend == _Backend.TORCH_SDPA:
|
||||
logger.info("Using Torch SDPA backend.")
|
||||
return "fastvideo.v1.attention.backends.sdpa.SDPABackend"
|
||||
elif selected_backend == _Backend.FLASH_ATTN or selected_backend is None:
|
||||
pass
|
||||
@@ -174,9 +187,7 @@ class CudaPlatformBase(Platform):
|
||||
target_backend = _Backend.TORCH_SDPA
|
||||
|
||||
if target_backend == _Backend.TORCH_SDPA:
|
||||
logger.info(
|
||||
"Using torch.nn.functional.scaled_dot_product_attention backend."
|
||||
)
|
||||
logger.info("Using Torch SDPA backend.")
|
||||
return "fastvideo.v1.attention.backends.sdpa.SDPABackend"
|
||||
|
||||
logger.info("Using Flash Attention backend.")
|
||||
|
||||
@@ -17,7 +17,7 @@ class _Backend(enum.Enum):
|
||||
FLASH_ATTN = enum.auto()
|
||||
SLIDING_TILE_ATTN = enum.auto()
|
||||
TORCH_SDPA = enum.auto()
|
||||
# SAGE_ATTN = enum.auto()
|
||||
SAGE_ATTN = enum.auto()
|
||||
NO_ATTENTION = enum.auto()
|
||||
|
||||
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
# type: ignore
|
||||
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import os
|
||||
|
||||
@@ -14,6 +14,7 @@ from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.utils import maybe_download_model
|
||||
from fastvideo.v1.configs.models.encoders import CLIPTextConfig
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -39,7 +40,8 @@ def test_clip_encoder():
|
||||
- Produce nearly identical outputs for the same input prompts
|
||||
"""
|
||||
args = FastVideoArgs(model_path="openai/clip-vit-large-patch14",
|
||||
precision="float16")
|
||||
text_encoder_precision_2="fp16",
|
||||
text_encoder_config_2=CLIPTextConfig())
|
||||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
|
||||
logger.info("Loading models from %s", args.model_path)
|
||||
@@ -60,7 +62,7 @@ def test_clip_encoder():
|
||||
from fastvideo.v1.models.loader.component_loader import TextEncoderLoader
|
||||
loader = TextEncoderLoader()
|
||||
args.device_str = "cuda:0"
|
||||
model2 = loader.load_model(TEXT_ENCODER_PATH, hf_config, device)
|
||||
model2 = loader.load(TEXT_ENCODER_PATH, "", args)
|
||||
|
||||
# Load the HuggingFace implementation directly
|
||||
# model2 = CLIPTextModel(hf_config)
|
||||
|
||||
@@ -13,6 +13,7 @@ from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.loader.component_loader import TextEncoderLoader
|
||||
from fastvideo.v1.utils import maybe_download_model
|
||||
from fastvideo.v1.configs.models.encoders import LlamaConfig
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -39,7 +40,8 @@ def test_llama_encoder():
|
||||
- Produce nearly identical outputs for the same input prompts
|
||||
"""
|
||||
args = FastVideoArgs(model_path="meta-llama/Llama-2-7b-hf",
|
||||
precision="float16")
|
||||
precision="float16",
|
||||
text_encoder_config=LlamaConfig())
|
||||
|
||||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
|
||||
@@ -57,7 +59,7 @@ def test_llama_encoder():
|
||||
loader = TextEncoderLoader()
|
||||
args.device_str = "cuda:0"
|
||||
device = torch.device(args.device_str)
|
||||
model2 = loader.load_model(TEXT_ENCODER_PATH, hf_config, device)
|
||||
model2 = loader.load(TEXT_ENCODER_PATH, "", args)
|
||||
|
||||
# Convert to float16 and move to device
|
||||
model2 = model2.to(torch.float16)
|
||||
|
||||
@@ -10,6 +10,8 @@ from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.loader.component_loader import TextEncoderLoader
|
||||
from fastvideo.v1.utils import maybe_download_model
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.configs.models.encoders import T5Config
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -36,8 +38,9 @@ def test_t5_encoder():
|
||||
precision).to(device).eval()
|
||||
tokenizer = AutoTokenizer.from_pretrained(TOKENIZER_PATH)
|
||||
|
||||
args = FastVideoArgs(model_path=TEXT_ENCODER_PATH, text_encoder_config=T5Config(), device_str="cuda")
|
||||
loader = TextEncoderLoader()
|
||||
model2 = loader.load_model(TEXT_ENCODER_PATH, hf_config, device)
|
||||
model2 = loader.load(TEXT_ENCODER_PATH, "", args)
|
||||
|
||||
# Convert to float16 and move to device
|
||||
model2 = model2.to(precision)
|
||||
|
||||
@@ -52,7 +52,7 @@ def initialize_identical_weights(model, seed=42):
|
||||
return model
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("distributed_setup")
|
||||
@pytest.mark.skip(reason="Incompatible with the new config")
|
||||
def test_hunyuanvideo_distributed():
|
||||
# Get tensor parallel info
|
||||
sp_rank = get_sequence_model_parallel_rank()
|
||||
|
||||
@@ -10,11 +10,11 @@ from fastvideo.v1.distributed.parallel_state import (
|
||||
get_sequence_model_parallel_rank,
|
||||
get_sequence_model_parallel_world_size)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.dits.hunyuanvideo import (
|
||||
HunyuanVideoTransformer3DModel as HunyuanVideoDit)
|
||||
from fastvideo.v1.models.loader.fsdp_load import load_fsdp_model
|
||||
from fastvideo.v1.models.loader.component_loader import TransformerLoader
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.utils import maybe_download_model
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.configs.models.dits import HunyuanVideoConfig
|
||||
|
||||
|
||||
logger = init_logger(__name__)
|
||||
@@ -59,13 +59,15 @@ def test_hunyuanvideo_distributed():
|
||||
config.pop("_class_name")
|
||||
config.pop("_diffusers_version")
|
||||
|
||||
weight_dir_list = glob.glob(os.path.join(TRANSFORMER_PATH, "*.safetensors"))
|
||||
weight_dir_list = [str(path) for path in weight_dir_list]
|
||||
model = load_fsdp_model(HunyuanVideoDit,
|
||||
init_params=config,
|
||||
weight_dir_list=weight_dir_list,
|
||||
device=torch.device(f"cuda:{LOCAL_RANK}"),
|
||||
cpu_offload=False)
|
||||
precision_str = "bf16"
|
||||
args = FastVideoArgs(model_path=TRANSFORMER_PATH,
|
||||
use_cpu_offload=False,
|
||||
precision=precision_str)
|
||||
args.device = torch.device(f"cuda:{LOCAL_RANK}")
|
||||
args.dit_config = HunyuanVideoConfig()
|
||||
|
||||
loader = TransformerLoader()
|
||||
model = loader.load(TRANSFORMER_PATH, "", args)
|
||||
|
||||
model.eval()
|
||||
|
||||
|
||||
@@ -11,6 +11,7 @@ from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.loader.component_loader import TransformerLoader
|
||||
from fastvideo.v1.utils import maybe_download_model
|
||||
from fastvideo.v1.configs.models.dits import WanVideoConfig
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -33,6 +34,7 @@ def test_wan_transformer():
|
||||
use_cpu_offload=False,
|
||||
precision=precision_str)
|
||||
args.device = device
|
||||
args.dit_config = WanVideoConfig()
|
||||
|
||||
loader = TransformerLoader()
|
||||
model2 = loader.load(TRANSFORMER_PATH, "", args).to(device, dtype=precision)
|
||||
@@ -113,6 +115,8 @@ def test_wan_transformer():
|
||||
# Check if outputs are similar (allowing for small numerical differences)
|
||||
max_diff = torch.max(torch.abs(output1 - output2))
|
||||
mean_diff = torch.mean(torch.abs(output1 - output2))
|
||||
logger.info("Max Diff: %s", max_diff.item())
|
||||
logger.info("Mean Diff: %s", mean_diff.item())
|
||||
assert max_diff < 1e-1, f"Maximum difference between outputs: {max_diff.item()}"
|
||||
# mean diff
|
||||
assert mean_diff < 1e-2, f"Mean difference between outputs: {mean_diff.item()}"
|
||||
@@ -8,8 +8,11 @@ import torch
|
||||
from safetensors.torch import load_file
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.vaes.hunyuanvae import (
|
||||
AutoencoderKLHunyuanVideo as MyHunyuanVAE)
|
||||
# from fastvideo.v1.models.vaes.hunyuanvae import (
|
||||
# AutoencoderKLHunyuanVideo as MyHunyuanVAE)
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.models.loader.component_loader import VAELoader
|
||||
from fastvideo.v1.configs.models.vaes import HunyuanVAEConfig
|
||||
from fastvideo.v1.utils import maybe_download_model
|
||||
|
||||
logger = init_logger(__name__)
|
||||
@@ -31,21 +34,14 @@ REFERENCE_LATENT = -105.51324462890625
|
||||
@pytest.mark.usefixtures("distributed_setup")
|
||||
def test_hunyuan_vae():
|
||||
device = torch.device("cuda:0")
|
||||
# Initialize the two model implementations
|
||||
config = json.load(open(CONFIG_PATH))
|
||||
config.pop("_class_name")
|
||||
config.pop("_diffusers_version")
|
||||
model = MyHunyuanVAE(**config).to(torch.bfloat16)
|
||||
precision = torch.bfloat16
|
||||
precision_str = "bf16"
|
||||
args = FastVideoArgs(model_path=VAE_PATH, vae_precision=precision_str)
|
||||
args.device = device
|
||||
args.vae_config = HunyuanVAEConfig()
|
||||
|
||||
loaded = load_file(os.path.join(VAE_PATH,
|
||||
"diffusion_pytorch_model.safetensors"))
|
||||
model.load_state_dict(loaded)
|
||||
|
||||
# Set model to eval mode
|
||||
model.eval()
|
||||
|
||||
# Move to GPU
|
||||
model = model.to(device)
|
||||
loader = VAELoader()
|
||||
model = loader.load(VAE_PATH, "", args)
|
||||
|
||||
model.enable_tiling(tile_sample_min_height=32,
|
||||
tile_sample_min_width=32,
|
||||
|
||||
@@ -9,6 +9,7 @@ from diffusers import AutoencoderKLWan
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.loader.component_loader import VAELoader
|
||||
from fastvideo.v1.configs.models.vaes import WanVAEConfig
|
||||
from fastvideo.v1.utils import maybe_download_model
|
||||
|
||||
logger = init_logger(__name__)
|
||||
@@ -30,6 +31,7 @@ def test_wan_vae():
|
||||
precision_str = "bf16"
|
||||
args = FastVideoArgs(model_path=VAE_PATH, vae_precision=precision_str)
|
||||
args.device = device
|
||||
args.vae_config = WanVAEConfig()
|
||||
|
||||
loader = VAELoader()
|
||||
model2 = loader.load(VAE_PATH, "", args)
|
||||
@@ -77,23 +79,19 @@ def test_wan_vae():
|
||||
# Test decoding
|
||||
logger.info("Testing decoding...")
|
||||
latent1_tensor = latent1.mode()
|
||||
latents_mean = (torch.tensor(model1.config.latents_mean).view(
|
||||
mean1 = (torch.tensor(model1.config.latents_mean).view(
|
||||
1, model1.config.z_dim, 1, 1, 1).to(input_tensor.device,
|
||||
input_tensor.dtype))
|
||||
latents_std = 1.0 / torch.tensor(model1.config.latents_std).view(
|
||||
1, model1.config.z_dim, 1, 1, 1).to(input_tensor.device,
|
||||
std1 = (1.0 / torch.tensor(model1.config.latents_std).view(
|
||||
1, model1.config.z_dim, 1, 1, 1)).to(input_tensor.device,
|
||||
input_tensor.dtype)
|
||||
latent1_tensor = latent1_tensor / latents_std + latents_mean
|
||||
latent1_tensor = latent1_tensor / std1 + mean1
|
||||
output1 = model1.decode(latent1_tensor).sample
|
||||
|
||||
mean2 = model2.config.arch_config.shift_factor.to(input_tensor.device, input_tensor.dtype)
|
||||
std2 = model2.config.arch_config.scaling_factor.to(input_tensor.device, input_tensor.dtype)
|
||||
latent2_tensor = latent2.mode()
|
||||
latents_mean = (torch.tensor(model2.config.latents_mean).view(
|
||||
1, model2.config.z_dim, 1, 1, 1).to(input_tensor.device,
|
||||
input_tensor.dtype))
|
||||
latents_std = 1.0 / torch.tensor(model2.config.latents_std).view(
|
||||
1, model2.config.z_dim, 1, 1, 1).to(input_tensor.device,
|
||||
input_tensor.dtype)
|
||||
latent2_tensor = latent2_tensor / latents_std + latents_mean
|
||||
latent2_tensor = latent2_tensor / std2 + mean2
|
||||
output2 = model2.decode(latent2_tensor)
|
||||
# Check if outputs have the same shape
|
||||
assert output1.shape == output2.shape, f"Output shapes don't match: {output1.shape} vs {output2.shape}"
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user