Compare commits
30
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
1856f01c06 | ||
|
|
55442027f8 | ||
|
|
710b9c549b | ||
|
|
40120579f2 | ||
|
|
f8409886a1 | ||
|
|
f37e1fa18e | ||
|
|
4917d58d8e | ||
|
|
bbeb227bc5 | ||
|
|
163f9227ff | ||
|
|
98070d9bca | ||
|
|
dfb15828a8 | ||
|
|
ef667b9942 | ||
|
|
01282b7622 | ||
|
|
b14cc034e4 | ||
|
|
6e5cb4b8e5 | ||
|
|
87913376be | ||
|
|
ce22b1279d | ||
|
|
f1dba33df7 | ||
|
|
e8ce4290ff | ||
|
|
fb1cc737e3 | ||
|
|
f5074a5c31 | ||
|
|
0d4ae825c9 | ||
|
|
da66e0fcea | ||
|
|
97f9db9b71 | ||
|
|
e297fe3949 | ||
|
|
dd2cf8ad69 | ||
|
|
f5f5903aa4 | ||
|
|
20b03a45d7 | ||
|
|
2e4b523bee | ||
|
|
ea0332f62a |
@@ -1,29 +0,0 @@
|
||||
name: 🐞 Bug report
|
||||
description: Create a report to help us reproduce and fix the bug
|
||||
title: "[Bug] "
|
||||
labels: ['Bug']
|
||||
|
||||
body:
|
||||
- type: textarea
|
||||
attributes:
|
||||
label: Environment
|
||||
description: |
|
||||
Please share your environment with us. You can run the command **python fastvideo/utils/env_utils.py** and copy-paste its output below.
|
||||
placeholder: FastVideo version, platform, python version, cuda version...
|
||||
validations:
|
||||
required: true
|
||||
- type: textarea
|
||||
attributes:
|
||||
label: Describe the bug
|
||||
description: A clear and concise description of what the bug is.
|
||||
validations:
|
||||
required: true
|
||||
- type: textarea
|
||||
attributes:
|
||||
label: Reproduction
|
||||
description: |
|
||||
What command or script did you run? Which **model** are you using?
|
||||
placeholder: |
|
||||
A placeholder for the command.
|
||||
validations:
|
||||
required: true
|
||||
@@ -1,17 +0,0 @@
|
||||
name: 🚀 Feature request
|
||||
description: Suggest an idea for this project
|
||||
title: "[Feature] "
|
||||
|
||||
body:
|
||||
- type: textarea
|
||||
attributes:
|
||||
label: Motivation
|
||||
description: |
|
||||
A clear and concise description of the motivation of the feature.
|
||||
validations:
|
||||
required: true
|
||||
- type: textarea
|
||||
attributes:
|
||||
label: Related resources
|
||||
description: |
|
||||
If there is an official code release or third-party implementations, please also provide the information here, which would be very helpful.
|
||||
@@ -1,155 +1,162 @@
|
||||
# FastVideo
|
||||
|
||||
<div align="center">
|
||||
<img src=assets/logo.jpg width="30%"/>
|
||||
<a href=""><img src="https://img.shields.io/static/v1?label=API:H100&message=Replicate&color=pink"></a>  
|
||||
<a href=""><img src="https://img.shields.io/static/v1?label=Discuss&message=Discord&color=purple&logo=discord"></a>  
|
||||
</div>
|
||||
<br>
|
||||
<div align="center">
|
||||
<img src=assets/logo.png width="50%"/>
|
||||
</div>
|
||||
|
||||
FastVideo is a lightweight framework for accelerating large video diffusion models.
|
||||
FastVideo is a scalable framework for post-training video diffusion models, addressing the growing challenges of fine-tuning, distillation, and inference as model sizes and sequence lengths increase. As a first step, it provides an efficient script for distilling and fine-tuning the 10B Mochi model, with plans to expand features and support for more models.
|
||||
|
||||
### Features
|
||||
|
||||
https://github.com/user-attachments/assets/5fbc4596-56d6-43aa-98e0-da472cf8e26c
|
||||
|
||||
|
||||
|
||||
|
||||
<p align="center">
|
||||
🤗 <a href="https://huggingface.co/FastVideo/FastMochi-diffusers" target="_blank">FastMochi</a> | 🤗 <a href="https://huggingface.co/FastVideo/FastHunyuan" target="_blank">FastHunyuan</a> | 🎮 <a href="https://discord.gg/REBzDQTWWt" target="_blank"> Discord </a> | 🕹️ <a href="https://replicate.com/lucataco/fast-hunyuan-video" target="_blank"> Replicate </a>
|
||||
</p>
|
||||
|
||||
|
||||
FastVideo currently offers: (with more to come)
|
||||
|
||||
- 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.
|
||||
- Scalable training with FSDP, sequence parallelism, and selective activation checkpointing, with near linear scaling to 64 GPUs.
|
||||
- Memory efficient finetuning with LoRA, precomputed latent, and precomputed text embeddings.
|
||||
|
||||
Dev in progress and highly experimental.
|
||||
|
||||
## 🎥 More Demos
|
||||
|
||||
Fast-Hunyuan comparison with original Hunyuan, achieving an 8X diffusion speed boost with the FastVideo framework.
|
||||
|
||||
https://github.com/user-attachments/assets/064ac1d2-11ed-4a0c-955b-4d412a96ef30
|
||||
|
||||
Comparison between OpenAI Sora, original Hunyuan and FastHunyuan
|
||||
|
||||
https://github.com/user-attachments/assets/d323b712-3f68-42b2-952b-94f6a49c4836
|
||||
|
||||
Comparison between original FastHunyuan, LLM-INT8 quantized FastHunyuan and NF4 quantized FastHunyuan
|
||||
|
||||
https://github.com/user-attachments/assets/cf89efb5-5f68-4949-a085-f41c1ef26c94
|
||||
- FastMochi, a distilled Mochi model that can generate videos with merely 8 sampling steps.
|
||||
- Finetuning with FSDP (both master weight and ema weight), sequence parallelism, and selective gradient checkpointing.
|
||||
- LoRA coupled with pecomputed the latents and text embedding for minumum memory consumption.
|
||||
- Finetuning with both image and videos.
|
||||
|
||||
## Change Log
|
||||
- ```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/06```: `FastMochi` v0.0.1 is released.
|
||||
|
||||
|
||||
## Fast and High-Quality Text-to-video Generation
|
||||
|
||||
### 8-Step Results of FastMochi
|
||||
|
||||
<table class="center">
|
||||
<td><img src=assets/8steps/1.gif width="320"></td></td>
|
||||
<td><img src=assets/8steps/2.gif width="320"></td></td></td>
|
||||
<tr>
|
||||
<td style="text-align:center;" width="320">tmp</td>
|
||||
<td style="text-align:center;" width="320">tmp</td>
|
||||
<tr>
|
||||
</table >
|
||||
|
||||
|
||||
## Table of Contents
|
||||
|
||||
Jump to a specific section:
|
||||
|
||||
- [🔧 Installation](#-installation)
|
||||
- [🚀 Inference](#-inference)
|
||||
- [🎯 Distill](#-distill)
|
||||
- [⚡ Finetune](#-lora-finetune)
|
||||
|
||||
|
||||
## 🔧 Installation
|
||||
The code is tested on Python 3.10.0, CUDA 12.1 and H100.
|
||||
|
||||
```
|
||||
./env_setup.sh fastvideo
|
||||
conda create -n fastmochi python=3.10.0 -y && conda activate fastmochi
|
||||
pip3 install torch==2.5.0 torchvision --index-url https://download.pytorch.org/whl/cu121
|
||||
pip install packaging ninja && pip install flash-attn==2.7.0.post2 --no-build-isolation
|
||||
pip install "git+https://github.com/huggingface/diffusers.git@bf64b32652a63a1865a0528a73a13652b201698b"
|
||||
git clone https://github.com/hao-ai-lab/FastVideo.git
|
||||
cd FastVideo && pip install -e .
|
||||
```
|
||||
|
||||
|
||||
|
||||
|
||||
## 🚀 Inference
|
||||
|
||||
### 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.
|
||||
Use [scripts/download_hf.py](scripts/download_hf.py) to download the hugging-face style model to a local directory. Use it like this:
|
||||
```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_diffusers_hunyuan.sh
|
||||
python scripts/download_hf.py --repo_id=FastVideo/FastMochi --local_dir=data/FastMochi --repo_type=model
|
||||
```
|
||||
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
|
||||
Start the gradio UI with
|
||||
```
|
||||
python fastvideo/demo/gradio_web_demo.py --model_path data/FastMochi
|
||||
```
|
||||
You can also inference FastHunyuan in the [official Hunyuan github](https://github.com/Tencent/HunyuanVideo).
|
||||
|
||||
### FastMochi
|
||||
We also provide CLI inference script featured with sequence parallelism.
|
||||
|
||||
```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
|
||||
```
|
||||
export NUM_GPUS=4
|
||||
|
||||
torchrun --nnodes=1 --nproc_per_node=$NUM_GPUS \
|
||||
fastvideo/sample/sample_t2v_mochi.py \
|
||||
--model_path data/FastMochi \
|
||||
--prompt_path assets/prompt.txt \
|
||||
--num_frames 163 \
|
||||
--height 480 \
|
||||
--width 848 \
|
||||
--num_inference_steps 8 \
|
||||
--guidance_scale 1.5 \
|
||||
--output_path outputs_video/demo_video \
|
||||
--seed 12345 \
|
||||
--scheduler_type "pcm_linear_quadratic" \
|
||||
--linear_threshold 0.1 \
|
||||
--linear_range 0.75
|
||||
```
|
||||
|
||||
For the mochi style, simply following the scripts list in mochi repo.
|
||||
|
||||
```
|
||||
git clone https://github.com/genmoai/mochi.git
|
||||
cd mochi
|
||||
|
||||
# install env
|
||||
...
|
||||
|
||||
python3 ./demos/cli.py --model_dir weights/ --cpu_offload
|
||||
```
|
||||
|
||||
|
||||
## 🎯 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=genmo/mochi-1-preview --local_dir=data/mochi --repo_type=model # original mochi
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/hunyuan --local_dir=data/hunyuan --repo_type=model # original hunyuan
|
||||
```
|
||||
To launch the distillation process, use the following commands:
|
||||
```
|
||||
bash scripts/distill/distill_mochi.sh # for mochi
|
||||
bash scripts/distill/distill_hunyuan.sh # for hunyuan
|
||||
```
|
||||
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 specificed 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
|
||||
Currently, we only provide Lora Finetune for Mochi model, the command for Lora Finetune is
|
||||
```
|
||||
bash scripts/finetune/finetune_mochi_lora.sh
|
||||
```
|
||||
### 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.
|
||||
### Finetune with Both Image and Video
|
||||
Our codebase support finetuning with both image and video.
|
||||
## 💰Hardware requirement
|
||||
|
||||
- VRAM is required for both distill 10B mochi model
|
||||
|
||||
To launch distillation, you will first need to prepare data in the following formats
|
||||
|
||||
```bash
|
||||
bash scripts/finetune/finetune_hunyuan.sh
|
||||
bash scripts/finetune/finetune_mochi_lora_mix.sh
|
||||
asset/example_data
|
||||
├── AAA.txt
|
||||
├── AAA.png
|
||||
├── BCC.txt
|
||||
├── BCC.png
|
||||
├── ......
|
||||
├── CCC.txt
|
||||
└── CCC.png
|
||||
```
|
||||
For Image-Video Mixture Fine-tuning, make sure to enable the --group_frame option in your script.
|
||||
|
||||
## 📑 Development Plan
|
||||
We provide a dataset example here. First download testing data. Use [scripts/download_hf.py](scripts/download_hf.py) to download the data to a local directory. Use it like this:
|
||||
```bash
|
||||
python scripts/download_hf.py --repo_id=Stealths-Video/Merge-425-Data --local_dir=data/Merge-425-Data --repo_type=dataset
|
||||
python scripts/download_hf.py --repo_id=Stealths-Video/validation_embeddings --local_dir=data/validation_embeddings --repo_type=dataset
|
||||
```
|
||||
|
||||
- More distillation methods
|
||||
- [ ] Add Distribution Matching Distillation
|
||||
- More models support
|
||||
- [ ] Add CogvideoX model
|
||||
- Code update
|
||||
- [ ] fp8 support
|
||||
- [ ] faster load model and save model support
|
||||
Then the distillation can be launched by:
|
||||
|
||||
```
|
||||
bash scripts/distill_t2v.sh
|
||||
```
|
||||
|
||||
|
||||
## ⚡ Lora Finetune
|
||||
|
||||
|
||||
## 💰Hardware requirement
|
||||
|
||||
- VRAM is required for both distill 10B mochi model
|
||||
|
||||
To launch finetuning, you will first need to prepare data in the following formats.
|
||||
|
||||
|
||||
|
||||
Then the finetuning can be launched by:
|
||||
|
||||
```
|
||||
bash scripts/lora_finetune.sh
|
||||
```
|
||||
|
||||
## 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 thank MBZUAI and Anyscale for their support throughout this project.
|
||||
We learned from and reused code from the following projects: [PCM](https://github.com/G-U-N/Phased-Consistency-Model), [diffusers](https://github.com/huggingface/diffusers), and [OpenSoraPlan](https://github.com/PKU-YuanGroup/Open-Sora-Plan).
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
|
Before Width: | Height: | Size: 22 MiB |
Binary file not shown.
Binary file not shown.
|
Before Width: | Height: | Size: 149 KiB |
Binary file not shown.
|
After Width: | Height: | Size: 380 KiB |
+9
-8
@@ -1,8 +1,9 @@
|
||||
Will Smith casually eats noodles, his relaxed demeanor contrasting with the energetic background of a bustling street food market. The scene captures a mix of humor and authenticity. Mid-shot framing, vibrant lighting.
|
||||
A lone hiker stands atop a towering cliff, silhouetted against the vast horizon. The rugged landscape stretches endlessly beneath, its earthy tones blending into the soft blues of the sky. The scene captures the spirit of exploration and human resilience. High angle, dynamic framing, with soft natural lighting emphasizing the grandeur of nature.
|
||||
A hand with delicate fingers picks up a bright yellow lemon from a wooden bowl filled with lemons and sprigs of mint against a peach-colored background. The hand gently tosses the lemon up and catches it, showcasing its smooth texture. A beige string bag sits beside the bowl, adding a rustic touch to the scene. Additional lemons, one halved, are scattered around the base of the bowl. The even lighting enhances the vibrant colors and creates a fresh, inviting atmosphere.
|
||||
A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes wide with interest. The playful yet serene atmosphere is complemented by soft natural light filtering through the petals. Mid-shot, warm and cheerful tones.
|
||||
A superintelligent humanoid robot waking up. The robot has a sleek metallic body with futuristic design features. Its glowing red eyes are the focal point, emanating a sharp, intense light as it powers on. The scene is set in a dimly lit, high-tech laboratory filled with glowing control panels, robotic arms, and holographic screens. The setting emphasizes advanced technology and an atmosphere of mystery. The ambiance is eerie and dramatic, highlighting the moment of awakening and the robots immense intelligence. Photorealistic style with a cinematic, dark sci-fi aesthetic. Aspect ratio: 16:9 --v 6.1
|
||||
fox in the forest close-up quickly turned its head to the left
|
||||
Man walking his dog in the woods on a hot sunny day
|
||||
A majestic lion strides across the golden savanna, its powerful frame glistening under the warm afternoon sun. The tall grass ripples gently in the breeze, enhancing the lion's commanding presence. The tone is vibrant, embodying the raw energy of the wild. Low angle, steady tracking shot, cinematic.
|
||||
A hand enters the frame, pulling a sheet of plastic wrap over three balls of dough placed on a wooden surface. The plastic wrap is stretched to cover the dough more securely. The hand adjusts the wrap, ensuring that it is tight and smooth over the dough. The scene focuses on the hand's movements as it secures the edges of the plastic wrap. No new objects appear, and the camera remains stationary, focusing on the action of covering the dough.
|
||||
A vintage train snakes through the mountains, its plume of white steam rising dramatically against the jagged peaks. The cars glint in the late afternoon sun, their deep crimson and gold accents lending a touch of elegance. The tracks carve a precarious path along the cliffside, revealing glimpses of a roaring river far below. Inside, passengers peer out the large windows, their faces lit with awe as the landscape unfolds.
|
||||
A crowded rooftop bar buzzes with energy, the city skyline twinkling like a field of stars in the background. Strings of fairy lights hang above, casting a warm, golden glow over the scene. Groups of people gather around high tables, their laughter blending with the soft rhythm of live jazz. The aroma of freshly mixed cocktails and charred appetizers wafts through the air, mingling with the cool night breeze.
|
||||
En "The Matrix", Neo, interpretado por Keanu Reeves, personifica la lucha contra un sistema opresor a través de su icónica imagen, que incluye unos anteojos oscuros. Estos lentes no son solo un accesorio de moda; representan una barrera entre la realidad y la percepción. Al usar estos anteojos, Neo se sumerge en un mundo donde la verdad se oculta detrás de ilusiones y engaños. La oscuridad de los lentes simboliza la ignorancia y el control que las máquinas tienen sobre la humanidad, mientras que su propia búsqueda de la verdad lo lleva a descubrir sus auténticos poderes. La escena en que se los pone se convierte en un momento crucial, marcando su transformación de un simple programador a "El Elegido". Esta imagen se ha convertido en un ícono cultural, encapsulando el mensaje de que, al enfrentar la oscuridad, podemos encontrar la luz que nos guía hacia la libertad. Así, los anteojos de Neo se convierten en un símbolo de resistencia y autoconocimiento en un mundo manipulado.
|
||||
Medium close up. Low-angle shot. A woman in a 1950s retro dress sits in a diner bathed in neon light, surrounded by classic decor and lively chatter. The camera starts with a medium shot of her sitting at the counter, then slowly zooms in as she blows a shiny pink bubblegum bubble. The bubble swells dramatically before popping with a soft, playful burst. The scene is vibrant and nostalgic, evoking the fun and carefree spirit of the 1950s.
|
||||
Will Smith eats noodles.
|
||||
A short clip of the blonde woman taking a sip from her whiskey glass, her eyes locking with the camera as she smirks playfully. The background shows a group of people laughing and enjoying the party, with vibrant neon signs illuminating the space. The shot is taken in a way that conveys the feeling of a tipsy, carefree night out. The camera then zooms in on her face as she winks, creating a cheeky, flirtatious vibe.
|
||||
A superintelligent humanoid robot waking up. The robot has a sleek metallic body with futuristic design features. Its glowing red eyes are the focal point, emanating a sharp, intense light as it powers on. The scene is set in a dimly lit, high-tech laboratory filled with glowing control panels, robotic arms, and holographic screens. The setting emphasizes advanced technology and an atmosphere of mystery. The ambiance is eerie and dramatic, highlighting the moment of awakening and the robot's immense intelligence. Photorealistic style with a cinematic, dark sci-fi aesthetic. Aspect ratio: 16:9 --v 6.1
|
||||
A chimpanzee lead vocalist singing into a microphone on stage. The camera zooms in to show him singing. There is a spotlight on him.
|
||||
@@ -1,24 +0,0 @@
|
||||
# Configuration for Cog ⚙️
|
||||
# Reference: https://cog.run/yaml
|
||||
|
||||
build:
|
||||
gpu: true
|
||||
cuda: "12.1"
|
||||
python_version: "3.10"
|
||||
python_packages:
|
||||
- "torch==2.4.0"
|
||||
- "torchvision"
|
||||
- "ninja==1.11.1.3"
|
||||
- "transformers==4.46.1"
|
||||
- "git+https://github.com/huggingface/diffusers.git@bf64b32652a63a1865a0528a73a13652b201698b"
|
||||
- "accelerate==1.0.1"
|
||||
- "safetensors==0.4.5"
|
||||
- "peft==0.13.2"
|
||||
- "packaging==24.2"
|
||||
- "git+https://github.com/hao-ai-lab/FastVideo"
|
||||
|
||||
run:
|
||||
- FLASH_ATTENTION_SKIP_CUDA_BUILD=TRUE pip install flash-attn --no-build-isolation
|
||||
- curl -o /usr/local/bin/pget -L "https://github.com/replicate/pget/releases/latest/download/pget_$(uname -s)_$(uname -m)" && chmod +x /usr/local/bin/pget
|
||||
|
||||
predict: "predict.py:Predictor"
|
||||
@@ -1,68 +0,0 @@
|
||||
|
||||
|
||||
|
||||
## 🧱 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.
|
||||
|
||||
|
||||
We provide a sample dataset to help you get started. Download the source media using the following command:
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/Image-Vid-Finetune-Src --local_dir=data/Image-Vid-Finetune-Src --repo_type=dataset
|
||||
```
|
||||
To preprocess the dataset for fine-tuning or distillation, run:
|
||||
```
|
||||
bash scripts/preprocess/preprocess_mochi_data.sh # for mochi
|
||||
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
|
||||
|
||||
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
|
||||
│ ├── 1.mp4
|
||||
│ ├── 2.jpg
|
||||
├── video2caption.json
|
||||
└── merge.txt
|
||||
|
||||
Format the JSON file as a list, where each item represents a media source:
|
||||
|
||||
For image media,
|
||||
```
|
||||
{
|
||||
"path": "0.jpg",
|
||||
"cap": ["captions"]
|
||||
}
|
||||
```
|
||||
For video media,
|
||||
```
|
||||
{
|
||||
"path": "1.mp4",
|
||||
"resolution": {
|
||||
"width": 848,
|
||||
"height": 480
|
||||
},
|
||||
"fps": 30.0,
|
||||
"duration": 6.033333333333333,
|
||||
"cap": [
|
||||
"caption"
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
Use a txt file (merge.txt) to contain the source folder for media and the JSON file for meta information:
|
||||
|
||||
```
|
||||
path_to_media_source_foder,path_to_json_file
|
||||
```
|
||||
|
||||
Adjust the `DATA_MERGE_PATH` and `OUTPUT_DIR` in `scripts/preprocess/preprocess_****_data.sh` accordingly and run:
|
||||
```
|
||||
bash scripts/preprocess/preprocess_****_data.sh
|
||||
```
|
||||
The preprocessed data will be put into the `OUTPUT_DIR` and the `videos2caption.json` can be used in finetune and distill scripts.
|
||||
@@ -1,10 +0,0 @@
|
||||
#!/bin/bash
|
||||
|
||||
# install torch
|
||||
pip install torch==2.5.0 torchvision --index-url https://download.pytorch.org/whl/cu121
|
||||
|
||||
# install FA2 and diffusers
|
||||
pip install packaging ninja && pip install flash-attn==2.7.0.post2 --no-build-isolation
|
||||
|
||||
# install fastvideo
|
||||
pip install -e .
|
||||
@@ -14,11 +14,11 @@ from torch.utils.data import DataLoader
|
||||
from fastvideo.utils.load import load_text_encoder, load_vae
|
||||
from diffusers.video_processor import VideoProcessor
|
||||
from tqdm import tqdm
|
||||
|
||||
|
||||
class T5dataset(Dataset):
|
||||
def __init__(
|
||||
self, json_path, vae_debug,
|
||||
self,
|
||||
json_path,
|
||||
vae_debug,
|
||||
):
|
||||
self.json_path = json_path
|
||||
self.vae_debug = vae_debug
|
||||
@@ -50,7 +50,7 @@ def main(args):
|
||||
local_rank = int(os.getenv("RANK", 0))
|
||||
world_size = int(os.getenv("WORLD_SIZE", 1))
|
||||
print("world_size", world_size, "local rank", local_rank)
|
||||
|
||||
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
torch.cuda.set_device(local_rank)
|
||||
if not dist.is_initialized():
|
||||
@@ -65,9 +65,9 @@ def main(args):
|
||||
os.makedirs(os.path.join(args.output_dir, "prompt_embed"), exist_ok=True)
|
||||
os.makedirs(os.path.join(args.output_dir, "prompt_attention_mask"), exist_ok=True)
|
||||
|
||||
latents_json_path = os.path.join(args.output_dir, "videos2caption_temp.json")
|
||||
latents_json_path = os.path.join(args.output_dir, "videos2caption_temp_replace.json")
|
||||
train_dataset = T5dataset(latents_json_path, args.vae_debug)
|
||||
text_encoder = load_text_encoder(args.model_type, args.model_path, device=device)
|
||||
text_encoder = load_text_encoder(args.model_type,args.model_path, device=device)
|
||||
vae, autocast_type, fps = load_vae(args.model_type, args.model_path)
|
||||
vae.enable_tiling()
|
||||
sampler = DistributedSampler(
|
||||
|
||||
@@ -13,7 +13,6 @@ import torch.distributed as dist
|
||||
from torch.utils.data.distributed import DistributedSampler
|
||||
from fastvideo.utils.load import load_vae
|
||||
from tqdm import tqdm
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
@@ -38,7 +37,7 @@ def main(args):
|
||||
dist.init_process_group(
|
||||
backend="nccl", init_method="env://", world_size=world_size, rank=local_rank
|
||||
)
|
||||
vae, autocast_type, fps = load_vae(args.model_type, args.model_path)
|
||||
vae, autocast_type = load_vae(args.model_type, args.model_path)
|
||||
vae.enable_tiling()
|
||||
os.makedirs(args.output_dir, exist_ok=True)
|
||||
os.makedirs(os.path.join(args.output_dir, "latent"), exist_ok=True)
|
||||
|
||||
@@ -20,7 +20,7 @@ def main(args):
|
||||
local_rank = int(os.getenv("RANK", 0))
|
||||
world_size = int(os.getenv("WORLD_SIZE", 1))
|
||||
print("world_size", world_size, "local rank", local_rank)
|
||||
|
||||
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
torch.cuda.set_device(local_rank)
|
||||
if not dist.is_initialized():
|
||||
@@ -28,22 +28,17 @@ def main(args):
|
||||
backend="nccl", init_method="env://", world_size=world_size, rank=local_rank
|
||||
)
|
||||
|
||||
text_encoder = load_text_encoder(args.model_type, args.model_path, device=device)
|
||||
text_encoder = load_text_encoder(args.model_type,args.model_path, device=device)
|
||||
autocast_type = torch.float16 if args.model_type == "hunyuan" else torch.bfloat16
|
||||
# output_dir/validation/prompt_attention_mask
|
||||
# output_dir/validation/prompt_embed
|
||||
os.makedirs(os.path.join(args.output_dir, "validation"), exist_ok=True)
|
||||
os.makedirs(
|
||||
os.path.join(args.output_dir, "validation", "prompt_attention_mask"),
|
||||
exist_ok=True,
|
||||
)
|
||||
os.makedirs(
|
||||
os.path.join(args.output_dir, "validation", "prompt_embed"), exist_ok=True
|
||||
)
|
||||
os.makedirs(os.path.join(args.output_dir,"validation"), exist_ok=True)
|
||||
os.makedirs(os.path.join(args.output_dir,"validation", "prompt_attention_mask"), exist_ok=True)
|
||||
os.makedirs(os.path.join(args.output_dir,"validation", "prompt_embed"), exist_ok=True)
|
||||
json_data = []
|
||||
with open(args.validation_prompt_txt, "r", encoding="utf-8") as file:
|
||||
with open(args.validation_prompt_txt, 'r', encoding='utf-8') as file:
|
||||
lines = file.readlines()
|
||||
prompts = [line.strip() for line in lines]
|
||||
prompts = [line.strip() for line in lines]
|
||||
for prompt in prompts:
|
||||
with torch.inference_mode():
|
||||
with torch.autocast("cuda", dtype=autocast_type):
|
||||
@@ -51,15 +46,8 @@ def main(args):
|
||||
prompt
|
||||
)
|
||||
file_name = prompt.split(".")[0]
|
||||
prompt_embed_path = os.path.join(
|
||||
args.output_dir, "validation", "prompt_embed", f"{file_name}.pt"
|
||||
)
|
||||
prompt_attention_mask_path = os.path.join(
|
||||
args.output_dir,
|
||||
"validation",
|
||||
"prompt_attention_mask",
|
||||
f"{file_name}.pt",
|
||||
)
|
||||
prompt_embed_path = os.path.join(args.output_dir,"validation", "prompt_embed", f"{file_name}.pt")
|
||||
prompt_attention_mask_path = os.path.join(args.output_dir,"validation", "prompt_attention_mask", f"{file_name}.pt")
|
||||
torch.save(prompt_embeds[0], prompt_embed_path)
|
||||
torch.save(prompt_attention_mask[0], prompt_attention_mask_path)
|
||||
print(f"sample {file_name} saved")
|
||||
|
||||
@@ -7,7 +7,10 @@ import random
|
||||
|
||||
class LatentDataset(Dataset):
|
||||
def __init__(
|
||||
self, json_path, num_latent_t, cfg_rate,
|
||||
self,
|
||||
json_path,
|
||||
num_latent_t,
|
||||
cfg_rate,
|
||||
):
|
||||
# data_merge_path: video_dir, latent_dir, prompt_embed_dir, json_path
|
||||
self.json_path = json_path
|
||||
@@ -107,7 +110,7 @@ def latent_collate_function(batch):
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
dataset = LatentDataset("data/HD-Mixkit-Finetune-Hunyuan/videos2caption.json", num_latent_t=8, cfg_rate=6)
|
||||
dataset = LatentDataset("data/Mochi-Synthetic-Data/merge.txt", num_latent_t=28)
|
||||
dataloader = torch.utils.data.DataLoader(
|
||||
dataset, batch_size=2, shuffle=False, collate_fn=latent_collate_function
|
||||
)
|
||||
|
||||
@@ -105,7 +105,6 @@ class T2V_dataset(Dataset):
|
||||
dataset_prog.set_cap_list(args.dataloader_num_workers, cap_list, n_elements)
|
||||
|
||||
print(f"video length: {len(dataset_prog.cap_list)}", flush=True)
|
||||
|
||||
def set_checkpoint(self, n_used_elements):
|
||||
for i in range(len(dataset_prog.n_used_elements)):
|
||||
dataset_prog.n_used_elements[i] = n_used_elements
|
||||
@@ -197,7 +196,7 @@ class T2V_dataset(Dataset):
|
||||
caps = [random.choice(caps)]
|
||||
text = caps
|
||||
input_ids, cond_mask = [], []
|
||||
text = text[0] if random.random() > self.cfg else ""
|
||||
text = text if random.random() > self.cfg else ""
|
||||
text_tokens_and_mask = self.tokenizer(
|
||||
text,
|
||||
max_length=self.text_max_length,
|
||||
@@ -270,8 +269,10 @@ class T2V_dataset(Dataset):
|
||||
# import ipdb;ipdb.set_trace()
|
||||
i["num_frames"] = math.ceil(fps * duration)
|
||||
# max 5.0 and min 1.0 are just thresholds to filter some videos which have suitable duration.
|
||||
if i["num_frames"] / fps > self.video_length_tolerance_range * (
|
||||
self.num_frames / self.train_fps * self.speed_factor
|
||||
if (
|
||||
i["num_frames"] / fps
|
||||
> self.video_length_tolerance_range
|
||||
* (self.num_frames / self.train_fps * self.speed_factor)
|
||||
): # too long video is not suitable for this training stage (self.num_frames)
|
||||
cnt_too_long += 1
|
||||
continue
|
||||
|
||||
@@ -288,7 +288,10 @@ class LongSideResizeVideo:
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self, size, skip_low_resolution=False, interpolation_mode="bilinear",
|
||||
self,
|
||||
size,
|
||||
skip_low_resolution=False,
|
||||
interpolation_mode="bilinear",
|
||||
):
|
||||
self.size = size
|
||||
self.skip_low_resolution = skip_low_resolution
|
||||
@@ -327,7 +330,10 @@ class CenterCropResizeVideo:
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self, size, top_crop=False, interpolation_mode="bilinear",
|
||||
self,
|
||||
size,
|
||||
top_crop=False,
|
||||
interpolation_mode="bilinear",
|
||||
):
|
||||
if len(size) != 2:
|
||||
raise ValueError(
|
||||
@@ -368,7 +374,9 @@ class UCFCenterCropVideo:
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self, size, interpolation_mode="bilinear",
|
||||
self,
|
||||
size,
|
||||
interpolation_mode="bilinear",
|
||||
):
|
||||
if isinstance(size, tuple):
|
||||
if len(size) != 2:
|
||||
@@ -405,7 +413,9 @@ class KineticsRandomCropResizeVideo:
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self, size, interpolation_mode="bilinear",
|
||||
self,
|
||||
size,
|
||||
interpolation_mode="bilinear",
|
||||
):
|
||||
if isinstance(size, tuple):
|
||||
if len(size) != 2:
|
||||
@@ -426,7 +436,9 @@ class KineticsRandomCropResizeVideo:
|
||||
|
||||
class CenterCropVideo:
|
||||
def __init__(
|
||||
self, size, interpolation_mode="bilinear",
|
||||
self,
|
||||
size,
|
||||
interpolation_mode="bilinear",
|
||||
):
|
||||
if isinstance(size, tuple):
|
||||
if len(size) != 2:
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
import gradio as gr
|
||||
import torch
|
||||
from fastvideo.models.mochi_hf.pipeline_mochi import MochiPipeline
|
||||
from fastvideo.models.mochi_hf.modeling_mochi import MochiTransformer3DModel
|
||||
from fastvideo.model.pipeline_mochi import MochiPipeline
|
||||
from fastvideo.model.modeling_mochi import MochiTransformer3DModel
|
||||
from diffusers import FlowMatchEulerDiscreteScheduler
|
||||
from diffusers.utils import export_to_video
|
||||
from fastvideo.distill.solver import PCMFMScheduler
|
||||
@@ -9,80 +9,57 @@ import tempfile
|
||||
import os
|
||||
import argparse
|
||||
|
||||
|
||||
def init_args():
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--prompts", nargs="+", default=[])
|
||||
parser.add_argument("--num_frames", type=int, default=25)
|
||||
parser.add_argument("--prompts", nargs='+', default=[])
|
||||
parser.add_argument("--num_frames", type=int, default=163)
|
||||
parser.add_argument("--height", type=int, default=480)
|
||||
parser.add_argument("--width", type=int, default=848)
|
||||
parser.add_argument("--num_inference_steps", type=int, default=8)
|
||||
parser.add_argument("--num_inference_steps", type=int, default=64)
|
||||
parser.add_argument("--guidance_scale", type=float, default=4.5)
|
||||
parser.add_argument("--model_path", type=str, default="data/mochi")
|
||||
parser.add_argument("--seed", type=int, default=12345)
|
||||
parser.add_argument("--seed", type=int, default=42)
|
||||
parser.add_argument("--transformer_path", type=str, default=None)
|
||||
parser.add_argument("--scheduler_type", type=str, default="pcm_linear_quadratic")
|
||||
parser.add_argument("--scheduler_type", type=str, default="euler")
|
||||
parser.add_argument("--lora_checkpoint_dir", type=str, default=None)
|
||||
parser.add_argument("--shift", type=float, default=8.0)
|
||||
parser.add_argument("--num_euler_timesteps", type=int, default=50)
|
||||
parser.add_argument("--linear_threshold", type=float, default=0.1)
|
||||
parser.add_argument("--linear_range", type=float, default=0.75)
|
||||
parser.add_argument("--num_euler_timesteps", type=int, default=100)
|
||||
parser.add_argument("--linear_threshold", type=float, default=0.025)
|
||||
parser.add_argument("--linear_range", type=float, default=0.5)
|
||||
parser.add_argument("--cpu_offload", action="store_true")
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def load_model(args):
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
if args.scheduler_type == "euler":
|
||||
scheduler = FlowMatchEulerDiscreteScheduler()
|
||||
else:
|
||||
linear_quadratic = True if "linear_quadratic" in args.scheduler_type else False
|
||||
scheduler = PCMFMScheduler(
|
||||
1000,
|
||||
args.shift,
|
||||
args.num_euler_timesteps,
|
||||
linear_quadratic,
|
||||
args.linear_threshold,
|
||||
args.linear_range,
|
||||
)
|
||||
|
||||
scheduler = PCMFMScheduler(1000, args.shift, args.num_euler_timesteps, False, args.linear_threshold, args.linear_range)
|
||||
|
||||
if args.transformer_path:
|
||||
transformer = MochiTransformer3DModel.from_pretrained(args.transformer_path)
|
||||
else:
|
||||
transformer = MochiTransformer3DModel.from_pretrained(
|
||||
args.model_path, subfolder="transformer/"
|
||||
)
|
||||
|
||||
pipe = MochiPipeline.from_pretrained(
|
||||
args.model_path, transformer=transformer, scheduler=scheduler
|
||||
)
|
||||
transformer = MochiTransformer3DModel.from_pretrained(args.model_path, subfolder='transformer/')
|
||||
|
||||
pipe = MochiPipeline.from_pretrained(args.model_path, transformer=transformer, scheduler=scheduler)
|
||||
pipe.enable_vae_tiling()
|
||||
# pipe.to(device)
|
||||
# if args.cpu_offload:
|
||||
pipe.enable_sequential_cpu_offload()
|
||||
pipe.to(device)
|
||||
if args.cpu_offload:
|
||||
pipe.enable_model_cpu_offload()
|
||||
return pipe
|
||||
|
||||
|
||||
def generate_video(
|
||||
prompt,
|
||||
negative_prompt,
|
||||
use_negative_prompt,
|
||||
seed,
|
||||
guidance_scale,
|
||||
num_frames,
|
||||
height,
|
||||
width,
|
||||
num_inference_steps,
|
||||
randomize_seed=False,
|
||||
):
|
||||
def generate_video(prompt, negative_prompt, use_negative_prompt, seed, guidance_scale,
|
||||
num_frames, height, width, num_inference_steps, randomize_seed=False):
|
||||
if randomize_seed:
|
||||
seed = torch.randint(0, 1000000, (1,)).item()
|
||||
|
||||
|
||||
pipe = load_model(args)
|
||||
print("load model successfully")
|
||||
generator = torch.Generator(device="cuda").manual_seed(seed)
|
||||
|
||||
|
||||
if not use_negative_prompt:
|
||||
negative_prompt = None
|
||||
|
||||
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
output = pipe(
|
||||
prompt=[prompt],
|
||||
@@ -94,24 +71,22 @@ def generate_video(
|
||||
guidance_scale=guidance_scale,
|
||||
generator=generator,
|
||||
).frames[0]
|
||||
|
||||
|
||||
output_path = os.path.join(tempfile.mkdtemp(), "output.mp4")
|
||||
export_to_video(output, output_path, fps=30)
|
||||
return output_path, seed
|
||||
|
||||
|
||||
examples = [
|
||||
"A hand enters the frame, pulling a sheet of plastic wrap over three balls of dough placed on a wooden surface. The plastic wrap is stretched to cover the dough more securely. The hand adjusts the wrap, ensuring that it is tight and smooth over the dough. The scene focuses on the hand’s movements as it secures the edges of the plastic wrap. No new objects appear, and the camera remains stationary, focusing on the action of covering the dough.",
|
||||
"A vintage train snakes through the mountains, its plume of white steam rising dramatically against the jagged peaks. The cars glint in the late afternoon sun, their deep crimson and gold accents lending a touch of elegance. The tracks carve a precarious path along the cliffside, revealing glimpses of a roaring river far below. Inside, passengers peer out the large windows, their faces lit with awe as the landscape unfolds.",
|
||||
"A crowded rooftop bar buzzes with energy, the city skyline twinkling like a field of stars in the background. Strings of fairy lights hang above, casting a warm, golden glow over the scene. Groups of people gather around high tables, their laughter blending with the soft rhythm of live jazz. The aroma of freshly mixed cocktails and charred appetizers wafts through the air, mingling with the cool night breeze.",
|
||||
"A crowded rooftop bar buzzes with energy, the city skyline twinkling like a field of stars in the background. Strings of fairy lights hang above, casting a warm, golden glow over the scene. Groups of people gather around high tables, their laughter blending with the soft rhythm of live jazz. The aroma of freshly mixed cocktails and charred appetizers wafts through the air, mingling with the cool night breeze."
|
||||
]
|
||||
|
||||
args = init_args()
|
||||
pipe = load_model(args)
|
||||
print("load model successfully")
|
||||
with gr.Blocks() as demo:
|
||||
gr.Markdown("# Fastvideo Mochi Video Generation Demo")
|
||||
|
||||
with gr.Blocks() as demo:
|
||||
gr.Markdown("# Mochi Video Generation Demo")
|
||||
|
||||
with gr.Group():
|
||||
with gr.Row():
|
||||
prompt = gr.Text(
|
||||
@@ -123,60 +98,33 @@ with gr.Blocks() as demo:
|
||||
)
|
||||
run_button = gr.Button("Run", scale=0)
|
||||
result = gr.Video(label="Result", show_label=False)
|
||||
|
||||
|
||||
with gr.Accordion("Advanced options", open=False):
|
||||
with gr.Group():
|
||||
with gr.Row():
|
||||
height = gr.Slider(
|
||||
label="Height",
|
||||
minimum=256,
|
||||
maximum=1024,
|
||||
step=32,
|
||||
value=args.height,
|
||||
)
|
||||
width = gr.Slider(
|
||||
label="Width", minimum=256, maximum=1024, step=32, value=args.width
|
||||
)
|
||||
|
||||
height = gr.Slider(label="Height", minimum=256, maximum=1024, step=32, value=args.height)
|
||||
width = gr.Slider(label="Width", minimum=256, maximum=1024, step=32, value=args.width)
|
||||
|
||||
with gr.Row():
|
||||
num_frames = gr.Slider(
|
||||
label="Number of Frames",
|
||||
minimum=21,
|
||||
maximum=163,
|
||||
value=args.num_frames,
|
||||
)
|
||||
guidance_scale = gr.Slider(
|
||||
label="Guidance Scale",
|
||||
minimum=1,
|
||||
maximum=12,
|
||||
value=args.guidance_scale,
|
||||
)
|
||||
num_inference_steps = gr.Slider(
|
||||
label="Inference Steps",
|
||||
minimum=4,
|
||||
maximum=100,
|
||||
value=args.num_inference_steps,
|
||||
)
|
||||
|
||||
num_frames = gr.Slider(label="Number of Frames", minimum=8, maximum=256, value=args.num_frames)
|
||||
guidance_scale = gr.Slider(label="Guidance Scale", minimum=1, maximum=20, value=args.guidance_scale)
|
||||
num_inference_steps = gr.Slider(label="Inference Steps", minimum=10, maximum=100, value=args.num_inference_steps)
|
||||
|
||||
with gr.Row():
|
||||
use_negative_prompt = gr.Checkbox(
|
||||
label="Use negative prompt", value=False
|
||||
)
|
||||
use_negative_prompt = gr.Checkbox(label="Use negative prompt", value=False)
|
||||
negative_prompt = gr.Text(
|
||||
label="Negative prompt",
|
||||
max_lines=1,
|
||||
placeholder="Enter a negative prompt",
|
||||
visible=False,
|
||||
)
|
||||
|
||||
seed = gr.Slider(
|
||||
label="Seed", minimum=0, maximum=1000000, step=1, value=args.seed
|
||||
visible=False
|
||||
)
|
||||
|
||||
seed = gr.Slider(label="Seed", minimum=0, maximum=1000000, step=1, value=args.seed)
|
||||
randomize_seed = gr.Checkbox(label="Randomize seed", value=True)
|
||||
seed_output = gr.Number(label="Used Seed")
|
||||
|
||||
|
||||
gr.Examples(examples=examples, inputs=prompt)
|
||||
|
||||
|
||||
use_negative_prompt.change(
|
||||
fn=lambda x: gr.update(visible=x),
|
||||
inputs=use_negative_prompt,
|
||||
@@ -185,19 +133,9 @@ with gr.Blocks() as demo:
|
||||
|
||||
run_button.click(
|
||||
fn=generate_video,
|
||||
inputs=[
|
||||
prompt,
|
||||
negative_prompt,
|
||||
use_negative_prompt,
|
||||
seed,
|
||||
guidance_scale,
|
||||
num_frames,
|
||||
height,
|
||||
width,
|
||||
num_inference_steps,
|
||||
randomize_seed,
|
||||
],
|
||||
outputs=[result, seed_output],
|
||||
inputs=[prompt, negative_prompt, use_negative_prompt, seed, guidance_scale,
|
||||
num_frames, height, width, num_inference_steps, randomize_seed],
|
||||
outputs=[result, seed_output]
|
||||
)
|
||||
|
||||
if __name__ == "__main__":
|
||||
+39
-37
@@ -27,8 +27,10 @@ import wandb
|
||||
from accelerate.utils import set_seed
|
||||
from tqdm.auto import tqdm
|
||||
from fastvideo.utils.fsdp_util import get_dit_fsdp_kwargs, apply_fsdp_checkpointing
|
||||
from diffusers import FlowMatchEulerDiscreteScheduler
|
||||
from fastvideo.utils.load import load_transformer
|
||||
from diffusers import (
|
||||
FlowMatchEulerDiscreteScheduler,
|
||||
)
|
||||
from fastvideo.utils.load import get_no_split_modules, load_transformer
|
||||
from fastvideo.distill.solver import EulerSolver, extract_into_tensor
|
||||
from copy import deepcopy
|
||||
from diffusers.optimization import get_scheduler
|
||||
@@ -37,7 +39,9 @@ from fastvideo.dataset.latent_datasets import LatentDataset, latent_collate_func
|
||||
import torch.distributed as dist
|
||||
from safetensors.torch import save_file
|
||||
from peft import LoraConfig
|
||||
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
|
||||
from torch.distributed.fsdp import (
|
||||
FullyShardedDataParallel as FSDP,
|
||||
)
|
||||
from fastvideo.utils.checkpoint import (
|
||||
save_checkpoint,
|
||||
save_lora_checkpoint,
|
||||
@@ -49,7 +53,6 @@ check_min_version("0.31.0")
|
||||
import time
|
||||
from collections import deque
|
||||
|
||||
|
||||
def main_print(content):
|
||||
if int(os.environ["LOCAL_RANK"]) <= 0:
|
||||
print(content)
|
||||
@@ -71,8 +74,7 @@ def save_checkpoint(transformer, rank, output_dir, step):
|
||||
weight_path = os.path.join(save_dir, "diffusion_pytorch_model.safetensors")
|
||||
save_file(cpu_state, weight_path)
|
||||
config_dict = dict(transformer.config)
|
||||
if "dtype" in config_dict:
|
||||
del config_dict["dtype"] # TODO
|
||||
if 'dtype' in config_dict: del config_dict['dtype'] # TODO
|
||||
config_path = os.path.join(save_dir, "config.json")
|
||||
# save dict as json
|
||||
with open(config_path, "w") as f:
|
||||
@@ -127,7 +129,7 @@ def distill_one_step(
|
||||
ema_decay,
|
||||
pred_decay_weight,
|
||||
pred_decay_type,
|
||||
hunyuan_teacher_disable_cfg,
|
||||
hunyuan_student_cfg_embed
|
||||
):
|
||||
total_loss = 0.0
|
||||
optimizer.zero_grad()
|
||||
@@ -166,18 +168,16 @@ def distill_one_step(
|
||||
noisy_model_input = sigmas * noise + (1.0 - sigmas) * model_input
|
||||
# Predict the noise residual
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
teacher_kwargs = {
|
||||
student_kwargs = {
|
||||
"hidden_states": noisy_model_input,
|
||||
"encoder_hidden_states": encoder_hidden_states,
|
||||
"timestep": timesteps,
|
||||
"encoder_attention_mask": encoder_attention_mask, # B, L
|
||||
"return_dict": False,
|
||||
}
|
||||
if hunyuan_teacher_disable_cfg:
|
||||
teacher_kwargs["guidance"] = torch.tensor(
|
||||
[1000.0], device=noisy_model_input.device, dtype=torch.bfloat16
|
||||
)
|
||||
model_pred = transformer(**teacher_kwargs)[0]
|
||||
if hunyuan_student_cfg_embed:
|
||||
student_kwargs["guidance"] = torch.tensor([hunyuan_student_cfg_embed], device=noisy_model_input.device, dtype=torch.bfloat16) * 1000
|
||||
model_pred = transformer(**student_kwargs)[0]
|
||||
|
||||
# if accelerator.is_main_process:
|
||||
model_pred, end_index = solver.euler_style_multiphase_pred(
|
||||
@@ -238,7 +238,7 @@ def distill_one_step(
|
||||
# loss = loss.mean()
|
||||
loss = (
|
||||
torch.mean(
|
||||
torch.sqrt((model_pred.float() - target.float()) ** 2 + huber_c ** 2)
|
||||
torch.sqrt((model_pred.float() - target.float()) ** 2 + huber_c**2)
|
||||
- huber_c
|
||||
)
|
||||
/ gradient_accumulation_steps
|
||||
@@ -318,13 +318,9 @@ def main(args):
|
||||
# Create model:
|
||||
|
||||
main_print(f"--> loading model from {args.pretrained_model_name_or_path}")
|
||||
|
||||
transformer = load_transformer(
|
||||
args.model_type,
|
||||
args.dit_model_name_or_path,
|
||||
args.pretrained_model_name_or_path,
|
||||
torch.float32 if args.master_weight_type == "fp32" else torch.bfloat16,
|
||||
)
|
||||
|
||||
|
||||
transformer = load_transformer(args.model_type,args.dit_model_name_or_path, args.pretrained_model_name_or_path,torch.float32 if args.master_weight_type == "fp32" else torch.bfloat16)
|
||||
|
||||
teacher_transformer = deepcopy(transformer)
|
||||
if args.use_ema:
|
||||
@@ -364,23 +360,26 @@ def main(args):
|
||||
transformer._no_split_modules = no_split_modules
|
||||
fsdp_kwargs["auto_wrap_policy"] = fsdp_kwargs["auto_wrap_policy"](transformer)
|
||||
|
||||
transformer = FSDP(transformer, **fsdp_kwargs,)
|
||||
teacher_transformer = FSDP(teacher_transformer, **fsdp_kwargs,)
|
||||
transformer = FSDP(
|
||||
transformer,
|
||||
**fsdp_kwargs,
|
||||
)
|
||||
teacher_transformer = FSDP(
|
||||
teacher_transformer,
|
||||
**fsdp_kwargs,
|
||||
)
|
||||
if args.use_ema:
|
||||
ema_transformer = FSDP(ema_transformer, **fsdp_kwargs,)
|
||||
ema_transformer = FSDP(
|
||||
ema_transformer,
|
||||
**fsdp_kwargs,
|
||||
)
|
||||
main_print(f"--> model loaded")
|
||||
|
||||
if args.gradient_checkpointing:
|
||||
apply_fsdp_checkpointing(
|
||||
transformer, no_split_modules, args.selective_checkpointing
|
||||
)
|
||||
apply_fsdp_checkpointing(
|
||||
teacher_transformer, no_split_modules, args.selective_checkpointing
|
||||
)
|
||||
apply_fsdp_checkpointing(transformer, no_split_modules, args.selective_checkpointing)
|
||||
apply_fsdp_checkpointing(teacher_transformer, no_split_modules, args.selective_checkpointing)
|
||||
if args.use_ema:
|
||||
apply_fsdp_checkpointing(
|
||||
ema_transformer, no_split_modules, args.selective_checkpointing
|
||||
)
|
||||
apply_fsdp_checkpointing(ema_transformer, no_split_modules, args.selective_checkpointing)
|
||||
# Set model as trainable.
|
||||
transformer.train()
|
||||
teacher_transformer.requires_grad_(False)
|
||||
@@ -566,7 +565,7 @@ def main(args):
|
||||
args.ema_decay,
|
||||
args.pred_decay_weight,
|
||||
args.pred_decay_type,
|
||||
args.hunyuan_teacher_disable_cfg,
|
||||
args.hunyuan_student_cfg_embed
|
||||
)
|
||||
|
||||
step_time = time.time() - start_time
|
||||
@@ -652,9 +651,12 @@ def main(args):
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
|
||||
|
||||
parser.add_argument(
|
||||
"--model_type", type=str, default="mochi", help="The type of model to train."
|
||||
"--model_type",
|
||||
type=str,
|
||||
default="mochi",
|
||||
help="The type of model to train."
|
||||
)
|
||||
|
||||
# dataset & dataloader
|
||||
@@ -892,7 +894,7 @@ if __name__ == "__main__":
|
||||
parser.add_argument("--multi_phased_distill_schedule", type=str, default=None)
|
||||
parser.add_argument("--pred_decay_weight", type=float, default=0.0)
|
||||
parser.add_argument("--pred_decay_type", default="l1")
|
||||
parser.add_argument("--hunyuan_teacher_disable_cfg", action="store_true")
|
||||
parser.add_argument("--hunyuan_student_cfg_embed", type=float)
|
||||
parser.add_argument(
|
||||
"--master_weight_type",
|
||||
type=str,
|
||||
|
||||
@@ -24,7 +24,7 @@ logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
|
||||
class DiscriminatorHead(nn.Module):
|
||||
def __init__(self, input_channel, output_channel=1, args=None):
|
||||
def __init__(self, input_channel, output_channel=1):
|
||||
super().__init__()
|
||||
inner_channel = 1024
|
||||
self.conv1 = nn.Sequential(
|
||||
@@ -43,21 +43,13 @@ class DiscriminatorHead(nn.Module):
|
||||
)
|
||||
|
||||
self.conv_out = nn.Conv2d(inner_channel, output_channel, 1, 1, 0)
|
||||
|
||||
vae_spatial_scale_factor = 8
|
||||
|
||||
self.patch_height = args.num_height // vae_spatial_scale_factor // 2
|
||||
self.patch_width = args.num_width // vae_spatial_scale_factor // 2
|
||||
print("## DiscriminatorHead: patch_height: ", self.patch_height)
|
||||
print("## DiscriminatorHead: patch_width: ", self.patch_width)
|
||||
|
||||
def forward(self, x):
|
||||
b, twh, c = x.shape
|
||||
|
||||
t = twh // (self.patch_height * self.patch_width)
|
||||
x = x.view(-1, self.patch_height * self.patch_width, c)
|
||||
t = twh // (30 * 53)
|
||||
x = x.view(-1, 30 * 53, c)
|
||||
x = x.permute(0, 2, 1)
|
||||
x = x.view(b * t, c, self.patch_height, self.patch_width)
|
||||
x = x.view(b * t, c, 30, 53)
|
||||
x = self.conv1(x)
|
||||
x = self.conv2(x) + x
|
||||
x = self.conv_out(x)
|
||||
@@ -70,11 +62,9 @@ class Discriminator(nn.Module):
|
||||
stride=8,
|
||||
num_h_per_head=1,
|
||||
adapter_channel_dims=[3072],
|
||||
total_layers = 48,
|
||||
args=None,
|
||||
):
|
||||
super().__init__()
|
||||
adapter_channel_dims = adapter_channel_dims * (total_layers // stride)
|
||||
adapter_channel_dims = adapter_channel_dims * (48 // stride)
|
||||
self.stride = stride
|
||||
self.num_h_per_head = num_h_per_head
|
||||
self.head_num = len(adapter_channel_dims)
|
||||
@@ -82,10 +72,7 @@ class Discriminator(nn.Module):
|
||||
[
|
||||
nn.ModuleList(
|
||||
[
|
||||
DiscriminatorHead(
|
||||
adapter_channel,
|
||||
args=args
|
||||
)
|
||||
DiscriminatorHead(adapter_channel)
|
||||
for _ in range(self.num_h_per_head)
|
||||
]
|
||||
)
|
||||
@@ -102,9 +89,9 @@ class Discriminator(nn.Module):
|
||||
|
||||
return custom_forward
|
||||
|
||||
assert len(features) == len(self.heads)
|
||||
for i in range(0, len(features)):
|
||||
for h in self.heads[i]:
|
||||
assert len(features) // self.stride == len(self.heads)
|
||||
for i in range(0, len(features), self.stride):
|
||||
for h in self.heads[i // self.stride]:
|
||||
# out = torch.utils.checkpoint.checkpoint(
|
||||
# create_custom_forward(h),
|
||||
# features[i],
|
||||
@@ -113,25 +100,3 @@ class Discriminator(nn.Module):
|
||||
out = h(features[i])
|
||||
outputs.append(out)
|
||||
return outputs
|
||||
|
||||
|
||||
class DMDiscriminator(nn.Module):
|
||||
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
|
||||
self.cls_pred_branch = nn.Sequential(
|
||||
nn.Conv2d(kernel_size=4, in_channels=1280, out_channels=1280, stride=2, padding=1), # 8x8 -> 4x4
|
||||
nn.GroupNorm(num_groups=32, num_channels=1280),
|
||||
nn.SiLU(),
|
||||
nn.Conv2d(kernel_size=4, in_channels=1280, out_channels=1280, stride=4, padding=0), # 4x4 -> 1x1
|
||||
nn.GroupNorm(num_groups=32, num_channels=1280),
|
||||
nn.SiLU(),
|
||||
nn.Conv2d(kernel_size=1, in_channels=1280, out_channels=1, stride=1, padding=0), # 1x1 -> 1x1
|
||||
)
|
||||
|
||||
self.cls_pred_branch.requires_grad_(True)
|
||||
|
||||
def forward(self, features):
|
||||
print("## features shape: ", features.shape)
|
||||
return self.cls_pred_branch(features)
|
||||
|
||||
@@ -275,7 +275,12 @@ class EulerSolver:
|
||||
return x_prev
|
||||
|
||||
def euler_style_multiphase_pred(
|
||||
self, sample, model_pred, timestep_index, multiphase, is_target=False,
|
||||
self,
|
||||
sample,
|
||||
model_pred,
|
||||
timestep_index,
|
||||
multiphase,
|
||||
is_target=False,
|
||||
):
|
||||
inference_indices = np.linspace(
|
||||
0, len(self.euler_timesteps), num=multiphase, endpoint=False
|
||||
|
||||
+63
-113
@@ -22,7 +22,6 @@ from torch.distributed.fsdp import (
|
||||
StateDictType,
|
||||
FullStateDictConfig,
|
||||
)
|
||||
from fastvideo.utils.load import load_transformer
|
||||
|
||||
from fastvideo.models.mochi_hf.pipeline_mochi import linear_quadratic_schedule
|
||||
import json
|
||||
@@ -53,41 +52,14 @@ from torch.distributed.fsdp import (
|
||||
FullyShardedDataParallel as FSDP,
|
||||
)
|
||||
from fastvideo.utils.checkpoint import (
|
||||
save_checkpoint,
|
||||
save_lora_checkpoint,
|
||||
resume_lora_optimizer,
|
||||
resume_training,
|
||||
save_checkpoint_generator_discriminator,
|
||||
resume_training_generator_discriminator,
|
||||
)
|
||||
# from fastvideo.utils.checkpoint import save_checkpoint
|
||||
from fastvideo.utils.logging_ import main_print
|
||||
from torch.distributed.fsdp import FullOptimStateDictConfig
|
||||
from safetensors.torch import save_file
|
||||
|
||||
def save_checkpoint(model, rank, output_dir, step, discriminator=False):
|
||||
with FSDP.state_dict_type(
|
||||
model,
|
||||
StateDictType.FULL_STATE_DICT,
|
||||
FullStateDictConfig(offload_to_cpu=True, rank0_only=True),
|
||||
FullOptimStateDictConfig(offload_to_cpu=True, rank0_only=True),
|
||||
):
|
||||
cpu_state = model.state_dict()
|
||||
|
||||
# todo move to get_state_dict
|
||||
save_dir = os.path.join(output_dir, f"checkpoint-{step}")
|
||||
os.makedirs(save_dir, exist_ok=True)
|
||||
# save using safetensors
|
||||
if rank <= 0 and not discriminator:
|
||||
weight_path = os.path.join(save_dir, "diffusion_pytorch_model.safetensors")
|
||||
save_file(cpu_state, weight_path)
|
||||
config_dict = dict(model.config)
|
||||
config_path = os.path.join(save_dir, "config.json")
|
||||
# save dict as json
|
||||
with open(config_path, "w") as f:
|
||||
json.dump(config_dict, f, indent=4)
|
||||
else:
|
||||
weight_path = os.path.join(save_dir, "discriminator_pytorch_model.safetensors")
|
||||
save_file(cpu_state, weight_path)
|
||||
|
||||
# Will error if the minimal version of diffusers is not installed. Remove at your own risks.
|
||||
check_min_version("0.31.0")
|
||||
@@ -104,7 +76,6 @@ def gan_d_loss(
|
||||
encoder_hidden_states,
|
||||
encoder_attention_mask,
|
||||
weight,
|
||||
discriminator_head_stride
|
||||
):
|
||||
loss = 0.0
|
||||
# collate sample_fake and sample_real
|
||||
@@ -114,8 +85,7 @@ def gan_d_loss(
|
||||
encoder_hidden_states,
|
||||
timestep,
|
||||
encoder_attention_mask,
|
||||
output_features=True,
|
||||
output_features_stride=discriminator_head_stride,
|
||||
output_attn=True,
|
||||
return_dict=False,
|
||||
)[1]
|
||||
real_features = teacher_transformer(
|
||||
@@ -123,8 +93,7 @@ def gan_d_loss(
|
||||
encoder_hidden_states,
|
||||
timestep,
|
||||
encoder_attention_mask,
|
||||
output_features=True,
|
||||
output_features_stride=discriminator_head_stride,
|
||||
output_attn=True,
|
||||
return_dict=False,
|
||||
)[1]
|
||||
|
||||
@@ -146,7 +115,6 @@ def gan_g_loss(
|
||||
encoder_hidden_states,
|
||||
encoder_attention_mask,
|
||||
weight,
|
||||
discriminator_head_stride
|
||||
):
|
||||
loss = 0.0
|
||||
features = teacher_transformer(
|
||||
@@ -154,8 +122,7 @@ def gan_g_loss(
|
||||
encoder_hidden_states,
|
||||
timestep,
|
||||
encoder_attention_mask,
|
||||
output_features=True,
|
||||
output_features_stride=discriminator_head_stride,
|
||||
output_attn=True,
|
||||
return_dict=False,
|
||||
)[1]
|
||||
fake_outputs = discriminator(
|
||||
@@ -168,19 +135,20 @@ def gan_g_loss(
|
||||
return loss
|
||||
|
||||
|
||||
def distill_one_step_adv(
|
||||
def train_one_step_mochi(
|
||||
transformer,
|
||||
model_type,
|
||||
teacher_transformer,
|
||||
optimizer,
|
||||
discriminator,
|
||||
discriminator_optimizer,
|
||||
global_step,
|
||||
lr_scheduler,
|
||||
loader,
|
||||
noise_scheduler,
|
||||
solver,
|
||||
noise_random_generator,
|
||||
sp_size,
|
||||
precondition_outputs,
|
||||
max_grad_norm,
|
||||
uncond_prompt_embed,
|
||||
uncond_prompt_mask,
|
||||
@@ -189,7 +157,6 @@ def distill_one_step_adv(
|
||||
not_apply_cfg_solver,
|
||||
distill_cfg,
|
||||
adv_weight,
|
||||
discriminator_head_stride
|
||||
):
|
||||
optimizer.zero_grad()
|
||||
discriminator_optimizer.zero_grad()
|
||||
@@ -200,7 +167,7 @@ def distill_one_step_adv(
|
||||
latents_attention_mask,
|
||||
encoder_attention_mask,
|
||||
) = next(loader)
|
||||
model_input = normalize_dit_input(model_type, latents)
|
||||
model_input = normalize_mochi_dit_input(latents)
|
||||
noise = torch.randn_like(model_input)
|
||||
bsz = model_input.shape[0]
|
||||
index = torch.randint(
|
||||
@@ -317,7 +284,6 @@ def distill_one_step_adv(
|
||||
encoder_hidden_states.float(),
|
||||
encoder_attention_mask,
|
||||
1.0,
|
||||
discriminator_head_stride
|
||||
)
|
||||
g_loss += g_gan_loss
|
||||
g_loss.backward()
|
||||
@@ -342,7 +308,6 @@ def distill_one_step_adv(
|
||||
encoder_hidden_states,
|
||||
encoder_attention_mask,
|
||||
1.0,
|
||||
discriminator_head_stride,
|
||||
)
|
||||
|
||||
d_loss.backward()
|
||||
@@ -382,14 +347,21 @@ def main(args):
|
||||
|
||||
main_print(f"--> loading model from {args.pretrained_model_name_or_path}")
|
||||
# keep the master weight to float32
|
||||
transformer = load_transformer(
|
||||
args.model_type,
|
||||
args.dit_model_name_or_path,
|
||||
args.pretrained_model_name_or_path,
|
||||
torch.float32 if args.master_weight_type == "fp32" else torch.bfloat16,
|
||||
)
|
||||
if args.dit_model_name_or_path:
|
||||
transformer = transformer = MochiTransformer3DModel.from_pretrained(
|
||||
args.dit_model_name_or_path,
|
||||
torch_dtype=torch.float32,
|
||||
# torch_dtype=torch.bfloat16 if args.use_lora else torch.float32,
|
||||
)
|
||||
else:
|
||||
transformer = MochiTransformer3DModel.from_pretrained(
|
||||
args.pretrained_model_name_or_path,
|
||||
subfolder="transformer",
|
||||
torch_dtype=torch.float32,
|
||||
# torch_dtype=torch.bfloat16 if args.use_lora else torch.float32,
|
||||
)
|
||||
teacher_transformer = deepcopy(transformer)
|
||||
discriminator = Discriminator(args.discriminator_head_stride, total_layers = 48 if args.model_type =="mochi" else 40)
|
||||
discriminator = Discriminator(args.discriminator_head_stride)
|
||||
|
||||
if args.use_lora:
|
||||
transformer.requires_grad_(False)
|
||||
@@ -411,20 +383,15 @@ def main(args):
|
||||
main_print(
|
||||
f"--> Initializing FSDP with sharding strategy: {args.fsdp_sharding_startegy}"
|
||||
)
|
||||
fsdp_kwargs, no_split_modules = get_dit_fsdp_kwargs(
|
||||
transformer,
|
||||
args.fsdp_sharding_startegy,
|
||||
args.use_lora,
|
||||
args.use_cpu_offload,
|
||||
args.master_weight_type,
|
||||
fsdp_kwargs = get_dit_fsdp_kwargs(
|
||||
args.fsdp_sharding_startegy, args.use_lora, args.use_cpu_offload
|
||||
)
|
||||
discriminator_fsdp_kwargs = get_discriminator_fsdp_kwargs(args.master_weight_type)
|
||||
if args.use_lora:
|
||||
assert args.model_type == "mochi", "LoRA is only supported for Mochi model."
|
||||
transformer.config.lora_rank = args.lora_rank
|
||||
transformer.config.lora_alpha = args.lora_alpha
|
||||
transformer.config.lora_target_modules = ["to_k", "to_q", "to_v", "to_out.0"]
|
||||
transformer._no_split_modules = no_split_modules
|
||||
transformer._no_split_modules = ["MochiTransformerBlock"]
|
||||
fsdp_kwargs["auto_wrap_policy"] = fsdp_kwargs["auto_wrap_policy"](transformer)
|
||||
|
||||
transformer = FSDP(
|
||||
@@ -442,12 +409,8 @@ def main(args):
|
||||
main_print(f"--> model loaded")
|
||||
|
||||
if args.gradient_checkpointing:
|
||||
apply_fsdp_checkpointing(
|
||||
transformer, no_split_modules, args.selective_checkpointing
|
||||
)
|
||||
apply_fsdp_checkpointing(
|
||||
teacher_transformer, no_split_modules, args.selective_checkpointing
|
||||
)
|
||||
apply_fsdp_checkpointing(transformer, args.selective_checkpointing)
|
||||
apply_fsdp_checkpointing(teacher_transformer, args.selective_checkpointing)
|
||||
# Set model as trainable.
|
||||
transformer.train()
|
||||
teacher_transformer.requires_grad_(False)
|
||||
@@ -599,49 +562,39 @@ def main(args):
|
||||
|
||||
step_times = deque(maxlen=100)
|
||||
# log_validation(args, transformer, device,
|
||||
# torch.bfloat16, 0, scheduler_type=args.scheduler_type, shift=args.shift, num_euler_timesteps=args.num_euler_timesteps, linear_quadratic_threshold=args.linear_quadratic_threshold,ema=False)
|
||||
def get_num_phases(multi_phased_distill_schedule, step):
|
||||
# step-phase,step-phase
|
||||
multi_phases = multi_phased_distill_schedule.split(",")
|
||||
phase = multi_phases[-1].split("-")[-1]
|
||||
for step_phases in multi_phases:
|
||||
phase_step, phase = step_phases.split("-")
|
||||
if step <= int(phase_step):
|
||||
return int(phase)
|
||||
return phase
|
||||
# torch.bfloat16, init_steps, scheduler_type=args.scheduler_type, shift=args.shift, num_euler_timesteps=args.num_euler_timesteps, linear_quadratic_threshold=args.linear_quadratic_threshold, ema=False)
|
||||
|
||||
for i in range(init_steps):
|
||||
_ = next(loader)
|
||||
for step in range(init_steps + 1, args.max_train_steps + 1):
|
||||
assert args.multi_phased_distill_schedule is not None
|
||||
num_phases = get_num_phases(args.multi_phased_distill_schedule, step)
|
||||
start_time = time.time()
|
||||
(
|
||||
generator_loss,
|
||||
generator_grad_norm,
|
||||
discriminator_loss,
|
||||
discriminator_grad_norm,
|
||||
) = distill_one_step_adv(
|
||||
) = train_one_step_mochi(
|
||||
transformer,
|
||||
args.model_type,
|
||||
teacher_transformer,
|
||||
optimizer,
|
||||
discriminator,
|
||||
discriminator_optimizer,
|
||||
step,
|
||||
lr_scheduler,
|
||||
loader,
|
||||
noise_scheduler,
|
||||
solver,
|
||||
noise_random_generator,
|
||||
args.sp_size,
|
||||
args.precondition_outputs,
|
||||
args.max_grad_norm,
|
||||
uncond_prompt_embed,
|
||||
uncond_prompt_mask,
|
||||
args.num_euler_timesteps,
|
||||
num_phases,
|
||||
args.validation_sampling_steps,
|
||||
args.not_apply_cfg_solver,
|
||||
args.distill_cfg,
|
||||
args.adv_weight,
|
||||
args.discriminator_head_stride
|
||||
)
|
||||
|
||||
step_time = time.time() - start_time
|
||||
@@ -680,17 +633,15 @@ def main(args):
|
||||
)
|
||||
else:
|
||||
# Your existing checkpoint saving code
|
||||
# TODO
|
||||
# save_checkpoint_generator_discriminator(
|
||||
# transformer,
|
||||
# optimizer,
|
||||
# discriminator,
|
||||
# discriminator_optimizer,
|
||||
# rank,
|
||||
# args.output_dir,
|
||||
# step,
|
||||
# )
|
||||
save_checkpoint(transformer, rank, args.output_dir, step, discriminator)
|
||||
save_checkpoint_generator_discriminator(
|
||||
transformer,
|
||||
optimizer,
|
||||
discriminator,
|
||||
discriminator_optimizer,
|
||||
rank,
|
||||
args.output_dir,
|
||||
step,
|
||||
)
|
||||
main_print(f"--> checkpoint saved at step {step}")
|
||||
dist.barrier()
|
||||
if args.log_validation and step % args.validation_steps == 0:
|
||||
@@ -704,17 +655,25 @@ def main(args):
|
||||
shift=args.shift,
|
||||
num_euler_timesteps=args.num_euler_timesteps,
|
||||
linear_quadratic_threshold=args.linear_quadratic_threshold,
|
||||
linear_range=args.linear_range,
|
||||
ema=False,
|
||||
)
|
||||
|
||||
|
||||
if args.use_lora:
|
||||
save_lora_checkpoint(
|
||||
transformer, optimizer, rank, args.output_dir, args.max_train_steps
|
||||
)
|
||||
else:
|
||||
save_checkpoint(transformer, rank, args.output_dir, args.max_train_steps)
|
||||
save_checkpoint(
|
||||
transformer, optimizer, rank, args.output_dir, args.max_train_steps
|
||||
)
|
||||
save_checkpoint(
|
||||
discriminator,
|
||||
discriminator_optimizer,
|
||||
rank,
|
||||
args.output_dir,
|
||||
step,
|
||||
discriminator=True,
|
||||
)
|
||||
|
||||
if get_sequence_parallel_state():
|
||||
destroy_sequence_parallel_group()
|
||||
@@ -723,13 +682,8 @@ def main(args):
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
|
||||
parser.add_argument(
|
||||
"--model_type", type=str, default="mochi", help="The type of model to train."
|
||||
)
|
||||
# dataset & dataloader
|
||||
parser.add_argument("--data_json_path", type=str, required=True)
|
||||
parser.add_argument("--num_height", type=int, default=480)
|
||||
parser.add_argument("--num_width", type=int, default=848)
|
||||
parser.add_argument("--num_frames", type=int, default=163)
|
||||
parser.add_argument(
|
||||
"--dataloader_num_workers",
|
||||
@@ -758,9 +712,16 @@ if __name__ == "__main__":
|
||||
parser.add_argument("--ema_decay", type=float, default=0.999)
|
||||
parser.add_argument("--ema_start_step", type=int, default=0)
|
||||
parser.add_argument("--cfg", type=float, default=0.1)
|
||||
parser.add_argument(
|
||||
"--precondition_outputs",
|
||||
action="store_true",
|
||||
help="Whether to precondition the outputs of the model.",
|
||||
)
|
||||
|
||||
# validation & logs
|
||||
parser.add_argument("--validation_sampling_steps", type=str, default="64")
|
||||
parser.add_argument("--validation_guidance_scale", type=str, default="4.5")
|
||||
parser.add_argument("--validation_prompt_dir", type=str)
|
||||
parser.add_argument("--validation_sampling_steps", type=int, default=64)
|
||||
parser.add_argument("--validation_guidance_scale", type=float, default=4.5)
|
||||
parser.add_argument("--validation_steps", type=float, default=64)
|
||||
parser.add_argument("--log_validation", action="store_true")
|
||||
parser.add_argument("--tracker_project_name", type=str, default=None)
|
||||
@@ -789,7 +750,6 @@ if __name__ == "__main__":
|
||||
" training using `--resume_from_checkpoint`."
|
||||
),
|
||||
)
|
||||
parser.add_argument("--validation_prompt_dir", type=str)
|
||||
parser.add_argument("--shift", type=float, default=1.0)
|
||||
parser.add_argument(
|
||||
"--resume_from_checkpoint",
|
||||
@@ -906,7 +866,6 @@ if __name__ == "__main__":
|
||||
"--lora_rank", type=int, default=128, help="LoRA rank parameter. "
|
||||
)
|
||||
parser.add_argument("--fsdp_sharding_startegy", default="full")
|
||||
parser.add_argument("--multi_phased_distill_schedule", type=str, default=None)
|
||||
parser.add_argument(
|
||||
"--gradient_accumulation_steps",
|
||||
type=int,
|
||||
@@ -967,15 +926,6 @@ if __name__ == "__main__":
|
||||
default=0.025,
|
||||
help="The threshold of the linear quadratic scheduler.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--linear_range",
|
||||
type=float,
|
||||
default=0.5,
|
||||
help="Range for linear quadratic scheduler.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--weight_decay", type=float, default=0.001, help="Weight decay to apply."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--master_weight_type",
|
||||
type=str,
|
||||
@@ -983,4 +933,4 @@ if __name__ == "__main__":
|
||||
help="Weight type to use - fp32 or bf16.",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
main(args)
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -31,4 +31,4 @@ def flash_attn_no_pad(
|
||||
"b s (h d) -> b s h d",
|
||||
h=nheads,
|
||||
)
|
||||
return output
|
||||
return output
|
||||
@@ -17,9 +17,9 @@ __all__ = [
|
||||
]
|
||||
|
||||
PRECISION_TO_TYPE = {
|
||||
"fp32": torch.float32,
|
||||
"fp16": torch.float16,
|
||||
"bf16": torch.bfloat16,
|
||||
'fp32': torch.float32,
|
||||
'fp16': torch.float16,
|
||||
'bf16': torch.bfloat16,
|
||||
}
|
||||
|
||||
# =================== Constant Values =====================
|
||||
@@ -34,7 +34,7 @@ PROMPT_TEMPLATE_ENCODE = (
|
||||
"<|start_header_id|>system<|end_header_id|>\n\nDescribe the image by detailing the color, shape, size, texture, "
|
||||
"quantity, text, spatial relationships of the objects and background:<|eot_id|>"
|
||||
"<|start_header_id|>user<|end_header_id|>\n\n{}<|eot_id|>"
|
||||
)
|
||||
)
|
||||
PROMPT_TEMPLATE_ENCODE_VIDEO = (
|
||||
"<|start_header_id|>system<|end_header_id|>\n\nDescribe the video by detailing the following aspects: "
|
||||
"1. The main content and theme of the video."
|
||||
@@ -43,12 +43,15 @@ PROMPT_TEMPLATE_ENCODE_VIDEO = (
|
||||
"4. background environment, light, style and atmosphere."
|
||||
"5. camera angles, movements, and transitions used in the video:<|eot_id|>"
|
||||
"<|start_header_id|>user<|end_header_id|>\n\n{}<|eot_id|>"
|
||||
)
|
||||
)
|
||||
|
||||
NEGATIVE_PROMPT = "Aerial view, aerial view, overexposed, low quality, deformation, a poor composition, bad hands, bad teeth, bad eyes, bad limbs, distortion"
|
||||
|
||||
PROMPT_TEMPLATE = {
|
||||
"dit-llm-encode": {"template": PROMPT_TEMPLATE_ENCODE, "crop_start": 36,},
|
||||
"dit-llm-encode": {
|
||||
"template": PROMPT_TEMPLATE_ENCODE,
|
||||
"crop_start": 36,
|
||||
},
|
||||
"dit-llm-encode-video": {
|
||||
"template": PROMPT_TEMPLATE_ENCODE_VIDEO,
|
||||
"crop_start": 95,
|
||||
|
||||
@@ -52,7 +52,6 @@ from einops import rearrange
|
||||
from fastvideo.utils.parallel_states import get_sequence_parallel_state, nccl_info
|
||||
from fastvideo.utils.communications import all_gather, all_to_all_4D
|
||||
import torch.nn.functional as F
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
EXAMPLE_DOC_STRING = """"""
|
||||
@@ -371,6 +370,10 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
bs_embed * num_videos_per_prompt, seq_len, -1
|
||||
)
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
return (
|
||||
prompt_embeds,
|
||||
negative_prompt_embeds,
|
||||
@@ -484,6 +487,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
f" {negative_prompt_embeds.shape}."
|
||||
)
|
||||
|
||||
|
||||
def prepare_latents(
|
||||
self,
|
||||
batch_size,
|
||||
@@ -676,7 +680,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
negative_prompt_embeds (`torch.Tensor`, *optional*):
|
||||
Pre-generated negative text embeddings. Can be used to easily tweak text inputs (prompt weighting). If
|
||||
not provided, `negative_prompt_embeds` are generated from the `negative_prompt` input argument.
|
||||
|
||||
|
||||
output_type (`str`, *optional*, defaults to `"pil"`):
|
||||
The output format of the generated image. Choose between `PIL.Image` or `np.array`.
|
||||
return_dict (`bool`, *optional*, defaults to `True`):
|
||||
@@ -763,11 +767,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
else:
|
||||
batch_size = prompt_embeds.shape[0]
|
||||
|
||||
device = (
|
||||
torch.device(f"cuda:{dist.get_rank()}")
|
||||
if dist.is_initialized()
|
||||
else self._execution_device
|
||||
)
|
||||
device = torch.device(f"cuda:{dist.get_rank()}") if dist.is_initialized() else self._execution_device
|
||||
|
||||
# 3. Encode input prompt
|
||||
lora_scale = (
|
||||
@@ -834,6 +834,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
if prompt_mask_2 is not None:
|
||||
prompt_mask_2 = torch.cat([negative_prompt_mask_2, prompt_mask_2])
|
||||
|
||||
|
||||
# 4. Prepare timesteps
|
||||
extra_set_timesteps_kwargs = self.prepare_extra_func_kwargs(
|
||||
self.scheduler.set_timesteps, {"n_tokens": n_tokens}
|
||||
@@ -866,7 +867,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
generator,
|
||||
latents,
|
||||
)
|
||||
|
||||
|
||||
world_size, rank = nccl_info.sp_size, nccl_info.rank_within_group
|
||||
if get_sequence_parallel_state():
|
||||
latents = rearrange(
|
||||
@@ -876,7 +877,8 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
|
||||
# 6. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline
|
||||
extra_step_kwargs = self.prepare_extra_func_kwargs(
|
||||
self.scheduler.step, {"generator": generator, "eta": eta},
|
||||
self.scheduler.step,
|
||||
{"generator": generator, "eta": eta},
|
||||
)
|
||||
|
||||
target_dtype = PRECISION_TO_TYPE[self.args.precision]
|
||||
@@ -910,12 +912,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
|
||||
t_expand = t.repeat(latent_model_input.shape[0])
|
||||
guidance_expand = (
|
||||
torch.tensor(
|
||||
[embedded_guidance_scale] * latent_model_input.shape[0],
|
||||
dtype=torch.float32,
|
||||
device=device,
|
||||
).to(target_dtype)
|
||||
* 1000.0
|
||||
torch.tensor([embedded_guidance_scale] * latent_model_input.shape[0],dtype=torch.float32,device=device,).to(target_dtype)* 1000.0
|
||||
if embedded_guidance_scale is not None
|
||||
else None
|
||||
)
|
||||
@@ -930,9 +927,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
(0, prompt_embeds.shape[2] - prompt_embeds_2.shape[1]),
|
||||
value=0,
|
||||
).unsqueeze(1)
|
||||
encoder_hidden_states = torch.cat(
|
||||
[prompt_embeds_2, prompt_embeds], dim=1
|
||||
)
|
||||
encoder_hidden_states= torch.cat([prompt_embeds_2, prompt_embeds], dim=1)
|
||||
noise_pred = self.transformer( # For an input image (129, 192, 336) (1, 256, 256)
|
||||
latent_model_input, # [2, 16, 33, 24, 42]
|
||||
encoder_hidden_states,
|
||||
@@ -940,9 +935,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
prompt_mask, # [2, 256]fpdb
|
||||
guidance=guidance_expand,
|
||||
return_dict=False,
|
||||
)[
|
||||
0
|
||||
]
|
||||
)[0]
|
||||
|
||||
# perform guidance
|
||||
if self.do_classifier_free_guidance:
|
||||
@@ -985,7 +978,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
if callback is not None and i % callback_steps == 0:
|
||||
step_idx = i // getattr(self.scheduler, "order", 1)
|
||||
callback(step_idx, t, latents)
|
||||
|
||||
|
||||
if get_sequence_parallel_state():
|
||||
latents = all_gather(latents, dim=2)
|
||||
|
||||
|
||||
@@ -140,7 +140,7 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin):
|
||||
Number of tokens in the input sequence.
|
||||
"""
|
||||
self.num_inference_steps = num_inference_steps
|
||||
|
||||
|
||||
sigmas = torch.linspace(1, 0, num_inference_steps + 1)
|
||||
sigmas = self.sd3_time_shift(sigmas)
|
||||
|
||||
|
||||
@@ -195,7 +195,10 @@ def add_denoise_schedule_args(parser: argparse.ArgumentParser):
|
||||
help="If reverse, learning/sampling from t=1 -> t=0.",
|
||||
)
|
||||
group.add_argument(
|
||||
"--flow-solver", type=str, default="euler", help="Solver for flow matching.",
|
||||
"--flow-solver",
|
||||
type=str,
|
||||
default="euler",
|
||||
help="Solver for flow matching.",
|
||||
)
|
||||
group.add_argument(
|
||||
"--use-linear-quadratic-schedule",
|
||||
@@ -357,10 +360,16 @@ def add_parallel_args(parser: argparse.ArgumentParser):
|
||||
|
||||
# ======================== Model loads ========================
|
||||
group.add_argument(
|
||||
"--ulysses-degree", type=int, default=1, help="Ulysses degree.",
|
||||
"--ulysses-degree",
|
||||
type=int,
|
||||
default=1,
|
||||
help="Ulysses degree.",
|
||||
)
|
||||
group.add_argument(
|
||||
"--ring-degree", type=int, default=1, help="Ulysses degree.",
|
||||
"--ring-degree",
|
||||
type=int,
|
||||
default=1,
|
||||
help="Ulysses degree.",
|
||||
)
|
||||
|
||||
return parser
|
||||
|
||||
@@ -9,18 +9,14 @@ from loguru import logger
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from fastvideo.models.hunyuan.constants import (
|
||||
PROMPT_TEMPLATE,
|
||||
NEGATIVE_PROMPT,
|
||||
PRECISION_TO_TYPE,
|
||||
)
|
||||
from fastvideo.models.hunyuan.constants import PROMPT_TEMPLATE, NEGATIVE_PROMPT, PRECISION_TO_TYPE
|
||||
from fastvideo.models.hunyuan.vae import load_vae
|
||||
from fastvideo.models.hunyuan.modules import load_model
|
||||
from fastvideo.models.hunyuan.text_encoder import TextEncoder
|
||||
from fastvideo.models.hunyuan.utils.data_utils import align_to
|
||||
from fastvideo.models.hunyuan.diffusion.schedulers import FlowMatchDiscreteScheduler
|
||||
from fastvideo.models.hunyuan.diffusion.pipelines import HunyuanVideoPipeline
|
||||
from safetensors.torch import load_file as safetensors_load_file
|
||||
|
||||
from fastvideo.utils.parallel_states import (
|
||||
initialize_sequence_parallel_state,
|
||||
nccl_info,
|
||||
@@ -75,14 +71,14 @@ class Inference(object):
|
||||
"""
|
||||
# ========================================================================
|
||||
logger.info(f"Got text-to-video model root path: {pretrained_model_path}")
|
||||
|
||||
|
||||
# ==================== Initialize Distributed Environment ================
|
||||
if nccl_info.sp_size > 1:
|
||||
device = torch.device(f"cuda:{os.environ['LOCAL_RANK']}")
|
||||
if device is None:
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
|
||||
parallel_args = None # {"ulysses_degree": args.ulysses_degree, "ring_degree": args.ring_degree}
|
||||
parallel_args = None #{"ulysses_degree": args.ulysses_degree, "ring_degree": args.ring_degree}
|
||||
|
||||
# ======================== Get the args path =============================
|
||||
|
||||
@@ -175,7 +171,7 @@ class Inference(object):
|
||||
use_cpu_offload=args.use_cpu_offload,
|
||||
device=device,
|
||||
logger=logger,
|
||||
parallel_args=parallel_args,
|
||||
parallel_args=parallel_args
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
@@ -240,16 +236,7 @@ class Inference(object):
|
||||
if not model_path.exists():
|
||||
raise ValueError(f"model_path not exists: {model_path}")
|
||||
logger.info(f"Loading torch model {model_path}...")
|
||||
if model_path.suffix == ".safetensors":
|
||||
# Use safetensors library for .safetensors files
|
||||
state_dict = safetensors_load_file(model_path)
|
||||
elif model_path.suffix == ".pt":
|
||||
# Use torch for .pt files
|
||||
state_dict = torch.load(
|
||||
model_path, map_location=lambda storage, loc: storage
|
||||
)
|
||||
else:
|
||||
raise ValueError(f"Unsupported file format: {model_path}")
|
||||
state_dict = torch.load(model_path, map_location=lambda storage, loc: storage)
|
||||
|
||||
if bare_model == "unknown" and ("ema" in state_dict or "module" in state_dict):
|
||||
bare_model = False
|
||||
@@ -290,7 +277,7 @@ class HunyuanVideoSampler(Inference):
|
||||
use_cpu_offload=False,
|
||||
device=0,
|
||||
logger=None,
|
||||
parallel_args=None,
|
||||
parallel_args=None
|
||||
):
|
||||
super().__init__(
|
||||
args,
|
||||
@@ -303,7 +290,7 @@ class HunyuanVideoSampler(Inference):
|
||||
use_cpu_offload=use_cpu_offload,
|
||||
device=device,
|
||||
logger=logger,
|
||||
parallel_args=parallel_args,
|
||||
parallel_args=parallel_args
|
||||
)
|
||||
|
||||
self.pipeline = self.load_diffusion_pipeline(
|
||||
@@ -425,8 +412,7 @@ class HunyuanVideoSampler(Inference):
|
||||
raise ValueError(
|
||||
f"Seed must be an integer, a list of integers, or None, got {seed}."
|
||||
)
|
||||
# Peiyuan: using GPU seed will cause A100 and H100 to generate different results...
|
||||
generator = [torch.Generator("cpu").manual_seed(seed) for seed in seeds]
|
||||
generator = [torch.Generator(self.device).manual_seed(seed) for seed in seeds]
|
||||
out_dict["seeds"] = seeds
|
||||
|
||||
# ========================================================================
|
||||
@@ -473,10 +459,11 @@ class HunyuanVideoSampler(Inference):
|
||||
scheduler = FlowMatchDiscreteScheduler(
|
||||
shift=flow_shift,
|
||||
reverse=self.args.flow_reverse,
|
||||
solver=self.args.flow_solver,
|
||||
solver=self.args.flow_solver
|
||||
)
|
||||
self.pipeline.scheduler = scheduler
|
||||
|
||||
|
||||
if "884" in self.args.vae:
|
||||
latents_size = [(video_length - 1) // 4 + 1, height // 8, width // 8]
|
||||
elif "888" in self.args.vae:
|
||||
|
||||
@@ -17,6 +17,7 @@ def load_model(args, in_channels, out_channels, factor_kwargs):
|
||||
model = HYVideoDiffusionTransformer(
|
||||
in_channels=in_channels,
|
||||
out_channels=out_channels,
|
||||
|
||||
**HUNYUAN_VIDEO_CONFIG[args.model],
|
||||
**factor_kwargs,
|
||||
)
|
||||
|
||||
@@ -5,48 +5,56 @@ import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
|
||||
|
||||
from fastvideo.utils.parallel_states import get_sequence_parallel_state, nccl_info
|
||||
from fastvideo.utils.communications import all_gather, all_to_all_4D
|
||||
from fastvideo.models.flash_attn_no_pad import flash_attn_no_pad
|
||||
|
||||
|
||||
|
||||
def attention(
|
||||
q, k, v, drop_rate=0, attn_mask=None, causal=False,
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
drop_rate=0,
|
||||
attn_mask=None,
|
||||
causal=False,
|
||||
):
|
||||
|
||||
qkv = torch.stack([q, k, v], dim=2)
|
||||
|
||||
if attn_mask is not None and attn_mask.dtype != torch.bool:
|
||||
attn_mask = attn_mask.bool()
|
||||
|
||||
x = flash_attn_no_pad(
|
||||
qkv, attn_mask, causal=causal, dropout_p=drop_rate, softmax_scale=None
|
||||
)
|
||||
|
||||
x = flash_attn_no_pad(qkv, attn_mask, causal=causal, dropout_p=drop_rate, softmax_scale=None)
|
||||
|
||||
b, s, a, d = x.shape
|
||||
out = x.reshape(b, s, -1)
|
||||
return out
|
||||
|
||||
|
||||
def parallel_attention(q, k, v, img_q_len, img_kv_len, text_mask):
|
||||
def parallel_attention(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
img_q_len,
|
||||
img_kv_len,
|
||||
text_mask
|
||||
):
|
||||
# 1GPU torch.Size([1, 11264, 24, 128]) tensor([ 0, 11275, 11520], device='cuda:0', dtype=torch.int32)
|
||||
# 2GPU torch.Size([1, 5632, 24, 128]) tensor([ 0, 5643, 5888], device='cuda:0', dtype=torch.int32)
|
||||
query, encoder_query = q
|
||||
key, encoder_key = k
|
||||
value, encoder_value = v
|
||||
if get_sequence_parallel_state():
|
||||
# batch_size, seq_len, attn_heads, head_dim
|
||||
# batch_size, seq_len, attn_heads, head_dim
|
||||
query = all_to_all_4D(query, scatter_dim=2, gather_dim=1)
|
||||
key = all_to_all_4D(key, scatter_dim=2, gather_dim=1)
|
||||
key = all_to_all_4D(key, scatter_dim=2, gather_dim=1)
|
||||
value = all_to_all_4D(value, scatter_dim=2, gather_dim=1)
|
||||
|
||||
|
||||
def shrink_head(encoder_state, dim):
|
||||
local_heads = encoder_state.shape[dim] // nccl_info.sp_size
|
||||
return encoder_state.narrow(
|
||||
dim, nccl_info.rank_within_group * local_heads, local_heads
|
||||
)
|
||||
|
||||
return encoder_state.narrow(dim, nccl_info.rank_within_group * local_heads, local_heads)
|
||||
encoder_query = shrink_head(encoder_query, dim=2)
|
||||
encoder_key = shrink_head(encoder_key, dim=2)
|
||||
encoder_value = shrink_head(encoder_value, dim=2)
|
||||
@@ -61,12 +69,10 @@ def parallel_attention(q, k, v, img_q_len, img_kv_len, text_mask):
|
||||
value = torch.cat([value, encoder_value], dim=1)
|
||||
# B, S, 3, H, D
|
||||
qkv = torch.stack([query, key, value], dim=2)
|
||||
|
||||
attn_mask = F.pad(text_mask, (sequence_length, 0), value=True)
|
||||
hidden_states = flash_attn_no_pad(
|
||||
qkv, attn_mask, causal=False, dropout_p=0.0, softmax_scale=None
|
||||
)
|
||||
|
||||
|
||||
attn_mask = F.pad(text_mask, (sequence_length, 0), value=True)
|
||||
hidden_states = flash_attn_no_pad(qkv, attn_mask, causal=False, dropout_p=0.0, softmax_scale=None)
|
||||
|
||||
hidden_states, encoder_hidden_states = hidden_states.split_with_sizes(
|
||||
(sequence_length, encoder_sequence_length), dim=1
|
||||
)
|
||||
@@ -75,8 +81,9 @@ def parallel_attention(q, k, v, img_q_len, img_kv_len, text_mask):
|
||||
encoder_hidden_states = all_gather(encoder_hidden_states, dim=2).contiguous()
|
||||
hidden_states = hidden_states.to(query.dtype)
|
||||
encoder_hidden_states = encoder_hidden_states.to(query.dtype)
|
||||
|
||||
|
||||
attn = torch.cat([hidden_states, encoder_hidden_states], dim=1)
|
||||
|
||||
|
||||
b, s, a, d = attn.shape
|
||||
attn = attn.reshape(b, s, -1)
|
||||
|
||||
@@ -43,7 +43,7 @@ class PatchEmbed(nn.Module):
|
||||
kernel_size=patch_size,
|
||||
stride=patch_size,
|
||||
bias=bias,
|
||||
**factory_kwargs,
|
||||
**factory_kwargs
|
||||
)
|
||||
nn.init.xavier_uniform_(self.proj.weight.view(self.proj.weight.size(0), -1))
|
||||
if bias:
|
||||
@@ -73,14 +73,14 @@ class TextProjection(nn.Module):
|
||||
in_features=in_channels,
|
||||
out_features=hidden_size,
|
||||
bias=True,
|
||||
**factory_kwargs,
|
||||
**factory_kwargs
|
||||
)
|
||||
self.act_1 = act_layer()
|
||||
self.linear_2 = nn.Linear(
|
||||
in_features=hidden_size,
|
||||
out_features=hidden_size,
|
||||
bias=True,
|
||||
**factory_kwargs,
|
||||
**factory_kwargs
|
||||
)
|
||||
|
||||
def forward(self, caption):
|
||||
|
||||
@@ -59,10 +59,9 @@ class MLP(nn.Module):
|
||||
return x
|
||||
|
||||
|
||||
#
|
||||
#
|
||||
class MLPEmbedder(nn.Module):
|
||||
"""copied from https://github.com/black-forest-labs/flux/blob/main/src/flux/modules/layers.py"""
|
||||
|
||||
def __init__(self, in_dim: int, hidden_dim: int, device=None, dtype=None):
|
||||
factory_kwargs = {"device": device, "dtype": dtype}
|
||||
super().__init__()
|
||||
@@ -92,7 +91,7 @@ class FinalLayer(nn.Module):
|
||||
hidden_size,
|
||||
patch_size * patch_size * out_channels,
|
||||
bias=True,
|
||||
**factory_kwargs,
|
||||
**factory_kwargs
|
||||
)
|
||||
else:
|
||||
self.linear = nn.Linear(
|
||||
|
||||
@@ -11,15 +11,16 @@ from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from .activation_layers import get_activation_layer
|
||||
from .norm_layers import get_norm_layer
|
||||
from .embed_layers import TimestepEmbedder, PatchEmbed, TextProjection
|
||||
from .attenion import parallel_attention
|
||||
from .attenion import parallel_attention
|
||||
from .posemb_layers import apply_rotary_emb
|
||||
from .mlp_layers import MLP, MLPEmbedder, FinalLayer
|
||||
from .modulate_layers import ModulateDiT, modulate, apply_gate
|
||||
from .token_refiner import SingleTokenRefiner
|
||||
from fastvideo.models.hunyuan.modules.posemb_layers import get_nd_rotary_pos_embed
|
||||
|
||||
from fastvideo.utils.parallel_states import nccl_info
|
||||
|
||||
from fastvideo.utils.parallel_states import (
|
||||
nccl_info,
|
||||
)
|
||||
|
||||
class MMDoubleStreamBlock(nn.Module):
|
||||
"""
|
||||
@@ -132,6 +133,7 @@ class MMDoubleStreamBlock(nn.Module):
|
||||
def disable_deterministic(self):
|
||||
self.deterministic = False
|
||||
|
||||
|
||||
def forward(
|
||||
self,
|
||||
img: torch.Tensor,
|
||||
@@ -172,18 +174,15 @@ class MMDoubleStreamBlock(nn.Module):
|
||||
|
||||
# Apply RoPE if needed.
|
||||
if freqs_cis is not None:
|
||||
|
||||
|
||||
def shrink_head(encoder_state, dim):
|
||||
local_heads = encoder_state.shape[dim] // nccl_info.sp_size
|
||||
return encoder_state.narrow(
|
||||
dim, nccl_info.rank_within_group * local_heads, local_heads
|
||||
)
|
||||
|
||||
return encoder_state.narrow(dim, nccl_info.rank_within_group * local_heads, local_heads)
|
||||
freqs_cis = (
|
||||
shrink_head(freqs_cis[0], dim=0),
|
||||
shrink_head(freqs_cis[1], dim=0),
|
||||
shrink_head(freqs_cis[0], dim=0),
|
||||
shrink_head(freqs_cis[1], dim=0)
|
||||
)
|
||||
|
||||
|
||||
img_qq, img_kk = apply_rotary_emb(img_q, img_k, freqs_cis, head_first=False)
|
||||
assert (
|
||||
img_qq.shape == img_q.shape and img_kk.shape == img_k.shape
|
||||
@@ -202,6 +201,7 @@ class MMDoubleStreamBlock(nn.Module):
|
||||
# Apply QK-Norm if needed.
|
||||
txt_q = self.txt_attn_q_norm(txt_q).to(txt_v)
|
||||
txt_k = self.txt_attn_k_norm(txt_k).to(txt_v)
|
||||
|
||||
|
||||
attn = parallel_attention(
|
||||
(img_q, txt_q),
|
||||
@@ -209,9 +209,10 @@ class MMDoubleStreamBlock(nn.Module):
|
||||
(img_v, txt_v),
|
||||
img_q_len=img_q.shape[1],
|
||||
img_kv_len=img_k.shape[1],
|
||||
text_mask=text_mask,
|
||||
text_mask=text_mask
|
||||
)
|
||||
|
||||
|
||||
# attention computation end
|
||||
|
||||
img_attn, txt_attn = attn[:, : img.shape[1]], attn[:, img.shape[1] :]
|
||||
@@ -332,14 +333,15 @@ class MMSingleStreamBlock(nn.Module):
|
||||
q = self.q_norm(q).to(v)
|
||||
k = self.k_norm(k).to(v)
|
||||
|
||||
|
||||
def shrink_head(encoder_state, dim):
|
||||
local_heads = encoder_state.shape[dim] // nccl_info.sp_size
|
||||
return encoder_state.narrow(
|
||||
dim, nccl_info.rank_within_group * local_heads, local_heads
|
||||
)
|
||||
|
||||
freqs_cis = (shrink_head(freqs_cis[0], dim=0), shrink_head(freqs_cis[1], dim=0))
|
||||
|
||||
return encoder_state.narrow(dim, nccl_info.rank_within_group * local_heads, local_heads)
|
||||
freqs_cis = (
|
||||
shrink_head(freqs_cis[0], dim=0),
|
||||
shrink_head(freqs_cis[1], dim=0)
|
||||
)
|
||||
|
||||
img_q, txt_q = q[:, :-txt_len, :, :], q[:, -txt_len:, :, :]
|
||||
img_k, txt_k = k[:, :-txt_len, :, :], k[:, -txt_len:, :, :]
|
||||
img_v, txt_v = v[:, :-txt_len, :, :], v[:, -txt_len:, :, :]
|
||||
@@ -348,6 +350,10 @@ class MMSingleStreamBlock(nn.Module):
|
||||
img_qq.shape == img_q.shape and img_kk.shape == img_k.shape
|
||||
), f"img_kk: {img_qq.shape}, img_q: {img_q.shape}, img_kk: {img_kk.shape}, img_k: {img_k.shape}"
|
||||
img_q, img_k = img_qq, img_kk
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
attn = parallel_attention(
|
||||
(img_q, txt_q),
|
||||
@@ -355,9 +361,10 @@ class MMSingleStreamBlock(nn.Module):
|
||||
(img_v, txt_v),
|
||||
img_q_len=img_q.shape[1],
|
||||
img_kv_len=img_k.shape[1],
|
||||
text_mask=text_mask,
|
||||
text_mask=text_mask
|
||||
)
|
||||
|
||||
|
||||
# attention computation end
|
||||
|
||||
# Compute activation in mlp stream, cat again and run second linear layer.
|
||||
@@ -439,8 +446,8 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
device: Optional[torch.device] = None,
|
||||
text_states_dim: int = 4096,
|
||||
text_states_dim_2: int = 768,
|
||||
rope_theta: int = 256,
|
||||
text_states_dim_2: int = 768,
|
||||
rope_theta:int = 256,
|
||||
):
|
||||
factory_kwargs = {"device": device, "dtype": dtype}
|
||||
super().__init__()
|
||||
@@ -451,12 +458,13 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
||||
self.unpatchify_channels = self.out_channels
|
||||
self.guidance_embed = guidance_embed
|
||||
self.rope_dim_list = rope_dim_list
|
||||
self.rope_theta = rope_theta
|
||||
self.rope_theta = rope_theta
|
||||
# Text projection. Default to linear projection.
|
||||
# Alternative: TokenRefiner. See more details (LI-DiT): http://arxiv.org/abs/2406.11831
|
||||
self.use_attention_mask = use_attention_mask
|
||||
self.text_projection = text_projection
|
||||
|
||||
|
||||
if hidden_size % heads_num != 0:
|
||||
raise ValueError(
|
||||
f"Hidden size {hidden_size} must be divisible by heads_num {heads_num}"
|
||||
@@ -484,11 +492,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
||||
)
|
||||
elif self.text_projection == "single_refiner":
|
||||
self.txt_in = SingleTokenRefiner(
|
||||
self.config.text_states_dim,
|
||||
hidden_size,
|
||||
heads_num,
|
||||
depth=2,
|
||||
**factory_kwargs,
|
||||
self.config.text_states_dim, hidden_size, heads_num, depth=2, **factory_kwargs
|
||||
)
|
||||
else:
|
||||
raise NotImplementedError(
|
||||
@@ -592,29 +596,25 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
||||
# text_states_2: Optional[torch.Tensor] = None, # Text embedding for modulation.
|
||||
# guidance: torch.Tensor = None, # Guidance for modulation, should be cfg_scale x 1000.
|
||||
# return_dict: bool = True,
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
timestep: torch.LongTensor,
|
||||
encoder_attention_mask: torch.Tensor,
|
||||
output_features=False,
|
||||
output_features_stride=8,
|
||||
output_attn=False,
|
||||
attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||
return_dict: bool = False,
|
||||
guidance=None,
|
||||
guidance = None,
|
||||
) -> Union[torch.Tensor, Dict[str, torch.Tensor]]:
|
||||
if guidance == None:
|
||||
guidance = torch.tensor(
|
||||
[6016.0], device=hidden_states.device, dtype=torch.bfloat16
|
||||
)
|
||||
guidance = torch.tensor([6016.], device=hidden_states.device, dtype=torch.bfloat16)
|
||||
out = {}
|
||||
img = x = hidden_states
|
||||
text_mask = encoder_attention_mask
|
||||
t = timestep
|
||||
txt = encoder_hidden_states[:, 1:]
|
||||
text_states_2 = encoder_hidden_states[:, 0, : self.config.text_states_dim_2]
|
||||
text_states_2 = encoder_hidden_states[:, 0, :self.config.text_states_dim_2]
|
||||
_, _, ot, oh, ow = x.shape
|
||||
tt, th, tw = (
|
||||
ot // self.patch_size[0],
|
||||
@@ -653,17 +653,23 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
||||
txt_seq_len = txt.shape[1]
|
||||
img_seq_len = img.shape[1]
|
||||
|
||||
|
||||
|
||||
freqs_cis = (freqs_cos, freqs_sin) if freqs_cos is not None else None
|
||||
# --------------------- Pass through DiT blocks ------------------------
|
||||
for _, block in enumerate(self.double_blocks):
|
||||
double_block_args = [img, txt, vec, freqs_cis, text_mask]
|
||||
double_block_args = [
|
||||
img,
|
||||
txt,
|
||||
vec,
|
||||
freqs_cis,
|
||||
text_mask
|
||||
]
|
||||
|
||||
img, txt = block(*double_block_args)
|
||||
|
||||
# Merge txt and img to pass through single stream blocks.
|
||||
x = torch.cat((img, txt), 1)
|
||||
if output_features:
|
||||
features_list = []
|
||||
if len(self.single_blocks) > 0:
|
||||
for _, block in enumerate(self.single_blocks):
|
||||
single_block_args = [
|
||||
@@ -671,12 +677,10 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
||||
vec,
|
||||
txt_seq_len,
|
||||
(freqs_cos, freqs_sin),
|
||||
text_mask,
|
||||
text_mask
|
||||
]
|
||||
|
||||
x = block(*single_block_args)
|
||||
if output_features and _ % output_features_stride == 0:
|
||||
features_list.append(x[:, :img_seq_len, ...])
|
||||
|
||||
img = x[:, :img_seq_len, ...]
|
||||
|
||||
@@ -684,12 +688,10 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
||||
img = self.final_layer(img, vec) # (N, T, patch_size ** 2 * out_channels)
|
||||
|
||||
img = self.unpatchify(img, tt, th, tw)
|
||||
assert return_dict == False, "return_dict is not supported."
|
||||
if output_features:
|
||||
features_list = torch.stack(features_list, dim=0)
|
||||
else:
|
||||
features_list = None
|
||||
return (img, features_list)
|
||||
if return_dict:
|
||||
out["x"] = img
|
||||
return out
|
||||
return (img, )
|
||||
|
||||
def unpatchify(self, x, t, h, w):
|
||||
"""
|
||||
|
||||
@@ -6,7 +6,6 @@ import torch.nn as nn
|
||||
|
||||
class ModulateDiT(nn.Module):
|
||||
"""Modulation layer for DiT."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int,
|
||||
|
||||
@@ -135,7 +135,10 @@ class IndividualTokenRefiner(nn.Module):
|
||||
)
|
||||
|
||||
def forward(
|
||||
self, x: torch.Tensor, c: torch.LongTensor, mask: Optional[torch.Tensor] = None,
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
c: torch.LongTensor,
|
||||
mask: Optional[torch.Tensor] = None,
|
||||
):
|
||||
mask = mask.clone().bool()
|
||||
# avoid attention weight become NaN
|
||||
@@ -149,7 +152,6 @@ class SingleTokenRefiner(nn.Module):
|
||||
"""
|
||||
A single token refiner block for llm text embedding refine.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_channels,
|
||||
|
||||
@@ -35,7 +35,6 @@ Given Input:
|
||||
input: "{input}"
|
||||
"""
|
||||
|
||||
|
||||
def get_rewrite_prompt(ori_prompt, mode="Normal"):
|
||||
if mode == "Normal":
|
||||
prompt = normal_mode_prompt.format(input=ori_prompt)
|
||||
@@ -45,9 +44,8 @@ def get_rewrite_prompt(ori_prompt, mode="Normal"):
|
||||
raise Exception("Only supports Normal and Normal", mode)
|
||||
return prompt
|
||||
|
||||
|
||||
ori_prompt = "一只小狗在草地上奔跑。"
|
||||
normal_prompt = get_rewrite_prompt(ori_prompt, mode="Normal")
|
||||
master_prompt = get_rewrite_prompt(ori_prompt, mode="Master")
|
||||
|
||||
# Then you can use the normal_prompt or master_prompt to access the hunyuan-large rewrite model to get the final prompt.
|
||||
# Then you can use the normal_prompt or master_prompt to access the hunyuan-large rewrite model to get the final prompt.
|
||||
@@ -44,7 +44,6 @@ def safe_file(path):
|
||||
path.parent.mkdir(exist_ok=True, parents=True)
|
||||
return path
|
||||
|
||||
|
||||
def save_videos_grid(videos: torch.Tensor, path: str, rescale=False, n_rows=1, fps=24):
|
||||
"""save videos by video tensor
|
||||
copy from https://github.com/guoyww/AnimateDiff/blob/e92bd5671ba62c0d774a32951453e328018b7c5b/animatediff/utils/util.py#L61
|
||||
|
||||
@@ -11,7 +11,6 @@ def _ntuple(n):
|
||||
x = tuple(repeat(x[0], n))
|
||||
return x
|
||||
return tuple(repeat(x, n))
|
||||
|
||||
return parse
|
||||
|
||||
|
||||
|
||||
@@ -10,12 +10,17 @@ def preprocess_text_encoder_tokenizer(args):
|
||||
|
||||
processor = AutoProcessor.from_pretrained(args.input_dir)
|
||||
model = LlavaForConditionalGeneration.from_pretrained(
|
||||
args.input_dir, torch_dtype=torch.float16, low_cpu_mem_usage=True,
|
||||
args.input_dir,
|
||||
torch_dtype=torch.float16,
|
||||
low_cpu_mem_usage=True,
|
||||
).to(0)
|
||||
|
||||
model.language_model.save_pretrained(f"{args.output_dir}")
|
||||
processor.tokenizer.save_pretrained(f"{args.output_dir}")
|
||||
|
||||
model.language_model.save_pretrained(
|
||||
f"{args.output_dir}"
|
||||
)
|
||||
processor.tokenizer.save_pretrained(
|
||||
f"{args.output_dir}"
|
||||
)
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
|
||||
@@ -5,15 +5,13 @@ import torch
|
||||
from .autoencoder_kl_causal_3d import AutoencoderKLCausal3D
|
||||
from ..constants import VAE_PATH, PRECISION_TO_TYPE
|
||||
|
||||
|
||||
def load_vae(
|
||||
vae_type: str = "884-16c-hy",
|
||||
vae_precision: str = None,
|
||||
sample_size: tuple = None,
|
||||
vae_path: str = None,
|
||||
logger=None,
|
||||
device=None,
|
||||
):
|
||||
def load_vae(vae_type: str="884-16c-hy",
|
||||
vae_precision: str=None,
|
||||
sample_size: tuple=None,
|
||||
vae_path: str=None,
|
||||
logger=None,
|
||||
device=None
|
||||
):
|
||||
"""the fucntion to load the 3D VAE model
|
||||
|
||||
Args:
|
||||
@@ -26,7 +24,7 @@ def load_vae(
|
||||
"""
|
||||
if vae_path is None:
|
||||
vae_path = VAE_PATH[vae_type]
|
||||
|
||||
|
||||
if logger is not None:
|
||||
logger.info(f"Loading 3D VAE model ({vae_type}) from: {vae_path}")
|
||||
config = AutoencoderKLCausal3D.load_config(vae_path)
|
||||
@@ -34,22 +32,20 @@ def load_vae(
|
||||
vae = AutoencoderKLCausal3D.from_config(config, sample_size=sample_size)
|
||||
else:
|
||||
vae = AutoencoderKLCausal3D.from_config(config)
|
||||
|
||||
|
||||
vae_ckpt = Path(vae_path) / "pytorch_model.pt"
|
||||
assert vae_ckpt.exists(), f"VAE checkpoint not found: {vae_ckpt}"
|
||||
|
||||
|
||||
ckpt = torch.load(vae_ckpt, map_location=vae.device)
|
||||
if "state_dict" in ckpt:
|
||||
ckpt = ckpt["state_dict"]
|
||||
if any(k.startswith("vae.") for k in ckpt.keys()):
|
||||
ckpt = {
|
||||
k.replace("vae.", ""): v for k, v in ckpt.items() if k.startswith("vae.")
|
||||
}
|
||||
ckpt = {k.replace("vae.", ""): v for k, v in ckpt.items() if k.startswith("vae.")}
|
||||
vae.load_state_dict(ckpt)
|
||||
|
||||
spatial_compression_ratio = vae.config.spatial_compression_ratio
|
||||
time_compression_ratio = vae.config.time_compression_ratio
|
||||
|
||||
|
||||
if vae_precision is not None:
|
||||
vae = vae.to(dtype=PRECISION_TO_TYPE[vae_precision])
|
||||
|
||||
|
||||
@@ -29,9 +29,7 @@ try:
|
||||
from diffusers.loaders import FromOriginalVAEMixin
|
||||
except ImportError:
|
||||
# Use this to be compatible with the original diffusers.
|
||||
from diffusers.loaders.single_file_model import (
|
||||
FromOriginalModelMixin as FromOriginalVAEMixin,
|
||||
)
|
||||
from diffusers.loaders.single_file_model import FromOriginalModelMixin as FromOriginalVAEMixin
|
||||
from diffusers.utils.accelerate_utils import apply_forward_hook
|
||||
from diffusers.models.attention_processor import (
|
||||
ADDED_KV_ATTENTION_PROCESSORS,
|
||||
@@ -43,13 +41,7 @@ from diffusers.models.attention_processor import (
|
||||
)
|
||||
from diffusers.models.modeling_outputs import AutoencoderKLOutput
|
||||
from diffusers.models.modeling_utils import ModelMixin
|
||||
from .vae import (
|
||||
DecoderCausal3D,
|
||||
BaseOutput,
|
||||
DecoderOutput,
|
||||
DiagonalGaussianDistribution,
|
||||
EncoderCausal3D,
|
||||
)
|
||||
from .vae import DecoderCausal3D, BaseOutput, DecoderOutput, DiagonalGaussianDistribution, EncoderCausal3D
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -119,12 +111,8 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
mid_block_add_attention=mid_block_add_attention,
|
||||
)
|
||||
|
||||
self.quant_conv = nn.Conv3d(
|
||||
2 * latent_channels, 2 * latent_channels, kernel_size=1
|
||||
)
|
||||
self.post_quant_conv = nn.Conv3d(
|
||||
latent_channels, latent_channels, kernel_size=1
|
||||
)
|
||||
self.quant_conv = nn.Conv3d(2 * latent_channels, 2 * latent_channels, kernel_size=1)
|
||||
self.post_quant_conv = nn.Conv3d(latent_channels, latent_channels, kernel_size=1)
|
||||
|
||||
self.use_slicing = False
|
||||
self.use_spatial_tiling = False
|
||||
@@ -140,9 +128,7 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
if isinstance(self.config.sample_size, (list, tuple))
|
||||
else self.config.sample_size
|
||||
)
|
||||
self.tile_latent_min_size = int(
|
||||
sample_size / (2 ** (len(self.config.block_out_channels) - 1))
|
||||
)
|
||||
self.tile_latent_min_size = int(sample_size / (2 ** (len(self.config.block_out_channels) - 1)))
|
||||
self.tile_overlap_factor = 0.25
|
||||
|
||||
def _set_gradient_checkpointing(self, module, value=False):
|
||||
@@ -203,15 +189,9 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
# set recursively
|
||||
processors = {}
|
||||
|
||||
def fn_recursive_add_processors(
|
||||
name: str,
|
||||
module: torch.nn.Module,
|
||||
processors: Dict[str, AttentionProcessor],
|
||||
):
|
||||
def fn_recursive_add_processors(name: str, module: torch.nn.Module, processors: Dict[str, AttentionProcessor]):
|
||||
if hasattr(module, "get_processor"):
|
||||
processors[f"{name}.processor"] = module.get_processor(
|
||||
return_deprecated_lora=True
|
||||
)
|
||||
processors[f"{name}.processor"] = module.get_processor(return_deprecated_lora=True)
|
||||
|
||||
for sub_name, child in module.named_children():
|
||||
fn_recursive_add_processors(f"{name}.{sub_name}", child, processors)
|
||||
@@ -225,9 +205,7 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
|
||||
# Copied from diffusers.models.unet_2d_condition.UNet2DConditionModel.set_attn_processor
|
||||
def set_attn_processor(
|
||||
self,
|
||||
processor: Union[AttentionProcessor, Dict[str, AttentionProcessor]],
|
||||
_remove_lora=False,
|
||||
self, processor: Union[AttentionProcessor, Dict[str, AttentionProcessor]], _remove_lora=False
|
||||
):
|
||||
r"""
|
||||
Sets the attention processor to use to compute attention.
|
||||
@@ -254,9 +232,7 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
if not isinstance(processor, dict):
|
||||
module.set_processor(processor, _remove_lora=_remove_lora)
|
||||
else:
|
||||
module.set_processor(
|
||||
processor.pop(f"{name}.processor"), _remove_lora=_remove_lora
|
||||
)
|
||||
module.set_processor(processor.pop(f"{name}.processor"), _remove_lora=_remove_lora)
|
||||
|
||||
for sub_name, child in module.named_children():
|
||||
fn_recursive_attn_processor(f"{name}.{sub_name}", child, processor)
|
||||
@@ -269,15 +245,9 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
"""
|
||||
Disables custom attention processors and sets the default attention implementation.
|
||||
"""
|
||||
if all(
|
||||
proc.__class__ in ADDED_KV_ATTENTION_PROCESSORS
|
||||
for proc in self.attn_processors.values()
|
||||
):
|
||||
if all(proc.__class__ in ADDED_KV_ATTENTION_PROCESSORS for proc in self.attn_processors.values()):
|
||||
processor = AttnAddedKVProcessor()
|
||||
elif all(
|
||||
proc.__class__ in CROSS_ATTENTION_PROCESSORS
|
||||
for proc in self.attn_processors.values()
|
||||
):
|
||||
elif all(proc.__class__ in CROSS_ATTENTION_PROCESSORS for proc in self.attn_processors.values()):
|
||||
processor = AttnProcessor()
|
||||
else:
|
||||
raise ValueError(
|
||||
@@ -307,10 +277,7 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
if self.use_temporal_tiling and x.shape[2] > self.tile_sample_min_tsize:
|
||||
return self.temporal_tiled_encode(x, return_dict=return_dict)
|
||||
|
||||
if self.use_spatial_tiling and (
|
||||
x.shape[-1] > self.tile_sample_min_size
|
||||
or x.shape[-2] > self.tile_sample_min_size
|
||||
):
|
||||
if self.use_spatial_tiling and (x.shape[-1] > self.tile_sample_min_size or x.shape[-2] > self.tile_sample_min_size):
|
||||
return self.spatial_tiled_encode(x, return_dict=return_dict)
|
||||
|
||||
if self.use_slicing and x.shape[0] > 1:
|
||||
@@ -327,18 +294,13 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
|
||||
return AutoencoderKLOutput(latent_dist=posterior)
|
||||
|
||||
def _decode(
|
||||
self, z: torch.FloatTensor, return_dict: bool = True
|
||||
) -> Union[DecoderOutput, torch.FloatTensor]:
|
||||
def _decode(self, z: torch.FloatTensor, return_dict: bool = True) -> Union[DecoderOutput, torch.FloatTensor]:
|
||||
assert len(z.shape) == 5, "The input tensor should have 5 dimensions."
|
||||
|
||||
if self.use_temporal_tiling and z.shape[2] > self.tile_latent_min_tsize:
|
||||
return self.temporal_tiled_decode(z, return_dict=return_dict)
|
||||
|
||||
if self.use_spatial_tiling and (
|
||||
z.shape[-1] > self.tile_latent_min_size
|
||||
or z.shape[-2] > self.tile_latent_min_size
|
||||
):
|
||||
if self.use_spatial_tiling and (z.shape[-1] > self.tile_latent_min_size or z.shape[-2] > self.tile_latent_min_size):
|
||||
return self.spatial_tiled_decode(z, return_dict=return_dict)
|
||||
|
||||
z = self.post_quant_conv(z)
|
||||
@@ -378,42 +340,25 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
|
||||
return DecoderOutput(sample=decoded)
|
||||
|
||||
def blend_v(
|
||||
self, a: torch.Tensor, b: torch.Tensor, blend_extent: int
|
||||
) -> torch.Tensor:
|
||||
def blend_v(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor:
|
||||
blend_extent = min(a.shape[-2], b.shape[-2], blend_extent)
|
||||
for y in range(blend_extent):
|
||||
b[:, :, :, y, :] = a[:, :, :, -blend_extent + y, :] * (
|
||||
1 - y / blend_extent
|
||||
) + b[:, :, :, y, :] * (y / blend_extent)
|
||||
b[:, :, :, y, :] = a[:, :, :, -blend_extent + y, :] * (1 - y / blend_extent) + b[:, :, :, y, :] * (y / blend_extent)
|
||||
return b
|
||||
|
||||
def blend_h(
|
||||
self, a: torch.Tensor, b: torch.Tensor, blend_extent: int
|
||||
) -> torch.Tensor:
|
||||
def blend_h(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor:
|
||||
blend_extent = min(a.shape[-1], b.shape[-1], blend_extent)
|
||||
for x in range(blend_extent):
|
||||
b[:, :, :, :, x] = a[:, :, :, :, -blend_extent + x] * (
|
||||
1 - x / blend_extent
|
||||
) + b[:, :, :, :, x] * (x / blend_extent)
|
||||
b[:, :, :, :, x] = a[:, :, :, :, -blend_extent + x] * (1 - x / blend_extent) + b[:, :, :, :, x] * (x / blend_extent)
|
||||
return b
|
||||
|
||||
def blend_t(
|
||||
self, a: torch.Tensor, b: torch.Tensor, blend_extent: int
|
||||
) -> torch.Tensor:
|
||||
def blend_t(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor:
|
||||
blend_extent = min(a.shape[-3], b.shape[-3], blend_extent)
|
||||
for x in range(blend_extent):
|
||||
b[:, :, x, :, :] = a[:, :, -blend_extent + x, :, :] * (
|
||||
1 - x / blend_extent
|
||||
) + b[:, :, x, :, :] * (x / blend_extent)
|
||||
b[:, :, x, :, :] = a[:, :, -blend_extent + x, :, :] * (1 - x / blend_extent) + b[:, :, x, :, :] * (x / blend_extent)
|
||||
return b
|
||||
|
||||
def spatial_tiled_encode(
|
||||
self,
|
||||
x: torch.FloatTensor,
|
||||
return_dict: bool = True,
|
||||
return_moments: bool = False,
|
||||
) -> AutoencoderKLOutput:
|
||||
def spatial_tiled_encode(self, x: torch.FloatTensor, return_dict: bool = True, return_moments: bool = False) -> AutoencoderKLOutput:
|
||||
r"""Encode a batch of images/videos using a tiled encoder.
|
||||
|
||||
When this option is enabled, the VAE will split the input tensor into tiles to compute encoding in several
|
||||
@@ -441,13 +386,7 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
for i in range(0, x.shape[-2], overlap_size):
|
||||
row = []
|
||||
for j in range(0, x.shape[-1], overlap_size):
|
||||
tile = x[
|
||||
:,
|
||||
:,
|
||||
:,
|
||||
i : i + self.tile_sample_min_size,
|
||||
j : j + self.tile_sample_min_size,
|
||||
]
|
||||
tile = x[:, :, :, i: i + self.tile_sample_min_size, j: j + self.tile_sample_min_size]
|
||||
tile = self.encoder(tile)
|
||||
tile = self.quant_conv(tile)
|
||||
row.append(tile)
|
||||
@@ -475,9 +414,7 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
|
||||
return AutoencoderKLOutput(latent_dist=posterior)
|
||||
|
||||
def spatial_tiled_decode(
|
||||
self, z: torch.FloatTensor, return_dict: bool = True
|
||||
) -> Union[DecoderOutput, torch.FloatTensor]:
|
||||
def spatial_tiled_decode(self, z: torch.FloatTensor, return_dict: bool = True) -> Union[DecoderOutput, torch.FloatTensor]:
|
||||
r"""
|
||||
Decode a batch of images/videos using a tiled decoder.
|
||||
|
||||
@@ -501,13 +438,7 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
for i in range(0, z.shape[-2], overlap_size):
|
||||
row = []
|
||||
for j in range(0, z.shape[-1], overlap_size):
|
||||
tile = z[
|
||||
:,
|
||||
:,
|
||||
:,
|
||||
i : i + self.tile_latent_min_size,
|
||||
j : j + self.tile_latent_min_size,
|
||||
]
|
||||
tile = z[:, :, :, i: i + self.tile_latent_min_size, j: j + self.tile_latent_min_size]
|
||||
tile = self.post_quant_conv(tile)
|
||||
decoded = self.decoder(tile)
|
||||
row.append(decoded)
|
||||
@@ -531,9 +462,7 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
|
||||
return DecoderOutput(sample=dec)
|
||||
|
||||
def temporal_tiled_encode(
|
||||
self, x: torch.FloatTensor, return_dict: bool = True
|
||||
) -> AutoencoderKLOutput:
|
||||
def temporal_tiled_encode(self, x: torch.FloatTensor, return_dict: bool = True) -> AutoencoderKLOutput:
|
||||
|
||||
B, C, T, H, W = x.shape
|
||||
overlap_size = int(self.tile_sample_min_tsize * (1 - self.tile_overlap_factor))
|
||||
@@ -543,11 +472,8 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
# Split the video into tiles and encode them separately.
|
||||
row = []
|
||||
for i in range(0, T, overlap_size):
|
||||
tile = x[:, :, i : i + self.tile_sample_min_tsize + 1, :, :]
|
||||
if self.use_spatial_tiling and (
|
||||
tile.shape[-1] > self.tile_sample_min_size
|
||||
or tile.shape[-2] > self.tile_sample_min_size
|
||||
):
|
||||
tile = x[:, :, i: i + self.tile_sample_min_tsize + 1, :, :]
|
||||
if self.use_spatial_tiling and (tile.shape[-1] > self.tile_sample_min_size or tile.shape[-2] > self.tile_sample_min_size):
|
||||
tile = self.spatial_tiled_encode(tile, return_moments=True)
|
||||
else:
|
||||
tile = self.encoder(tile)
|
||||
@@ -561,7 +487,7 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
tile = self.blend_t(row[i - 1], tile, blend_extent)
|
||||
result_row.append(tile[:, :, :t_limit, :, :])
|
||||
else:
|
||||
result_row.append(tile[:, :, : t_limit + 1, :, :])
|
||||
result_row.append(tile[:, :, :t_limit + 1, :, :])
|
||||
|
||||
moments = torch.cat(result_row, dim=2)
|
||||
posterior = DiagonalGaussianDistribution(moments)
|
||||
@@ -571,9 +497,7 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
|
||||
return AutoencoderKLOutput(latent_dist=posterior)
|
||||
|
||||
def temporal_tiled_decode(
|
||||
self, z: torch.FloatTensor, return_dict: bool = True
|
||||
) -> Union[DecoderOutput, torch.FloatTensor]:
|
||||
def temporal_tiled_decode(self, z: torch.FloatTensor, return_dict: bool = True) -> Union[DecoderOutput, torch.FloatTensor]:
|
||||
# Split z into overlapping tiles and decode them separately.
|
||||
|
||||
B, C, T, H, W = z.shape
|
||||
@@ -583,11 +507,8 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
|
||||
row = []
|
||||
for i in range(0, T, overlap_size):
|
||||
tile = z[:, :, i : i + self.tile_latent_min_tsize + 1, :, :]
|
||||
if self.use_spatial_tiling and (
|
||||
tile.shape[-1] > self.tile_latent_min_size
|
||||
or tile.shape[-2] > self.tile_latent_min_size
|
||||
):
|
||||
tile = z[:, :, i: i + self.tile_latent_min_tsize + 1, :, :]
|
||||
if self.use_spatial_tiling and (tile.shape[-1] > self.tile_latent_min_size or tile.shape[-2] > self.tile_latent_min_size):
|
||||
decoded = self.spatial_tiled_decode(tile, return_dict=True).sample
|
||||
else:
|
||||
tile = self.post_quant_conv(tile)
|
||||
@@ -601,7 +522,7 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
tile = self.blend_t(row[i - 1], tile, blend_extent)
|
||||
result_row.append(tile[:, :, :t_limit, :, :])
|
||||
else:
|
||||
result_row.append(tile[:, :, : t_limit + 1, :, :])
|
||||
result_row.append(tile[:, :, :t_limit + 1, :, :])
|
||||
|
||||
dec = torch.cat(result_row, dim=2)
|
||||
if not return_dict:
|
||||
@@ -659,9 +580,7 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
|
||||
for _, attn_processor in self.attn_processors.items():
|
||||
if "Added" in str(attn_processor.__class__.__name__):
|
||||
raise ValueError(
|
||||
"`fuse_qkv_projections()` is not supported for models having added KV projections."
|
||||
)
|
||||
raise ValueError("`fuse_qkv_projections()` is not supported for models having added KV projections.")
|
||||
|
||||
self.original_attn_processors = self.attn_processors
|
||||
|
||||
|
||||
@@ -34,9 +34,7 @@ from diffusers.models.normalization import RMSNorm
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
|
||||
def prepare_causal_attention_mask(
|
||||
n_frame: int, n_hw: int, dtype, device, batch_size: int = None
|
||||
):
|
||||
def prepare_causal_attention_mask(n_frame: int, n_hw: int, dtype, device, batch_size: int = None):
|
||||
seq_len = n_frame * n_hw
|
||||
mask = torch.full((seq_len, seq_len), float("-inf"), dtype=dtype, device=device)
|
||||
for i in range(seq_len):
|
||||
@@ -60,25 +58,16 @@ class CausalConv3d(nn.Module):
|
||||
kernel_size: Union[int, Tuple[int, int, int]],
|
||||
stride: Union[int, Tuple[int, int, int]] = 1,
|
||||
dilation: Union[int, Tuple[int, int, int]] = 1,
|
||||
pad_mode="replicate",
|
||||
**kwargs,
|
||||
pad_mode='replicate',
|
||||
**kwargs
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.pad_mode = pad_mode
|
||||
padding = (
|
||||
kernel_size // 2,
|
||||
kernel_size // 2,
|
||||
kernel_size // 2,
|
||||
kernel_size // 2,
|
||||
kernel_size - 1,
|
||||
0,
|
||||
) # W, H, T
|
||||
padding = (kernel_size // 2, kernel_size // 2, kernel_size // 2, kernel_size // 2, kernel_size - 1, 0) # W, H, T
|
||||
self.time_causal_padding = padding
|
||||
|
||||
self.conv = nn.Conv3d(
|
||||
chan_in, chan_out, kernel_size, stride=stride, dilation=dilation, **kwargs
|
||||
)
|
||||
self.conv = nn.Conv3d(chan_in, chan_out, kernel_size, stride=stride, dilation=dilation, **kwargs)
|
||||
|
||||
def forward(self, x):
|
||||
x = F.pad(x, self.time_causal_padding, mode=self.pad_mode)
|
||||
@@ -130,9 +119,7 @@ class UpsampleCausal3D(nn.Module):
|
||||
elif use_conv:
|
||||
if kernel_size is None:
|
||||
kernel_size = 3
|
||||
conv = CausalConv3d(
|
||||
self.channels, self.out_channels, kernel_size=kernel_size, bias=bias
|
||||
)
|
||||
conv = CausalConv3d(self.channels, self.out_channels, kernel_size=kernel_size, bias=bias)
|
||||
|
||||
if name == "conv":
|
||||
self.conv = conv
|
||||
@@ -169,14 +156,10 @@ class UpsampleCausal3D(nn.Module):
|
||||
first_h, other_h = hidden_states.split((1, T - 1), dim=2)
|
||||
if output_size is None:
|
||||
if T > 1:
|
||||
other_h = F.interpolate(
|
||||
other_h, scale_factor=self.upsample_factor, mode="nearest"
|
||||
)
|
||||
other_h = F.interpolate(other_h, scale_factor=self.upsample_factor, mode="nearest")
|
||||
|
||||
first_h = first_h.squeeze(2)
|
||||
first_h = F.interpolate(
|
||||
first_h, scale_factor=self.upsample_factor[1:], mode="nearest"
|
||||
)
|
||||
first_h = F.interpolate(first_h, scale_factor=self.upsample_factor[1:], mode="nearest")
|
||||
first_h = first_h.unsqueeze(2)
|
||||
else:
|
||||
raise NotImplementedError
|
||||
@@ -237,11 +220,7 @@ class DownsampleCausal3D(nn.Module):
|
||||
|
||||
if use_conv:
|
||||
conv = CausalConv3d(
|
||||
self.channels,
|
||||
self.out_channels,
|
||||
kernel_size=kernel_size,
|
||||
stride=stride,
|
||||
bias=bias,
|
||||
self.channels, self.out_channels, kernel_size=kernel_size, stride=stride, bias=bias
|
||||
)
|
||||
else:
|
||||
raise NotImplementedError
|
||||
@@ -254,15 +233,11 @@ class DownsampleCausal3D(nn.Module):
|
||||
else:
|
||||
self.conv = conv
|
||||
|
||||
def forward(
|
||||
self, hidden_states: torch.FloatTensor, scale: float = 1.0
|
||||
) -> torch.FloatTensor:
|
||||
def forward(self, hidden_states: torch.FloatTensor, scale: float = 1.0) -> torch.FloatTensor:
|
||||
assert hidden_states.shape[1] == self.channels
|
||||
|
||||
if self.norm is not None:
|
||||
hidden_states = self.norm(hidden_states.permute(0, 2, 3, 1)).permute(
|
||||
0, 3, 1, 2
|
||||
)
|
||||
hidden_states = self.norm(hidden_states.permute(0, 2, 3, 1)).permute(0, 3, 1, 2)
|
||||
|
||||
assert hidden_states.shape[1] == self.channels
|
||||
|
||||
@@ -323,9 +298,7 @@ class ResnetBlockCausal3D(nn.Module):
|
||||
elif self.time_embedding_norm == "spatial":
|
||||
self.norm1 = SpatialNorm(in_channels, temb_channels)
|
||||
else:
|
||||
self.norm1 = torch.nn.GroupNorm(
|
||||
num_groups=groups, num_channels=in_channels, eps=eps, affine=True
|
||||
)
|
||||
self.norm1 = torch.nn.GroupNorm(num_groups=groups, num_channels=in_channels, eps=eps, affine=True)
|
||||
|
||||
self.conv1 = CausalConv3d(in_channels, out_channels, kernel_size=3, stride=1)
|
||||
|
||||
@@ -334,15 +307,10 @@ class ResnetBlockCausal3D(nn.Module):
|
||||
self.time_emb_proj = linear_cls(temb_channels, out_channels)
|
||||
elif self.time_embedding_norm == "scale_shift":
|
||||
self.time_emb_proj = linear_cls(temb_channels, 2 * out_channels)
|
||||
elif (
|
||||
self.time_embedding_norm == "ada_group"
|
||||
or self.time_embedding_norm == "spatial"
|
||||
):
|
||||
elif self.time_embedding_norm == "ada_group" or self.time_embedding_norm == "spatial":
|
||||
self.time_emb_proj = None
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Unknown time_embedding_norm : {self.time_embedding_norm} "
|
||||
)
|
||||
raise ValueError(f"Unknown time_embedding_norm : {self.time_embedding_norm} ")
|
||||
else:
|
||||
self.time_emb_proj = None
|
||||
|
||||
@@ -351,15 +319,11 @@ class ResnetBlockCausal3D(nn.Module):
|
||||
elif self.time_embedding_norm == "spatial":
|
||||
self.norm2 = SpatialNorm(out_channels, temb_channels)
|
||||
else:
|
||||
self.norm2 = torch.nn.GroupNorm(
|
||||
num_groups=groups_out, num_channels=out_channels, eps=eps, affine=True
|
||||
)
|
||||
self.norm2 = torch.nn.GroupNorm(num_groups=groups_out, num_channels=out_channels, eps=eps, affine=True)
|
||||
|
||||
self.dropout = torch.nn.Dropout(dropout)
|
||||
conv_3d_out_channels = conv_3d_out_channels or out_channels
|
||||
self.conv2 = CausalConv3d(
|
||||
out_channels, conv_3d_out_channels, kernel_size=3, stride=1
|
||||
)
|
||||
self.conv2 = CausalConv3d(out_channels, conv_3d_out_channels, kernel_size=3, stride=1)
|
||||
|
||||
self.nonlinearity = get_activation(non_linearity)
|
||||
|
||||
@@ -369,11 +333,7 @@ class ResnetBlockCausal3D(nn.Module):
|
||||
elif self.down:
|
||||
self.downsample = DownsampleCausal3D(in_channels, use_conv=False, name="op")
|
||||
|
||||
self.use_in_shortcut = (
|
||||
self.in_channels != conv_3d_out_channels
|
||||
if use_in_shortcut is None
|
||||
else use_in_shortcut
|
||||
)
|
||||
self.use_in_shortcut = self.in_channels != conv_3d_out_channels if use_in_shortcut is None else use_in_shortcut
|
||||
|
||||
self.conv_shortcut = None
|
||||
if self.use_in_shortcut:
|
||||
@@ -393,10 +353,7 @@ class ResnetBlockCausal3D(nn.Module):
|
||||
) -> torch.FloatTensor:
|
||||
hidden_states = input_tensor
|
||||
|
||||
if (
|
||||
self.time_embedding_norm == "ada_group"
|
||||
or self.time_embedding_norm == "spatial"
|
||||
):
|
||||
if self.time_embedding_norm == "ada_group" or self.time_embedding_norm == "spatial":
|
||||
hidden_states = self.norm1(hidden_states, temb)
|
||||
else:
|
||||
hidden_states = self.norm1(hidden_states)
|
||||
@@ -408,26 +365,33 @@ class ResnetBlockCausal3D(nn.Module):
|
||||
if hidden_states.shape[0] >= 64:
|
||||
input_tensor = input_tensor.contiguous()
|
||||
hidden_states = hidden_states.contiguous()
|
||||
input_tensor = self.upsample(input_tensor, scale=scale)
|
||||
hidden_states = self.upsample(hidden_states, scale=scale)
|
||||
input_tensor = (
|
||||
self.upsample(input_tensor, scale=scale)
|
||||
)
|
||||
hidden_states = (
|
||||
self.upsample(hidden_states, scale=scale)
|
||||
)
|
||||
elif self.downsample is not None:
|
||||
input_tensor = self.downsample(input_tensor, scale=scale)
|
||||
hidden_states = self.downsample(hidden_states, scale=scale)
|
||||
input_tensor = (
|
||||
self.downsample(input_tensor, scale=scale)
|
||||
)
|
||||
hidden_states = (
|
||||
self.downsample(hidden_states, scale=scale)
|
||||
)
|
||||
|
||||
hidden_states = self.conv1(hidden_states)
|
||||
|
||||
if self.time_emb_proj is not None:
|
||||
if not self.skip_time_act:
|
||||
temb = self.nonlinearity(temb)
|
||||
temb = self.time_emb_proj(temb, scale)[:, :, None, None]
|
||||
temb = (
|
||||
self.time_emb_proj(temb, scale)[:, :, None, None]
|
||||
)
|
||||
|
||||
if temb is not None and self.time_embedding_norm == "default":
|
||||
hidden_states = hidden_states + temb
|
||||
|
||||
if (
|
||||
self.time_embedding_norm == "ada_group"
|
||||
or self.time_embedding_norm == "spatial"
|
||||
):
|
||||
if self.time_embedding_norm == "ada_group" or self.time_embedding_norm == "spatial":
|
||||
hidden_states = self.norm2(hidden_states, temb)
|
||||
else:
|
||||
hidden_states = self.norm2(hidden_states)
|
||||
@@ -442,7 +406,9 @@ class ResnetBlockCausal3D(nn.Module):
|
||||
hidden_states = self.conv2(hidden_states)
|
||||
|
||||
if self.conv_shortcut is not None:
|
||||
input_tensor = self.conv_shortcut(input_tensor)
|
||||
input_tensor = (
|
||||
self.conv_shortcut(input_tensor)
|
||||
)
|
||||
|
||||
output_tensor = (input_tensor + hidden_states) / self.output_scale_factor
|
||||
|
||||
@@ -484,11 +450,7 @@ def get_down_block3d(
|
||||
)
|
||||
attention_head_dim = num_attention_heads
|
||||
|
||||
down_block_type = (
|
||||
down_block_type[7:]
|
||||
if down_block_type.startswith("UNetRes")
|
||||
else down_block_type
|
||||
)
|
||||
down_block_type = down_block_type[7:] if down_block_type.startswith("UNetRes") else down_block_type
|
||||
if down_block_type == "DownEncoderBlockCausal3D":
|
||||
return DownEncoderBlockCausal3D(
|
||||
num_layers=num_layers,
|
||||
@@ -542,9 +504,7 @@ def get_up_block3d(
|
||||
)
|
||||
attention_head_dim = num_attention_heads
|
||||
|
||||
up_block_type = (
|
||||
up_block_type[7:] if up_block_type.startswith("UNetRes") else up_block_type
|
||||
)
|
||||
up_block_type = up_block_type[7:] if up_block_type.startswith("UNetRes") else up_block_type
|
||||
if up_block_type == "UpDecoderBlockCausal3D":
|
||||
return UpDecoderBlockCausal3D(
|
||||
num_layers=num_layers,
|
||||
@@ -585,15 +545,11 @@ class UNetMidBlockCausal3D(nn.Module):
|
||||
output_scale_factor: float = 1.0,
|
||||
):
|
||||
super().__init__()
|
||||
resnet_groups = (
|
||||
resnet_groups if resnet_groups is not None else min(in_channels // 4, 32)
|
||||
)
|
||||
resnet_groups = resnet_groups if resnet_groups is not None else min(in_channels // 4, 32)
|
||||
self.add_attention = add_attention
|
||||
|
||||
if attn_groups is None:
|
||||
attn_groups = (
|
||||
resnet_groups if resnet_time_scale_shift == "default" else None
|
||||
)
|
||||
attn_groups = resnet_groups if resnet_time_scale_shift == "default" else None
|
||||
|
||||
# there is always at least one resnet
|
||||
resnets = [
|
||||
@@ -628,11 +584,7 @@ class UNetMidBlockCausal3D(nn.Module):
|
||||
rescale_output_factor=output_scale_factor,
|
||||
eps=resnet_eps,
|
||||
norm_num_groups=attn_groups,
|
||||
spatial_norm_dim=(
|
||||
temb_channels
|
||||
if resnet_time_scale_shift == "spatial"
|
||||
else None
|
||||
),
|
||||
spatial_norm_dim=temb_channels if resnet_time_scale_shift == "spatial" else None,
|
||||
residual_connection=True,
|
||||
bias=True,
|
||||
upcast_softmax=True,
|
||||
@@ -660,9 +612,7 @@ class UNetMidBlockCausal3D(nn.Module):
|
||||
self.attentions = nn.ModuleList(attentions)
|
||||
self.resnets = nn.ModuleList(resnets)
|
||||
|
||||
def forward(
|
||||
self, hidden_states: torch.FloatTensor, temb: Optional[torch.FloatTensor] = None
|
||||
) -> torch.FloatTensor:
|
||||
def forward(self, hidden_states: torch.FloatTensor, temb: Optional[torch.FloatTensor] = None) -> torch.FloatTensor:
|
||||
hidden_states = self.resnets[0](hidden_states, temb)
|
||||
for attn, resnet in zip(self.attentions, self.resnets[1:]):
|
||||
if attn is not None:
|
||||
@@ -671,12 +621,8 @@ class UNetMidBlockCausal3D(nn.Module):
|
||||
attention_mask = prepare_causal_attention_mask(
|
||||
T, H * W, hidden_states.dtype, hidden_states.device, batch_size=B
|
||||
)
|
||||
hidden_states = attn(
|
||||
hidden_states, temb=temb, attention_mask=attention_mask
|
||||
)
|
||||
hidden_states = rearrange(
|
||||
hidden_states, "b (f h w) c -> b c f h w", f=T, h=H, w=W
|
||||
)
|
||||
hidden_states = attn(hidden_states, temb=temb, attention_mask=attention_mask)
|
||||
hidden_states = rearrange(hidden_states, "b (f h w) c -> b c f h w", f=T, h=H, w=W)
|
||||
hidden_states = resnet(hidden_states, temb)
|
||||
|
||||
return hidden_states
|
||||
@@ -737,9 +683,7 @@ class DownEncoderBlockCausal3D(nn.Module):
|
||||
else:
|
||||
self.downsamplers = None
|
||||
|
||||
def forward(
|
||||
self, hidden_states: torch.FloatTensor, scale: float = 1.0
|
||||
) -> torch.FloatTensor:
|
||||
def forward(self, hidden_states: torch.FloatTensor, scale: float = 1.0) -> torch.FloatTensor:
|
||||
for resnet in self.resnets:
|
||||
hidden_states = resnet(hidden_states, temb=None, scale=scale)
|
||||
|
||||
@@ -808,10 +752,7 @@ class UpDecoderBlockCausal3D(nn.Module):
|
||||
self.resolution_idx = resolution_idx
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.FloatTensor,
|
||||
temb: Optional[torch.FloatTensor] = None,
|
||||
scale: float = 1.0,
|
||||
self, hidden_states: torch.FloatTensor, temb: Optional[torch.FloatTensor] = None, scale: float = 1.0
|
||||
) -> torch.FloatTensor:
|
||||
for resnet in self.resnets:
|
||||
hidden_states = resnet(hidden_states, temb=temb, scale=scale)
|
||||
|
||||
@@ -51,9 +51,7 @@ class EncoderCausal3D(nn.Module):
|
||||
super().__init__()
|
||||
self.layers_per_block = layers_per_block
|
||||
|
||||
self.conv_in = CausalConv3d(
|
||||
in_channels, block_out_channels[0], kernel_size=3, stride=1
|
||||
)
|
||||
self.conv_in = CausalConv3d(in_channels, block_out_channels[0], kernel_size=3, stride=1)
|
||||
self.mid_block = None
|
||||
self.down_blocks = nn.ModuleList([])
|
||||
|
||||
@@ -73,9 +71,7 @@ class EncoderCausal3D(nn.Module):
|
||||
and not is_final_block
|
||||
)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Unsupported time_compression_ratio: {time_compression_ratio}."
|
||||
)
|
||||
raise ValueError(f"Unsupported time_compression_ratio: {time_compression_ratio}.")
|
||||
|
||||
downsample_stride_HW = (2, 2) if add_spatial_downsample else (1, 1)
|
||||
downsample_stride_T = (2,) if add_time_downsample else (1,)
|
||||
@@ -110,15 +106,11 @@ class EncoderCausal3D(nn.Module):
|
||||
)
|
||||
|
||||
# out
|
||||
self.conv_norm_out = nn.GroupNorm(
|
||||
num_channels=block_out_channels[-1], num_groups=norm_num_groups, eps=1e-6
|
||||
)
|
||||
self.conv_norm_out = nn.GroupNorm(num_channels=block_out_channels[-1], num_groups=norm_num_groups, eps=1e-6)
|
||||
self.conv_act = nn.SiLU()
|
||||
|
||||
conv_out_channels = 2 * out_channels if double_z else out_channels
|
||||
self.conv_out = CausalConv3d(
|
||||
block_out_channels[-1], conv_out_channels, kernel_size=3
|
||||
)
|
||||
self.conv_out = CausalConv3d(block_out_channels[-1], conv_out_channels, kernel_size=3)
|
||||
|
||||
def forward(self, sample: torch.FloatTensor) -> torch.FloatTensor:
|
||||
r"""The forward method of the `EncoderCausal3D` class."""
|
||||
@@ -163,9 +155,7 @@ class DecoderCausal3D(nn.Module):
|
||||
super().__init__()
|
||||
self.layers_per_block = layers_per_block
|
||||
|
||||
self.conv_in = CausalConv3d(
|
||||
in_channels, block_out_channels[-1], kernel_size=3, stride=1
|
||||
)
|
||||
self.conv_in = CausalConv3d(in_channels, block_out_channels[-1], kernel_size=3, stride=1)
|
||||
self.mid_block = None
|
||||
self.up_blocks = nn.ModuleList([])
|
||||
|
||||
@@ -201,15 +191,11 @@ class DecoderCausal3D(nn.Module):
|
||||
and not is_final_block
|
||||
)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Unsupported time_compression_ratio: {time_compression_ratio}."
|
||||
)
|
||||
raise ValueError(f"Unsupported time_compression_ratio: {time_compression_ratio}.")
|
||||
|
||||
upsample_scale_factor_HW = (2, 2) if add_spatial_upsample else (1, 1)
|
||||
upsample_scale_factor_T = (2,) if add_time_upsample else (1,)
|
||||
upsample_scale_factor = tuple(
|
||||
upsample_scale_factor_T + upsample_scale_factor_HW
|
||||
)
|
||||
upsample_scale_factor = tuple(upsample_scale_factor_T + upsample_scale_factor_HW)
|
||||
up_block = get_up_block3d(
|
||||
up_block_type,
|
||||
num_layers=self.layers_per_block + 1,
|
||||
@@ -232,9 +218,7 @@ class DecoderCausal3D(nn.Module):
|
||||
if norm_type == "spatial":
|
||||
self.conv_norm_out = SpatialNorm(block_out_channels[0], temb_channels)
|
||||
else:
|
||||
self.conv_norm_out = nn.GroupNorm(
|
||||
num_channels=block_out_channels[0], num_groups=norm_num_groups, eps=1e-6
|
||||
)
|
||||
self.conv_norm_out = nn.GroupNorm(num_channels=block_out_channels[0], num_groups=norm_num_groups, eps=1e-6)
|
||||
self.conv_act = nn.SiLU()
|
||||
self.conv_out = CausalConv3d(block_out_channels[0], out_channels, kernel_size=3)
|
||||
|
||||
@@ -286,9 +270,7 @@ class DecoderCausal3D(nn.Module):
|
||||
|
||||
# up
|
||||
for up_block in self.up_blocks:
|
||||
sample = torch.utils.checkpoint.checkpoint(
|
||||
create_custom_forward(up_block), sample, latent_embeds
|
||||
)
|
||||
sample = torch.utils.checkpoint.checkpoint(create_custom_forward(up_block), sample, latent_embeds)
|
||||
else:
|
||||
# middle
|
||||
sample = self.mid_block(sample, latent_embeds)
|
||||
@@ -359,14 +341,13 @@ class DiagonalGaussianDistribution(object):
|
||||
dim=reduce_dim,
|
||||
)
|
||||
|
||||
def nll(
|
||||
self, sample: torch.Tensor, dims: Tuple[int, ...] = [1, 2, 3]
|
||||
) -> torch.Tensor:
|
||||
def nll(self, sample: torch.Tensor, dims: Tuple[int, ...] = [1, 2, 3]) -> torch.Tensor:
|
||||
if self.deterministic:
|
||||
return torch.Tensor([0.0])
|
||||
logtwopi = np.log(2.0 * np.pi)
|
||||
return 0.5 * torch.sum(
|
||||
logtwopi + self.logvar + torch.pow(sample - self.mean, 2) / self.var,
|
||||
logtwopi + self.logvar +
|
||||
torch.pow(sample - self.mean, 2) / self.var,
|
||||
dim=dims,
|
||||
)
|
||||
|
||||
|
||||
@@ -5,80 +5,43 @@ import os
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--diffusers_path", required=True, type=str)
|
||||
parser.add_argument(
|
||||
"--transformer_path", type=str, default=None, help="Path to save transformer model"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--vae_encoder_path", type=str, default=None, help="Path to save VAE encoder model"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--vae_decoder_path", type=str, default=None, help="Path to save VAE decoder model"
|
||||
)
|
||||
parser.add_argument("--transformer_path", type=str, default=None, help="Path to save transformer model")
|
||||
parser.add_argument("--vae_encoder_path", type=str, default=None, help="Path to save VAE encoder model")
|
||||
parser.add_argument("--vae_decoder_path", type=str, default=None, help="Path to save VAE decoder model")
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
|
||||
def reverse_scale_shift(weight, dim):
|
||||
scale, shift = weight.chunk(2, dim=0)
|
||||
new_weight = torch.cat([shift, scale], dim=0)
|
||||
return new_weight
|
||||
|
||||
|
||||
def reverse_proj_gate(weight):
|
||||
gate, proj = weight.chunk(2, dim=0)
|
||||
new_weight = torch.cat([proj, gate], dim=0)
|
||||
return new_weight
|
||||
|
||||
|
||||
def convert_diffusers_transformer_to_mochi(state_dict):
|
||||
original_state_dict = state_dict.copy()
|
||||
new_state_dict = {}
|
||||
|
||||
# Convert patch_embed
|
||||
new_state_dict["x_embedder.proj.weight"] = original_state_dict.pop(
|
||||
"patch_embed.proj.weight"
|
||||
)
|
||||
new_state_dict["x_embedder.proj.bias"] = original_state_dict.pop(
|
||||
"patch_embed.proj.bias"
|
||||
)
|
||||
new_state_dict["x_embedder.proj.weight"] = original_state_dict.pop("patch_embed.proj.weight")
|
||||
new_state_dict["x_embedder.proj.bias"] = original_state_dict.pop("patch_embed.proj.bias")
|
||||
|
||||
# Convert time_embed
|
||||
new_state_dict["t_embedder.mlp.0.weight"] = original_state_dict.pop(
|
||||
"time_embed.timestep_embedder.linear_1.weight"
|
||||
)
|
||||
new_state_dict["t_embedder.mlp.0.bias"] = original_state_dict.pop(
|
||||
"time_embed.timestep_embedder.linear_1.bias"
|
||||
)
|
||||
new_state_dict["t_embedder.mlp.2.weight"] = original_state_dict.pop(
|
||||
"time_embed.timestep_embedder.linear_2.weight"
|
||||
)
|
||||
new_state_dict["t_embedder.mlp.2.bias"] = original_state_dict.pop(
|
||||
"time_embed.timestep_embedder.linear_2.bias"
|
||||
)
|
||||
new_state_dict["t5_y_embedder.to_kv.weight"] = original_state_dict.pop(
|
||||
"time_embed.pooler.to_kv.weight"
|
||||
)
|
||||
new_state_dict["t5_y_embedder.to_kv.bias"] = original_state_dict.pop(
|
||||
"time_embed.pooler.to_kv.bias"
|
||||
)
|
||||
new_state_dict["t5_y_embedder.to_q.weight"] = original_state_dict.pop(
|
||||
"time_embed.pooler.to_q.weight"
|
||||
)
|
||||
new_state_dict["t5_y_embedder.to_q.bias"] = original_state_dict.pop(
|
||||
"time_embed.pooler.to_q.bias"
|
||||
)
|
||||
new_state_dict["t5_y_embedder.to_out.weight"] = original_state_dict.pop(
|
||||
"time_embed.pooler.to_out.weight"
|
||||
)
|
||||
new_state_dict["t5_y_embedder.to_out.bias"] = original_state_dict.pop(
|
||||
"time_embed.pooler.to_out.bias"
|
||||
)
|
||||
new_state_dict["t5_yproj.weight"] = original_state_dict.pop(
|
||||
"time_embed.caption_proj.weight"
|
||||
)
|
||||
new_state_dict["t5_yproj.bias"] = original_state_dict.pop(
|
||||
"time_embed.caption_proj.bias"
|
||||
)
|
||||
new_state_dict["t_embedder.mlp.0.weight"] = original_state_dict.pop("time_embed.timestep_embedder.linear_1.weight")
|
||||
new_state_dict["t_embedder.mlp.0.bias"] = original_state_dict.pop("time_embed.timestep_embedder.linear_1.bias")
|
||||
new_state_dict["t_embedder.mlp.2.weight"] = original_state_dict.pop("time_embed.timestep_embedder.linear_2.weight")
|
||||
new_state_dict["t_embedder.mlp.2.bias"] = original_state_dict.pop("time_embed.timestep_embedder.linear_2.bias")
|
||||
new_state_dict["t5_y_embedder.to_kv.weight"] = original_state_dict.pop("time_embed.pooler.to_kv.weight")
|
||||
new_state_dict["t5_y_embedder.to_kv.bias"] = original_state_dict.pop("time_embed.pooler.to_kv.bias")
|
||||
new_state_dict["t5_y_embedder.to_q.weight"] = original_state_dict.pop("time_embed.pooler.to_q.weight")
|
||||
new_state_dict["t5_y_embedder.to_q.bias"] = original_state_dict.pop("time_embed.pooler.to_q.bias")
|
||||
new_state_dict["t5_y_embedder.to_out.weight"] = original_state_dict.pop("time_embed.pooler.to_out.weight")
|
||||
new_state_dict["t5_y_embedder.to_out.bias"] = original_state_dict.pop("time_embed.pooler.to_out.bias")
|
||||
new_state_dict["t5_yproj.weight"] = original_state_dict.pop("time_embed.caption_proj.weight")
|
||||
new_state_dict["t5_yproj.bias"] = original_state_dict.pop("time_embed.caption_proj.bias")
|
||||
|
||||
# Convert transformer blocks
|
||||
num_layers = 48
|
||||
@@ -87,12 +50,8 @@ def convert_diffusers_transformer_to_mochi(state_dict):
|
||||
new_prefix = f"blocks.{i}."
|
||||
|
||||
# norm1
|
||||
new_state_dict[new_prefix + "mod_x.weight"] = original_state_dict.pop(
|
||||
block_prefix + "norm1.linear.weight"
|
||||
)
|
||||
new_state_dict[new_prefix + "mod_x.bias"] = original_state_dict.pop(
|
||||
block_prefix + "norm1.linear.bias"
|
||||
)
|
||||
new_state_dict[new_prefix + "mod_x.weight"] = original_state_dict.pop(block_prefix + "norm1.linear.weight")
|
||||
new_state_dict[new_prefix + "mod_x.bias"] = original_state_dict.pop(block_prefix + "norm1.linear.bias")
|
||||
|
||||
if i < num_layers - 1:
|
||||
new_state_dict[new_prefix + "mod_y.weight"] = original_state_dict.pop(
|
||||
@@ -111,7 +70,7 @@ def convert_diffusers_transformer_to_mochi(state_dict):
|
||||
|
||||
# Visual attention
|
||||
q = original_state_dict.pop(block_prefix + "attn1.to_q.weight")
|
||||
k = original_state_dict.pop(block_prefix + "attn1.to_k.weight")
|
||||
k = original_state_dict.pop(block_prefix + "attn1.to_k.weight")
|
||||
v = original_state_dict.pop(block_prefix + "attn1.to_v.weight")
|
||||
qkv_weight = torch.cat([q, k, v], dim=0)
|
||||
new_state_dict[new_prefix + "attn.qkv_x.weight"] = qkv_weight
|
||||
@@ -154,9 +113,7 @@ def convert_diffusers_transformer_to_mochi(state_dict):
|
||||
new_state_dict[new_prefix + "mlp_x.w1.weight"] = reverse_proj_gate(
|
||||
original_state_dict.pop(block_prefix + "ff.net.0.proj.weight")
|
||||
)
|
||||
new_state_dict[new_prefix + "mlp_x.w2.weight"] = original_state_dict.pop(
|
||||
block_prefix + "ff.net.2.weight"
|
||||
)
|
||||
new_state_dict[new_prefix + "mlp_x.w2.weight"] = original_state_dict.pop(block_prefix + "ff.net.2.weight")
|
||||
if i < num_layers - 1:
|
||||
new_state_dict[new_prefix + "mlp_y.w1.weight"] = reverse_proj_gate(
|
||||
original_state_dict.pop(block_prefix + "ff_context.net.0.proj.weight")
|
||||
@@ -172,9 +129,7 @@ def convert_diffusers_transformer_to_mochi(state_dict):
|
||||
new_state_dict["final_layer.mod.bias"] = reverse_scale_shift(
|
||||
original_state_dict.pop("norm_out.linear.bias"), dim=0
|
||||
)
|
||||
new_state_dict["final_layer.linear.weight"] = original_state_dict.pop(
|
||||
"proj_out.weight"
|
||||
)
|
||||
new_state_dict["final_layer.linear.weight"] = original_state_dict.pop("proj_out.weight")
|
||||
new_state_dict["final_layer.linear.bias"] = original_state_dict.pop("proj_out.bias")
|
||||
|
||||
new_state_dict["pos_frequencies"] = original_state_dict.pop("pos_frequencies")
|
||||
@@ -183,7 +138,6 @@ def convert_diffusers_transformer_to_mochi(state_dict):
|
||||
|
||||
return new_state_dict
|
||||
|
||||
|
||||
def convert_diffusers_vae_to_mochi(state_dict):
|
||||
original_state_dict = state_dict.copy()
|
||||
encoder_state_dict = {}
|
||||
@@ -192,12 +146,8 @@ def convert_diffusers_vae_to_mochi(state_dict):
|
||||
# Convert encoder
|
||||
prefix = "encoder."
|
||||
|
||||
encoder_state_dict["layers.0.weight"] = original_state_dict.pop(
|
||||
f"{prefix}proj_in.weight"
|
||||
)
|
||||
encoder_state_dict["layers.0.bias"] = original_state_dict.pop(
|
||||
f"{prefix}proj_in.bias"
|
||||
)
|
||||
encoder_state_dict["layers.0.weight"] = original_state_dict.pop(f"{prefix}proj_in.weight")
|
||||
encoder_state_dict["layers.0.bias"] = original_state_dict.pop(f"{prefix}proj_in.bias")
|
||||
|
||||
# Convert block_in
|
||||
for i in range(3):
|
||||
@@ -229,89 +179,57 @@ def convert_diffusers_vae_to_mochi(state_dict):
|
||||
# Convert down_blocks
|
||||
down_block_layers = [3, 4, 6]
|
||||
for block in range(3):
|
||||
encoder_state_dict[
|
||||
f"layers.{block+4}.layers.0.weight"
|
||||
] = original_state_dict.pop(f"{prefix}down_blocks.{block}.conv_in.conv.weight")
|
||||
encoder_state_dict[f"layers.{block+4}.layers.0.weight"] = original_state_dict.pop(
|
||||
f"{prefix}down_blocks.{block}.conv_in.conv.weight"
|
||||
)
|
||||
encoder_state_dict[f"layers.{block+4}.layers.0.bias"] = original_state_dict.pop(
|
||||
f"{prefix}down_blocks.{block}.conv_in.conv.bias"
|
||||
)
|
||||
|
||||
for i in range(down_block_layers[block]):
|
||||
# Convert resnets
|
||||
encoder_state_dict[
|
||||
f"layers.{block+4}.layers.{i+1}.stack.0.weight"
|
||||
] = original_state_dict.pop(
|
||||
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.stack.0.weight"] = original_state_dict.pop(
|
||||
f"{prefix}down_blocks.{block}.resnets.{i}.norm1.norm_layer.weight"
|
||||
)
|
||||
encoder_state_dict[
|
||||
f"layers.{block+4}.layers.{i+1}.stack.0.bias"
|
||||
] = original_state_dict.pop(
|
||||
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.stack.0.bias"] = original_state_dict.pop(
|
||||
f"{prefix}down_blocks.{block}.resnets.{i}.norm1.norm_layer.bias"
|
||||
)
|
||||
encoder_state_dict[
|
||||
f"layers.{block+4}.layers.{i+1}.stack.2.weight"
|
||||
] = original_state_dict.pop(
|
||||
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.stack.2.weight"] = original_state_dict.pop(
|
||||
f"{prefix}down_blocks.{block}.resnets.{i}.conv1.conv.weight"
|
||||
)
|
||||
encoder_state_dict[
|
||||
f"layers.{block+4}.layers.{i+1}.stack.2.bias"
|
||||
] = original_state_dict.pop(
|
||||
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.stack.2.bias"] = original_state_dict.pop(
|
||||
f"{prefix}down_blocks.{block}.resnets.{i}.conv1.conv.bias"
|
||||
)
|
||||
encoder_state_dict[
|
||||
f"layers.{block+4}.layers.{i+1}.stack.3.weight"
|
||||
] = original_state_dict.pop(
|
||||
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.stack.3.weight"] = original_state_dict.pop(
|
||||
f"{prefix}down_blocks.{block}.resnets.{i}.norm2.norm_layer.weight"
|
||||
)
|
||||
encoder_state_dict[
|
||||
f"layers.{block+4}.layers.{i+1}.stack.3.bias"
|
||||
] = original_state_dict.pop(
|
||||
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.stack.3.bias"] = original_state_dict.pop(
|
||||
f"{prefix}down_blocks.{block}.resnets.{i}.norm2.norm_layer.bias"
|
||||
)
|
||||
encoder_state_dict[
|
||||
f"layers.{block+4}.layers.{i+1}.stack.5.weight"
|
||||
] = original_state_dict.pop(
|
||||
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.stack.5.weight"] = original_state_dict.pop(
|
||||
f"{prefix}down_blocks.{block}.resnets.{i}.conv2.conv.weight"
|
||||
)
|
||||
encoder_state_dict[
|
||||
f"layers.{block+4}.layers.{i+1}.stack.5.bias"
|
||||
] = original_state_dict.pop(
|
||||
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.stack.5.bias"] = original_state_dict.pop(
|
||||
f"{prefix}down_blocks.{block}.resnets.{i}.conv2.conv.bias"
|
||||
)
|
||||
|
||||
# Convert attentions
|
||||
q = original_state_dict.pop(
|
||||
f"{prefix}down_blocks.{block}.attentions.{i}.to_q.weight"
|
||||
)
|
||||
k = original_state_dict.pop(
|
||||
f"{prefix}down_blocks.{block}.attentions.{i}.to_k.weight"
|
||||
)
|
||||
v = original_state_dict.pop(
|
||||
f"{prefix}down_blocks.{block}.attentions.{i}.to_v.weight"
|
||||
)
|
||||
q = original_state_dict.pop(f"{prefix}down_blocks.{block}.attentions.{i}.to_q.weight")
|
||||
k = original_state_dict.pop(f"{prefix}down_blocks.{block}.attentions.{i}.to_k.weight")
|
||||
v = original_state_dict.pop(f"{prefix}down_blocks.{block}.attentions.{i}.to_v.weight")
|
||||
qkv_weight = torch.cat([q, k, v], dim=0)
|
||||
encoder_state_dict[
|
||||
f"layers.{block+4}.layers.{i+1}.attn_block.attn.qkv.weight"
|
||||
] = qkv_weight
|
||||
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.attn_block.attn.qkv.weight"] = qkv_weight
|
||||
|
||||
encoder_state_dict[
|
||||
f"layers.{block+4}.layers.{i+1}.attn_block.attn.out.weight"
|
||||
] = original_state_dict.pop(
|
||||
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.attn_block.attn.out.weight"] = original_state_dict.pop(
|
||||
f"{prefix}down_blocks.{block}.attentions.{i}.to_out.0.weight"
|
||||
)
|
||||
encoder_state_dict[
|
||||
f"layers.{block+4}.layers.{i+1}.attn_block.attn.out.bias"
|
||||
] = original_state_dict.pop(
|
||||
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.attn_block.attn.out.bias"] = original_state_dict.pop(
|
||||
f"{prefix}down_blocks.{block}.attentions.{i}.to_out.0.bias"
|
||||
)
|
||||
encoder_state_dict[
|
||||
f"layers.{block+4}.layers.{i+1}.attn_block.norm.weight"
|
||||
] = original_state_dict.pop(
|
||||
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.attn_block.norm.weight"] = original_state_dict.pop(
|
||||
f"{prefix}down_blocks.{block}.norms.{i}.norm_layer.weight"
|
||||
)
|
||||
encoder_state_dict[
|
||||
f"layers.{block+4}.layers.{i+1}.attn_block.norm.bias"
|
||||
] = original_state_dict.pop(
|
||||
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.attn_block.norm.bias"] = original_state_dict.pop(
|
||||
f"{prefix}down_blocks.{block}.norms.{i}.norm_layer.bias"
|
||||
)
|
||||
|
||||
@@ -348,39 +266,29 @@ def convert_diffusers_vae_to_mochi(state_dict):
|
||||
qkv_weight = torch.cat([q, k, v], dim=0)
|
||||
encoder_state_dict[f"layers.{i+7}.attn_block.attn.qkv.weight"] = qkv_weight
|
||||
|
||||
encoder_state_dict[
|
||||
f"layers.{i+7}.attn_block.attn.out.weight"
|
||||
] = original_state_dict.pop(f"{prefix}block_out.attentions.{i}.to_out.0.weight")
|
||||
encoder_state_dict[
|
||||
f"layers.{i+7}.attn_block.attn.out.bias"
|
||||
] = original_state_dict.pop(f"{prefix}block_out.attentions.{i}.to_out.0.bias")
|
||||
encoder_state_dict[
|
||||
f"layers.{i+7}.attn_block.norm.weight"
|
||||
] = original_state_dict.pop(f"{prefix}block_out.norms.{i}.norm_layer.weight")
|
||||
encoder_state_dict[
|
||||
f"layers.{i+7}.attn_block.norm.bias"
|
||||
] = original_state_dict.pop(f"{prefix}block_out.norms.{i}.norm_layer.bias")
|
||||
encoder_state_dict[f"layers.{i+7}.attn_block.attn.out.weight"] = original_state_dict.pop(
|
||||
f"{prefix}block_out.attentions.{i}.to_out.0.weight"
|
||||
)
|
||||
encoder_state_dict[f"layers.{i+7}.attn_block.attn.out.bias"] = original_state_dict.pop(
|
||||
f"{prefix}block_out.attentions.{i}.to_out.0.bias"
|
||||
)
|
||||
encoder_state_dict[f"layers.{i+7}.attn_block.norm.weight"] = original_state_dict.pop(
|
||||
f"{prefix}block_out.norms.{i}.norm_layer.weight"
|
||||
)
|
||||
encoder_state_dict[f"layers.{i+7}.attn_block.norm.bias"] = original_state_dict.pop(
|
||||
f"{prefix}block_out.norms.{i}.norm_layer.bias"
|
||||
)
|
||||
|
||||
# Convert output layers
|
||||
encoder_state_dict["output_norm.weight"] = original_state_dict.pop(
|
||||
f"{prefix}norm_out.norm_layer.weight"
|
||||
)
|
||||
encoder_state_dict["output_norm.bias"] = original_state_dict.pop(
|
||||
f"{prefix}norm_out.norm_layer.bias"
|
||||
)
|
||||
encoder_state_dict["output_proj.weight"] = original_state_dict.pop(
|
||||
f"{prefix}proj_out.weight"
|
||||
)
|
||||
encoder_state_dict["output_norm.weight"] = original_state_dict.pop(f"{prefix}norm_out.norm_layer.weight")
|
||||
encoder_state_dict["output_norm.bias"] = original_state_dict.pop(f"{prefix}norm_out.norm_layer.bias")
|
||||
encoder_state_dict["output_proj.weight"] = original_state_dict.pop(f"{prefix}proj_out.weight")
|
||||
|
||||
# Convert decoder
|
||||
prefix = "decoder."
|
||||
|
||||
decoder_state_dict["blocks.0.0.weight"] = original_state_dict.pop(
|
||||
f"{prefix}conv_in.weight"
|
||||
)
|
||||
decoder_state_dict["blocks.0.0.bias"] = original_state_dict.pop(
|
||||
f"{prefix}conv_in.bias"
|
||||
)
|
||||
decoder_state_dict["blocks.0.0.weight"] = original_state_dict.pop(f"{prefix}conv_in.weight")
|
||||
decoder_state_dict["blocks.0.0.bias"] = original_state_dict.pop(f"{prefix}conv_in.bias")
|
||||
|
||||
# Convert block_in
|
||||
for i in range(3):
|
||||
@@ -413,44 +321,28 @@ def convert_diffusers_vae_to_mochi(state_dict):
|
||||
up_block_layers = [6, 4, 3]
|
||||
for block in range(3):
|
||||
for i in range(up_block_layers[block]):
|
||||
decoder_state_dict[
|
||||
f"blocks.{block+1}.blocks.{i}.stack.0.weight"
|
||||
] = original_state_dict.pop(
|
||||
decoder_state_dict[f"blocks.{block+1}.blocks.{i}.stack.0.weight"] = original_state_dict.pop(
|
||||
f"{prefix}up_blocks.{block}.resnets.{i}.norm1.norm_layer.weight"
|
||||
)
|
||||
decoder_state_dict[
|
||||
f"blocks.{block+1}.blocks.{i}.stack.0.bias"
|
||||
] = original_state_dict.pop(
|
||||
decoder_state_dict[f"blocks.{block+1}.blocks.{i}.stack.0.bias"] = original_state_dict.pop(
|
||||
f"{prefix}up_blocks.{block}.resnets.{i}.norm1.norm_layer.bias"
|
||||
)
|
||||
decoder_state_dict[
|
||||
f"blocks.{block+1}.blocks.{i}.stack.2.weight"
|
||||
] = original_state_dict.pop(
|
||||
decoder_state_dict[f"blocks.{block+1}.blocks.{i}.stack.2.weight"] = original_state_dict.pop(
|
||||
f"{prefix}up_blocks.{block}.resnets.{i}.conv1.conv.weight"
|
||||
)
|
||||
decoder_state_dict[
|
||||
f"blocks.{block+1}.blocks.{i}.stack.2.bias"
|
||||
] = original_state_dict.pop(
|
||||
decoder_state_dict[f"blocks.{block+1}.blocks.{i}.stack.2.bias"] = original_state_dict.pop(
|
||||
f"{prefix}up_blocks.{block}.resnets.{i}.conv1.conv.bias"
|
||||
)
|
||||
decoder_state_dict[
|
||||
f"blocks.{block+1}.blocks.{i}.stack.3.weight"
|
||||
] = original_state_dict.pop(
|
||||
decoder_state_dict[f"blocks.{block+1}.blocks.{i}.stack.3.weight"] = original_state_dict.pop(
|
||||
f"{prefix}up_blocks.{block}.resnets.{i}.norm2.norm_layer.weight"
|
||||
)
|
||||
decoder_state_dict[
|
||||
f"blocks.{block+1}.blocks.{i}.stack.3.bias"
|
||||
] = original_state_dict.pop(
|
||||
decoder_state_dict[f"blocks.{block+1}.blocks.{i}.stack.3.bias"] = original_state_dict.pop(
|
||||
f"{prefix}up_blocks.{block}.resnets.{i}.norm2.norm_layer.bias"
|
||||
)
|
||||
decoder_state_dict[
|
||||
f"blocks.{block+1}.blocks.{i}.stack.5.weight"
|
||||
] = original_state_dict.pop(
|
||||
decoder_state_dict[f"blocks.{block+1}.blocks.{i}.stack.5.weight"] = original_state_dict.pop(
|
||||
f"{prefix}up_blocks.{block}.resnets.{i}.conv2.conv.weight"
|
||||
)
|
||||
decoder_state_dict[
|
||||
f"blocks.{block+1}.blocks.{i}.stack.5.bias"
|
||||
] = original_state_dict.pop(
|
||||
decoder_state_dict[f"blocks.{block+1}.blocks.{i}.stack.5.bias"] = original_state_dict.pop(
|
||||
f"{prefix}up_blocks.{block}.resnets.{i}.conv2.conv.bias"
|
||||
)
|
||||
decoder_state_dict[f"blocks.{block+1}.proj.weight"] = original_state_dict.pop(
|
||||
@@ -487,32 +379,25 @@ def convert_diffusers_vae_to_mochi(state_dict):
|
||||
f"{prefix}block_out.resnets.{i}.conv2.conv.bias"
|
||||
)
|
||||
|
||||
# Convert output layers
|
||||
decoder_state_dict["output_proj.weight"] = original_state_dict.pop(
|
||||
f"{prefix}proj_out.weight"
|
||||
)
|
||||
decoder_state_dict["output_proj.bias"] = original_state_dict.pop(
|
||||
f"{prefix}proj_out.bias"
|
||||
)
|
||||
# Convert output layers
|
||||
decoder_state_dict["output_proj.weight"] = original_state_dict.pop(f"{prefix}proj_out.weight")
|
||||
decoder_state_dict["output_proj.bias"] = original_state_dict.pop(f"{prefix}proj_out.bias")
|
||||
|
||||
return encoder_state_dict, decoder_state_dict
|
||||
|
||||
|
||||
def ensure_safetensors_extension(path):
|
||||
if not path.endswith(".safetensors"):
|
||||
path = path + ".safetensors"
|
||||
if not path.endswith('.safetensors'):
|
||||
path = path + '.safetensors'
|
||||
return path
|
||||
|
||||
|
||||
def ensure_directory_exists(path):
|
||||
directory = os.path.dirname(path)
|
||||
if directory:
|
||||
os.makedirs(directory, exist_ok=True)
|
||||
|
||||
|
||||
def main(args):
|
||||
from diffusers import MochiPipeline
|
||||
|
||||
|
||||
pipe = MochiPipeline.from_pretrained(args.diffusers_path)
|
||||
|
||||
if args.transformer_path:
|
||||
@@ -520,9 +405,7 @@ def main(args):
|
||||
ensure_directory_exists(transformer_path)
|
||||
|
||||
print(f"Converting transformer model...")
|
||||
transformer_state_dict = convert_diffusers_transformer_to_mochi(
|
||||
pipe.transformer.state_dict()
|
||||
)
|
||||
transformer_state_dict = convert_diffusers_transformer_to_mochi(pipe.transformer.state_dict())
|
||||
save_file(transformer_state_dict, transformer_path)
|
||||
print(f"Saved transformer to {transformer_path}")
|
||||
|
||||
@@ -534,9 +417,7 @@ def main(args):
|
||||
ensure_directory_exists(decoder_path)
|
||||
|
||||
print(f"Converting VAE models...")
|
||||
encoder_state_dict, decoder_state_dict = convert_diffusers_vae_to_mochi(
|
||||
pipe.vae.state_dict()
|
||||
)
|
||||
encoder_state_dict, decoder_state_dict = convert_diffusers_vae_to_mochi(pipe.vae.state_dict())
|
||||
|
||||
save_file(encoder_state_dict, encoder_path)
|
||||
print(f"Saved VAE encoder to {encoder_path}")
|
||||
@@ -544,10 +425,7 @@ def main(args):
|
||||
save_file(decoder_state_dict, decoder_path)
|
||||
print(f"Saved VAE decoder to {decoder_path}")
|
||||
elif args.vae_encoder_path or args.vae_decoder_path:
|
||||
print(
|
||||
"Warning: Both VAE encoder and decoder paths must be specified to convert VAE models."
|
||||
)
|
||||
|
||||
print("Warning: Both VAE encoder and decoder paths must be specified to convert VAE models.")
|
||||
|
||||
if __name__ == "__main__":
|
||||
main(args)
|
||||
main(args)
|
||||
@@ -42,6 +42,7 @@ def normalize_dit_input(model_type, latents):
|
||||
latents = (latents - latents_mean) / latents_std
|
||||
return latents
|
||||
elif model_type == "hunyuan":
|
||||
return latents * 0.476986
|
||||
return latents * 0.476986
|
||||
else:
|
||||
raise NotImplementedError(f"model_type {model_type} not supported")
|
||||
|
||||
@@ -80,6 +80,9 @@ class FeedForward(HF_FeedForward):
|
||||
return self.net[2](LigerSiLUMulFunction.apply(gate, hidden_states))
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
class MochiAttention(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
@@ -621,8 +624,7 @@ class MochiTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
timestep: torch.LongTensor,
|
||||
encoder_attention_mask: torch.Tensor,
|
||||
output_features=False,
|
||||
output_features_stride=8,
|
||||
output_attn=False,
|
||||
attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||
return_dict: bool = False,
|
||||
) -> torch.Tensor:
|
||||
@@ -698,7 +700,7 @@ class MochiTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
|
||||
encoder_attention_mask,
|
||||
temb,
|
||||
image_rotary_emb,
|
||||
output_features,
|
||||
output_attn,
|
||||
**ckpt_kwargs,
|
||||
)
|
||||
else:
|
||||
@@ -708,10 +710,9 @@ class MochiTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
|
||||
encoder_attention_mask=encoder_attention_mask,
|
||||
temb=temb,
|
||||
image_rotary_emb=image_rotary_emb,
|
||||
output_attn=output_features,
|
||||
output_attn=output_attn,
|
||||
)
|
||||
if i % output_features_stride == 0:
|
||||
attn_outputs_list.append(attn_outputs)
|
||||
attn_outputs_list.append(attn_outputs)
|
||||
|
||||
hidden_states = self.norm_out(hidden_states, temb)
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
@@ -726,7 +727,7 @@ class MochiTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
|
||||
# remove `lora_scale` from each PEFT layer
|
||||
unscale_lora_layers(self, lora_scale)
|
||||
|
||||
if not output_features:
|
||||
if not output_attn:
|
||||
attn_outputs_list = None
|
||||
else:
|
||||
attn_outputs_list = torch.stack(attn_outputs_list, dim=0)
|
||||
|
||||
@@ -67,7 +67,11 @@ class MochiRMSNorm(nn.Module):
|
||||
|
||||
class MochiLayerNormContinuous(nn.Module):
|
||||
def __init__(
|
||||
self, embedding_dim: int, conditioning_embedding_dim: int, eps=1e-5, bias=True,
|
||||
self,
|
||||
embedding_dim: int,
|
||||
conditioning_embedding_dim: int,
|
||||
eps=1e-5,
|
||||
bias=True,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
@@ -77,7 +81,9 @@ class MochiLayerNormContinuous(nn.Module):
|
||||
self.norm = MochiModulatedRMSNorm(eps=eps)
|
||||
|
||||
def forward(
|
||||
self, x: torch.Tensor, conditioning_embedding: torch.Tensor,
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
conditioning_embedding: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
input_dtype = x.dtype
|
||||
|
||||
|
||||
@@ -36,7 +36,7 @@ from diffusers.pipelines.mochi.pipeline_output import MochiPipelineOutput
|
||||
from einops import rearrange
|
||||
from fastvideo.utils.parallel_states import get_sequence_parallel_state, nccl_info
|
||||
from fastvideo.utils.communications import all_gather
|
||||
from diffusers.loaders import Mochi1LoraLoaderMixin
|
||||
|
||||
if is_torch_xla_available():
|
||||
import torch_xla.core.xla_model as xm
|
||||
|
||||
@@ -85,13 +85,13 @@ def linear_quadratic_schedule(num_steps, threshold_noise, linear_steps=None):
|
||||
]
|
||||
threshold_noise_step_diff = linear_steps - threshold_noise * num_steps
|
||||
quadratic_steps = num_steps - linear_steps
|
||||
quadratic_coef = threshold_noise_step_diff / (linear_steps * quadratic_steps ** 2)
|
||||
quadratic_coef = threshold_noise_step_diff / (linear_steps * quadratic_steps**2)
|
||||
linear_coef = threshold_noise / linear_steps - 2 * threshold_noise_step_diff / (
|
||||
quadratic_steps ** 2
|
||||
quadratic_steps**2
|
||||
)
|
||||
const = quadratic_coef * (linear_steps ** 2)
|
||||
const = quadratic_coef * (linear_steps**2)
|
||||
quadratic_sigma_schedule = [
|
||||
quadratic_coef * (i ** 2) + linear_coef * i + const
|
||||
quadratic_coef * (i**2) + linear_coef * i + const
|
||||
for i in range(linear_steps, num_steps)
|
||||
]
|
||||
sigma_schedule = linear_sigma_schedule + quadratic_sigma_schedule
|
||||
@@ -165,7 +165,7 @@ def retrieve_timesteps(
|
||||
return timesteps, num_inference_steps
|
||||
|
||||
|
||||
class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
|
||||
class MochiPipeline(DiffusionPipeline):
|
||||
r"""
|
||||
The mochi pipeline for text-to-video generation.
|
||||
|
||||
@@ -502,8 +502,7 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
|
||||
f" size of {batch_size}. Make sure the batch size matches the length of the generators."
|
||||
)
|
||||
|
||||
latents = randn_tensor(shape, generator=generator, device=device, dtype=torch.float32)
|
||||
latents = latents.to(dtype)
|
||||
latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
|
||||
return latents
|
||||
|
||||
@property
|
||||
@@ -534,8 +533,8 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
|
||||
negative_prompt: Optional[Union[str, List[str]]] = None,
|
||||
height: Optional[int] = None,
|
||||
width: Optional[int] = None,
|
||||
num_frames: int = 19,
|
||||
num_inference_steps: int = 64,
|
||||
num_frames: int = 16,
|
||||
num_inference_steps: int = 28,
|
||||
timesteps: List[int] = None,
|
||||
guidance_scale: float = 4.5,
|
||||
num_videos_per_prompt: Optional[int] = 1,
|
||||
@@ -712,11 +711,17 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
|
||||
# check if of type FlowMatchEulerDiscreteScheduler
|
||||
if isinstance(self.scheduler, FlowMatchEulerDiscreteScheduler):
|
||||
timesteps, num_inference_steps = retrieve_timesteps(
|
||||
self.scheduler, num_inference_steps, device, timesteps, sigmas,
|
||||
self.scheduler,
|
||||
num_inference_steps,
|
||||
device,
|
||||
timesteps,
|
||||
sigmas,
|
||||
)
|
||||
else:
|
||||
timesteps, num_inference_steps = retrieve_timesteps(
|
||||
self.scheduler, num_inference_steps, device,
|
||||
self.scheduler,
|
||||
num_inference_steps,
|
||||
device,
|
||||
)
|
||||
num_warmup_steps = max(
|
||||
len(timesteps) - num_inference_steps * self.scheduler.order, 0
|
||||
@@ -724,7 +729,6 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
|
||||
self._num_timesteps = len(timesteps)
|
||||
|
||||
# 6. Denoising loop
|
||||
self._progress_bar_config = {"disable": nccl_info.rank_within_group != 0}
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
if self.interrupt:
|
||||
|
||||
@@ -1,240 +0,0 @@
|
||||
import torch
|
||||
from diffusers import HunyuanVideoPipeline, HunyuanVideoTransformer3DModel, BitsAndBytesConfig
|
||||
import imageio as iio
|
||||
import math
|
||||
import numpy as np
|
||||
import io
|
||||
import time
|
||||
import argparse
|
||||
import os
|
||||
|
||||
def export_to_video_bytes(fps, frames):
|
||||
request = iio.core.Request("<bytes>", mode="w", extension=".mp4")
|
||||
pyavobject = iio.plugins.pyav.PyAVPlugin(request)
|
||||
if isinstance(frames, np.ndarray):
|
||||
frames = (np.array(frames) * 255).astype('uint8')
|
||||
else:
|
||||
frames = np.array(frames)
|
||||
new_bytes = pyavobject.write(frames, codec="libx264", fps=fps)
|
||||
out_bytes = io.BytesIO(new_bytes)
|
||||
return out_bytes
|
||||
|
||||
def export_to_video(frames, path, fps):
|
||||
video_bytes = export_to_video_bytes(fps, frames)
|
||||
video_bytes.seek(0)
|
||||
with open(path, "wb") as f:
|
||||
f.write(video_bytes.getbuffer())
|
||||
|
||||
def main(args):
|
||||
torch.manual_seed(args.seed)
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
prompt_template = {
|
||||
"template": (
|
||||
"<|start_header_cid|>system<|end_header_id|>\n\nDescribe the video by detailing the following aspects: "
|
||||
"1. The main content and theme of the video."
|
||||
"2. The color, shape, size, texture, quantity, text, and spatial relationships of the contents, including objects, people, and anything else."
|
||||
"3. Actions, events, behaviors temporal relationships, physical movement changes of the contents."
|
||||
"4. Background environment, light, style, atmosphere, and qualities."
|
||||
"5. Camera angles, movements, and transitions used in the video."
|
||||
"6. Thematic and aesthetic concepts associated with the scene, i.e. realistic, futuristic, fairy tale, etc<|eot_id|>"
|
||||
"<|start_header_id|>user<|end_header_id|>\n\n{}<|eot_id|>"
|
||||
),
|
||||
"crop_start": 95,
|
||||
}
|
||||
|
||||
|
||||
model_id = args.model_path
|
||||
|
||||
if args.quantization == "nf4":
|
||||
quantization_config = BitsAndBytesConfig(load_in_4bit=True, bnb_4bit_compute_dtype=torch.bfloat16, bnb_4bit_quant_type="nf4", llm_int8_skip_modules=["proj_out", "norm_out"])
|
||||
transformer = HunyuanVideoTransformer3DModel.from_pretrained(
|
||||
model_id, subfolder="transformer/" ,torch_dtype=torch.bfloat16, quantization_config=quantization_config
|
||||
)
|
||||
if args.quantization == "int8":
|
||||
quantization_config = BitsAndBytesConfig(load_in_8bit=True, llm_int8_skip_modules=["proj_out", "norm_out"])
|
||||
transformer = HunyuanVideoTransformer3DModel.from_pretrained(
|
||||
model_id, subfolder="transformer/" ,torch_dtype=torch.bfloat16, quantization_config=quantization_config
|
||||
)
|
||||
elif not args.quantization:
|
||||
transformer = HunyuanVideoTransformer3DModel.from_pretrained(
|
||||
model_id, subfolder="transformer/" ,torch_dtype=torch.bfloat16
|
||||
).to(device)
|
||||
|
||||
print("Max vram for read transofrmer:", round(torch.cuda.max_memory_allocated(device="cuda") / 1024 ** 3, 3), "GiB")
|
||||
torch.cuda.reset_max_memory_allocated(device)
|
||||
|
||||
if not args.cpu_offload:
|
||||
pipe = HunyuanVideoPipeline.from_pretrained(model_id, torch_dtype=torch.bfloat16).to(device)
|
||||
pipe.transformer = transformer
|
||||
else:
|
||||
pipe = HunyuanVideoPipeline.from_pretrained(model_id, transformer=transformer, torch_dtype=torch.bfloat16)
|
||||
torch.cuda.reset_max_memory_allocated(device)
|
||||
pipe.scheduler._shift = args.flow_shift
|
||||
pipe.vae.enable_tiling()
|
||||
if args.cpu_offload:
|
||||
pipe.enable_model_cpu_offload()
|
||||
print("Max vram for init pipeline:", round(torch.cuda.max_memory_allocated(device="cuda") / 1024 ** 3, 3), "GiB")
|
||||
with open(args.prompt) as f:
|
||||
prompts = f.readlines()
|
||||
|
||||
generator = torch.Generator("cpu").manual_seed(args.seed)
|
||||
os.makedirs(os.path.dirname(args.output_path), exist_ok=True)
|
||||
torch.cuda.reset_max_memory_allocated(device)
|
||||
for prompt in prompts:
|
||||
start_time = time.perf_counter()
|
||||
output = pipe(
|
||||
prompt=prompt,
|
||||
height = args.height,
|
||||
width = args.width,
|
||||
num_frames = args.num_frames,
|
||||
prompt_template=prompt_template,
|
||||
num_inference_steps = args.num_inference_steps,
|
||||
generator=generator,
|
||||
).frames[0]
|
||||
export_to_video(output, os.path.join(args.output_path, f"{prompt[:100]}.mp4"), fps=args.fps)
|
||||
print("Time:", round(time.perf_counter() - start_time, 2), "seconds")
|
||||
print("Max vram for denoise:", round(torch.cuda.max_memory_allocated(device="cuda") / 1024 ** 3, 3), "GiB")
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
|
||||
# Basic parameters
|
||||
parser.add_argument("--prompt", type=str, help="prompt file for inference")
|
||||
parser.add_argument("--num_frames", type=int, default=16)
|
||||
parser.add_argument("--height", type=int, default=256)
|
||||
parser.add_argument("--width", type=int, default=256)
|
||||
parser.add_argument("--num_inference_steps", type=int, default=50)
|
||||
parser.add_argument("--model_path", type=str, default="data/hunyuan")
|
||||
parser.add_argument("--output_path", type=str, default="./outputs/video")
|
||||
parser.add_argument("--fps", type=int, default=24)
|
||||
parser.add_argument("--quantization", type=str, default=None)
|
||||
parser.add_argument("--cpu_offload", action="store_true")
|
||||
# Additional parameters
|
||||
parser.add_argument(
|
||||
"--denoise-type",
|
||||
type=str,
|
||||
default="flow",
|
||||
help="Denoise type for noised inputs.",
|
||||
)
|
||||
parser.add_argument("--seed", type=int, default=None, help="Seed for evaluation.")
|
||||
parser.add_argument(
|
||||
"--neg_prompt", type=str, default=None, help="Negative prompt for sampling."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--guidance_scale",
|
||||
type=float,
|
||||
default=1.0,
|
||||
help="Classifier free guidance scale.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--embedded_cfg_scale",
|
||||
type=float,
|
||||
default=6.0,
|
||||
help="Embedded classifier free guidance scale.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--flow_shift", type=int, default=7, help="Flow shift parameter."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--batch_size", type=int, default=1, help="Batch size for inference."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--num_videos",
|
||||
type=int,
|
||||
default=1,
|
||||
help="Number of videos to generate per prompt.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--load-key",
|
||||
type=str,
|
||||
default="module",
|
||||
help="Key to load the model states. 'module' for the main model, 'ema' for the EMA model.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--use-cpu-offload",
|
||||
action="store_true",
|
||||
help="Use CPU offload for the model load.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--dit-weight",
|
||||
type=str,
|
||||
default="data/hunyuan/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--reproduce",
|
||||
action="store_true",
|
||||
help="Enable reproducibility by setting random seeds and deterministic algorithms.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--disable-autocast",
|
||||
action="store_true",
|
||||
help="Disable autocast for denoising loop and vae decoding in pipeline sampling.",
|
||||
)
|
||||
|
||||
# Flow Matching
|
||||
parser.add_argument(
|
||||
"--flow-reverse",
|
||||
action="store_true",
|
||||
help="If reverse, learning/sampling from t=1 -> t=0.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--flow-solver", type=str, default="euler", help="Solver for flow matching."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--use-linear-quadratic-schedule",
|
||||
action="store_true",
|
||||
help="Use linear quadratic schedule for flow matching. Following MovieGen (https://ai.meta.com/static-resource/movie-gen-research-paper)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--linear-schedule-end",
|
||||
type=int,
|
||||
default=25,
|
||||
help="End step for linear quadratic schedule for flow matching.",
|
||||
)
|
||||
|
||||
# Model parameters
|
||||
parser.add_argument("--model", type=str, default="HYVideo-T/2-cfgdistill")
|
||||
parser.add_argument("--latent-channels", type=int, default=16)
|
||||
parser.add_argument(
|
||||
"--precision", type=str, default="bf16", choices=["fp32", "fp16", "bf16", "fp8"]
|
||||
)
|
||||
parser.add_argument(
|
||||
"--rope-theta", type=int, default=256, help="Theta used in RoPE."
|
||||
)
|
||||
|
||||
parser.add_argument("--vae", type=str, default="884-16c-hy")
|
||||
parser.add_argument(
|
||||
"--vae-precision", type=str, default="fp16", choices=["fp32", "fp16", "bf16"]
|
||||
)
|
||||
parser.add_argument("--vae-tiling", action="store_true", default=True)
|
||||
|
||||
parser.add_argument("--text-encoder", type=str, default="llm")
|
||||
parser.add_argument(
|
||||
"--text-encoder-precision",
|
||||
type=str,
|
||||
default="fp16",
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
)
|
||||
parser.add_argument("--text-states-dim", type=int, default=4096)
|
||||
parser.add_argument("--text-len", type=int, default=256)
|
||||
parser.add_argument("--tokenizer", type=str, default="llm")
|
||||
parser.add_argument("--prompt-template", type=str, default="dit-llm-encode")
|
||||
parser.add_argument(
|
||||
"--prompt-template-video", type=str, default="dit-llm-encode-video"
|
||||
)
|
||||
parser.add_argument("--hidden-state-skip-layer", type=int, default=2)
|
||||
parser.add_argument("--apply-final-norm", action="store_true")
|
||||
|
||||
parser.add_argument("--text-encoder-2", type=str, default="clipL")
|
||||
parser.add_argument(
|
||||
"--text-encoder-precision-2",
|
||||
type=str,
|
||||
default="fp16",
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
)
|
||||
parser.add_argument("--text-states-dim-2", type=int, default=768)
|
||||
parser.add_argument("--tokenizer-2", type=str, default="clipL")
|
||||
parser.add_argument("--text-len-2", type=int, default=77)
|
||||
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
@@ -22,7 +22,6 @@ from fastvideo.utils.parallel_states import (
|
||||
nccl_info,
|
||||
)
|
||||
|
||||
|
||||
def initialize_distributed():
|
||||
local_rank = int(os.getenv("RANK", 0))
|
||||
world_size = int(os.getenv("WORLD_SIZE", 1))
|
||||
@@ -38,33 +37,27 @@ def main(args):
|
||||
initialize_distributed()
|
||||
print(nccl_info.sp_size)
|
||||
device = torch.cuda.current_device()
|
||||
|
||||
|
||||
print(args)
|
||||
models_root_path = Path(args.model_path)
|
||||
if not models_root_path.exists():
|
||||
raise ValueError(f"`models_root` not exists: {models_root_path}")
|
||||
|
||||
|
||||
# Create save folder to save the samples
|
||||
save_path = args.output_path
|
||||
os.makedirs(os.path.dirname(save_path), exist_ok=True)
|
||||
|
||||
# Load models
|
||||
hunyuan_video_sampler = HunyuanVideoSampler.from_pretrained(
|
||||
models_root_path, args=args
|
||||
)
|
||||
|
||||
hunyuan_video_sampler = HunyuanVideoSampler.from_pretrained(models_root_path, args=args)
|
||||
|
||||
# Get the updated args
|
||||
args = hunyuan_video_sampler.args
|
||||
|
||||
# Start sampling
|
||||
samples = []
|
||||
|
||||
with open(args.prompt) as f:
|
||||
prompts = f.readlines()
|
||||
|
||||
for prompt in prompts:
|
||||
for prompt in args.prompts:
|
||||
outputs = hunyuan_video_sampler.predict(
|
||||
prompt=prompt,
|
||||
prompt=prompt,
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
video_length=args.num_frames,
|
||||
@@ -75,25 +68,25 @@ def main(args):
|
||||
num_videos_per_prompt=args.num_videos,
|
||||
flow_shift=args.flow_shift,
|
||||
batch_size=args.batch_size,
|
||||
embedded_guidance_scale=args.embedded_cfg_scale,
|
||||
embedded_guidance_scale=args.embedded_cfg_scale
|
||||
)
|
||||
videos = rearrange(outputs["samples"], "b c t h w -> t b c h w")
|
||||
samples.append(outputs['samples'][0])
|
||||
|
||||
for prompt, video in zip(args.prompts, samples):
|
||||
videos = rearrange(video.unsqueeze(0), "b c t h w -> t b c h w")
|
||||
outputs = []
|
||||
for x in videos:
|
||||
x = torchvision.utils.make_grid(x, nrow=6)
|
||||
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
|
||||
outputs.append((x * 255).numpy().astype(np.uint8))
|
||||
os.makedirs(os.path.dirname(args.output_path), exist_ok=True)
|
||||
imageio.mimsave(
|
||||
os.path.join(args.output_path, f"{prompt[:100]}.mp4"), outputs, fps=args.fps
|
||||
)
|
||||
|
||||
imageio.mimsave(args.output_path + f"{prompt[:100]}.mp4", outputs, fps=args.fps)
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
|
||||
|
||||
# Basic parameters
|
||||
parser.add_argument("--prompt", type=str, help="prompt file for inference")
|
||||
parser.add_argument("--prompts", nargs="+", default=[])
|
||||
parser.add_argument("--num_frames", type=int, default=16)
|
||||
parser.add_argument("--height", type=int, default=256)
|
||||
parser.add_argument("--width", type=int, default=256)
|
||||
@@ -101,133 +94,55 @@ if __name__ == "__main__":
|
||||
parser.add_argument("--model_path", type=str, default="data/hunyuan")
|
||||
parser.add_argument("--output_path", type=str, default="./outputs/video")
|
||||
parser.add_argument("--fps", type=int, default=24)
|
||||
|
||||
|
||||
# Additional parameters
|
||||
parser.add_argument(
|
||||
"--denoise-type",
|
||||
type=str,
|
||||
default="flow",
|
||||
help="Denoise type for noised inputs.",
|
||||
)
|
||||
parser.add_argument("--denoise-type", type=str, default="flow", help="Denoise type for noised inputs.")
|
||||
parser.add_argument("--seed", type=int, default=None, help="Seed for evaluation.")
|
||||
parser.add_argument(
|
||||
"--neg_prompt", type=str, default=None, help="Negative prompt for sampling."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--guidance_scale",
|
||||
type=float,
|
||||
default=1.0,
|
||||
help="Classifier free guidance scale.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--embedded_cfg_scale",
|
||||
type=float,
|
||||
default=6.0,
|
||||
help="Embedded classifier free guidance scale.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--flow_shift", type=int, default=7, help="Flow shift parameter."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--batch_size", type=int, default=1, help="Batch size for inference."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--num_videos",
|
||||
type=int,
|
||||
default=1,
|
||||
help="Number of videos to generate per prompt.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--load-key",
|
||||
type=str,
|
||||
default="module",
|
||||
help="Key to load the model states. 'module' for the main model, 'ema' for the EMA model.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--use-cpu-offload",
|
||||
action="store_true",
|
||||
help="Use CPU offload for the model load.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--dit-weight",
|
||||
type=str,
|
||||
default="data/hunyuan/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--reproduce",
|
||||
action="store_true",
|
||||
help="Enable reproducibility by setting random seeds and deterministic algorithms.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--disable-autocast",
|
||||
action="store_true",
|
||||
help="Disable autocast for denoising loop and vae decoding in pipeline sampling.",
|
||||
)
|
||||
|
||||
parser.add_argument("--neg_prompt", type=str, default=None, help="Negative prompt for sampling.")
|
||||
parser.add_argument("--guidance_scale", type=float, default=1.0, help="Classifier free guidance scale.")
|
||||
parser.add_argument("--embedded_cfg_scale", type=float, default=6.0, help="Embedded classifier free guidance scale.")
|
||||
parser.add_argument("--flow_shift", type=int, default=7, help="Flow shift parameter.")
|
||||
parser.add_argument("--batch_size", type=int, default=1, help="Batch size for inference.")
|
||||
parser.add_argument("--num_videos", type=int, default=1, help="Number of videos to generate per prompt.")
|
||||
parser.add_argument("--load-key", type=str, default="module", help="Key to load the model states. 'module' for the main model, 'ema' for the EMA model.")
|
||||
parser.add_argument("--use-cpu-offload", action="store_true", help="Use CPU offload for the model load.")
|
||||
parser.add_argument("--dit-weight", type=str, default="data/hunyuan/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt")
|
||||
parser.add_argument("--reproduce", action="store_true", help="Enable reproducibility by setting random seeds and deterministic algorithms.")
|
||||
parser.add_argument("--disable-autocast", action="store_true", help="Disable autocast for denoising loop and vae decoding in pipeline sampling.")
|
||||
|
||||
# Flow Matching
|
||||
parser.add_argument(
|
||||
"--flow-reverse",
|
||||
action="store_true",
|
||||
help="If reverse, learning/sampling from t=1 -> t=0.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--flow-solver", type=str, default="euler", help="Solver for flow matching."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--use-linear-quadratic-schedule",
|
||||
action="store_true",
|
||||
help="Use linear quadratic schedule for flow matching. Following MovieGen (https://ai.meta.com/static-resource/movie-gen-research-paper)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--linear-schedule-end",
|
||||
type=int,
|
||||
default=25,
|
||||
help="End step for linear quadratic schedule for flow matching.",
|
||||
)
|
||||
|
||||
parser.add_argument("--flow-reverse", action="store_true", help="If reverse, learning/sampling from t=1 -> t=0.")
|
||||
parser.add_argument("--flow-solver", type=str, default="euler", help="Solver for flow matching.")
|
||||
parser.add_argument("--use-linear-quadratic-schedule", action="store_true",
|
||||
help="Use linear quadratic schedule for flow matching. Following MovieGen (https://ai.meta.com/static-resource/movie-gen-research-paper)")
|
||||
parser.add_argument("--linear-schedule-end", type=int, default=25,
|
||||
help="End step for linear quadratic schedule for flow matching.")
|
||||
|
||||
# Model parameters
|
||||
parser.add_argument("--model", type=str, default="HYVideo-T/2-cfgdistill")
|
||||
parser.add_argument("--latent-channels", type=int, default=16)
|
||||
parser.add_argument(
|
||||
"--precision", type=str, default="bf16", choices=["fp32", "fp16", "bf16"]
|
||||
)
|
||||
parser.add_argument(
|
||||
"--rope-theta", type=int, default=256, help="Theta used in RoPE."
|
||||
)
|
||||
|
||||
parser.add_argument("--precision", type=str, default="bf16", choices=["fp32", "fp16", "bf16"])
|
||||
parser.add_argument("--rope-theta", type=int, default=256, help="Theta used in RoPE.")
|
||||
|
||||
parser.add_argument("--vae", type=str, default="884-16c-hy")
|
||||
parser.add_argument(
|
||||
"--vae-precision", type=str, default="fp16", choices=["fp32", "fp16", "bf16"]
|
||||
)
|
||||
parser.add_argument("--vae-precision", type=str, default="fp16", choices=["fp32", "fp16", "bf16"])
|
||||
parser.add_argument("--vae-tiling", action="store_true", default=True)
|
||||
|
||||
parser.add_argument("--text-encoder", type=str, default="llm")
|
||||
parser.add_argument(
|
||||
"--text-encoder-precision",
|
||||
type=str,
|
||||
default="fp16",
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
)
|
||||
parser.add_argument("--text-encoder-precision", type=str, default="fp16", choices=["fp32", "fp16", "bf16"])
|
||||
parser.add_argument("--text-states-dim", type=int, default=4096)
|
||||
parser.add_argument("--text-len", type=int, default=256)
|
||||
parser.add_argument("--tokenizer", type=str, default="llm")
|
||||
parser.add_argument("--prompt-template", type=str, default="dit-llm-encode")
|
||||
parser.add_argument(
|
||||
"--prompt-template-video", type=str, default="dit-llm-encode-video"
|
||||
)
|
||||
parser.add_argument("--prompt-template-video", type=str, default="dit-llm-encode-video")
|
||||
parser.add_argument("--hidden-state-skip-layer", type=int, default=2)
|
||||
parser.add_argument("--apply-final-norm", action="store_true")
|
||||
|
||||
parser.add_argument("--text-encoder-2", type=str, default="clipL")
|
||||
parser.add_argument(
|
||||
"--text-encoder-precision-2",
|
||||
type=str,
|
||||
default="fp16",
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
)
|
||||
parser.add_argument("--text-encoder-precision-2", type=str, default="fp16", choices=["fp32", "fp16", "bf16"])
|
||||
parser.add_argument("--text-states-dim-2", type=int, default=768)
|
||||
parser.add_argument("--tokenizer-2", type=str, default="clipL")
|
||||
parser.add_argument("--text-len-2", type=int, default=77)
|
||||
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
main(args)
|
||||
@@ -0,0 +1,126 @@
|
||||
import os
|
||||
import imageio
|
||||
import time
|
||||
from einops import rearrange
|
||||
|
||||
import torch
|
||||
import torchvision
|
||||
import numpy as np
|
||||
from pathlib import Path
|
||||
from loguru import logger
|
||||
from datetime import datetime
|
||||
import argparse
|
||||
from diffusers.utils import export_to_video
|
||||
|
||||
from fastvideo.models.hunyuan.utils.file_utils import save_videos_grid
|
||||
from fastvideo.models.hunyuan.inference import HunyuanVideoSampler
|
||||
|
||||
def main(args):
|
||||
print(args)
|
||||
models_root_path = Path(args.model_path)
|
||||
if not models_root_path.exists():
|
||||
raise ValueError(f"`models_root` not exists: {models_root_path}")
|
||||
|
||||
# Create save folder to save the samples
|
||||
save_path = args.output_path
|
||||
os.makedirs(os.path.dirname(save_path), exist_ok=True)
|
||||
|
||||
# Load models
|
||||
hunyuan_video_sampler = HunyuanVideoSampler.from_pretrained(models_root_path, args=args)
|
||||
|
||||
# Get the updated args
|
||||
args = hunyuan_video_sampler.args
|
||||
|
||||
# Start sampling
|
||||
samples = []
|
||||
for prompt in args.prompts:
|
||||
outputs = hunyuan_video_sampler.predict(
|
||||
prompt=prompt,
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
video_length=args.num_frames,
|
||||
seed=args.seed,
|
||||
negative_prompt=args.neg_prompt,
|
||||
infer_steps=args.num_inference_steps,
|
||||
guidance_scale=args.guidance_scale,
|
||||
num_videos_per_prompt=args.num_videos,
|
||||
flow_shift=args.flow_shift,
|
||||
batch_size=args.batch_size,
|
||||
embedded_guidance_scale=args.embedded_cfg_scale
|
||||
)
|
||||
samples.append(outputs['samples'][0])
|
||||
|
||||
for prompt, video in zip(args.prompts, samples):
|
||||
videos = rearrange(video.unsqueeze(0), "b c t h w -> t b c h w")
|
||||
outputs = []
|
||||
for x in videos:
|
||||
x = torchvision.utils.make_grid(x, nrow=6)
|
||||
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
|
||||
outputs.append((x * 255).numpy().astype(np.uint8))
|
||||
os.makedirs(os.path.dirname(args.output_path), exist_ok=True)
|
||||
imageio.mimsave(args.output_path + f"{prompt[:100]}.mp4", outputs, fps=args.fps)
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
|
||||
# Basic parameters
|
||||
parser.add_argument("--prompts", nargs="+", default=[])
|
||||
parser.add_argument("--num_frames", type=int, default=16)
|
||||
parser.add_argument("--height", type=int, default=256)
|
||||
parser.add_argument("--width", type=int, default=256)
|
||||
parser.add_argument("--num_inference_steps", type=int, default=50)
|
||||
parser.add_argument("--model_path", type=str, default="data/hunyuan")
|
||||
parser.add_argument("--output_path", type=str, default="./outputs/video")
|
||||
parser.add_argument("--fps", type=int, default=24)
|
||||
|
||||
# Additional parameters
|
||||
parser.add_argument("--denoise-type", type=str, default="flow", help="Denoise type for noised inputs.")
|
||||
parser.add_argument("--seed", type=int, default=None, help="Seed for evaluation.")
|
||||
parser.add_argument("--neg_prompt", type=str, default=None, help="Negative prompt for sampling.")
|
||||
parser.add_argument("--guidance_scale", type=float, default=1.0, help="Classifier free guidance scale.")
|
||||
parser.add_argument("--embedded_cfg_scale", type=float, default=6.0, help="Embedded classifier free guidance scale.")
|
||||
parser.add_argument("--flow_shift", type=int, default=7, help="Flow shift parameter.")
|
||||
parser.add_argument("--batch_size", type=int, default=1, help="Batch size for inference.")
|
||||
parser.add_argument("--num_videos", type=int, default=1, help="Number of videos to generate per prompt.")
|
||||
parser.add_argument("--load-key", type=str, default="module", help="Key to load the model states. 'module' for the main model, 'ema' for the EMA model.")
|
||||
parser.add_argument("--use-cpu-offload", action="store_true", help="Use CPU offload for the model load.")
|
||||
parser.add_argument("--dit-weight", type=str, default="data/hunyuan/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt")
|
||||
parser.add_argument("--reproduce", action="store_true", help="Enable reproducibility by setting random seeds and deterministic algorithms.")
|
||||
parser.add_argument("--disable-autocast", action="store_true", help="Disable autocast for denoising loop and vae decoding in pipeline sampling.")
|
||||
|
||||
# Flow Matching
|
||||
parser.add_argument("--flow-reverse", action="store_true", help="If reverse, learning/sampling from t=1 -> t=0.")
|
||||
parser.add_argument("--flow-solver", type=str, default="euler", help="Solver for flow matching.")
|
||||
parser.add_argument("--use-linear-quadratic-schedule", action="store_true",
|
||||
help="Use linear quadratic schedule for flow matching. Following MovieGen (https://ai.meta.com/static-resource/movie-gen-research-paper)")
|
||||
parser.add_argument("--linear-schedule-end", type=int, default=25,
|
||||
help="End step for linear quadratic schedule for flow matching.")
|
||||
|
||||
# Model parameters
|
||||
parser.add_argument("--model", type=str, default="HYVideo-T/2-cfgdistill")
|
||||
parser.add_argument("--latent-channels", type=int, default=16)
|
||||
parser.add_argument("--precision", type=str, default="bf16", choices=["fp32", "fp16", "bf16"])
|
||||
parser.add_argument("--rope-theta", type=int, default=256, help="Theta used in RoPE.")
|
||||
|
||||
parser.add_argument("--vae", type=str, default="884-16c-hy")
|
||||
parser.add_argument("--vae-precision", type=str, default="fp16", choices=["fp32", "fp16", "bf16"])
|
||||
parser.add_argument("--vae-tiling", action="store_true", default=True)
|
||||
|
||||
parser.add_argument("--text-encoder", type=str, default="llm")
|
||||
parser.add_argument("--text-encoder-precision", type=str, default="fp16", choices=["fp32", "fp16", "bf16"])
|
||||
parser.add_argument("--text-states-dim", type=int, default=4096)
|
||||
parser.add_argument("--text-len", type=int, default=256)
|
||||
parser.add_argument("--tokenizer", type=str, default="llm")
|
||||
parser.add_argument("--prompt-template", type=str, default="dit-llm-encode")
|
||||
parser.add_argument("--prompt-template-video", type=str, default="dit-llm-encode-video")
|
||||
parser.add_argument("--hidden-state-skip-layer", type=int, default=2)
|
||||
parser.add_argument("--apply-final-norm", action="store_true")
|
||||
|
||||
parser.add_argument("--text-encoder-2", type=str, default="clipL")
|
||||
parser.add_argument("--text-encoder-precision-2", type=str, default="fp16", choices=["fp32", "fp16", "bf16"])
|
||||
parser.add_argument("--text-states-dim-2", type=int, default=768)
|
||||
parser.add_argument("--tokenizer-2", type=str, default="clipL")
|
||||
parser.add_argument("--text-len-2", type=int, default=77)
|
||||
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
@@ -39,7 +39,7 @@ def main(args):
|
||||
initialize_distributed()
|
||||
print(nccl_info.sp_size)
|
||||
device = torch.cuda.current_device()
|
||||
# Peiyuan: GPU seed will cause A100 and H100 to produce different results .....
|
||||
generator = torch.Generator(device).manual_seed(args.seed)
|
||||
weight_dtype = torch.bfloat16
|
||||
if args.scheduler_type == "euler":
|
||||
scheduler = FlowMatchEulerDiscreteScheduler()
|
||||
@@ -107,9 +107,9 @@ def main(args):
|
||||
encoder_attention_mask = None
|
||||
|
||||
if prompts is not None:
|
||||
videos = []
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
for prompt in prompts:
|
||||
generator = torch.Generator("cpu").manual_seed(args.seed)
|
||||
video = pipe(
|
||||
prompt=[prompt],
|
||||
height=args.height,
|
||||
@@ -119,17 +119,9 @@ def main(args):
|
||||
guidance_scale=args.guidance_scale,
|
||||
generator=generator,
|
||||
).frames
|
||||
if nccl_info.global_rank <= 0:
|
||||
os.makedirs(args.output_path, exist_ok=True)
|
||||
suffix = prompt.split(".")[0]
|
||||
export_to_video(
|
||||
video[0],
|
||||
os.path.join(args.output_path, f"{suffix}.mp4"),
|
||||
fps=30,
|
||||
)
|
||||
videos.append(video[0])
|
||||
else:
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
generator = torch.Generator("cpu").manual_seed(args.seed)
|
||||
videos = pipe(
|
||||
prompt_embeds=prompt_embeds,
|
||||
prompt_attention_mask=encoder_attention_mask,
|
||||
@@ -141,7 +133,16 @@ def main(args):
|
||||
generator=generator,
|
||||
).frames
|
||||
|
||||
if nccl_info.global_rank <= 0:
|
||||
if nccl_info.global_rank <= 0:
|
||||
if prompts is not None:
|
||||
# mkdir
|
||||
os.makedirs(args.output_path, exist_ok=True)
|
||||
for video, prompt in zip(videos, prompts):
|
||||
suffix = prompt.split(".")[0]
|
||||
export_to_video(
|
||||
video, os.path.join(args.output_path, f"{suffix}.mp4"), fps=30
|
||||
)
|
||||
else:
|
||||
export_to_video(videos[0], args.output_path + ".mp4", fps=30)
|
||||
|
||||
|
||||
|
||||
+22
-27
@@ -30,8 +30,9 @@ from accelerate.utils import set_seed
|
||||
from tqdm.auto import tqdm
|
||||
from fastvideo.utils.fsdp_util import get_dit_fsdp_kwargs, apply_fsdp_checkpointing
|
||||
from diffusers.utils import convert_unet_state_dict_to_peft
|
||||
from diffusers import FlowMatchEulerDiscreteScheduler
|
||||
from fastvideo.utils.load import load_transformer
|
||||
from diffusers import (
|
||||
FlowMatchEulerDiscreteScheduler,
|
||||
)
|
||||
from diffusers.optimization import get_scheduler
|
||||
from fastvideo.models.mochi_hf.modeling_mochi import MochiTransformer3DModel
|
||||
from diffusers.utils import check_min_version
|
||||
@@ -39,7 +40,9 @@ from fastvideo.dataset.latent_datasets import LatentDataset, latent_collate_func
|
||||
import torch.distributed as dist
|
||||
from safetensors.torch import save_file, load_file
|
||||
from peft import LoraConfig, get_peft_model_state_dict, set_peft_model_state_dict
|
||||
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
|
||||
from torch.distributed.fsdp import (
|
||||
FullyShardedDataParallel as FSDP,
|
||||
)
|
||||
from fastvideo.utils.checkpoint import (
|
||||
save_checkpoint,
|
||||
save_lora_checkpoint,
|
||||
@@ -47,7 +50,6 @@ from fastvideo.utils.checkpoint import (
|
||||
)
|
||||
from fastvideo.utils.logging_ import main_print
|
||||
from fastvideo.models.mochi_hf.pipeline_mochi import MochiPipeline
|
||||
|
||||
# Will error if the minimal version of diffusers is not installed. Remove at your own risks.
|
||||
check_min_version("0.31.0")
|
||||
import time
|
||||
@@ -99,9 +101,8 @@ def get_sigmas(noise_scheduler, device, timesteps, n_dim=4, dtype=torch.float32)
|
||||
return sigma
|
||||
|
||||
|
||||
def train_one_step(
|
||||
def train_one_step_mochi(
|
||||
transformer,
|
||||
model_type,
|
||||
optimizer,
|
||||
lr_scheduler,
|
||||
loader,
|
||||
@@ -125,7 +126,7 @@ def train_one_step(
|
||||
latents_attention_mask,
|
||||
encoder_attention_mask,
|
||||
) = next(loader)
|
||||
latents = normalize_dit_input(model_type, latents)
|
||||
latents = normalize_mochi_dit_input(latents)
|
||||
batch_size = latents.shape[0]
|
||||
noise = torch.randn_like(latents)
|
||||
u = compute_density_for_timestep_sampling(
|
||||
@@ -216,15 +217,15 @@ def main(args):
|
||||
|
||||
main_print(f"--> loading model from {args.pretrained_model_name_or_path}")
|
||||
# keep the master weight to float32
|
||||
transformer = load_transformer(
|
||||
args.model_type,
|
||||
args.dit_model_name_or_path,
|
||||
transformer = MochiTransformer3DModel.from_pretrained(
|
||||
args.pretrained_model_name_or_path,
|
||||
torch.float32 if args.master_weight_type == "fp32" else torch.bfloat16,
|
||||
subfolder="transformer",
|
||||
torch_dtype=torch.float32
|
||||
if args.master_weight_type == "fp32"
|
||||
else torch.bfloat16,
|
||||
)
|
||||
|
||||
if args.use_lora:
|
||||
assert args.model_type == "mochi", "LoRA is only supported for Mochi model."
|
||||
transformer.requires_grad_(False)
|
||||
transformer_lora_config = LoraConfig(
|
||||
r=args.lora_rank,
|
||||
@@ -262,8 +263,7 @@ def main(args):
|
||||
main_print(
|
||||
f"--> Initializing FSDP with sharding strategy: {args.fsdp_sharding_startegy}"
|
||||
)
|
||||
fsdp_kwargs, no_split_modules = get_dit_fsdp_kwargs(
|
||||
transformer,
|
||||
fsdp_kwargs = get_dit_fsdp_kwargs(
|
||||
args.fsdp_sharding_startegy,
|
||||
args.use_lora,
|
||||
args.use_cpu_offload,
|
||||
@@ -274,18 +274,17 @@ def main(args):
|
||||
transformer.config.lora_rank = args.lora_rank
|
||||
transformer.config.lora_alpha = args.lora_alpha
|
||||
transformer.config.lora_target_modules = ["to_k", "to_q", "to_v", "to_out.0"]
|
||||
transformer._no_split_modules = [
|
||||
no_split_module.__name__ for no_split_module in no_split_modules
|
||||
]
|
||||
transformer._no_split_modules = ["MochiTransformerBlock"]
|
||||
fsdp_kwargs["auto_wrap_policy"] = fsdp_kwargs["auto_wrap_policy"](transformer)
|
||||
|
||||
transformer = FSDP(transformer, **fsdp_kwargs,)
|
||||
transformer = FSDP(
|
||||
transformer,
|
||||
**fsdp_kwargs,
|
||||
)
|
||||
main_print(f"--> model loaded")
|
||||
|
||||
if args.gradient_checkpointing:
|
||||
apply_fsdp_checkpointing(
|
||||
transformer, no_split_modules, args.selective_checkpointing
|
||||
)
|
||||
apply_fsdp_checkpointing(transformer, args.selective_checkpointing)
|
||||
|
||||
# Set model as trainable.
|
||||
transformer.train()
|
||||
@@ -411,9 +410,8 @@ def main(args):
|
||||
next(loader)
|
||||
for step in range(init_steps + 1, args.max_train_steps + 1):
|
||||
start_time = time.time()
|
||||
loss, grad_norm = train_one_step(
|
||||
loss, grad_norm = train_one_step_mochi(
|
||||
transformer,
|
||||
args.model_type,
|
||||
optimizer,
|
||||
lr_scheduler,
|
||||
loader,
|
||||
@@ -480,9 +478,7 @@ def main(args):
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument(
|
||||
"--model_type", type=str, default="mochi", help="The type of model to train."
|
||||
)
|
||||
|
||||
# dataset & dataloader
|
||||
parser.add_argument("--data_json_path", type=str, required=True)
|
||||
parser.add_argument("--num_frames", type=int, default=163)
|
||||
@@ -506,7 +502,6 @@ if __name__ == "__main__":
|
||||
|
||||
# text encoder & vae & diffusion model
|
||||
parser.add_argument("--pretrained_model_name_or_path", type=str)
|
||||
parser.add_argument("--dit_model_name_or_path", type=str, default=None)
|
||||
parser.add_argument("--cache_dir", type=str, default="./cache_dir")
|
||||
|
||||
# diffusion setting
|
||||
|
||||
@@ -28,7 +28,10 @@ def save_checkpoint(model, optimizer, rank, output_dir, step, discriminator=Fals
|
||||
FullOptimStateDictConfig(offload_to_cpu=True, rank0_only=True),
|
||||
):
|
||||
cpu_state = model.state_dict()
|
||||
optim_state = FSDP.optim_state_dict(model, optimizer,)
|
||||
optim_state = FSDP.optim_state_dict(
|
||||
model,
|
||||
optimizer,
|
||||
)
|
||||
|
||||
# todo move to get_state_dict
|
||||
save_dir = os.path.join(output_dir, f"checkpoint-{step}")
|
||||
@@ -49,10 +52,16 @@ def save_checkpoint(model, optimizer, rank, output_dir, step, discriminator=Fals
|
||||
save_file(cpu_state, weight_path)
|
||||
optimizer_path = os.path.join(save_dir, "discriminator_optimizer.pt")
|
||||
torch.save(optim_state, optimizer_path)
|
||||
|
||||
|
||||
|
||||
def save_checkpoint_generator_discriminator(
|
||||
model, optimizer, discriminator, discriminator_optimizer, rank, output_dir, step,
|
||||
model,
|
||||
optimizer,
|
||||
discriminator,
|
||||
discriminator_optimizer,
|
||||
rank,
|
||||
output_dir,
|
||||
step,
|
||||
):
|
||||
with FSDP.state_dict_type(
|
||||
model,
|
||||
@@ -183,36 +192,6 @@ def resume_training_generator_discriminator(
|
||||
)
|
||||
return model, optimizer, discriminator, discriminator_optimizer, step
|
||||
|
||||
def resume_training_generator_fake_transformer(
|
||||
model,
|
||||
optimizer,
|
||||
fake_transformer,
|
||||
guidance_optimizer,
|
||||
discriminator,
|
||||
discriminator_optimizer,
|
||||
checkpoint_dir,
|
||||
rank,
|
||||
):
|
||||
step = int(checkpoint_dir.split("-")[-1])
|
||||
model_weight_dir = os.path.join(checkpoint_dir, "model_weights_state")
|
||||
model_optimizer_dir = os.path.join(checkpoint_dir, "model_optimizer_state")
|
||||
fake_model_weight_dir = os.path.join(checkpoint_dir, "fake_model_weights_state")
|
||||
fake_model_optimizer_dir = os.path.join(checkpoint_dir, "fake_model_optimizer_state")
|
||||
model, optimizer = load_sharded_model(
|
||||
model, optimizer, model_weight_dir, model_optimizer_dir
|
||||
)
|
||||
fake_transformer, guidance_optimizer = load_sharded_model(
|
||||
fake_transformer, guidance_optimizer, fake_model_weight_dir, fake_model_optimizer_dir
|
||||
)
|
||||
discriminator_ckpt_file = os.path.join(
|
||||
checkpoint_dir, "discriminator_fsdp_state", "discriminator_state.pt"
|
||||
)
|
||||
discriminator, discriminator_optimizer = load_full_state_model(
|
||||
discriminator, discriminator_optimizer, discriminator_ckpt_file, rank
|
||||
)
|
||||
return model, optimizer, fake_transformer, guidance_optimizer, discriminator, discriminator_optimizer, step
|
||||
|
||||
|
||||
|
||||
def resume_training(model, optimizer, checkpoint_dir, discriminator=False):
|
||||
weight_path = os.path.join(checkpoint_dir, "diffusion_pytorch_model.safetensors")
|
||||
@@ -251,7 +230,10 @@ def save_lora_checkpoint(transformer, optimizer, rank, output_dir, step):
|
||||
FullStateDictConfig(offload_to_cpu=True, rank0_only=True),
|
||||
):
|
||||
full_state_dict = transformer.state_dict()
|
||||
lora_optim_state = FSDP.optim_state_dict(transformer, optimizer,)
|
||||
lora_optim_state = FSDP.optim_state_dict(
|
||||
transformer,
|
||||
optimizer,
|
||||
)
|
||||
|
||||
if rank <= 0:
|
||||
save_dir = os.path.join(output_dir, f"lora-checkpoint-{step}")
|
||||
|
||||
@@ -132,7 +132,9 @@ class SeqAllToAll4D(torch.autograd.Function):
|
||||
|
||||
|
||||
def all_to_all_4D(
|
||||
input_: torch.Tensor, scatter_dim: int = 2, gather_dim: int = 1,
|
||||
input_: torch.Tensor,
|
||||
scatter_dim: int = 2,
|
||||
gather_dim: int = 1,
|
||||
):
|
||||
return SeqAllToAll4D.apply(nccl_info.group, input_, scatter_dim, gather_dim)
|
||||
|
||||
@@ -191,7 +193,9 @@ class _AllToAll(torch.autograd.Function):
|
||||
|
||||
|
||||
def all_to_all(
|
||||
input_: torch.Tensor, scatter_dim: int = 2, gather_dim: int = 1,
|
||||
input_: torch.Tensor,
|
||||
scatter_dim: int = 2,
|
||||
gather_dim: int = 1,
|
||||
):
|
||||
return _AllToAll.apply(input_, nccl_info.group, scatter_dim, gather_dim)
|
||||
|
||||
|
||||
@@ -55,7 +55,6 @@ def pad_to_multiple(number, ds_stride):
|
||||
padding = ds_stride - remainder
|
||||
return number + padding
|
||||
|
||||
|
||||
# TODO
|
||||
class Collate:
|
||||
def __init__(self, args):
|
||||
|
||||
@@ -1,40 +0,0 @@
|
||||
import platform
|
||||
|
||||
import accelerate
|
||||
import peft
|
||||
import torch
|
||||
import transformers
|
||||
from transformers.utils import is_torch_cuda_available, is_torch_npu_available
|
||||
|
||||
|
||||
VERSION = "1.2.0"
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
info = {
|
||||
"FastVideo version": VERSION,
|
||||
"Platform": platform.platform(),
|
||||
"Python version": platform.python_version(),
|
||||
"PyTorch version": torch.__version__,
|
||||
"Transformers version": transformers.__version__,
|
||||
"Accelerate version": accelerate.__version__,
|
||||
"PEFT version": peft.__version__,
|
||||
}
|
||||
|
||||
if is_torch_cuda_available():
|
||||
info["PyTorch version"] += " (GPU)"
|
||||
info["GPU type"] = torch.cuda.get_device_name()
|
||||
|
||||
if is_torch_npu_available():
|
||||
info["PyTorch version"] += " (NPU)"
|
||||
info["NPU type"] = torch.npu.get_device_name()
|
||||
info["CANN version"] = torch.version.cann
|
||||
|
||||
try:
|
||||
import bitsandbytes
|
||||
|
||||
info["Bitsandbytes version"] = bitsandbytes.__version__
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
print("\n" + "\n".join([f"- {key}: {value}" for key, value in info.items()]) + "\n")
|
||||
@@ -29,13 +29,14 @@ import functools
|
||||
|
||||
|
||||
non_reentrant_wrapper = partial(
|
||||
checkpoint_wrapper, checkpoint_impl=CheckpointImpl.NO_REENTRANT,
|
||||
checkpoint_wrapper,
|
||||
checkpoint_impl=CheckpointImpl.NO_REENTRANT,
|
||||
)
|
||||
|
||||
check_fn = lambda submodule: isinstance(submodule, MochiTransformerBlock)
|
||||
|
||||
|
||||
def apply_fsdp_checkpointing(model, no_split_modules, p=1):
|
||||
def apply_fsdp_checkpointing(model,no_split_modules, p=1):
|
||||
# https://github.com/foundation-model-stack/fms-fsdp/blob/408c7516d69ea9b6bcd4c0f5efab26c0f64b3c2d/fms_fsdp/policies/ac_handler.py#L16
|
||||
"""apply activation checkpointing to model
|
||||
returns None as model is updated directly
|
||||
@@ -79,18 +80,15 @@ def get_mixed_precision(master_weight_type="fp32"):
|
||||
|
||||
|
||||
def get_dit_fsdp_kwargs(
|
||||
transformer,
|
||||
sharding_strategy,
|
||||
use_lora=False,
|
||||
cpu_offload=False,
|
||||
master_weight_type="fp32",
|
||||
transformer, sharding_strategy, use_lora=False, cpu_offload=False, master_weight_type="fp32"
|
||||
):
|
||||
no_split_modules = get_no_split_modules(transformer)
|
||||
if use_lora:
|
||||
auto_wrap_policy = fsdp_auto_wrap_policy
|
||||
else:
|
||||
auto_wrap_policy = functools.partial(
|
||||
transformer_auto_wrap_policy, transformer_layer_cls=no_split_modules,
|
||||
transformer_auto_wrap_policy,
|
||||
transformer_layer_cls=no_split_modules,
|
||||
)
|
||||
|
||||
# we use float32 for fsdp but autocast during training
|
||||
|
||||
+46
-75
@@ -1,26 +1,17 @@
|
||||
import torch
|
||||
from fastvideo.models.mochi_hf.modeling_mochi import (
|
||||
MochiTransformer3DModel,
|
||||
MochiTransformerBlock,
|
||||
)
|
||||
from fastvideo.models.hunyuan.modules.models import (
|
||||
HYVideoDiffusionTransformer,
|
||||
MMDoubleStreamBlock,
|
||||
MMSingleStreamBlock,
|
||||
)
|
||||
from fastvideo.models.mochi_hf.modeling_mochi import MochiTransformer3DModel, MochiTransformerBlock
|
||||
from fastvideo.models.hunyuan.modules.models import HYVideoDiffusionTransformer, MMDoubleStreamBlock, MMSingleStreamBlock
|
||||
from fastvideo.models.hunyuan.vae.autoencoder_kl_causal_3d import AutoencoderKLCausal3D
|
||||
from diffusers import AutoencoderKLMochi
|
||||
from transformers import T5EncoderModel, AutoTokenizer
|
||||
import os
|
||||
import os
|
||||
from torch import nn
|
||||
|
||||
# Path
|
||||
from pathlib import Path
|
||||
import torch.nn.functional as F
|
||||
from fastvideo.models.hunyuan.text_encoder import TextEncoder
|
||||
from fastvideo.utils.logging_ import main_print
|
||||
|
||||
hunyuan_config = {
|
||||
hunyuan_config = {
|
||||
"mm_double_blocks_depth": 20,
|
||||
"mm_single_blocks_depth": 40,
|
||||
"rope_dim_list": [16, 56, 56],
|
||||
@@ -35,7 +26,7 @@ PROMPT_TEMPLATE_ENCODE = (
|
||||
"<|start_header_id|>system<|end_header_id|>\n\nDescribe the image by detailing the color, shape, size, texture, "
|
||||
"quantity, text, spatial relationships of the objects and background:<|eot_id|>"
|
||||
"<|start_header_id|>user<|end_header_id|>\n\n{}<|eot_id|>"
|
||||
)
|
||||
)
|
||||
PROMPT_TEMPLATE_ENCODE_VIDEO = (
|
||||
"<|start_header_id|>system<|end_header_id|>\n\nDescribe the video by detailing the following aspects: "
|
||||
"1. The main content and theme of the video."
|
||||
@@ -44,31 +35,34 @@ PROMPT_TEMPLATE_ENCODE_VIDEO = (
|
||||
"4. background environment, light, style and atmosphere."
|
||||
"5. camera angles, movements, and transitions used in the video:<|eot_id|>"
|
||||
"<|start_header_id|>user<|end_header_id|>\n\n{}<|eot_id|>"
|
||||
)
|
||||
)
|
||||
|
||||
NEGATIVE_PROMPT = "Aerial view, aerial view, overexposed, low quality, deformation, a poor composition, bad hands, bad teeth, bad eyes, bad limbs, distortion"
|
||||
|
||||
PROMPT_TEMPLATE = {
|
||||
"dit-llm-encode": {"template": PROMPT_TEMPLATE_ENCODE, "crop_start": 36,},
|
||||
"dit-llm-encode": {
|
||||
"template": PROMPT_TEMPLATE_ENCODE,
|
||||
"crop_start": 36,
|
||||
},
|
||||
"dit-llm-encode-video": {
|
||||
"template": PROMPT_TEMPLATE_ENCODE_VIDEO,
|
||||
"crop_start": 95,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
|
||||
class HunyuanTextEncoderWrapper(nn.Module):
|
||||
def __init__(self, pretrained_model_name_or_path, device):
|
||||
super().__init__()
|
||||
|
||||
text_len = 256
|
||||
crop_start = PROMPT_TEMPLATE["dit-llm-encode-video"].get("crop_start", 0)
|
||||
|
||||
|
||||
text_len = 256
|
||||
crop_start = PROMPT_TEMPLATE["dit-llm-encode-video"].get("crop_start", 0 )
|
||||
|
||||
max_length = text_len + crop_start
|
||||
|
||||
# prompt_template
|
||||
prompt_template = PROMPT_TEMPLATE["dit-llm-encode"]
|
||||
|
||||
|
||||
# prompt_template_video
|
||||
prompt_template_video = PROMPT_TEMPLATE["dit-llm-encode-video"]
|
||||
text_encoder_path = os.path.join(pretrained_model_name_or_path, "text_encoder")
|
||||
@@ -86,9 +80,7 @@ class HunyuanTextEncoderWrapper(nn.Module):
|
||||
logger=None,
|
||||
device=device,
|
||||
)
|
||||
text_encoder_path_2 = os.path.join(
|
||||
pretrained_model_name_or_path, "text_encoder_2"
|
||||
)
|
||||
text_encoder_path_2 = os.path.join(pretrained_model_name_or_path, "text_encoder_2")
|
||||
self.text_encoder_2 = TextEncoder(
|
||||
text_encoder_type="clipL",
|
||||
text_encoder_path=text_encoder_path_2,
|
||||
@@ -157,33 +149,26 @@ class HunyuanTextEncoderWrapper(nn.Module):
|
||||
bs_embed * num_videos_per_prompt, seq_len, -1
|
||||
)
|
||||
return (prompt_embeds, attention_mask)
|
||||
|
||||
|
||||
def encode_prompt(self, prompt):
|
||||
prompt_embeds, attention_mask = self.encode_(prompt, self.text_encoder)
|
||||
prompt_embeds_2, attention_mask_2 = self.encode_(prompt, self.text_encoder_2)
|
||||
prompt_embeds_2 = F.pad(
|
||||
prompt_embeds_2,
|
||||
(0, prompt_embeds.shape[2] - prompt_embeds_2.shape[1]),
|
||||
value=0,
|
||||
).unsqueeze(1)
|
||||
prompt_embeds_2 = F.pad(prompt_embeds_2, (0, prompt_embeds.shape[2] - prompt_embeds_2.shape[1]), value=0).unsqueeze(1)
|
||||
prompt_embeds = torch.cat([prompt_embeds_2, prompt_embeds], dim=1)
|
||||
return prompt_embeds, attention_mask
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
class MochiTextEncoderWrapper(nn.Module):
|
||||
def __init__(self, pretrained_model_name_or_path, device):
|
||||
super().__init__()
|
||||
self.text_encoder = T5EncoderModel.from_pretrained(
|
||||
os.path.join(pretrained_model_name_or_path, "text_encoder")
|
||||
).to(device)
|
||||
self.tokenizer = AutoTokenizer.from_pretrained(
|
||||
os.path.join(pretrained_model_name_or_path, "tokenizer")
|
||||
)
|
||||
self.text_encoder = T5EncoderModel.from_pretrained(os.path.join(pretrained_model_name_or_path, "text_encoder")).to(device)
|
||||
self.tokenizer = AutoTokenizer.from_pretrained(os.path.join(pretrained_model_name_or_path, "text_encoder"))
|
||||
self.max_sequence_length = 256
|
||||
|
||||
def encode_prompt(self, prompt):
|
||||
device = self.text_encoder.device
|
||||
dtype = self.text_encoder.dtype
|
||||
device = self.text_encoder.device
|
||||
dtype = self.dtype
|
||||
|
||||
prompt = [prompt] if isinstance(prompt, str) else prompt
|
||||
batch_size = len(prompt)
|
||||
@@ -220,20 +205,19 @@ class MochiTextEncoderWrapper(nn.Module):
|
||||
|
||||
# duplicate text embeddings for each generation per prompt, using mps friendly method
|
||||
_, seq_len, _ = prompt_embeds.shape
|
||||
prompt_embeds = prompt_embeds.view(batch_size, seq_len, -1)
|
||||
prompt_embeds = prompt_embeds.view(
|
||||
batch_size , seq_len, -1
|
||||
)
|
||||
prompt_attention_mask = prompt_attention_mask.view(batch_size, -1)
|
||||
|
||||
return prompt_embeds, prompt_attention_mask
|
||||
|
||||
|
||||
def load_hunyuan_state_dict(model, dit_model_name_or_path):
|
||||
load_key = "module"
|
||||
model_path = dit_model_name_or_path
|
||||
bare_model = "unknown"
|
||||
|
||||
state_dict = torch.load(
|
||||
model_path, map_location=lambda storage, loc: storage, weights_only=True
|
||||
)
|
||||
|
||||
state_dict = torch.load(model_path, map_location=lambda storage, loc: storage, weights_only=True)
|
||||
|
||||
if bare_model == "unknown" and ("ema" in state_dict or "module" in state_dict):
|
||||
bare_model = False
|
||||
@@ -247,14 +231,8 @@ def load_hunyuan_state_dict(model, dit_model_name_or_path):
|
||||
)
|
||||
model.load_state_dict(state_dict, strict=True)
|
||||
return model
|
||||
|
||||
|
||||
def load_transformer(
|
||||
model_type,
|
||||
dit_model_name_or_path,
|
||||
pretrained_model_name_or_path,
|
||||
master_weight_type,
|
||||
):
|
||||
|
||||
def load_transformer(model_type,dit_model_name_or_path, pretrained_model_name_or_path, master_weight_type):
|
||||
if model_type == "mochi":
|
||||
if dit_model_name_or_path:
|
||||
transformer = MochiTransformer3DModel.from_pretrained(
|
||||
@@ -271,16 +249,16 @@ def load_transformer(
|
||||
)
|
||||
elif model_type == "hunyuan":
|
||||
transformer = HYVideoDiffusionTransformer(
|
||||
in_channels=16, out_channels=16, **hunyuan_config, dtype=master_weight_type,
|
||||
in_channels=16,
|
||||
out_channels=16,
|
||||
**hunyuan_config,
|
||||
dtype=master_weight_type,
|
||||
)
|
||||
transformer = load_hunyuan_state_dict(transformer, dit_model_name_or_path)
|
||||
if master_weight_type == torch.bfloat16:
|
||||
transformer = transformer.bfloat16()
|
||||
else:
|
||||
raise ValueError(f"Unsupported model type: {model_type}")
|
||||
return transformer
|
||||
|
||||
|
||||
def load_vae(model_type, pretrained_model_name_or_path):
|
||||
weight_dtype = torch.float32
|
||||
if model_type == "mochi":
|
||||
@@ -291,25 +269,19 @@ def load_vae(model_type, pretrained_model_name_or_path):
|
||||
fps = 30
|
||||
elif model_type == "hunyuan":
|
||||
vae_precision = torch.float32
|
||||
vae_path = os.path.join(
|
||||
pretrained_model_name_or_path, "hunyuan-video-t2v-720p/vae"
|
||||
)
|
||||
|
||||
vae_path = os.path.join(pretrained_model_name_or_path, "hunyuan-video-t2v-720p/vae")
|
||||
|
||||
config = AutoencoderKLCausal3D.load_config(vae_path)
|
||||
vae = AutoencoderKLCausal3D.from_config(config)
|
||||
|
||||
|
||||
vae_ckpt = Path(vae_path) / "pytorch_model.pt"
|
||||
assert vae_ckpt.exists(), f"VAE checkpoint not found: {vae_ckpt}"
|
||||
|
||||
|
||||
ckpt = torch.load(vae_ckpt, map_location=vae.device, weights_only=True)
|
||||
if "state_dict" in ckpt:
|
||||
ckpt = ckpt["state_dict"]
|
||||
if any(k.startswith("vae.") for k in ckpt.keys()):
|
||||
ckpt = {
|
||||
k.replace("vae.", ""): v
|
||||
for k, v in ckpt.items()
|
||||
if k.startswith("vae.")
|
||||
}
|
||||
ckpt = {k.replace("vae.", ""): v for k, v in ckpt.items() if k.startswith("vae.")}
|
||||
vae.load_state_dict(ckpt)
|
||||
vae = vae.to(dtype=vae_precision)
|
||||
vae.requires_grad_(False)
|
||||
@@ -318,6 +290,7 @@ def load_vae(model_type, pretrained_model_name_or_path):
|
||||
autocast_type = torch.float32
|
||||
fps = 24
|
||||
return vae, autocast_type, fps
|
||||
|
||||
|
||||
|
||||
def load_text_encoder(model_type, pretrained_model_name_or_path, device):
|
||||
@@ -329,7 +302,6 @@ def load_text_encoder(model_type, pretrained_model_name_or_path, device):
|
||||
raise ValueError(f"Unsupported model type: {model_type}")
|
||||
return text_encoder
|
||||
|
||||
|
||||
def get_no_split_modules(transformer):
|
||||
# if of type MochiTransformer3DModel
|
||||
if isinstance(transformer, MochiTransformer3DModel):
|
||||
@@ -338,12 +310,11 @@ def get_no_split_modules(transformer):
|
||||
return (MMDoubleStreamBlock, MMSingleStreamBlock)
|
||||
else:
|
||||
raise ValueError(f"Unsupported transformer type: {type(transformer)}")
|
||||
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# test encode prompt
|
||||
device = torch.cuda.current_device()
|
||||
pretrained_model_name_or_path = "data/hunyuan"
|
||||
text_encoder = load_text_encoder("hunyuan", pretrained_model_name_or_path, device)
|
||||
prompt = "A man on stage claps his hands together while facing the audience. The audience, visible in the foreground, holds up mobile devices to record the event, capturing the moment from various angles. The background features a large banner with text identifying the man on stage. Throughout the sequence, the man's expression remains engaged and directed towards the audience. The camera angle remains constant, focusing on capturing the interaction between the man on stage and the audience."
|
||||
prompt = "A man on stage claps his hands together while facing the audience. The audience, visible in the foreground, holds up mobile devices to record the event, capturing the moment from various angles. The background features a large banner with text identifying the man on stage. Throughout the sequence, the man's expression remains engaged and directed towards the audience. The camera angle remains constant, focusing on capturing the interaction between the man on stage and the audience."
|
||||
prompt_embeds, attention_mask = text_encoder.encode_prompt(prompt)
|
||||
|
||||
@@ -1,23 +1,20 @@
|
||||
from accelerate.logging import get_logger
|
||||
import torch
|
||||
|
||||
def get_optimizer(
|
||||
params_to_optimize,
|
||||
args,
|
||||
lr=1e-5,
|
||||
betas=(0.9, 0.999),
|
||||
weight_decay=1e-3,
|
||||
eps=1e-8,
|
||||
):
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
def get_optimizer(args, params_to_optimize, use_deepspeed: bool = False):
|
||||
# Optimizer creation
|
||||
supported_optimizers = ["adam", "adamw"]
|
||||
supported_optimizers = ["adam", "adamw", "prodigy"]
|
||||
if args.optimizer not in supported_optimizers:
|
||||
print(
|
||||
logger.warning(
|
||||
f"Unsupported choice of optimizer: {args.optimizer}. Supported optimizers include {supported_optimizers}. Defaulting to AdamW"
|
||||
)
|
||||
args.optimizer = "adamw"
|
||||
|
||||
|
||||
if args.use_8bit_adam and not (args.optimizer.lower() not in ["adam", "adamw"]):
|
||||
print(
|
||||
logger.warning(
|
||||
f"use_8bit_adam is ignored when optimizer is not set to 'Adam' or 'AdamW'. Optimizer was "
|
||||
f"set to {args.optimizer.lower()}"
|
||||
)
|
||||
@@ -37,20 +34,44 @@ def get_optimizer(
|
||||
|
||||
optimizer = optimizer_class(
|
||||
params_to_optimize,
|
||||
lr=lr,
|
||||
betas=betas,
|
||||
eps=eps,
|
||||
weight_decay=weight_decay,
|
||||
betas=(args.adam_beta1, args.adam_beta2),
|
||||
eps=args.adam_epsilon,
|
||||
weight_decay=args.adam_weight_decay,
|
||||
)
|
||||
elif args.optimizer.lower() == "adam":
|
||||
optimizer_class = bnb.optim.Adam8bit if args.use_8bit_adam else torch.optim.Adam
|
||||
|
||||
optimizer = optimizer_class(
|
||||
params_to_optimize,
|
||||
lr=lr,
|
||||
betas=betas,
|
||||
eps=eps,
|
||||
weight_decay=weight_decay,
|
||||
betas=(args.adam_beta1, args.adam_beta2),
|
||||
eps=args.adam_epsilon,
|
||||
weight_decay=args.adam_weight_decay,
|
||||
)
|
||||
elif args.optimizer.lower() == "prodigy":
|
||||
try:
|
||||
import prodigyopt
|
||||
except ImportError:
|
||||
raise ImportError(
|
||||
"To use Prodigy, please install the prodigyopt library: `pip install prodigyopt`"
|
||||
)
|
||||
|
||||
optimizer_class = prodigyopt.Prodigy
|
||||
|
||||
if args.learning_rate <= 0.1:
|
||||
logger.warning(
|
||||
"Learning rate is too low. When using prodigy, it's generally better to set learning rate around 1.0"
|
||||
)
|
||||
|
||||
optimizer = optimizer_class(
|
||||
params_to_optimize,
|
||||
lr=args.learning_rate,
|
||||
betas=(args.adam_beta1, args.adam_beta2),
|
||||
beta3=args.prodigy_beta3,
|
||||
weight_decay=args.adam_weight_decay,
|
||||
eps=args.adam_epsilon,
|
||||
decouple=args.prodigy_decouple,
|
||||
use_bias_correction=args.prodigy_use_bias_correction,
|
||||
safeguard_warmup=args.prodigy_safeguard_warmup,
|
||||
)
|
||||
|
||||
return optimizer
|
||||
|
||||
@@ -23,7 +23,6 @@ import wandb
|
||||
import gc
|
||||
from fastvideo.utils.load import load_vae
|
||||
|
||||
|
||||
def prepare_latents(
|
||||
batch_size,
|
||||
num_channels_latents,
|
||||
@@ -46,7 +45,6 @@ def prepare_latents(
|
||||
return latents
|
||||
|
||||
|
||||
|
||||
def sample_validation_video(
|
||||
transformer,
|
||||
vae,
|
||||
@@ -67,7 +65,7 @@ def sample_validation_video(
|
||||
output_type: Optional[str] = "pil",
|
||||
vae_spatial_scale_factor=8,
|
||||
vae_temporal_scale_factor=6,
|
||||
num_channels_latents=12,
|
||||
num_channels_latents=12
|
||||
):
|
||||
device = vae.device
|
||||
|
||||
@@ -108,11 +106,17 @@ def sample_validation_video(
|
||||
sigmas = np.array(sigmas)
|
||||
if scheduler_type == "euler":
|
||||
timesteps, num_inference_steps = retrieve_timesteps(
|
||||
scheduler, num_inference_steps, device, timesteps, sigmas,
|
||||
scheduler,
|
||||
num_inference_steps,
|
||||
device,
|
||||
timesteps,
|
||||
sigmas,
|
||||
)
|
||||
else:
|
||||
timesteps, num_inference_steps = retrieve_timesteps(
|
||||
scheduler, num_inference_steps, device,
|
||||
scheduler,
|
||||
num_inference_steps,
|
||||
device,
|
||||
)
|
||||
num_warmup_steps = max(len(timesteps) - num_inference_steps * scheduler.order, 0)
|
||||
|
||||
@@ -208,7 +212,7 @@ def log_validation(
|
||||
args,
|
||||
transformer,
|
||||
device,
|
||||
weight_dtype, # TODO
|
||||
weight_dtype, # TODO
|
||||
global_step,
|
||||
scheduler_type="euler",
|
||||
shift=1.0,
|
||||
@@ -220,18 +224,16 @@ def log_validation(
|
||||
# TODO
|
||||
print(f"Running validation....\n")
|
||||
if args.model_type == "mochi":
|
||||
vae_spatial_scale_factor = 8
|
||||
vae_temporal_scale_factor = 6
|
||||
num_channels_latents = 12
|
||||
vae_spatial_scale_factor=8
|
||||
vae_temporal_scale_factor=6
|
||||
num_channels_latents=12
|
||||
elif args.model_type == "hunyuan":
|
||||
vae_spatial_scale_factor = 8
|
||||
vae_temporal_scale_factor = 4
|
||||
num_channels_latents = 16
|
||||
vae_spatial_scale_factor=8
|
||||
vae_temporal_scale_factor=4
|
||||
num_channels_latents=16
|
||||
else:
|
||||
raise ValueError(f"Model type {args.model_type} not supported")
|
||||
vae, autocast_type, fps = load_vae(
|
||||
args.model_type, args.pretrained_model_name_or_path
|
||||
)
|
||||
vae, autocast_type, fps = load_vae(args.model_type, args.pretrained_model_name_or_path)
|
||||
vae.enable_tiling()
|
||||
if scheduler_type == "euler":
|
||||
scheduler = FlowMatchEulerDiscreteScheduler()
|
||||
@@ -266,19 +268,20 @@ def log_validation(
|
||||
num_sp_groups = int(os.getenv("WORLD_SIZE", "1")) // nccl_info.sp_size
|
||||
# pad to multiple of groups
|
||||
if num_embeds % num_sp_groups != 0:
|
||||
validation_prompt_ids += [0] * (
|
||||
num_sp_groups - num_embeds % num_sp_groups
|
||||
)
|
||||
validation_prompt_ids += [0] * (num_sp_groups - num_embeds % num_sp_groups)
|
||||
num_embeds_per_group = len(validation_prompt_ids) // num_sp_groups
|
||||
local_prompt_ids = validation_prompt_ids[
|
||||
nccl_info.group_id
|
||||
* num_embeds_per_group : (nccl_info.group_id + 1)
|
||||
nccl_info.group_id * num_embeds_per_group : (nccl_info.group_id + 1)
|
||||
* num_embeds_per_group
|
||||
]
|
||||
|
||||
for i in local_prompt_ids:
|
||||
prompt_embed_path = os.path.join(embe_dir, f"{embeds[i]}")
|
||||
prompt_mask_path = os.path.join(mask_dir, f"{masks[i]}")
|
||||
prompt_embed_path = os.path.join(
|
||||
embe_dir, f"{embeds[i]}"
|
||||
)
|
||||
prompt_mask_path = os.path.join(
|
||||
mask_dir, f"{masks[i]}"
|
||||
)
|
||||
prompt_embeds = (
|
||||
torch.load(prompt_embed_path, map_location="cpu", weights_only=True)
|
||||
.to(device)
|
||||
@@ -289,7 +292,9 @@ def log_validation(
|
||||
.to(device)
|
||||
.unsqueeze(0)
|
||||
)
|
||||
negative_prompt_embeds = torch.zeros(256, 4096).to(device).unsqueeze(0)
|
||||
negative_prompt_embeds = (
|
||||
torch.zeros(256, 4096).to(device).unsqueeze(0)
|
||||
)
|
||||
negative_prompt_attention_mask = (
|
||||
torch.zeros(256).bool().to(device).unsqueeze(0)
|
||||
)
|
||||
@@ -312,7 +317,7 @@ def log_validation(
|
||||
negative_prompt_attention_mask=negative_prompt_attention_mask,
|
||||
vae_spatial_scale_factor=vae_spatial_scale_factor,
|
||||
vae_temporal_scale_factor=vae_temporal_scale_factor,
|
||||
num_channels_latents=num_channels_latents,
|
||||
num_channels_latents=num_channels_latents
|
||||
)[0]
|
||||
if nccl_info.rank_within_group == 0:
|
||||
videos.append(video[0])
|
||||
|
||||
@@ -0,0 +1,12 @@
|
||||
python3 fastvideo/sample/sample_t2v_hunyuan_no_sp.py \
|
||||
--height 500 \
|
||||
--width 700 \
|
||||
--num_frames 29 \
|
||||
--num_inference_steps 50 \
|
||||
--guidance_scale 1 \
|
||||
--embedded_cfg_scale 6 \
|
||||
--flow-reverse \
|
||||
--prompts "A cat walks on the grass, realistic style." \
|
||||
--prompts "A dog runs in the park, realistic style." \
|
||||
--seed 42 \
|
||||
--output_path outputs_video/hunyuan/
|
||||
@@ -0,0 +1,39 @@
|
||||
num_gpus=4
|
||||
|
||||
torchrun --nnodes=1 --nproc_per_node=$num_gpus --master_port 29503 \
|
||||
fastvideo/sample/sample_t2v_hunyuan.py \
|
||||
--height 512 \
|
||||
--width 512 \
|
||||
--num_frames 29 \
|
||||
--num_inference_steps 4 \
|
||||
--guidance_scale 1 \
|
||||
--embedded_cfg_scale 6 \
|
||||
--flow_shift 17 \
|
||||
--flow-reverse \
|
||||
--prompts "A man on stage claps his hands together while facing the audience. The audience, visible in the foreground, holds up mobile devices to record the event, capturing the moment from various angles. The background features a large banner with text identifying the man on stage. Throughout the sequence, the man's expression remains engaged and directed towards the audience. The camera angle remains constant, focusing on capturing the interaction between the man on stage and the audience."\
|
||||
--seed 12345 \
|
||||
--output_path outputs_video/hunyuan/
|
||||
|
||||
|
||||
tensor(-0.1065, device='cuda:0', dtype=torch.float16)
|
||||
tensor(-0.0034, device='cuda:2', dtype=torch.float16)
|
||||
|
||||
tensor(-0.0230, device='cuda:0', dtype=torch.float16)
|
||||
>>> weight[0, :768].mean()
|
||||
tensor(-0.0367, device='cuda:0', dtype=torch.float16)
|
||||
|
||||
num_gpus=1
|
||||
|
||||
torchrun --nnodes=1 --nproc_per_node=$num_gpus --master_port 29503 \
|
||||
fastvideo/sample/sample_t2v_hunyuan.py \
|
||||
--height 480 \
|
||||
--width 848 \
|
||||
--num_frames 93 \
|
||||
--num_inference_steps 50 \
|
||||
--guidance_scale 1 \
|
||||
--embedded_cfg_scale 6 \
|
||||
--flow_shift 17 \
|
||||
--flow-reverse \
|
||||
--prompts "A man on stage claps his hands together while facing the audience. The audience, visible in the foreground, holds up mobile devices to record the event, capturing the moment from various angles. The background features a large banner with text identifying the man on stage. Throughout the sequence, the man's expression remains engaged and directed towards the audience. The camera angle remains constant, focusing on capturing the interaction between the man on stage and the audience."\
|
||||
--seed 12345 \
|
||||
--output_path outputs_video/hunyuan/
|
||||
@@ -0,0 +1,20 @@
|
||||
|
||||
|
||||
|
||||
num_gpus=4
|
||||
|
||||
torchrun --nnodes=1 --nproc_per_node=$num_gpus --master_port 29503 \
|
||||
fastvideo/sample/sample_t2v_mochi.py \
|
||||
--model_path data/mochi \
|
||||
--prompt_embed_path "data/synthetic_debug2/prompt_embed/2.pt" \
|
||||
--encoder_attention_mask_path "data/synthetic_debug2/prompt_attention_mask/1.pt" \
|
||||
--num_frames 163 \
|
||||
--height 480 \
|
||||
--width 848 \
|
||||
--num_inference_steps 32 \
|
||||
--guidance_scale 4.5 \
|
||||
--output_path outputs_video/debug \
|
||||
--shift 8 \
|
||||
--seed 12345 \
|
||||
--scheduler_type "pcm_linear_quadratic"
|
||||
|
||||
@@ -0,0 +1,11 @@
|
||||
python fastvideo/sample/sample_t2v_mochi_no_sp.py \
|
||||
--model_path data/mochi \
|
||||
--prompts "A hand enters the frame, pulling a sheet of plastic wrap over three balls of dough placed on a wooden surface. The plastic wrap is stretched to cover the dough more securely. The hand adjusts the wrap, ensuring that it is tight and smooth over the dough. The scene focuses on the hand's movements as it secures the edges of the plastic wrap. No new objects appear, and the camera remains stationary, focusing on the action of covering the dough." \
|
||||
--num_frames 163 \
|
||||
--height 480 \
|
||||
--width 848 \
|
||||
--num_inference_steps 64 \
|
||||
--guidance_scale 0.0 \
|
||||
--seed 12346 \
|
||||
--transformer_path data/outputs/debug/checkpoint-100/transformer \
|
||||
--output_path outputs_video/single_no_guidance
|
||||
-130
@@ -1,130 +0,0 @@
|
||||
# Prediction interface for Cog ⚙️
|
||||
# https://cog.run/python
|
||||
|
||||
from cog import BasePredictor, Input, Path
|
||||
import os
|
||||
import time
|
||||
import torch
|
||||
import imageio
|
||||
import argparse
|
||||
import subprocess
|
||||
import torchvision
|
||||
import numpy as np
|
||||
from einops import rearrange
|
||||
|
||||
MODEL_CACHE = 'FastHunyuan'
|
||||
os.environ['MODEL_BASE'] = './'+MODEL_CACHE
|
||||
from fastvideo.models.hunyuan.inference import HunyuanVideoSampler
|
||||
MODEL_URL = "https://weights.replicate.delivery/default/FastVideo/FastHunyuan/model.tar"
|
||||
|
||||
def download_weights(url, dest):
|
||||
start = time.time()
|
||||
print("downloading url: ", url)
|
||||
print("downloading to: ", dest)
|
||||
subprocess.check_call(["pget", "-xf", url, dest], close_fds=False)
|
||||
print("downloading took: ", time.time() - start)
|
||||
|
||||
class Predictor(BasePredictor):
|
||||
def setup(self):
|
||||
"""Load the model into memory"""
|
||||
print("Model Base: " + os.environ['MODEL_BASE'])
|
||||
# Download weights
|
||||
if not os.path.exists(MODEL_CACHE):
|
||||
download_weights(MODEL_URL, MODEL_CACHE)
|
||||
|
||||
self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
args = argparse.Namespace(
|
||||
num_frames=125,
|
||||
height=720,
|
||||
width=1280,
|
||||
num_inference_steps=6,
|
||||
fps=24,
|
||||
denoise_type='flow',
|
||||
seed=1024,
|
||||
neg_prompt=None,
|
||||
guidance_scale=1.0,
|
||||
embedded_cfg_scale=6.0,
|
||||
flow_shift=17,
|
||||
batch_size=1,
|
||||
num_videos=1,
|
||||
load_key='module',
|
||||
use_cpu_offload=False,
|
||||
dit_weight='FastHunyuan/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt',
|
||||
reproduce=True,
|
||||
disable_autocast=False,
|
||||
flow_reverse=True,
|
||||
flow_solver='euler',
|
||||
use_linear_quadratic_schedule=False,
|
||||
linear_schedule_end=25,
|
||||
model='HYVideo-T/2-cfgdistill',
|
||||
latent_channels=16,
|
||||
precision='bf16',
|
||||
rope_theta=256,
|
||||
vae='884-16c-hy',
|
||||
vae_precision='fp16',
|
||||
vae_tiling=True,
|
||||
text_encoder='llm',
|
||||
text_encoder_precision='fp16',
|
||||
text_states_dim=4096,
|
||||
text_len=256,
|
||||
tokenizer='llm',
|
||||
prompt_template='dit-llm-encode',
|
||||
prompt_template_video='dit-llm-encode-video',
|
||||
hidden_state_skip_layer=2,
|
||||
apply_final_norm=False,
|
||||
text_encoder_2='clipL',
|
||||
text_encoder_precision_2='fp16',
|
||||
text_states_dim_2=768,
|
||||
tokenizer_2='clipL',
|
||||
text_len_2=77,
|
||||
model_path=MODEL_CACHE,
|
||||
)
|
||||
self.model = HunyuanVideoSampler.from_pretrained(MODEL_CACHE, args=args)
|
||||
|
||||
def predict(
|
||||
self,
|
||||
prompt: str = Input(description="Text prompt for video generation", default="A cat walks on the grass, realistic style."),
|
||||
negative_prompt: str = Input(description="Text prompt to specify what you don't want in the video.", default=""),
|
||||
width: int = Input(description="Width of output video", default=1280, ge=256),
|
||||
height: int = Input(description="Height of output video", default=720, ge=256),
|
||||
num_frames: int = Input(description="Number of frames to generate", default=125, ge=16),
|
||||
num_inference_steps: int = Input(description="Number of denoising steps", default=6, ge=1, le=50),
|
||||
guidance_scale: float = Input(description="Classifier free guidance scale", default=1.0, ge=0.1, le=10.0),
|
||||
embedded_cfg_scale: float = Input(description="Embedded classifier free guidance scale", default=6.0, ge=0.1, le=10.0),
|
||||
flow_shift: int = Input(description="Flow shift parameter", default=17, ge=1, le=20),
|
||||
fps: int = Input(description="Frames per second of output video", default=24, ge=1, le=60),
|
||||
seed: int = Input(description="0 for Random seed. Set for reproducible generation", default=0),
|
||||
) -> Path:
|
||||
"""Run video generation"""
|
||||
if seed <=0:
|
||||
seed = int.from_bytes(os.urandom(2), "big")
|
||||
print(f"Using seed: {seed}")
|
||||
|
||||
outputs = self.model.predict(
|
||||
prompt=prompt,
|
||||
height=height,
|
||||
width=width,
|
||||
video_length=num_frames,
|
||||
seed=seed,
|
||||
negative_prompt=negative_prompt,
|
||||
infer_steps=num_inference_steps,
|
||||
guidance_scale=guidance_scale,
|
||||
embedded_guidance_scale=embedded_cfg_scale,
|
||||
flow_shift=flow_shift,
|
||||
flow_reverse=True,
|
||||
batch_size=1,
|
||||
num_videos_per_prompt=1,
|
||||
)
|
||||
|
||||
# Process output video
|
||||
videos = rearrange(outputs["samples"], "b c t h w -> t b c h w")
|
||||
frames = []
|
||||
for x in videos:
|
||||
x = torchvision.utils.make_grid(x, nrow=6)
|
||||
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
|
||||
frames.append((x * 255).numpy().astype(np.uint8))
|
||||
|
||||
# Save video
|
||||
output_path = Path("/tmp/output.mp4")
|
||||
imageio.mimsave(str(output_path), frames, fps=fps)
|
||||
return Path(output_path)
|
||||
+6
-1
@@ -21,9 +21,14 @@ dependencies = [
|
||||
"timm==1.0.11", "torchdiffeq==0.2.4", "torchmetrics==1.5.1", "tqdm==4.66.5", "urllib3==2.2.0", "uvicorn==0.32.0",
|
||||
"scikit-video==1.1.11", "imageio-ffmpeg==0.5.1", "sentencepiece==0.2.0", "beautifulsoup4==4.12.3", "ftfy==6.3.0",
|
||||
"moviepy==1.0.3", "wandb==0.18.5", "tensorboard==2.18.0", "pydantic==2.9.2", "gradio==5.3.0", "huggingface_hub==0.26.1", "protobuf==5.28.3",
|
||||
"watch", "gpustat", "peft==0.13.2", "liger_kernel==0.4.1", "einops==0.8.0", "wheel==0.44.0", "loguru", "diffusers==0.32.0", "bitsandbytes"]
|
||||
"watch", "gpustat", "peft==0.13.2", "liger_kernel==0.4.1", "einops==0.8.0", "wheel==0.44.0"]
|
||||
|
||||
|
||||
[project.optional-dependencies]
|
||||
hunyuan = [
|
||||
"loguru"
|
||||
]
|
||||
|
||||
[tool.setuptools.packages.find]
|
||||
exclude = ["assets*", "docker*", "docs", "scripts*"]
|
||||
|
||||
|
||||
@@ -1,49 +0,0 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export WANDB_API_KEY=7abd9763ff10869f9e88526d919c2639766139a8
|
||||
|
||||
DATA_DIR=./data
|
||||
|
||||
torchrun --nnodes 1 --nproc_per_node 4\
|
||||
fastvideo/distill_dmd.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/hunyuan\
|
||||
--dit_model_name_or_path $DATA_DIR/hunyuan/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt\
|
||||
--model_type "hunyuan" \
|
||||
--cache_dir "$DATA_DIR/.cache"\
|
||||
--data_json_path "$DATA_DIR/HD-Mixkit-Finetune-Hunyuan/videos2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/HD-Mixkit-Finetune-Hunyuan/validation"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 8\
|
||||
--sp_size 4\
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=30000\
|
||||
--learning_rate=1e-5\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=1\
|
||||
--validation_steps 1\
|
||||
--validation_sampling_steps "2,4,8" \
|
||||
--num_inference_steps 4 \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--log_validation\
|
||||
--output_dir="$DATA_DIR/outputs/hy_phase1_shift17_bs_32"\
|
||||
--tracker_project_name Hunyuan_Distill \
|
||||
--num_height 784 \
|
||||
--num_width 1280 \
|
||||
--num_frames 29 \
|
||||
--shift 17 \
|
||||
--validation_guidance_scale "1.0" \
|
||||
--num_euler_timesteps 50 \
|
||||
--multi_phased_distill_schedule "4000-1" \
|
||||
--not_apply_cfg_solver \
|
||||
--master_weight_type "bf16" \
|
||||
--optimizer "AdamW" \
|
||||
--use_8bit_adam \
|
||||
--generator_update_steps 5 \
|
||||
--run_pod
|
||||
@@ -1,38 +0,0 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
|
||||
torchrun --nnodes 1 --nproc_per_node 4 \
|
||||
fastvideo/distill.py \
|
||||
--seed 42 \
|
||||
--pretrained_model_name_or_path data/mochi \
|
||||
--model_type "mochi" \
|
||||
--cache_dir data/.cache \
|
||||
--data_json_path data/Merge-30k-Data/video2caption.json \
|
||||
--validation_prompt_dir data/Image-Vid-Finetune-Mochi/validation \
|
||||
--gradient_checkpointing \
|
||||
--train_batch_size=1 \
|
||||
--num_latent_t 28 \
|
||||
--sp_size 4 \
|
||||
--train_sp_batch_size 2 \
|
||||
--dataloader_num_workers 4 \
|
||||
--gradient_accumulation_steps=1 \
|
||||
--max_train_steps=4000 \
|
||||
--learning_rate=1e-6 \
|
||||
--mixed_precision=bf16 \
|
||||
--checkpointing_steps=64 \
|
||||
--validation_steps 1 \
|
||||
--validation_sampling_steps 8 \
|
||||
--checkpoints_total_limit 3 \
|
||||
--allow_tf32 \
|
||||
--ema_start_step 0 \
|
||||
--cfg 0.0 \
|
||||
--log_validation \
|
||||
--output_dir="data/outputs/lq_euler_50_thres0.1_lrg_0.75_phase1_lr1e-6_repro" \
|
||||
--tracker_project_name PCM \
|
||||
--num_frames 163 \
|
||||
--scheduler_type pcm_linear_quadratic \
|
||||
--validation_guidance_scale 0.5,1.5,2.5 \
|
||||
--num_euler_timesteps 50 \
|
||||
--linear_quadratic_threshold 0.1 \
|
||||
--linear_range 0.75 \
|
||||
--multi_phased_distill_schedule 4000-1
|
||||
@@ -1,39 +0,0 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
|
||||
torchrun --nnodes 1 --nproc_per_node 4 \
|
||||
fastvideo/distill_dmd.py \
|
||||
--seed 42 \
|
||||
--pretrained_model_name_or_path data/mochi \
|
||||
--model_type "mochi" \
|
||||
--cache_dir data/.cache \
|
||||
--data_json_path data/Mochi-425-Data/videos2caption.json \
|
||||
--validation_prompt_dir data/Mochi-425-Data/validation \
|
||||
--gradient_checkpointing \
|
||||
--train_batch_size=1 \
|
||||
--num_latent_t 16 \
|
||||
--sp_size 4 \
|
||||
--train_sp_batch_size 1 \
|
||||
--generator_update_steps=1 \
|
||||
--dataloader_num_workers 4 \
|
||||
--gradient_accumulation_steps=1 \
|
||||
--max_train_steps=4000 \
|
||||
--learning_rate=1e-6 \
|
||||
--mixed_precision=bf16 \
|
||||
--checkpointing_steps=64 \
|
||||
--validation_steps=64 \
|
||||
--validation_sampling_steps 6 \
|
||||
--checkpoints_total_limit 3 \
|
||||
--allow_tf32 \
|
||||
--ema_start_step 0 \
|
||||
--cfg 0.0 \
|
||||
--log_validation \
|
||||
--output_dir="data/outputs/lq_euler_50_thres0.1_lrg_0.75_phase1_lr1e-6_repro" \
|
||||
--tracker_project_name PCM \
|
||||
--num_frames 93 \
|
||||
--scheduler_type pcm_linear_quadratic \
|
||||
--validation_guidance_scale 0.5,1.5,2.5 \
|
||||
--num_euler_timesteps 50 \
|
||||
--linear_quadratic_threshold 0.1 \
|
||||
--linear_range 0.75 \
|
||||
--multi_phased_distill_schedule 4000-1
|
||||
@@ -0,0 +1,7 @@
|
||||
## How to Distill Hunyuan
|
||||
|
||||
python scripts/download_hf.py --repo_id FastVideo/hunyuan --local_dir data/hunyuan --repo_type model
|
||||
|
||||
|
||||
python scripts/download_hf.py --repo_id FastVideo/Hunyuan-Distill-Data --local_dir data/Hunyuan-Distill-Data --repo_type=dataset
|
||||
|
||||
@@ -1,16 +1,18 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
|
||||
|
||||
|
||||
DATA_DIR=./data
|
||||
|
||||
torchrun --nnodes 1 --nproc_per_node 8\
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/hunyuan\
|
||||
--pretrained_model_name_or_path data/hunyuan\
|
||||
--dit_model_name_or_path $DATA_DIR/hunyuan/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt\
|
||||
--model_type "hunyuan" \
|
||||
--cache_dir "$DATA_DIR/.cache"\
|
||||
--data_json_path "$DATA_DIR/Hunyuan-30K-Distill-Data/videos2caption.json"\
|
||||
--data_json_path "$DATA_DIR/Hunyuan-Distill-Data/videos2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/Hunyuan-Distill-Data/validation"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
@@ -18,8 +20,8 @@ torchrun --nnodes 1 --nproc_per_node 8\
|
||||
--sp_size 1\
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=2000\
|
||||
--gradient_accumulation_steps=2\
|
||||
--max_train_steps=480\
|
||||
--learning_rate=1e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=64\
|
||||
@@ -31,7 +33,7 @@ torchrun --nnodes 1 --nproc_per_node 8\
|
||||
--cfg 0.0\
|
||||
--log_validation\
|
||||
--output_dir="$DATA_DIR/outputs/hy_phase1_shift17_bs_32"\
|
||||
--tracker_project_name Hunyuan_Distill \
|
||||
--tracker_project_name PCM \
|
||||
--num_frames 93 \
|
||||
--shift 17 \
|
||||
--validation_guidance_scale "1.0" \
|
||||
@@ -1,36 +0,0 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
|
||||
torchrun --nnodes 1 --nproc_per_node 8 \
|
||||
fastvideo/train.py \
|
||||
--seed 42 \
|
||||
--pretrained_model_name_or_path data/hunyuan \
|
||||
--dit_model_name_or_path data/hunyuan/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt\
|
||||
--model_type "hunyuan" \
|
||||
--cache_dir data/.cache \
|
||||
--data_json_path data/Image-Vid-Finetune-HunYuan/videos2caption.json \
|
||||
--validation_prompt_dir data/Image-Vid-Finetune-HunYuan/validation \
|
||||
--gradient_checkpointing \
|
||||
--train_batch_size=1 \
|
||||
--num_latent_t 24 \
|
||||
--sp_size 4 \
|
||||
--train_sp_batch_size 1 \
|
||||
--dataloader_num_workers 4 \
|
||||
--gradient_accumulation_steps=1 \
|
||||
--max_train_steps=2000 \
|
||||
--learning_rate=5e-6 \
|
||||
--mixed_precision=bf16 \
|
||||
--checkpointing_steps=200 \
|
||||
--validation_steps 100 \
|
||||
--validation_sampling_steps 64 \
|
||||
--checkpoints_total_limit 3 \
|
||||
--allow_tf32 \
|
||||
--ema_start_step 0 \
|
||||
--cfg 0.0 \
|
||||
--ema_decay 0.999 \
|
||||
--log_validation \
|
||||
--output_dir=data/outputs/HSH-Taylor-Finetune-Hunyuan \
|
||||
--tracker_project_name HSH-Taylor-Finetune-Hunyuan \
|
||||
--num_frames 93 \
|
||||
--validation_guidance_scale "1.0" \
|
||||
--group_frame
|
||||
@@ -1,32 +0,0 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
|
||||
torchrun --nnodes 1 --nproc_per_node 4 \
|
||||
fastvideo/train.py \
|
||||
--seed 42 \
|
||||
--pretrained_model_name_or_path data/mochi \
|
||||
--cache_dir data/.cache \
|
||||
--data_json_path data/Mochi-Black-Myth/videos2caption.json \
|
||||
--validation_prompt_dir data/Mochi-Black-Myth/validation \
|
||||
--gradient_checkpointing \
|
||||
--train_batch_size=1 \
|
||||
--num_latent_t 16 \
|
||||
--sp_size 1 \
|
||||
--train_sp_batch_size 1 \
|
||||
--dataloader_num_workers 4 \
|
||||
--gradient_accumulation_steps=1 \
|
||||
--max_train_steps=2000 \
|
||||
--learning_rate=5e-6 \
|
||||
--mixed_precision=bf16 \
|
||||
--checkpointing_steps=200 \
|
||||
--validation_steps 100 \
|
||||
--validation_sampling_steps 64 \
|
||||
--checkpoints_total_limit 3 \
|
||||
--allow_tf32 \
|
||||
--ema_start_step 0 \
|
||||
--cfg 0.0 \
|
||||
--ema_decay 0.999 \
|
||||
--log_validation \
|
||||
--output_dir=data/outputs/Black-Myth-Finetune \
|
||||
--tracker_project_name Black-Myth-Finetune \
|
||||
--num_frames 93
|
||||
@@ -1,37 +0,0 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
|
||||
torchrun --nnodes 1 --nproc_per_node 2 \
|
||||
fastvideo/train.py \
|
||||
--seed 42 \
|
||||
--pretrained_model_name_or_path data/mochi \
|
||||
--cache_dir data/.cache \
|
||||
--data_json_path data/Mochi-Black-Myth/videos2caption.json \
|
||||
--validation_prompt_dir data/Mochi-Black-Myth/validation \
|
||||
--gradient_checkpointing \
|
||||
--train_batch_size 1 \
|
||||
--num_latent_t 14 \
|
||||
--sp_size 2 \
|
||||
--train_sp_batch_size 1 \
|
||||
--dataloader_num_workers 1 \
|
||||
--gradient_accumulation_steps 2 \
|
||||
--max_train_steps 2000 \
|
||||
--learning_rate 5e-6 \
|
||||
--mixed_precision bf16 \
|
||||
--checkpointing_steps 200 \
|
||||
--validation_steps 100 \
|
||||
--validation_sampling_steps 64 \
|
||||
--checkpoints_total_limit 3 \
|
||||
--allow_tf32 \
|
||||
--ema_start_step 0 \
|
||||
--cfg 0.0 \
|
||||
--ema_decay 0.999 \
|
||||
--log_validation \
|
||||
--output_dir=data/outputs/Black-Myth-Lora-FT \
|
||||
--tracker_project_name Black-Myth-Lora-Finetune \
|
||||
--num_frames 91 \
|
||||
--lora_rank 128 \
|
||||
--lora_alpha 256 \
|
||||
--master_weight_type "bf16" \
|
||||
--use_lora \
|
||||
--use_cpu_offload
|
||||
@@ -1,37 +0,0 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
|
||||
CUDA_VISIBLE_DEVICES=5 torchrun --nnodes 1 --nproc_per_node 1 \
|
||||
fastvideo/train.py \
|
||||
--seed 42 \
|
||||
--pretrained_model_name_or_path data/mochi \
|
||||
--cache_dir data/.cache \
|
||||
--data_json_path data/Image-Vid-Finetune-Mochi/videos2caption.json \
|
||||
--validation_prompt_dir data/Image-Vid-Finetune-Mochi/validation \
|
||||
--gradient_checkpointing \
|
||||
--train_batch_size=1 \
|
||||
--num_latent_t 14 \
|
||||
--sp_size 1 \
|
||||
--train_sp_batch_size 1 \
|
||||
--dataloader_num_workers 1 \
|
||||
--gradient_accumulation_steps=1 \
|
||||
--max_train_steps=2000 \
|
||||
--learning_rate=5e-6 \
|
||||
--mixed_precision=bf16 \
|
||||
--checkpointing_steps=200 \
|
||||
--validation_steps 100 \
|
||||
--validation_sampling_steps 64 \
|
||||
--checkpoints_total_limit 3 \
|
||||
--allow_tf32 \
|
||||
--ema_start_step 0 \
|
||||
--cfg 0.0 \
|
||||
--ema_decay 0.999 \
|
||||
--log_validation \
|
||||
--output_dir=data/outputs/HSH-Taylor-Finetune-Lora \
|
||||
--tracker_project_name HSH-Taylor-Finetune-Lora \
|
||||
--num_frames 91 \
|
||||
--group_frame \
|
||||
--lora_rank 128 \
|
||||
--lora_alpha 256 \
|
||||
--master_weight_type "bf16" \
|
||||
--use_lora
|
||||
@@ -0,0 +1,36 @@
|
||||
torchrun --nnodes 4 --nproc_per_node 4 \
|
||||
--node_rank=3 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=172.23.30.15:29500 \
|
||||
fastvideo/train.py \
|
||||
--seed 42 \
|
||||
--pretrained_model_name_or_path data/mochi \
|
||||
--cache_dir "data/.cache" \
|
||||
--data_json_path "data/BLACK-MYTH-Finetune-Dataset/videos2caption.json" \
|
||||
--validation_prompt_dir "data/BLACK-MYTH-Finetune-Dataset/validation_prompt_embed_mask" \
|
||||
--uncond_prompt_dir "data/BLACK-MYTH-Finetune-Dataset/uncond_prompt_embed_mask" \
|
||||
--gradient_checkpointing \
|
||||
--train_batch_size=1 \
|
||||
--num_latent_t 16 \
|
||||
--sp_size 4 \
|
||||
--train_sp_batch_size 4 \
|
||||
--dataloader_num_workers 4 \
|
||||
--gradient_accumulation_steps=1 \
|
||||
--max_train_steps=4000 \
|
||||
--learning_rate=2e-5 \
|
||||
--mixed_precision="bf16" \
|
||||
--checkpointing_steps=500 \
|
||||
--validation_steps=100 \
|
||||
--validation_sampling_steps 64 \
|
||||
--checkpoints_total_limit 3 \
|
||||
--allow_tf32 \
|
||||
--ema_start_step 0 \
|
||||
--cfg 0.0 \
|
||||
--ema_decay 0.999 \
|
||||
--log_validation \
|
||||
--output_dir="data/outputs/black_myth_correct_video_mask" \
|
||||
--weighting_scheme "uniform" \
|
||||
--num_frames 91 \
|
||||
--selective_checkpointing 1.0
|
||||
|
||||
@@ -0,0 +1,13 @@
|
||||
num_gpus=1
|
||||
|
||||
torchrun --nproc_per_node=$num_gpus fastvideo/sample/generate_synthetic.py \
|
||||
--model_path data/mochi \
|
||||
--num_frames 1 \
|
||||
--height 480 \
|
||||
--width 848 \
|
||||
--num_inference_steps 4 \
|
||||
--guidance_scale 4.5 \
|
||||
--prompt_path "data/prompt.txt" \
|
||||
--dataset_output_dir data/synthetic_debug2
|
||||
|
||||
|
||||
@@ -1,9 +0,0 @@
|
||||
from huggingface_hub import HfApi
|
||||
|
||||
api = HfApi()
|
||||
|
||||
api.upload_folder(
|
||||
folder_path="data/Black-Myth-Taylor-Src",
|
||||
repo_id="FastVideo/Image-Vid-Finetune-Src",
|
||||
repo_type="dataset",
|
||||
)
|
||||
@@ -0,0 +1,92 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_DIR="$HOME"
|
||||
export WANDB_MODE=online
|
||||
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
|
||||
export LD_LIBRARY_PATH=/opt/amazon/efa/lib:/opt/aws-ofi-nccl/lib:$LD_LIBRARY_PATH
|
||||
export FI_PROVIDER=efa
|
||||
export FI_EFA_USE_DEVICE_RDMA=1
|
||||
export NCCL_PROTO=simple
|
||||
|
||||
DATA_DIR=/data
|
||||
IP=10.4.139.86
|
||||
|
||||
torchrun --nnodes 2 --nproc_per_node 8\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=$IP:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/hunyuan\
|
||||
--dit_model_name_or_path $DATA_DIR/hunyuan/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt\
|
||||
--model_type "hunyuan" \
|
||||
--cache_dir "$DATA_DIR/.cache"\
|
||||
--data_json_path "$DATA_DIR/Hunyuan-Distill-Data/videos2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/Hunyuan-Distill-Data/validation"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 24\
|
||||
--sp_size 2\
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=480\
|
||||
--learning_rate=1e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=64\
|
||||
--validation_steps 64\
|
||||
--validation_sampling_steps "2,4,8" \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--log_validation\
|
||||
--output_dir="$DATA_DIR/outputs/hy_phase1_shift17"\
|
||||
--tracker_project_name Hunyuan_Distill \
|
||||
--num_frames 93 \
|
||||
--shift 17 \
|
||||
--validation_guidance_scale "1.0" \
|
||||
--num_euler_timesteps 50 \
|
||||
--multi_phased_distill_schedule "4000-1" \
|
||||
--not_apply_cfg_solver
|
||||
|
||||
|
||||
torchrun --nnodes 2 --nproc_per_node 8\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=$IP:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/hunyuan\
|
||||
--dit_model_name_or_path $DATA_DIR/hunyuan/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt\
|
||||
--model_type "hunyuan" \
|
||||
--cache_dir "$DATA_DIR/.cache"\
|
||||
--data_json_path "$DATA_DIR/Hunyuan-Distill-Data/videos2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/Hunyuan-Distill-Data/validation"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 24\
|
||||
--sp_size 2\
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=480\
|
||||
--learning_rate=3e-7\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=64\
|
||||
--validation_steps 64\
|
||||
--validation_sampling_steps "2,4,8" \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--log_validation\
|
||||
--output_dir="$DATA_DIR/outputs/hy_phase1_shift17_lr_3e-7"\
|
||||
--tracker_project_name Hunyuan_Distill \
|
||||
--num_frames 93 \
|
||||
--shift 17 \
|
||||
--validation_guidance_scale "1.0" \
|
||||
--num_euler_timesteps 50 \
|
||||
--multi_phased_distill_schedule "4000-1" \
|
||||
--not_apply_cfg_solver
|
||||
@@ -0,0 +1,51 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_DIR="$HOME"
|
||||
export WANDB_MODE=online
|
||||
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
|
||||
export LD_LIBRARY_PATH=/opt/amazon/efa/lib:/opt/aws-ofi-nccl/lib:$LD_LIBRARY_PATH
|
||||
export FI_PROVIDER=efa
|
||||
export FI_EFA_USE_DEVICE_RDMA=1
|
||||
export NCCL_PROTO=simple
|
||||
|
||||
DATA_DIR=/data
|
||||
IP=10.4.139.86
|
||||
|
||||
torchrun --nnodes 2 --nproc_per_node 8\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=$IP:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/hunyuan\
|
||||
--dit_model_name_or_path $DATA_DIR/hunyuan/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt\
|
||||
--model_type "hunyuan" \
|
||||
--cache_dir "$DATA_DIR/.cache"\
|
||||
--data_json_path "$DATA_DIR/Hunyuan-Distill-Data/videos2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/Hunyuan-Distill-Data/validation"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 24\
|
||||
--sp_size 1\
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=2\
|
||||
--max_train_steps=640\
|
||||
--learning_rate=1e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=64\
|
||||
--validation_steps 64\
|
||||
--validation_sampling_steps "2,4,8" \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--log_validation\
|
||||
--output_dir="$DATA_DIR/outputs/hy_phase1_shift25_bs_32"\
|
||||
--tracker_project_name Hunyuan_Distill \
|
||||
--num_frames 93 \
|
||||
--shift 25 \
|
||||
--validation_guidance_scale "1.0" \
|
||||
--num_euler_timesteps 50 \
|
||||
--multi_phased_distill_schedule "4000-1" \
|
||||
--not_apply_cfg_solver
|
||||
@@ -0,0 +1,51 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_DIR="$HOME"
|
||||
export WANDB_MODE=online
|
||||
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
|
||||
export LD_LIBRARY_PATH=/opt/amazon/efa/lib:/opt/aws-ofi-nccl/lib:$LD_LIBRARY_PATH
|
||||
export FI_PROVIDER=efa
|
||||
export FI_EFA_USE_DEVICE_RDMA=1
|
||||
export NCCL_PROTO=simple
|
||||
|
||||
DATA_DIR=/data
|
||||
IP=10.4.139.86
|
||||
|
||||
torchrun --nnodes 2 --nproc_per_node 8\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=$IP:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/hunyuan\
|
||||
--dit_model_name_or_path $DATA_DIR/hunyuan/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt\
|
||||
--model_type "hunyuan" \
|
||||
--cache_dir "$DATA_DIR/.cache"\
|
||||
--data_json_path "$DATA_DIR/Hunyuan-Distill-Data/videos2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/Hunyuan-Distill-Data/validation"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 24\
|
||||
--sp_size 1\
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=2\
|
||||
--max_train_steps=640\
|
||||
--learning_rate=1e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=64\
|
||||
--validation_steps 64\
|
||||
--validation_sampling_steps "2,4,8" \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--log_validation\
|
||||
--output_dir="$DATA_DIR/outputs/hy_phase1_shift33_bs_32"\
|
||||
--tracker_project_name Hunyuan_Distill \
|
||||
--num_frames 93 \
|
||||
--shift 33 \
|
||||
--validation_guidance_scale "1.0" \
|
||||
--num_euler_timesteps 50 \
|
||||
--multi_phased_distill_schedule "4000-1" \
|
||||
--not_apply_cfg_solver
|
||||
@@ -0,0 +1,53 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_DIR="$HOME"
|
||||
export WANDB_MODE=online
|
||||
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
|
||||
export LD_LIBRARY_PATH=/opt/amazon/efa/lib:/opt/aws-ofi-nccl/lib:$LD_LIBRARY_PATH
|
||||
export FI_PROVIDER=efa
|
||||
export FI_EFA_USE_DEVICE_RDMA=1
|
||||
export NCCL_PROTO=simple
|
||||
|
||||
DATA_DIR=/data
|
||||
IP=10.4.139.86
|
||||
|
||||
torchrun --nnodes 2 --nproc_per_node 8\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=$IP:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/hunyuan\
|
||||
--dit_model_name_or_path $DATA_DIR/hunyuan/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt\
|
||||
--model_type "hunyuan" \
|
||||
--cache_dir "$DATA_DIR/.cache"\
|
||||
--data_json_path "$DATA_DIR/Hunyuan-Distill-Data/videos2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/Hunyuan-Distill-Data/validation"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 24\
|
||||
--sp_size 1\
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=2\
|
||||
--max_train_steps=640\
|
||||
--learning_rate=1e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=64\
|
||||
--validation_steps 64\
|
||||
--validation_sampling_steps "2,4,8" \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--log_validation\
|
||||
--output_dir="$DATA_DIR/outputs/hy_phase1_shift33_student3_batchsize32"\
|
||||
--tracker_project_name Hunyuan_Distill \
|
||||
--num_frames 93 \
|
||||
--shift 33 \
|
||||
--validation_guidance_scale "1.0" \
|
||||
--num_euler_timesteps 50 \
|
||||
--multi_phased_distill_schedule "4000-1" \
|
||||
--not_apply_cfg_solver \
|
||||
--hunyuan_student_cfg_embed 3
|
||||
|
||||
@@ -0,0 +1,92 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_DIR="$HOME"
|
||||
export WANDB_MODE=online
|
||||
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
|
||||
export LD_LIBRARY_PATH=/opt/amazon/efa/lib:/opt/aws-ofi-nccl/lib:$LD_LIBRARY_PATH
|
||||
export FI_PROVIDER=efa
|
||||
export FI_EFA_USE_DEVICE_RDMA=1
|
||||
export NCCL_PROTO=simple
|
||||
|
||||
DATA_DIR=/data
|
||||
IP=10.4.139.86
|
||||
|
||||
torchrun --nnodes 2 --nproc_per_node 8\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=$IP:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/hunyuan\
|
||||
--dit_model_name_or_path $DATA_DIR/hunyuan/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt\
|
||||
--model_type "hunyuan" \
|
||||
--cache_dir "$DATA_DIR/.cache"\
|
||||
--data_json_path "$DATA_DIR/Hunyuan-Distill-Data/videos2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/Hunyuan-Distill-Data/validation"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 24\
|
||||
--sp_size 2\
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=480\
|
||||
--learning_rate=1e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=64\
|
||||
--validation_steps 64\
|
||||
--validation_sampling_steps "2,4,8" \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--log_validation\
|
||||
--output_dir="$DATA_DIR/outputs/hy_phase1_shift27"\
|
||||
--tracker_project_name Hunyuan_Distill \
|
||||
--num_frames 93 \
|
||||
--shift 27 \
|
||||
--validation_guidance_scale "1.0" \
|
||||
--num_euler_timesteps 50 \
|
||||
--multi_phased_distill_schedule "4000-1" \
|
||||
--not_apply_cfg_solver
|
||||
|
||||
|
||||
torchrun --nnodes 2 --nproc_per_node 8\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=$IP:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/hunyuan\
|
||||
--dit_model_name_or_path $DATA_DIR/hunyuan/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt\
|
||||
--model_type "hunyuan" \
|
||||
--cache_dir "$DATA_DIR/.cache"\
|
||||
--data_json_path "$DATA_DIR/Hunyuan-Distill-Data/videos2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/Hunyuan-Distill-Data/validation"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 24\
|
||||
--sp_size 2\
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=480\
|
||||
--learning_rate=3e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=64\
|
||||
--validation_steps 64\
|
||||
--validation_sampling_steps "2,4,8" \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--log_validation\
|
||||
--output_dir="$DATA_DIR/outputs/hy_phase1_shift17_lr_3e-6"\
|
||||
--tracker_project_name Hunyuan_Distill \
|
||||
--num_frames 93 \
|
||||
--shift 17 \
|
||||
--validation_guidance_scale "1.0" \
|
||||
--num_euler_timesteps 50 \
|
||||
--multi_phased_distill_schedule "4000-1" \
|
||||
--not_apply_cfg_solver
|
||||
@@ -0,0 +1,95 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_DIR="$HOME"
|
||||
export WANDB_MODE=online
|
||||
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
|
||||
export LD_LIBRARY_PATH=/opt/amazon/efa/lib:/opt/aws-ofi-nccl/lib:$LD_LIBRARY_PATH
|
||||
export FI_PROVIDER=efa
|
||||
export FI_EFA_USE_DEVICE_RDMA=1
|
||||
export NCCL_PROTO=simple
|
||||
|
||||
DATA_DIR=/data
|
||||
IP=10.4.139.86
|
||||
|
||||
torchrun --nnodes 2 --nproc_per_node 8\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=$IP:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/hunyuan\
|
||||
--dit_model_name_or_path $DATA_DIR/hunyuan/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt\
|
||||
--model_type "hunyuan" \
|
||||
--cache_dir "$DATA_DIR/.cache"\
|
||||
--data_json_path "$DATA_DIR/Hunyuan-Distill-Data/videos2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/Hunyuan-Distill-Data/validation"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 24\
|
||||
--sp_size 2\
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=480\
|
||||
--learning_rate=1e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=64\
|
||||
--validation_steps 64\
|
||||
--validation_sampling_steps "2,4,8" \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--log_validation\
|
||||
--output_dir="$DATA_DIR/outputs/hy_phase1_shift17_teacher_no_cfg"\
|
||||
--tracker_project_name Hunyuan_Distill \
|
||||
--num_frames 93 \
|
||||
--shift 17 \
|
||||
--validation_guidance_scale "1.0" \
|
||||
--num_euler_timesteps 50 \
|
||||
--multi_phased_distill_schedule "4000-1" \
|
||||
--not_apply_cfg_solver \
|
||||
--hunyuan_teacher_disable_cfg
|
||||
|
||||
|
||||
torchrun --nnodes 2 --nproc_per_node 8\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=$IP:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/hunyuan\
|
||||
--dit_model_name_or_path $DATA_DIR/hunyuan/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt\
|
||||
--model_type "hunyuan" \
|
||||
--cache_dir "$DATA_DIR/.cache"\
|
||||
--data_json_path "$DATA_DIR/Hunyuan-Distill-Data/videos2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/Hunyuan-Distill-Data/validation"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 24\
|
||||
--sp_size 2\
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=480\
|
||||
--learning_rate=1e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=64\
|
||||
--validation_steps 64\
|
||||
--validation_sampling_steps "2,4,8" \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--log_validation\
|
||||
--output_dir="$DATA_DIR/outputs/hy_phase1_shift17_ema"\
|
||||
--tracker_project_name Hunyuan_Distill \
|
||||
--num_frames 93 \
|
||||
--shift 17 \
|
||||
--validation_guidance_scale "1.0" \
|
||||
--num_euler_timesteps 50 \
|
||||
--multi_phased_distill_schedule "4000-1" \
|
||||
--not_apply_cfg_solver \
|
||||
--use_ema
|
||||
|
||||
@@ -1,26 +1,36 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_DIR="$HOME"
|
||||
export WANDB_MODE=online
|
||||
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
|
||||
export LD_LIBRARY_PATH=/opt/amazon/efa/lib:/opt/aws-ofi-nccl/lib:$LD_LIBRARY_PATH
|
||||
export FI_PROVIDER=efa
|
||||
export FI_EFA_USE_DEVICE_RDMA=1
|
||||
export NCCL_PROTO=simple
|
||||
|
||||
DATA_DIR=./data
|
||||
DATA_DIR=/data
|
||||
IP=10.4.139.86
|
||||
|
||||
torchrun --nnodes 1 --nproc_per_node 4\
|
||||
fastvideo/distill_adv.py\
|
||||
torchrun --nnodes 2 --nproc_per_node 8\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=$IP:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/hunyuan\
|
||||
--dit_model_name_or_path $DATA_DIR/hunyuan/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt\
|
||||
--model_type "hunyuan" \
|
||||
--cache_dir "$DATA_DIR/.cache"\
|
||||
--data_json_path "$DATA_DIR/HD-Mixkit-Finetune-Hunyuan/videos2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/HD-Mixkit-Finetune-Hunyuan/validation"\
|
||||
--data_json_path "$DATA_DIR/Hunyuan-Distill-Data/videos2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/Hunyuan-Distill-Data/validation"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 8\
|
||||
--sp_size 4 \
|
||||
--num_latent_t 24\
|
||||
--sp_size 1\
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=2000\
|
||||
--gradient_accumulation_steps=2\
|
||||
--max_train_steps=480\
|
||||
--learning_rate=1e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=64\
|
||||
@@ -31,14 +41,11 @@ torchrun --nnodes 1 --nproc_per_node 4\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--log_validation\
|
||||
--output_dir="$DATA_DIR/outputs/hy_phase1_shift17_bs_8_adv_HD"\
|
||||
--output_dir="$DATA_DIR/outputs/hy_phase1_shift17_bs_32"\
|
||||
--tracker_project_name Hunyuan_Distill \
|
||||
--num_height 720 \
|
||||
--num_width 1280 \
|
||||
--num_frames 125 \
|
||||
--num_frames 93 \
|
||||
--shift 17 \
|
||||
--validation_guidance_scale "1.0" \
|
||||
--num_euler_timesteps 50 \
|
||||
--multi_phased_distill_schedule "4000-1" \
|
||||
--not_apply_cfg_solver \
|
||||
--master_weight_type "bf16"
|
||||
--not_apply_cfg_solver
|
||||
@@ -0,0 +1,92 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_DIR="$HOME"
|
||||
export WANDB_MODE=online
|
||||
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
|
||||
export LD_LIBRARY_PATH=/opt/amazon/efa/lib:/opt/aws-ofi-nccl/lib:$LD_LIBRARY_PATH
|
||||
export FI_PROVIDER=efa
|
||||
export FI_EFA_USE_DEVICE_RDMA=1
|
||||
export NCCL_PROTO=simple
|
||||
|
||||
DATA_DIR=/data
|
||||
IP=10.4.139.86
|
||||
|
||||
torchrun --nnodes 2 --nproc_per_node 8\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=$IP:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/hunyuan\
|
||||
--dit_model_name_or_path $DATA_DIR/hunyuan/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt\
|
||||
--model_type "hunyuan" \
|
||||
--cache_dir "$DATA_DIR/.cache"\
|
||||
--data_json_path "$DATA_DIR/Hunyuan-Distill-Data/videos2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/Hunyuan-Distill-Data/validation"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 24\
|
||||
--sp_size 1\
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=320\
|
||||
--learning_rate=1e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=64\
|
||||
--validation_steps 64\
|
||||
--validation_sampling_steps "2,4,8" \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--log_validation\
|
||||
--output_dir="$DATA_DIR/outputs/hy_phase1_shift17_bs_16_nosp"\
|
||||
--tracker_project_name Hunyuan_Distill \
|
||||
--num_frames 93 \
|
||||
--shift 17 \
|
||||
--validation_guidance_scale "1.0" \
|
||||
--num_euler_timesteps 50 \
|
||||
--multi_phased_distill_schedule "4000-1" \
|
||||
--not_apply_cfg_solver
|
||||
|
||||
|
||||
torchrun --nnodes 2 --nproc_per_node 8\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=$IP:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/hunyuan\
|
||||
--dit_model_name_or_path $DATA_DIR/hunyuan/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt\
|
||||
--model_type "hunyuan" \
|
||||
--cache_dir "$DATA_DIR/.cache"\
|
||||
--data_json_path "$DATA_DIR/Hunyuan-Distill-Data/videos2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/Hunyuan-Distill-Data/validation"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 24\
|
||||
--sp_size 2\
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=2\
|
||||
--max_train_steps=320\
|
||||
--learning_rate=1e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=64\
|
||||
--validation_steps 64\
|
||||
--validation_sampling_steps "2,4,8" \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--log_validation\
|
||||
--output_dir="$DATA_DIR/outputs/hy_phase1_shift17_bs_16_sp"\
|
||||
--tracker_project_name Hunyuan_Distill \
|
||||
--num_frames 93 \
|
||||
--shift 17 \
|
||||
--validation_guidance_scale "1.0" \
|
||||
--num_euler_timesteps 50 \
|
||||
--multi_phased_distill_schedule "4000-1" \
|
||||
--not_apply_cfg_solver
|
||||
@@ -1,48 +1,53 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_DIR="$HOME"
|
||||
export WANDB_MODE=online
|
||||
export WANDB_API_KEY=7abd9763ff10869f9e88526d919c2639766139a8
|
||||
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
|
||||
export LD_LIBRARY_PATH=/opt/amazon/efa/lib:/opt/aws-ofi-nccl/lib:$LD_LIBRARY_PATH
|
||||
export FI_PROVIDER=efa
|
||||
export FI_EFA_USE_DEVICE_RDMA=1
|
||||
export NCCL_PROTO=simple
|
||||
|
||||
DATA_DIR=./data
|
||||
DATA_DIR=/data
|
||||
IP=10.4.139.86
|
||||
|
||||
torchrun --nnodes 1 --nproc_per_node 4\
|
||||
fastvideo/distill_dmd.py\
|
||||
torchrun --nnodes 2 --nproc_per_node 8\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=$IP:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/hunyuan\
|
||||
--dit_model_name_or_path $DATA_DIR/hunyuan/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt\
|
||||
--model_type "hunyuan" \
|
||||
--cache_dir "$DATA_DIR/.cache"\
|
||||
--data_json_path "$DATA_DIR/HD-Mixkit-Finetune-Hunyuan/videos2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/HD-Mixkit-Finetune-Hunyuan/validation"\
|
||||
--data_json_path "$DATA_DIR/Hunyuan-Distill-Data/videos2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/Hunyuan-Distill-Data/validation"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 8\
|
||||
--sp_size 4\
|
||||
--num_latent_t 24\
|
||||
--sp_size 1\
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=30000\
|
||||
--learning_rate=1e-5\
|
||||
--gradient_accumulation_steps=2\
|
||||
--max_train_steps=640\
|
||||
--learning_rate=1e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=64\
|
||||
--validation_steps 64\
|
||||
--validation_sampling_steps "2,4,8" \
|
||||
--num_inference_steps 8 \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--log_validation\
|
||||
--output_dir="$DATA_DIR/outputs/hy_phase1_shift17_bs_32"\
|
||||
--output_dir="$DATA_DIR/outputs/hy_phase1_shift17_student3_batchsize32"\
|
||||
--tracker_project_name Hunyuan_Distill \
|
||||
--num_height 720 \
|
||||
--num_width 1280 \
|
||||
--num_frames 125 \
|
||||
--num_frames 93 \
|
||||
--shift 17 \
|
||||
--validation_guidance_scale "1.0" \
|
||||
--num_euler_timesteps 50 \
|
||||
--multi_phased_distill_schedule "4000-1" \
|
||||
--not_apply_cfg_solver \
|
||||
--master_weight_type "bf16" \
|
||||
--optimizer "AdamW" \
|
||||
--use_8bit_adam \
|
||||
--generator_update_steps 3
|
||||
--hunyuan_student_cfg_embed 3
|
||||
|
||||
@@ -1,48 +1,53 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_DIR="$HOME"
|
||||
export WANDB_MODE=online
|
||||
export WANDB_API_KEY=7abd9763ff10869f9e88526d919c2639766139a8
|
||||
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
|
||||
export LD_LIBRARY_PATH=/opt/amazon/efa/lib:/opt/aws-ofi-nccl/lib:$LD_LIBRARY_PATH
|
||||
export FI_PROVIDER=efa
|
||||
export FI_EFA_USE_DEVICE_RDMA=1
|
||||
export NCCL_PROTO=simple
|
||||
|
||||
DATA_DIR=./data
|
||||
DATA_DIR=/data
|
||||
IP=10.4.139.86
|
||||
|
||||
torchrun --nnodes 1 --nproc_per_node 4\
|
||||
fastvideo/distill_dmd.py\
|
||||
torchrun --nnodes 2 --nproc_per_node 8\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=$IP:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/hunyuan\
|
||||
--dit_model_name_or_path $DATA_DIR/hunyuan/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt\
|
||||
--model_type "hunyuan" \
|
||||
--cache_dir "$DATA_DIR/.cache"\
|
||||
--data_json_path "$DATA_DIR/HD-Mixkit-Finetune-Hunyuan/videos2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/HD-Mixkit-Finetune-Hunyuan/validation"\
|
||||
--data_json_path "$DATA_DIR/Hunyuan-Distill-Data/videos2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/Hunyuan-Distill-Data/validation"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 8\
|
||||
--sp_size 4\
|
||||
--num_latent_t 24\
|
||||
--sp_size 1\
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=30000\
|
||||
--learning_rate=1e-5\
|
||||
--gradient_accumulation_steps=2\
|
||||
--max_train_steps=640\
|
||||
--learning_rate=1e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=64\
|
||||
--validation_steps 64\
|
||||
--validation_sampling_steps "2,4,8" \
|
||||
--num_inference_steps 4 \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--log_validation\
|
||||
--output_dir="$DATA_DIR/outputs/hy_phase1_shift17_bs_32"\
|
||||
--output_dir="$DATA_DIR/outputs/hy_phase1_shift17_student4_batchsize32"\
|
||||
--tracker_project_name Hunyuan_Distill \
|
||||
--num_height 720 \
|
||||
--num_width 1280 \
|
||||
--num_frames 125 \
|
||||
--num_frames 93 \
|
||||
--shift 17 \
|
||||
--validation_guidance_scale "1.0" \
|
||||
--num_euler_timesteps 50 \
|
||||
--multi_phased_distill_schedule "4000-1" \
|
||||
--not_apply_cfg_solver \
|
||||
--master_weight_type "bf16" \
|
||||
--optimizer "AdamW" \
|
||||
--use_8bit_adam \
|
||||
--generator_update_steps 3
|
||||
--hunyuan_student_cfg_embed 4
|
||||
|
||||
@@ -0,0 +1,51 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_DIR="$HOME"
|
||||
export WANDB_MODE=online
|
||||
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
|
||||
export LD_LIBRARY_PATH=/opt/amazon/efa/lib:/opt/aws-ofi-nccl/lib:$LD_LIBRARY_PATH
|
||||
export FI_PROVIDER=efa
|
||||
export FI_EFA_USE_DEVICE_RDMA=1
|
||||
export NCCL_PROTO=simple
|
||||
|
||||
DATA_DIR=/data
|
||||
IP=10.4.139.86
|
||||
|
||||
torchrun --nnodes 2 --nproc_per_node 8\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=$IP:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/hunyuan\
|
||||
--dit_model_name_or_path $DATA_DIR/hunyuan/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt\
|
||||
--model_type "hunyuan" \
|
||||
--cache_dir "$DATA_DIR/.cache"\
|
||||
--data_json_path "$DATA_DIR/Hunyuan-Distill-Data/videos2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/Hunyuan-Distill-Data/validation"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 24\
|
||||
--sp_size 1\
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=2\
|
||||
--max_train_steps=640\
|
||||
--learning_rate=3e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=64\
|
||||
--validation_steps 64\
|
||||
--validation_sampling_steps "2,4,8" \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--log_validation\
|
||||
--output_dir="$DATA_DIR/outputs/hy_phase1_shift17_bs_32_lr_3e-6"\
|
||||
--tracker_project_name Hunyuan_Distill \
|
||||
--num_frames 93 \
|
||||
--shift 17 \
|
||||
--validation_guidance_scale "1.0" \
|
||||
--num_euler_timesteps 50 \
|
||||
--multi_phased_distill_schedule "4000-1" \
|
||||
--not_apply_cfg_solver
|
||||
@@ -0,0 +1,51 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_DIR="$HOME"
|
||||
export WANDB_MODE=online
|
||||
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
|
||||
export LD_LIBRARY_PATH=/opt/amazon/efa/lib:/opt/aws-ofi-nccl/lib:$LD_LIBRARY_PATH
|
||||
export FI_PROVIDER=efa
|
||||
export FI_EFA_USE_DEVICE_RDMA=1
|
||||
export NCCL_PROTO=simple
|
||||
|
||||
DATA_DIR=/data
|
||||
IP=10.4.139.86
|
||||
|
||||
torchrun --nnodes 2 --nproc_per_node 8\
|
||||
--node_rank=0 \
|
||||
--rdzv_id=456 \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint=$IP:29500 \
|
||||
fastvideo/distill.py\
|
||||
--seed 42\
|
||||
--pretrained_model_name_or_path $DATA_DIR/hunyuan\
|
||||
--dit_model_name_or_path $DATA_DIR/hunyuan/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt\
|
||||
--model_type "hunyuan" \
|
||||
--cache_dir "$DATA_DIR/.cache"\
|
||||
--data_json_path "$DATA_DIR/Hunyuan-Distill-Data/videos2caption.json"\
|
||||
--validation_prompt_dir "$DATA_DIR/Hunyuan-Distill-Data/validation"\
|
||||
--gradient_checkpointing\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 24\
|
||||
--sp_size 1\
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 4\
|
||||
--gradient_accumulation_steps=2\
|
||||
--max_train_steps=640\
|
||||
--learning_rate=1e-6\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=64\
|
||||
--validation_steps 64\
|
||||
--validation_sampling_steps "2,4,8" \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--log_validation\
|
||||
--output_dir="$DATA_DIR/outputs/hy_phase2_shift17_bs_32"\
|
||||
--tracker_project_name Hunyuan_Distill \
|
||||
--num_frames 93 \
|
||||
--shift 17 \
|
||||
--validation_guidance_scale "1.0" \
|
||||
--num_euler_timesteps 50 \
|
||||
--multi_phased_distill_schedule "4000-2" \
|
||||
--not_apply_cfg_solver
|
||||
@@ -1,20 +0,0 @@
|
||||
#!/bin/bash
|
||||
|
||||
num_gpus=1
|
||||
export MODEL_BASE="data/FastHunyuan-diffusers"
|
||||
torchrun --nnodes=1 --nproc_per_node=$num_gpus --master_port 12345 \
|
||||
fastvideo/sample/sample_t2v_diffusers_hunyuan.py \
|
||||
--height 720 \
|
||||
--width 1280 \
|
||||
--num_frames 45 \
|
||||
--num_inference_steps 6 \
|
||||
--guidance_scale 1 \
|
||||
--embedded_cfg_scale 6 \
|
||||
--flow_shift 17 \
|
||||
--flow-reverse \
|
||||
--prompt ./assets/prompt.txt \
|
||||
--seed 1024 \
|
||||
--output_path outputs_video/hunyuan_quant/nf4/ \
|
||||
--model_path $MODEL_BASE \
|
||||
--quantization "nf4" \
|
||||
--cpu_offload
|
||||
@@ -1,19 +0,0 @@
|
||||
#!/bin/bash
|
||||
|
||||
num_gpus=4
|
||||
export MODEL_BASE=data/FastHunyuan
|
||||
torchrun --nnodes=1 --nproc_per_node=$num_gpus --master_port 29503 \
|
||||
fastvideo/sample/sample_t2v_hunyuan.py \
|
||||
--height 720 \
|
||||
--width 1280 \
|
||||
--num_frames 125 \
|
||||
--num_inference_steps 6 \
|
||||
--guidance_scale 1 \
|
||||
--embedded_cfg_scale 6 \
|
||||
--flow_shift 17 \
|
||||
--flow-reverse \
|
||||
--prompt ./assets/prompt.txt \
|
||||
--seed 1024 \
|
||||
--output_path outputs_video/hunyuan/cfg6/ \
|
||||
--model_path $MODEL_BASE \
|
||||
--dit-weight ${MODEL_BASE}/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt
|
||||
@@ -1,19 +0,0 @@
|
||||
#!/bin/bash
|
||||
|
||||
num_gpus=4
|
||||
|
||||
torchrun --nnodes=1 --nproc_per_node=$num_gpus --master_port 29503 \
|
||||
fastvideo/sample/sample_t2v_mochi.py \
|
||||
--model_path data/FastMochi-diffusers \
|
||||
--prompt_path "assets/prompt.txt" \
|
||||
--num_frames 163 \
|
||||
--height 480 \
|
||||
--width 848 \
|
||||
--num_inference_steps 8 \
|
||||
--guidance_scale 1.5 \
|
||||
--output_path outputs_video/mochi_sp/ \
|
||||
--seed 1024 \
|
||||
--scheduler_type "pcm_linear_quadratic" \
|
||||
--linear_threshold 0.1 \
|
||||
--linear_range 0.75
|
||||
|
||||
@@ -1,33 +0,0 @@
|
||||
# export WANDB_MODE="offline"
|
||||
GPU_NUM=1 # 2,4,8
|
||||
MODEL_PATH="data/FastMochi-diffusers"
|
||||
MODEL_TYPE="mochi"
|
||||
DATA_MERGE_PATH="data/Image-Vid-Finetune-Src/merge.txt"
|
||||
OUTPUT_DIR="data/Image-Vid-Finetune-Mochi"
|
||||
VALIDATION_PATH="assets/prompt.txt"
|
||||
|
||||
torchrun --nproc_per_node=$GPU_NUM \
|
||||
fastvideo/data_preprocess/preprocess_vae_latents.py \
|
||||
--model_path $MODEL_PATH \
|
||||
--data_merge_path $DATA_MERGE_PATH \
|
||||
--train_batch_size=1 \
|
||||
--max_height=480 \
|
||||
--max_width=848 \
|
||||
--num_frames=93 \
|
||||
--dataloader_num_workers 1 \
|
||||
--output_dir=$OUTPUT_DIR \
|
||||
--model_type $MODEL_TYPE \
|
||||
--train_fps 24
|
||||
|
||||
torchrun --nproc_per_node=$GPU_NUM \
|
||||
fastvideo/data_preprocess/preprocess_text_embeddings.py \
|
||||
--model_type $MODEL_TYPE \
|
||||
--model_path $MODEL_PATH \
|
||||
--output_dir=$OUTPUT_DIR
|
||||
|
||||
torchrun --nproc_per_node=1 \
|
||||
fastvideo/data_preprocess/preprocess_validation_text_embeddings.py \
|
||||
--model_type $MODEL_TYPE \
|
||||
--model_path $MODEL_PATH \
|
||||
--output_dir=$OUTPUT_DIR \
|
||||
--validation_prompt_txt $VALIDATION_PATH
|
||||
+18
-15
@@ -1,23 +1,24 @@
|
||||
# export WANDB_MODE="offline"
|
||||
GPU_NUM=1 # 2,4,8
|
||||
GPU_NUM=8
|
||||
MODEL_PATH="data/hunyuan"
|
||||
MODEL_TYPE="hunyuan"
|
||||
DATA_MERGE_PATH="data/Image-Vid-Finetune-Src/merge.txt"
|
||||
OUTPUT_DIR="data/Image-Vid-Finetune-HunYuan"
|
||||
DATA_MERGE_PATH="data/Mixkit-All-Clips/merge.txt"
|
||||
OUTPUT_DIR="data/Hunyuan-Mixkit-Data"
|
||||
VALIDATION_PATH="assets/prompt.txt"
|
||||
# torchrun --nproc_per_node=$GPU_NUM \
|
||||
# fastvideo/data_preprocess/preprocess_vae_latents.py \
|
||||
# --model_path $MODEL_PATH \
|
||||
# --data_merge_path $DATA_MERGE_PATH \
|
||||
# --train_batch_size=1 \
|
||||
# --max_height=480 \
|
||||
# --max_width=848 \
|
||||
# --num_frames=93 \
|
||||
# --dataloader_num_workers 1 \
|
||||
# --output_dir=$OUTPUT_DIR \
|
||||
# --model_type $MODEL_TYPE \
|
||||
# --train_fps 24
|
||||
|
||||
|
||||
torchrun --nproc_per_node=$GPU_NUM \
|
||||
fastvideo/data_preprocess/preprocess_vae_latents.py \
|
||||
--model_path $MODEL_PATH \
|
||||
--data_merge_path $DATA_MERGE_PATH \
|
||||
--train_batch_size=1 \
|
||||
--max_height=480 \
|
||||
--max_width=848 \
|
||||
--num_frames=93 \
|
||||
--dataloader_num_workers 1 \
|
||||
--output_dir=$OUTPUT_DIR \
|
||||
--model_type $MODEL_TYPE \
|
||||
--train_fps 24
|
||||
|
||||
torchrun --nproc_per_node=$GPU_NUM \
|
||||
fastvideo/data_preprocess/preprocess_text_embeddings.py \
|
||||
@@ -25,6 +26,8 @@ torchrun --nproc_per_node=$GPU_NUM \
|
||||
--model_path $MODEL_PATH \
|
||||
--output_dir=$OUTPUT_DIR
|
||||
|
||||
|
||||
|
||||
torchrun --nproc_per_node=1 \
|
||||
fastvideo/data_preprocess/preprocess_validation_text_embeddings.py \
|
||||
--model_type $MODEL_TYPE \
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user