Compare commits

..
Author SHA1 Message Date
Your Name 1856f01c06 update 2024-12-12 20:59:03 -08:00
Your Name 55442027f8 update 2024-12-12 20:50:16 -08:00
Your Name 710b9c549b update 2024-12-12 18:33:49 -08:00
rlsu9 40120579f2 update 2024-12-12 22:28:41 +00:00
rlsu9 f8409886a1 update 2024-12-12 21:35:29 +00:00
rlsu9 f37e1fa18e update 2024-12-12 21:32:43 +00:00
rlsu9 4917d58d8e update 2024-12-12 21:30:45 +00:00
rlsu9 bbeb227bc5 add scripts 2024-12-12 21:26:36 +00:00
rlsu9 163f9227ff update 2024-12-12 21:26:03 +00:00
rlsu9 98070d9bca add val gen 2024-12-12 19:54:22 +00:00
rlsu9 dfb15828a8 rename preprocess 2024-12-12 19:53:57 +00:00
rlsu9 ef667b9942 move data preprocess 2024-12-12 19:12:58 +00:00
rlsu9 01282b7622 debug ac; disitll; move fsdp_util 2024-12-12 19:07:43 +00:00
rlsu9 b14cc034e4 update 2024-12-12 18:11:04 +00:00
Peiyuan Zhang 6e5cb4b8e5 update 2024-12-12 06:34:22 +00:00
Peiyuan Zhang 87913376be add training 2024-12-12 05:59:11 +00:00
Peiyuan Zhang ce22b1279d vae done 2024-12-12 00:41:02 +00:00
Peiyuan Zhang f1dba33df7 update 2024-12-11 23:56:13 +00:00
Peiyuan Zhang e8ce4290ff update 2024-12-11 20:27:18 +00:00
Peiyuan Zhang fb1cc737e3 the same input output format for transformer 2024-12-11 19:56:57 +00:00
Peiyuan Zhang f5074a5c31 refactor single card attention to no pad 2024-12-11 19:00:10 +00:00
Peiyuan Zhang 0d4ae825c9 add no pad; support batch size > 1 2024-12-11 18:39:30 +00:00
Peiyuan Zhang da66e0fcea move rope freq inside transformer 2024-12-11 18:18:44 +00:00
foreverpiano 97f9db9b71 update 2024-12-11 13:36:38 +00:00
foreverpiano e297fe3949 make distill a framework 2024-12-11 13:33:24 +00:00
foreverpiano dd2cf8ad69 update 2024-12-11 13:04:48 +00:00
foreverpiano f5f5903aa4 update sp 2024-12-10 13:59:33 +00:00
foreverpiano 20b03a45d7 update 2024-12-09 10:25:53 +00:00
foreverpiano 2e4b523bee fps=24 2024-12-09 10:16:26 +00:00
foreverpiano ea0332f62a inference hunyuan ok 2024-12-09 10:10:59 +00:00
59 changed files with 841 additions and 1936 deletions
+159 -16
View File
@@ -1,19 +1,162 @@
prepare the src data
```
python scripts/download_hf.py --repo_id=FastVideo/Distill-30K-Data-Src --local_dir=data/Distill-30K-Data-Src --repo_type=dataset
cd data/Distill-30K-Data-Src
cat archive_name.tar.gz.* > Distill-30K-Data-Src.tar.gz
rm archive_name.tar.gz.*
tar --use-compress-program="pigz --processes 64" -xvf Distill-30K-Data-Src.tar.gz
mv ephemeral/hao.zhang/codefolder/FastVideo-OSP/data/Distill-30K-Src/* .
rm -r ephemeral
rm Distill-30K-Data-Src.tar.gz
cd ../..
```
Now the src data is stored in folder data/Distill-30K-Data-Src, adjust the path in `data/Distill-30K-Src/merge.txt` correspondingly
# FastVideo
<div align="center">
<a href=""><img src="https://img.shields.io/static/v1?label=API:H100&message=Replicate&color=pink"></a> &ensp;
<a href=""><img src="https://img.shields.io/static/v1?label=Discuss&message=Discord&color=purple&logo=discord"></a> &ensp;
</div>
<br>
<div align="center">
<img src=assets/logo.png width="50%"/>
</div>
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
- 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/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
Adjust the `SHARD_NUM` (total shards num) and `SHARD_IDX` (idx of the shard to preprocess) correspoindingly in `./scripts/preprocess_hunyuan_data.sh` for each GPU node to preprocess different shards
```
bash ./scripts/preprocess_hunyuan_data.sh
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 .
```
The preprocessed data will be stored in `data/HD-Hunyuan-30K-Distill-Data_Shard*`
## 🚀 Inference
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
python scripts/download_hf.py --repo_id=FastVideo/FastMochi --local_dir=data/FastMochi --repo_type=model
```
Start the gradio UI with
```
python fastvideo/demo/gradio_web_demo.py --model_path data/FastMochi
```
We also provide CLI inference script featured with sequence parallelism.
```
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
## 💰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
asset/example_data
├── AAA.txt
├── AAA.png
├── BCC.txt
├── BCC.png
├── ......
├── CCC.txt
└── CCC.png
```
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
```
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 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.

Before

Width:  |  Height:  |  Size: 26 MiB

Binary file not shown.
@@ -14,8 +14,6 @@ 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,
@@ -52,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():
@@ -67,11 +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)
@@ -79,8 +78,6 @@ if __name__ == "__main__":
parser.add_argument("--model_type", type=str, default="mochi")
parser.add_argument("--data_merge_path", type=str, required=True)
parser.add_argument("--num_frames", type=int, default=163)
parser.add_argument("--shard_num", type=int)
parser.add_argument("--shard_idx", type=int)
parser.add_argument(
"--dataloader_num_workers",
type=int,
@@ -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")
+1 -1
View File
@@ -46,7 +46,7 @@ class LatentDataset(Dataset):
map_location="cpu",
weights_only=True,
)
latent = latent.squeeze(0)[:, :self.num_latent_t]
latent = latent.squeeze(0)[:, -self.num_latent_t :]
if random.random() < self.cfg_rate:
prompt_embed = self.uncond_prompt_embed
prompt_attention_mask = self.uncond_prompt_mask
+6 -11
View File
@@ -92,8 +92,6 @@ class T2V_dataset(Dataset):
self.v_decoder = DecordInit()
self.video_length_tolerance_range = args.video_length_tolerance_range
self.support_Chinese = True
self.shard_num = args.shard_num
self.shard_idx = args.shard_idx
if not ("mt5" in args.text_encoder_name):
self.support_Chinese = False
@@ -107,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
@@ -199,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,
@@ -272,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
@@ -342,11 +341,7 @@ class T2V_dataset(Dataset):
for i in range(len(sub_list)):
sub_list[i]["path"] = opj(folder, sub_list[i]["path"])
cap_lists += sub_list
total = len(cap_lists)
shard_size = (total + self.shard_num - 1) // self.shard_num # Ceiling division
start_idx = self.shard_idx * shard_size
end_idx = min(start_idx + shard_size, total)
return cap_lists[start_idx:end_idx]
return cap_lists
def get_cap_list(self):
cap_lists = self.read_jsons(self.data)
+32 -94
View File
@@ -9,10 +9,9 @@ import tempfile
import os
import argparse
def init_args():
parser = argparse.ArgumentParser()
parser.add_argument("--prompts", nargs="+", default=[])
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)
@@ -30,60 +29,37 @@ def init_args():
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:
scheduler = PCMFMScheduler(
1000,
args.shift,
args.num_euler_timesteps,
False,
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_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],
@@ -95,23 +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()
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=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,
)
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,20 +133,10 @@ 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__":
demo.queue(max_size=20).launch(server_name="0.0.0.0", server_port=7860)
demo.queue(max_size=20).launch(server_name="0.0.0.0", server_port=7860)
+43 -59
View File
@@ -30,7 +30,7 @@ from fastvideo.utils.fsdp_util import get_dit_fsdp_kwargs, apply_fsdp_checkpoint
from diffusers import (
FlowMatchEulerDiscreteScheduler,
)
from fastvideo.utils.load import load_transformer
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
@@ -53,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)
@@ -75,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:
@@ -131,6 +129,7 @@ def distill_one_step(
ema_decay,
pred_decay_weight,
pred_decay_type,
hunyuan_student_cfg_embed
):
total_loss = 0.0
optimizer.zero_grad()
@@ -140,13 +139,6 @@ def distill_one_step(
"absolute mean": 0.0,
"absolute max": 0.0,
}
distill_cfg = distill_cfg.split(",")
distill_cfg = [float(cfg) for cfg in distill_cfg]
# randomly pick
distill_cfg = distill_cfg[torch.randint(0, len(distill_cfg), (1,)).item()]
distill_cfg = torch.tensor([distill_cfg], device=transformer.device, dtype=torch.bfloat16) * 1000
if sp_size > 1:
broadcast(distill_cfg)
for _ in range(gradient_accumulation_steps):
(
latents,
@@ -176,16 +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):
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 model_type == "hunyuan":
kwargs["guidance"] = distill_cfg
model_pred = transformer(**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(
@@ -194,18 +186,14 @@ def distill_one_step(
with torch.no_grad():
w = distill_cfg
with torch.autocast("cuda", dtype=torch.bfloat16):
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 model_type == "hunyuan":
kwargs["guidance"] = distill_cfg
cond_teacher_output = teacher_transformer(**kwargs)[0].float()
if not_apply_cfg_solver or model_type == "hunyuan":
cond_teacher_output = teacher_transformer(
noisy_model_input,
encoder_hidden_states,
timesteps,
encoder_attention_mask, # B, L
return_dict=False,
)[0].float()
if not_apply_cfg_solver:
uncond_teacher_output = cond_teacher_output
else:
# Get teacher model prediction on noisy_latents and unconditional embedding
@@ -225,19 +213,22 @@ def distill_one_step(
# 20.4.12. Get target LCM prediction on x_prev, w, c, t_n
with torch.no_grad():
with torch.autocast("cuda", dtype=torch.bfloat16):
kwargs = {
"hidden_states": x_prev,
"encoder_hidden_states": encoder_hidden_states,
"timestep": timesteps_prev,
"encoder_attention_mask": encoder_attention_mask, # B, L
"return_dict": False,
}
if model_type == "hunyuan":
kwargs["guidance"] = distill_cfg
if ema_transformer is not None:
target_pred = ema_transformer(**kwargs)[0]
target_pred = ema_transformer(
x_prev.float(),
encoder_hidden_states,
timesteps_prev,
encoder_attention_mask, # B, L
return_dict=False,
)[0]
else:
target_pred = transformer(**kwargs)[0]
target_pred = transformer(
x_prev.float(),
encoder_hidden_states,
timesteps_prev,
encoder_attention_mask, # B, L
return_dict=False,
)[0]
target, end_index = solver.euler_style_multiphase_pred(
x_prev, target_pred, index, multiphase, True
@@ -327,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:
@@ -389,16 +376,10 @@ 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, 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)
@@ -584,6 +565,7 @@ def main(args):
args.ema_decay,
args.pred_decay_weight,
args.pred_decay_type,
args.hunyuan_student_cfg_embed
)
step_time = time.time() - start_time
@@ -669,15 +651,16 @@ 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
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",
@@ -886,7 +869,7 @@ if __name__ == "__main__":
help="Whether to apply the cfg_solver.",
)
parser.add_argument(
"--distill_cfg", type=str
"--distill_cfg", type=float, default=3.0, help="Distillation coefficient."
)
# ["euler_linear_quadratic", "pcm", "pcm_linear_qudratic"]
parser.add_argument(
@@ -911,6 +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_student_cfg_embed", type=float)
parser.add_argument(
"--master_weight_type",
type=str,
+7 -12
View File
@@ -46,14 +46,10 @@ class DiscriminatorHead(nn.Module):
def forward(self, x):
b, twh, c = x.shape
# from fastvideo.utils.logging_ import ForkedPdb
# ForkedPdb().set_trace()
after_patching_height = 49
after_patching_wight = 80
t = twh // (after_patching_height * after_patching_wight)
x = x.view(-1, after_patching_height * after_patching_wight, c)
t = twh // (30 * 53)
x = x.view(-1, 30 * 53, c)
x = x.permute(0, 2, 1)
x = x.view(b * t, c, after_patching_height, after_patching_wight)
x = x.view(b * t, c, 30, 53)
x = self.conv1(x)
x = self.conv2(x) + x
x = self.conv_out(x)
@@ -66,10 +62,9 @@ class Discriminator(nn.Module):
stride=8,
num_h_per_head=1,
adapter_channel_dims=[3072],
total_layers = 48,
):
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)
@@ -94,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],
+62 -112
View File
@@ -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,
+1 -1
View File
@@ -31,4 +31,4 @@ def flash_attn_no_pad(
"b s (h d) -> b s h d",
h=nheads,
)
return output
return output
+5 -5
View File
@@ -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,7 +43,7 @@ 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"
@@ -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(
@@ -911,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
)
@@ -931,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,
@@ -941,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:
@@ -986,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)
+11 -12
View File
@@ -9,11 +9,7 @@ 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
@@ -56,7 +52,9 @@ class Inference(object):
self.device = (
device
if device is not None
else "cuda" if torch.cuda.is_available() else "cpu"
else "cuda"
if torch.cuda.is_available()
else "cpu"
)
self.logger = logger
self.parallel_args = parallel_args
@@ -73,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 =============================
@@ -173,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
@@ -279,7 +277,7 @@ class HunyuanVideoSampler(Inference):
use_cpu_offload=False,
device=0,
logger=None,
parallel_args=None,
parallel_args=None
):
super().__init__(
args,
@@ -292,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(
@@ -461,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,
)
+22 -20
View File
@@ -5,12 +5,13 @@ 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,
@@ -24,34 +25,36 @@ def attention(
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)
@@ -66,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
)
@@ -80,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(
+47 -47
View File
@@ -11,7 +11,7 @@ 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
@@ -22,7 +22,6 @@ from fastvideo.utils.parallel_states import (
nccl_info,
)
class MMDoubleStreamBlock(nn.Module):
"""
A multimodal dit block with seperate modulation for
@@ -134,6 +133,7 @@ class MMDoubleStreamBlock(nn.Module):
def disable_deterministic(self):
self.deterministic = False
def forward(
self,
img: torch.Tensor,
@@ -174,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
@@ -204,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),
@@ -211,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] :]
@@ -272,7 +271,7 @@ class MMSingleStreamBlock(nn.Module):
head_dim = hidden_size // heads_num
mlp_hidden_dim = int(hidden_size * mlp_width_ratio)
self.mlp_hidden_dim = mlp_hidden_dim
self.scale = qk_scale or head_dim**-0.5
self.scale = qk_scale or head_dim ** -0.5
# qkv and mlp_in
self.linear1 = nn.Linear(
@@ -334,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:, :, :]
@@ -350,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),
@@ -357,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.
@@ -441,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__()
@@ -453,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}"
@@ -486,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(
@@ -594,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],
@@ -655,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 = [
@@ -673,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, ...]
@@ -686,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,
@@ -152,7 +152,6 @@ class SingleTokenRefiner(nn.Module):
"""
A single token refiner block for llm text embedding refine.
"""
def __init__(
self,
in_channels,
+1 -3
View File
@@ -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
@@ -15,9 +15,12 @@ def preprocess_text_encoder_tokenizer(args):
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__":
+12 -16
View File
@@ -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)
+12 -31
View File
@@ -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,8 +179,8 @@ 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"
@@ -238,80 +188,48 @@ def convert_diffusers_vae_to_mochi(state_dict):
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(
f"{prefix}down_blocks.{block}.resnets.{i}.norm1.norm_layer.weight"
)
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(
f"{prefix}down_blocks.{block}.resnets.{i}.norm1.norm_layer.bias"
)
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(
f"{prefix}down_blocks.{block}.resnets.{i}.conv1.conv.weight"
)
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(
f"{prefix}down_blocks.{block}.resnets.{i}.conv1.conv.bias"
)
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(
f"{prefix}down_blocks.{block}.resnets.{i}.norm2.norm_layer.weight"
)
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(
f"{prefix}down_blocks.{block}.resnets.{i}.norm2.norm_layer.bias"
)
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(
f"{prefix}down_blocks.{block}.resnets.{i}.conv2.conv.weight"
)
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(
f"{prefix}down_blocks.{block}.resnets.{i}.conv2.conv.bias"
)
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.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.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.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.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,45 +321,29 @@ 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(
f"{prefix}up_blocks.{block}.resnets.{i}.norm1.norm_layer.weight"
)
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(
f"{prefix}up_blocks.{block}.resnets.{i}.norm1.norm_layer.bias"
)
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(
f"{prefix}up_blocks.{block}.resnets.{i}.conv1.conv.weight"
)
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(
f"{prefix}up_blocks.{block}.resnets.{i}.conv1.conv.bias"
)
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(
f"{prefix}up_blocks.{block}.resnets.{i}.norm2.norm_layer.weight"
)
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(
f"{prefix}up_blocks.{block}.resnets.{i}.norm2.norm_layer.bias"
)
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(
f"{prefix}up_blocks.{block}.resnets.{i}.conv2.conv.weight"
)
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(
f"{prefix}up_blocks.{block}.resnets.{i}.conv2.conv.bias"
)
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(
f"{prefix}up_blocks.{block}.proj.weight"
@@ -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")
+8 -7
View File
@@ -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)
+38 -120
View File
@@ -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,21 +37,19 @@ 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
@@ -60,7 +57,7 @@ def main(args):
samples = []
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,
@@ -71,10 +68,10 @@ 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
)
samples.append(outputs["samples"][0])
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 = []
@@ -85,10 +82,9 @@ def main(args):
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)
@@ -98,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)
+37 -119
View File
@@ -15,22 +15,19 @@ 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
)
hunyuan_video_sampler = HunyuanVideoSampler.from_pretrained(models_root_path, args=args)
# Get the updated args
args = hunyuan_video_sampler.args
@@ -38,7 +35,7 @@ def main(args):
samples = []
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,
@@ -49,10 +46,10 @@ 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
)
samples.append(outputs["samples"][0])
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 = []
@@ -63,10 +60,9 @@ def main(args):
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)
@@ -76,133 +72,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)
+3 -4
View File
@@ -50,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
@@ -221,9 +220,9 @@ def main(args):
transformer = MochiTransformer3DModel.from_pretrained(
args.pretrained_model_name_or_path,
subfolder="transformer",
torch_dtype=(
torch.float32 if args.master_weight_type == "fp32" else torch.bfloat16
),
torch_dtype=torch.float32
if args.master_weight_type == "fp32"
else torch.bfloat16,
)
if args.use_lora:
-1
View File
@@ -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):
+2 -6
View File
@@ -36,7 +36,7 @@ non_reentrant_wrapper = partial(
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
@@ -80,11 +80,7 @@ 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:
+37 -70
View File
@@ -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,7 +35,7 @@ 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"
@@ -58,20 +49,20 @@ PROMPT_TEMPLATE = {
"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")
@@ -89,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,
@@ -160,32 +149,25 @@ 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, "text_encoder")
)
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
device = self.text_encoder.device
dtype = self.dtype
prompt = [prompt] if isinstance(prompt, str) else prompt
@@ -223,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
@@ -250,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(
@@ -284,7 +259,6 @@ def load_transformer(
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":
@@ -295,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)
@@ -322,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):
@@ -333,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):
@@ -342,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)
+32 -36
View File
@@ -23,7 +23,6 @@ import wandb
import gc
from fastvideo.utils.load import load_vae
def prepare_latents(
batch_size,
num_channels_latents,
@@ -57,7 +56,6 @@ def sample_validation_video(
num_inference_steps: int = 28,
timesteps: List[int] = None,
guidance_scale: float = 4.5,
embed_guidance_scale: float = None,
num_videos_per_prompt: Optional[int] = 1,
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
prompt_embeds: Optional[torch.Tensor] = None,
@@ -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
@@ -139,16 +137,14 @@ def sample_validation_video(
# broadcast to batch dimension in a way that's compatible with ONNX/Core ML
timestep = t.expand(latent_model_input.shape[0])
with torch.autocast("cuda", dtype=torch.bfloat16):
kwargs = {
"hidden_states": latent_model_input,
"encoder_hidden_states": prompt_embeds,
"timestep": timestep,
"encoder_attention_mask": prompt_attention_mask,
"return_dict": False,
}
if embed_guidance_scale > 0.0:
kwargs["guidance"] = torch.tensor([embed_guidance_scale], device=device, dtype=torch.bfloat16) * 1000
noise_pred = transformer(**kwargs)[0]
noise_pred = transformer(
hidden_states=latent_model_input,
encoder_hidden_states=prompt_embeds,
timestep=timestep,
encoder_attention_mask=prompt_attention_mask,
return_dict=False,
)[0]
# Mochi CFG + Sampling runs in FP32
noise_pred = noise_pred.to(torch.float32)
if do_classifier_free_guidance:
@@ -216,7 +212,7 @@ def log_validation(
args,
transformer,
device,
weight_dtype, # TODO
weight_dtype, # TODO
global_step,
scheduler_type="euler",
shift=1.0,
@@ -228,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()
@@ -274,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)
@@ -297,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)
)
@@ -309,11 +306,10 @@ def log_validation(
scheduler_type=scheduler_type,
num_frames=args.num_frames,
# Peiyuan TODO: remove hardcode
height=args.num_height,
width=args.num_width,
height=480,
width=848,
num_inference_steps=validation_sampling_step,
guidance_scale=0.0 if args.model_type == "hunyuan" else validation_guidance_scale,
embed_guidance_scale= validation_guidance_scale if args.model_type == "hunyuan" else 0.0,
guidance_scale=validation_guidance_scale,
generator=generator,
prompt_embeds=prompt_embeds,
prompt_attention_mask=prompt_attention_mask,
@@ -321,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])
+6 -1
View File
@@ -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"]
"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*"]
+2
View File
@@ -1,5 +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
+9 -7
View File
@@ -2,11 +2,13 @@ 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_adv.py\
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"\
@@ -15,11 +17,11 @@ torchrun --nnodes 1 --nproc_per_node 8\
--gradient_checkpointing\
--train_batch_size=1\
--num_latent_t 24\
--sp_size 4\
--sp_size 1\
--train_sp_batch_size 1\
--dataloader_num_workers 4\
--gradient_accumulation_steps=1\
--max_train_steps=640\
--gradient_accumulation_steps=2\
--max_train_steps=480\
--learning_rate=1e-6\
--mixed_precision="bf16"\
--checkpointing_steps=64\
@@ -30,8 +32,8 @@ torchrun --nnodes 1 --nproc_per_node 8\
--ema_start_step 0\
--cfg 0.0\
--log_validation\
--output_dir="$DATA_DIR/outputs/hy_phase1_shift17_bs_32_adv"\
--tracker_project_name Hunyuan_Distill \
--output_dir="$DATA_DIR/outputs/hy_phase1_shift17_bs_32"\
--tracker_project_name PCM \
--num_frames 93 \
--shift 17 \
--validation_guidance_scale "1.0" \
+6 -4
View File
@@ -31,7 +31,7 @@ torchrun --nnodes 2 --nproc_per_node 8\
--dataloader_num_workers 4\
--gradient_accumulation_steps=2\
--max_train_steps=640\
--learning_rate=3e-7\
--learning_rate=1e-6\
--mixed_precision="bf16"\
--checkpointing_steps=64\
--validation_steps 64\
@@ -41,11 +41,13 @@ torchrun --nnodes 2 --nproc_per_node 8\
--ema_start_step 0\
--cfg 0.0\
--log_validation\
--output_dir="$DATA_DIR/outputs/hy_phase1_shift17_bs_32_lr3e-7"\
--output_dir="$DATA_DIR/outputs/hy_phase1_shift33_student3_batchsize32"\
--tracker_project_name Hunyuan_Distill \
--num_frames 93 \
--shift 17 \
--shift 33 \
--validation_guidance_scale "1.0" \
--num_euler_timesteps 50 \
--multi_phased_distill_schedule "4000-1" \
--not_apply_cfg_solver
--not_apply_cfg_solver \
--hunyuan_student_cfg_embed 3
-51
View File
@@ -1,51 +0,0 @@
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_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-1" \
--not_apply_cfg_solver
-53
View File
@@ -1,53 +0,0 @@
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_adv.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 32\
--sp_size 4 \
--train_sp_batch_size 1\
--dataloader_num_workers 4\
--gradient_accumulation_steps=1\
--max_train_steps=2000\
--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_8_adv_HD"\
--tracker_project_name Hunyuan_Distill \
--num_height 720 \
--num_width 1280 \
--num_frames 125 \
--shift 17 \
--validation_guidance_scale "1.0" \
--num_euler_timesteps 50 \
--multi_phased_distill_schedule "4000-1" \
--not_apply_cfg_solver
-51
View File
@@ -1,51 +0,0 @@
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_shift17_bs_32_euler20"\
--tracker_project_name Hunyuan_Distill \
--num_frames 93 \
--shift 17 \
--validation_guidance_scale "1.0" \
--num_euler_timesteps 20 \
--multi_phased_distill_schedule "4000-1" \
--not_apply_cfg_solver
-51
View File
@@ -1,51 +0,0 @@
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 4 --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-30K-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=2000\
--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_32_moredata"\
--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
-53
View File
@@ -1,53 +0,0 @@
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 4 --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"\
--gradient_checkpointing\
--train_batch_size=1\
--num_latent_t 32 \
--sp_size 2 \
--train_sp_batch_size 1\
--dataloader_num_workers 4\
--gradient_accumulation_steps=1\
--max_train_steps=2000\
--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_HD"\
--tracker_project_name Hunyuan_Distill \
--num_height 720 \
--num_width 1280 \
--num_frames 125 \
--shift 17 \
--validation_guidance_scale "1.0" \
--num_euler_timesteps 50 \
--multi_phased_distill_schedule "4000-1" \
--not_apply_cfg_solver
-53
View File
@@ -1,53 +0,0 @@
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 4 --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"\
--gradient_checkpointing\
--train_batch_size=1\
--num_latent_t 32 \
--sp_size 2 \
--train_sp_batch_size 1\
--dataloader_num_workers 4\
--gradient_accumulation_steps=1\
--max_train_steps=2000\
--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_16_HD"\
--tracker_project_name Hunyuan_Distill \
--num_height 720 \
--num_width 1280 \
--num_frames 125 \
--shift 25 \
--validation_guidance_scale "1.0" \
--num_euler_timesteps 50 \
--multi_phased_distill_schedule "4000-1" \
--not_apply_cfg_solver
-53
View File
@@ -1,53 +0,0 @@
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 4 --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"\
--gradient_checkpointing\
--train_batch_size=1\
--num_latent_t 32 \
--sp_size 2 \
--train_sp_batch_size 1\
--dataloader_num_workers 4\
--gradient_accumulation_steps=1\
--max_train_steps=2000\
--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_16_HD"\
--tracker_project_name Hunyuan_Distill \
--num_height 720 \
--num_width 1280 \
--num_frames 125 \
--shift 33 \
--validation_guidance_scale "1.0" \
--num_euler_timesteps 50 \
--multi_phased_distill_schedule "4000-1" \
--not_apply_cfg_solver
-58
View File
@@ -1,58 +0,0 @@
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 4 --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"\
--gradient_checkpointing\
--train_batch_size=1\
--num_latent_t 32 \
--sp_size 2 \
--train_sp_batch_size 1\
--dataloader_num_workers 4\
--gradient_accumulation_steps=1\
--max_train_steps=2000\
--learning_rate=1e-6\
--mixed_precision="bf16"\
--checkpointing_steps=64\
--validation_steps 64\
--validation_sampling_steps "4,6,8" \
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--cfg 0.0\
--log_validation\
--output_dir="$DATA_DIR/outputs/hy_phase1_lq_0.1_0.67_bs_16_HD"\
--tracker_project_name Hunyuan_Distill \
--num_height 720 \
--num_width 1280 \
--num_frames 125 \
--scheduler_type pcm_linear_quadratic \
--linear_quadratic_threshold 0.1 \
--linear_range 0.67 \
--validation_guidance_scale "6.0" \
--num_euler_timesteps 50 \
--multi_phased_distill_schedule "4000-1" \
--not_apply_cfg_solver \
--distill_cfg "6.0" \
-58
View File
@@ -1,58 +0,0 @@
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 4 --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"\
--gradient_checkpointing\
--train_batch_size=1\
--num_latent_t 32 \
--sp_size 2 \
--train_sp_batch_size 1\
--dataloader_num_workers 4\
--gradient_accumulation_steps=1\
--max_train_steps=2000\
--learning_rate=1e-6\
--mixed_precision="bf16"\
--checkpointing_steps=64\
--validation_steps 64\
--validation_sampling_steps "4,6,8" \
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--cfg 0.0\
--log_validation\
--output_dir="$DATA_DIR/outputs/hy_phase1_lq_0.05_0.67_bs_16_HD"\
--tracker_project_name Hunyuan_Distill \
--num_height 720 \
--num_width 1280 \
--num_frames 125 \
--scheduler_type pcm_linear_quadratic \
--linear_quadratic_threshold 0.05 \
--linear_range 0.67 \
--validation_guidance_scale "6.0" \
--num_euler_timesteps 50 \
--multi_phased_distill_schedule "4000-1" \
--not_apply_cfg_solver \
--distill_cfg "6.0"
-56
View File
@@ -1,56 +0,0 @@
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 4 --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"\
--gradient_checkpointing\
--train_batch_size=1\
--num_latent_t 32 \
--sp_size 2 \
--train_sp_batch_size 1\
--dataloader_num_workers 4\
--gradient_accumulation_steps=1\
--max_train_steps=2000\
--learning_rate=1e-6\
--mixed_precision="bf16"\
--checkpointing_steps=64\
--validation_steps 64\
--validation_sampling_steps "4,6,8" \
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--cfg 0.0\
--log_validation\
--output_dir="$DATA_DIR/outputs/hy_phase1_lq_0.1_0.67_bs_16_HD_random_cfg"\
--tracker_project_name Hunyuan_Distill \
--num_height 720 \
--num_width 1280 \
--num_frames 125 \
--scheduler_type pcm_linear_quadratic \
--linear_quadratic_threshold 0.1 \
--linear_range 0.67 \
--validation_guidance_scale "3.0,4.0,5.0,6.0" \
--num_euler_timesteps 50 \
--multi_phased_distill_schedule "4000-1" \
--not_apply_cfg_solver \
--distill_cfg "1.0,2.0,3.0,4.0,5.0,6.0"
-54
View File
@@ -1,54 +0,0 @@
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 4 --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"\
--gradient_checkpointing\
--train_batch_size=1\
--num_latent_t 32 \
--sp_size 2 \
--train_sp_batch_size 1\
--dataloader_num_workers 4\
--gradient_accumulation_steps=1\
--max_train_steps=2000\
--learning_rate=1e-6\
--mixed_precision="bf16"\
--checkpointing_steps=64\
--validation_steps 64\
--validation_sampling_steps "4,6,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_HD_random_cfg"\
--tracker_project_name Hunyuan_Distill \
--num_height 720 \
--num_width 1280 \
--num_frames 125 \
--shift 17 \
--validation_guidance_scale "3.0,4.0,5.0,6.0" \
--num_euler_timesteps 50 \
--multi_phased_distill_schedule "4000-1" \
--not_apply_cfg_solver \
--distill_cfg "1.0,2.0,3.0,4.0,5.0,6.0"
+16 -21
View File
@@ -1,31 +1,26 @@
# export WANDB_MODE="offline"
GPU_NUM=8
SHARD_NUM=8
SHARD_IDX=0
MODEL_PATH="data/hunyuan"
MODEL_TYPE="hunyuan"
DATA_MERGE_PATH="data/Distill-30K-Src/merge.txt"
OUTPUT_DIR="data/HD-Hunyuan-30K-Distill-Data_Shard${SHARD_IDX}"
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=720 \
--max_width=1280 \
--num_frames=129 \
--dataloader_num_workers 1 \
--output_dir=$OUTPUT_DIR \
--model_type $MODEL_TYPE \
--train_fps 24 \
--shard_num=$SHARD_NUM \
--shard_idx=$SHARD_IDX
torchrun --nproc_per_node=$GPU_NUM \
fastvideo/data_preprocess/preprocess_text_embeddings.py \
--model_type $MODEL_TYPE \
--model_path $MODEL_PATH \
+14 -10
View File
@@ -1,22 +1,26 @@
import torch
from fastvideo.models.hunyuan.diffusion.pipelines.pipeline_hunyuan_video import (
HunyuanVideoPipeline,
)
from fastvideo.models.hunyuan.diffusion.pipelines.pipeline_hunyuan_video import HunyuanVideoPipeline
from fastvideo.models.hunyuan.modules.models import HYVideoDiffusionTransformer
from fastvideo.models.hunyuan.vae.autoencoder_kl_causal_3d import AutoencoderKLCausal3D
transformer = HYVideoDiffusionTransformer.from_pretrained(
"data/hyvideo-diffusers", torch_dtype=torch.bfloat16, subfolder="transformer"
transformer=HYVideoDiffusionTransformer.from_pretrained(
'data/hyvideo-diffusers',
torch_dtype=torch.bfloat16,
subfolder='transformer'
)
vae = AutoencoderKLCausal3D.from_pretrained(
"data/hyvideo-diffusers", torch_dtype=torch.float16, subfolder="vae"
vae=AutoencoderKLCausal3D.from_pretrained(
'data/hyvideo-diffusers',
torch_dtype=torch.float16,
subfolder='vae'
)
pipe = HunyuanVideoPipeline.from_pretrained(
"data/hyvideo-diffusers", transformer=transformer, vae=vae
'data/hyvideo-diffusers',
transformer=transformer,
vae=vae
)
pipe = pipe.to("cuda")
pipe = pipe.to('cuda')
pipe.vae.enable_tiling()
prompt = "Close-up, A little girl wearing a red hoodie in winter strikes a match. The sky is dark, there is a layer of snow on the ground, and it is still snowing lightly. The flame of the match flickers, illuminating the girl's face intermittently."
@@ -35,4 +39,4 @@ output = result.videos[0].permute(1, 2, 3, 0).detach().cpu().numpy()
output = (output * 255).clip(0, 255).astype("uint8")
output = [PIL.Image.fromarray(x) for x in output]
export_to_video(output, "output.mp4", fps=24)
export_to_video(output, "output.mp4", fps=24)