Compare commits

...
Author SHA1 Message Date
Will Lin e4ceadb5d5 fix denoising stage init 2025-06-06 11:44:36 -07:00
Zhang PeiyuanandWill Lin 8f8ce6d9e1 [bugfix] [training] Add negative prompt to preprocessing and validation (#479)
Co-authored-by: Will Lin <wlsaidhi@gmail.com>
2025-06-06 11:21:21 -07:00
KevinzzandBrianChen1129 e3d0cbe185 [STA] Implement mask search for V1's Wan2.1 (#415)
Co-authored-by: BrianChen1129 <yongqich@umich.edu>
2025-06-05 15:03:31 -04:00
eeecho d5ec468d43 [V0] [Distill] support distill in V0 for wan (#444) 2025-06-04 16:20:14 -07:00
Yongqi Chen e55fa6e5dc [Preprocess] I2V dataset (#473) 2025-06-04 16:19:20 -07:00
Wenxuan Tan 61b6ddeee1 [Issue template] Move env report to the end for readability (#476) 2025-06-04 16:18:22 -07:00
William Lin a9a000f45d [bugfix] fix bz >1 for training (#477) 2025-06-04 16:03:39 -07:00
Wenxuan Tan 66b8b8561e [LoRA] Support V1 LoRA inference (#451) 2025-06-04 16:19:29 -05:00
Wenxuan Tan 6684872616 [misc] Find unused port in distributed init (#475) 2025-06-04 14:15:12 -05:00
Wenxuan Tan 7f654e3332 [misc] Polish V1 training code (#469) 2025-06-03 22:02:03 -05:00
William Lin 8631c1b806 [Training] Support Multi-Node training with FSDP + SP (#459) 2025-06-03 17:40:22 -07:00
Wei Zhouand“BrianChen1129” 5357e12b5a Update v1 inference scripts (#467)
Co-authored-by: “BrianChen1129” <yongqich@umich.edu>
2025-06-03 02:54:12 -04:00
Kevin Lin bdfdf1dfee [Training] Add distributed checkpointing (#458) 2025-06-02 15:15:45 -07:00
Wei Zhou d156461785 Fix WanVideo (#461) 2025-06-01 23:08:11 -07:00
Yongqi Chen 6edf113838 [Misc] Bring back mask files under asset/ and update new Wan mask strategy file (#462) 2025-06-01 19:46:11 -07:00
William Lin dcf7738cbc [Misc] disable cast_forward_inputs (#460) 2025-05-31 22:06:29 -07:00
Wenxuan Tan b2ebaaf865 [Misc] Remove InferenceEngine (#455) 2025-05-31 11:12:05 -07:00
Wenxuan Tan 7768bb80f6 misc: add remote pdb for debugging workers (#456) 2025-05-30 20:44:27 -05:00
97 changed files with 1341264 additions and 1081 deletions
+8 -8
View File
@@ -4,14 +4,6 @@ title: "[Bug] "
labels: ['Bug']
body:
- type: textarea
attributes:
label: Environment
description: |
Please share your environment with us. You can run the command **python fastvideo/utils/collect_env.py** and copy-paste its output below.
placeholder: FastVideo version, platform, python version, cuda version...
validations:
required: true
- type: textarea
attributes:
label: Describe the bug
@@ -25,5 +17,13 @@ body:
What command or script did you run? Which **model** are you using?
placeholder: |
A placeholder for the command.
validations:
required: true
- type: textarea
attributes:
label: Environment
description: |
Please share your environment with us. You can run the command **python fastvideo/utils/collect_env.py** and copy-paste its output below.
placeholder: FastVideo version, platform, python version, cuda version...
validations:
required: true
-1
View File
@@ -27,7 +27,6 @@ env
**/build/
**.pyc
**.txt
**.json
# Distribution / packaging
build/
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
@@ -57,8 +57,9 @@ Run the script with:
python example.py
```
The generated video will be saved in the current directory under `my_videos/`.
The generated video will be saved in the current directory under `my_videos/`
More inference example scripts can be found in `scripts/inference/`
## Available Models
Please see the [support matrix](#support-matrix) for the list of supported models and their available optimizations.
@@ -79,7 +80,6 @@ def main():
sampling_param = SamplingParam.from_pretrained(model_name)
sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
sampling_param.num_frames = 107
sampling_param.image_strength = 0.8 # How much to preserve the original image (0-1)
# Generate video based on the image
prompt = "A photograph coming to life with gentle movement"
-6
View File
@@ -72,12 +72,6 @@ FastVideo will automatically detect and use `FA3` if it is installed when using
pip install st_attn==0.0.4
```
Then download STA mask strategy from Hugging Face
```bash
python scripts/huggingface/download_hf.py --repo_id=FastVideo/STA_Mask_Strategy --local_dir=assets/ --repo_type=dataset
```
Please see [this page](#sta-installation) for more installation instructions.
(optimizations-sage)=
+6 -4
View File
@@ -2,7 +2,7 @@ from fastvideo import VideoGenerator
# from fastvideo.v1.configs.sample import SamplingParam
OUTPUT_PATH = "video_samples"
def main():
# FastVideo will automatically use the optimal default arguments for the
# model.
@@ -11,7 +11,9 @@ def main():
generator = VideoGenerator.from_pretrained(
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
# if num_gpus > 1, FastVideo will automatically handle distributed setup
num_gpus=1,
num_gpus=2,
use_fsdp_inference=True,
use_cpu_offload=False
)
# sampling_param = SamplingParam.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
@@ -23,7 +25,7 @@ def main():
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
)
video = generator.generate_video(prompt)
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True)
# video = generator.generate_video(prompt, sampling_param=sampling_param, output_path="wan_t2v_videos/")
# Generate another video with a different prompt, without reloading the
@@ -34,7 +36,7 @@ def main():
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
"cinematic.")
video2 = generator.generate_video(prompt2)
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True)
if __name__ == "__main__":
+1 -1
View File
@@ -1,5 +1,5 @@
from fastvideo import VideoGenerator
from fastvideo.v1.configs.pipelines.base import PipelineConfig
def main():
@@ -0,0 +1,45 @@
from fastvideo import VideoGenerator
from fastvideo.v1.configs.sample import SamplingParam
OUTPUT_PATH = "./lora"
def main():
# Initialize VideoGenerator with the Wan model
generator = VideoGenerator.from_pretrained(
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
num_gpus=2,
lora_path="benjamin-paine/steamboat-willie-1.3b",
lora_nickname="steamboat"
)
kwargs = {
"height": 480,
"width": 832,
"num_frames": 81,
"guidance_scale": 5.0,
"num_inference_steps": 32,
}
# Generate video with LoRA style
prompt = "steamboat willie style, golden era animation, close-up of a short fluffy monster kneeling beside a melting red candle. the mood is one of wonder and curiosity, as the monster gazes at the flame with wide eyes and open mouth. Its pose and expression convey a sense of innocence and playfulness, as if it is exploring the world around it for the first time. The use of warm colors and dramatic lighting further enhances the cozy atmosphere of the image."
negative_prompt = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
video = generator.generate_video(
prompt,
# sampling_param=sampling_param,
output_path=OUTPUT_PATH,
save_video=True,
negative_prompt=negative_prompt,
**kwargs
)
generator.set_lora_adapter(lora_nickname="flat_color", lora_path="motimalu/wan-flat-color-1.3b-v2")
prompt = "flat color, no lineart, blending, negative space, artist:[john kafka|ponsuke kaikai|hara id 21|yoneyama mai|fuzichoco], 1girl, sakura miko, pink hair, cowboy shot, white shirt, floral print, off shoulder, outdoors, cherry blossom, tree shade, wariza, looking up, falling petals, half-closed eyes, white sky, clouds, live2d animation, upper body, high quality cinematic video of a woman sitting under a sakura tree. Dreamy and lonely, the camera close-ups on the face of the woman as she turns towards the viewer. The Camera is steady, This is a cowboy shot. The animation is smooth and fluid."
negative_prompt = "bad quality video,色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
video = generator.generate_video(
prompt,
output_path=OUTPUT_PATH,
save_video=True,
negative_prompt=negative_prompt,
**kwargs
)
if __name__ == "__main__":
main()
@@ -0,0 +1,5 @@
# STA Mask Search Examples
```bash
bash examples/inference/sta_mask_search/inference_wan_sta.sh
```
@@ -0,0 +1,39 @@
#!/bin/bash
export FASTVIDEO_ATTENTION_CONFIG=assets/mask_strategy_wan.json
export FASTVIDEO_ATTENTION_BACKEND=SLIDING_TILE_ATTN
export MODEL_BASE=Wan-AI/Wan2.1-T2V-14B-Diffusers
base_port=29503
num_gpu=$(nvidia-smi --query-gpu=gpu_name --format=csv,noheader | wc -l)
gpu_ids=$(seq 0 $((num_gpu-1)))
skip_time_steps=12
output_path="inference_results/sta/mask_search_full"
STA_mode="STA_searching"
for i in $gpu_ids; do
port=$((base_port+i))
CUDA_VISIBLE_DEVICES=$i MASTER_PORT=$port python examples/inference/sta_mask_search/wan_example.py \
--prompt_path ./assets/prompt_extend_${i}.txt \
--output_path $output_path \
--STA_mode $STA_mode &
sleep 1
done
wait
echo "STA searching completed"
output_path="inference_results/sta/mask_search_sparse"
STA_mode="STA_tuning"
for i in $gpu_ids; do
port=$((base_port+i))
CUDA_VISIBLE_DEVICES=$i MASTER_PORT=$port python examples/inference/sta_mask_search/wan_example.py \
--prompt_path ./assets/prompt_extend_${i}.txt \
--output_path $output_path \
--STA_mode $STA_mode \
--skip_time_steps $skip_time_steps &
sleep 1
done
wait
echo "STA tuning completed"
echo "All jobs completed"
@@ -0,0 +1,63 @@
import os
import argparse
from fastvideo import VideoGenerator, SamplingParam
def main(args):
os.makedirs(args.output_path, exist_ok=True)
# Create a video generator with a pre-trained model
generator = VideoGenerator.from_pretrained(
"Wan-AI/Wan2.1-T2V-14B-Diffusers",
num_gpus=args.num_gpus, # Adjust based on your hardware
STA_mode=args.STA_mode,
skip_time_steps=args.skip_time_steps
)
# Prompts for your video
prompt = args.prompt
prompt_path = args.prompt_path
negative_prompt = "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"
if prompt_path is not None:
with open(prompt_path, "r") as f:
prompts = f.readlines()
else:
prompts = [prompt]
params = SamplingParam(
height=args.height,
width=args.width,
num_frames=args.num_frames,
num_inference_steps=args.num_inference_steps,
fps=args.fps,
guidance_scale=args.guidance_scale,
seed=args.seed,
return_frames=True, # Also return frames from this call (defaults to False)
output_path=args.output_path, # Controls where videos are saved
save_video=True,
negative_prompt=negative_prompt
)
# Generate the video
for prompt in prompts:
video = generator.generate_video(
prompt,
sampling_param=params,
)
if __name__ == '__main__':
parser = argparse.ArgumentParser()
parser.add_argument("--prompt", type=str, default="A man is dancing.")
parser.add_argument("--prompt_path", type=str, default=None)
parser.add_argument("--height", type=int, default=768)
parser.add_argument("--width", type=int, default=1280)
parser.add_argument("--num_frames", type=int, default=69)
parser.add_argument("--num_inference_steps", type=int, default=50)
parser.add_argument("--fps", type=int, default=16)
parser.add_argument("--guidance_scale", type=float, default=5.0)
parser.add_argument("--seed", type=int, default=12345)
parser.add_argument("--output_path", type=str, default="my_videos/")
parser.add_argument("--num_gpus", type=int, default=1)
parser.add_argument("--STA_mode", type=str, default="STA_searching")
parser.add_argument("--skip_time_steps", type=int, default=12)
args = parser.parse_args()
main(args)
+4 -2
View File
@@ -11,7 +11,8 @@ from fastvideo.v1.distributed import init_distributed_environment, initialize_mo
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.configs.models.vaes import WanVAEConfig
from fastvideo import PipelineConfig
from fastvideo.v1.pipelines.preprocess_pipeline import PreprocessPipeline
from fastvideo.v1.pipelines.preprocess.preprocess_pipeline_i2v import PreprocessPipeline_I2V
from fastvideo.v1.pipelines.preprocess.preprocess_pipeline_t2v import PreprocessPipeline_T2V
logger = init_logger(__name__)
@@ -42,7 +43,7 @@ def main(args):
)
fastvideo_args.check_fastvideo_args()
fastvideo_args.device = torch.device(f"cuda:{local_rank}")
PreprocessPipeline = PreprocessPipeline_I2V if args.preprocess_task == "i2v" else PreprocessPipeline_T2V
pipeline = PreprocessPipeline(args.model_path, fastvideo_args)
pipeline.forward(batch=None, fastvideo_args=fastvideo_args, args=args)
@@ -91,6 +92,7 @@ if __name__ == "__main__":
parser.add_argument("--group_frame", action="store_true") # TODO
parser.add_argument("--group_resolution", action="store_true") # TODO
parser.add_argument("--dataset", default="t2v")
parser.add_argument("--preprocess_task", type=str, default="t2v")
parser.add_argument("--train_fps", type=int, default=30)
parser.add_argument("--use_image_num", type=int, default=0)
parser.add_argument("--text_max_length", type=int, default=256)
@@ -68,7 +68,8 @@ def main(args):
train_dataset = T5dataset(latents_json_path, args.vae_debug)
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()
if args.model_type != "wan":
vae.enable_tiling()
sampler = DistributedSampler(train_dataset, rank=local_rank, num_replicas=world_size, shuffle=True)
train_dataloader = DataLoader(
train_dataset,
@@ -33,7 +33,8 @@ def main(args):
if not dist.is_initialized():
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.enable_tiling()
if args.model_type != "wan":
vae.enable_tiling()
os.makedirs(args.output_dir, exist_ok=True)
os.makedirs(os.path.join(args.output_dir, "latent"), exist_ok=True)
+75 -41
View File
@@ -12,6 +12,7 @@ import torch.distributed as dist
import wandb
from accelerate.utils import set_seed
from diffusers import FlowMatchEulerDiscreteScheduler
from fastvideo.distill.solver import PCMFMScheduler
from diffusers.optimization import get_scheduler
from diffusers.utils import check_min_version
from peft import LoraConfig
@@ -23,7 +24,7 @@ from tqdm.auto import tqdm
from fastvideo.dataset.latent_datasets import (LatentDataset, latent_collate_function)
from fastvideo.distill.solver import EulerSolver, extract_into_tensor
from fastvideo.models.mochi_hf.mochi_latents_utils import normalize_dit_input
from fastvideo.utils.latents_utils import normalize_dit_input
from fastvideo.models.mochi_hf.pipeline_mochi import linear_quadratic_schedule
from fastvideo.utils.checkpoint import (resume_lora_optimizer, save_checkpoint, save_lora_checkpoint)
from fastvideo.utils.communications import (broadcast, sp_parallel_dataloader_wrapper)
@@ -123,13 +124,21 @@ def distill_one_step(
noisy_model_input = sigmas * noise + (1.0 - sigmas) * model_input
# Predict the noise residual
with torch.autocast("cuda", dtype=torch.bfloat16):
teacher_kwargs = {
"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 args.model_type == "wan":
teacher_kwargs = {
"hidden_states": noisy_model_input,
"encoder_hidden_states": encoder_hidden_states,
"timestep": timesteps,
"return_dict": True,
}
else:
teacher_kwargs = {
"hidden_states": noisy_model_input,
"encoder_hidden_states": encoder_hidden_states,
"timestep": timesteps,
"encoder_attention_mask": encoder_attention_mask, # B, L
"return_dict": False,
}
if hunyuan_teacher_disable_cfg:
teacher_kwargs["guidance"] = torch.tensor([1000.0],
device=noisy_model_input.device,
@@ -141,47 +150,70 @@ def distill_one_step(
with torch.no_grad():
w = distill_cfg
with torch.autocast("cuda", dtype=torch.bfloat16):
cond_teacher_output = teacher_transformer(
noisy_model_input,
encoder_hidden_states,
timesteps,
encoder_attention_mask, # B, L
return_dict=False,
)[0].float()
if args.model_type == "wan":
cond_teacher_kwargs ={
"hidden_states": noisy_model_input,
"encoder_hidden_states": encoder_hidden_states,
"timestep": timesteps,
"return_dict": True,
}
else:
cond_teacher_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,
}
cond_teacher_output = teacher_transformer(**cond_teacher_kwargs)[0].float()
if not_apply_cfg_solver:
uncond_teacher_output = cond_teacher_output
else:
# Get teacher model prediction on noisy_latents and unconditional embedding
with torch.autocast("cuda", dtype=torch.bfloat16):
uncond_teacher_output = teacher_transformer(
noisy_model_input,
uncond_prompt_embed.unsqueeze(0).expand(bsz, -1, -1),
timesteps,
uncond_prompt_mask.unsqueeze(0).expand(bsz, -1),
return_dict=False,
)[0].float()
if args.model_type == "wan":
uncond_teacher_kwargs = {
"hidden_states": noisy_model_input,
"encoder_hidden_states":uncond_prompt_embed.unsqueeze(0).expand(bsz, -1, -1),
"timestep": timesteps,
"return_dict": True,
}
else:
uncond_teacher_kwargs = {
"hidden_states": noisy_model_input,
"encoder_hidden_states":uncond_prompt_embed.unsqueeze(0).expand(bsz, -1, -1),
"timestep": timesteps,
"encoder_attention_mask": uncond_prompt_mask.unsqueeze(0).expand(bsz, -1),
"return_dict": False,
}
uncond_teacher_output = teacher_transformer(**uncond_teacher_kwargs)[0].float()
teacher_output = uncond_teacher_output + w * (cond_teacher_output - uncond_teacher_output)
x_prev = solver.euler_step(noisy_model_input, teacher_output, index)
# 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):
if ema_transformer is not None:
target_pred = ema_transformer(
x_prev.float(),
encoder_hidden_states,
timesteps_prev,
encoder_attention_mask, # B, L
return_dict=False,
)[0]
if args.model_type == "wan":
target_pred_kwargs = {
"hidden_states": x_prev.float(),
"encoder_hidden_states": encoder_hidden_states,
"timestep":timesteps_prev,
"return_dict":True,
}
else:
target_pred = transformer(
x_prev.float(),
encoder_hidden_states,
timesteps_prev,
encoder_attention_mask, # B, L
return_dict=False,
)[0]
target_pred_kwargs = {
"hidden_states": x_prev.float(),
"encoder_hidden_states": encoder_hidden_states,
"timestep":timesteps_prev,
"encoder_attention_mask":encoder_attention_mask,
"return_dict":False,
}
if ema_transformer is not None:
target_pred = ema_transformer(**target_pred_kwargs)[0]
else:
target_pred = transformer(**target_pred_kwargs)[0]
target, end_index = solver.euler_style_multiphase_pred(x_prev, target_pred, index, multiphase, True)
@@ -242,7 +274,7 @@ def main(args):
noise_random_generator = None
# Handle the repository creation
if rank <= 0 and args.output_dir is not None:
if rank == 0 and args.output_dir is not None:
os.makedirs(args.output_dir, exist_ok=True)
# For mixed precision training we cast all non-trainable weights to half-precision
@@ -319,7 +351,9 @@ def main(args):
teacher_transformer.requires_grad_(False)
if args.use_ema:
ema_transformer.requires_grad_(False)
noise_scheduler = FlowMatchEulerDiscreteScheduler(shift=args.shift)
noise_scheduler = FlowMatchEulerDiscreteScheduler()
if args.scheduler_type == "pcm_linear_quadratic":
linear_steps = int(noise_scheduler.config.num_train_timesteps * args.linear_range)
sigmas = linear_quadratic_schedule(
@@ -391,7 +425,7 @@ def main(args):
len(train_dataloader) / args.gradient_accumulation_steps * args.sp_size / args.train_sp_batch_size)
args.num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch)
if rank <= 0:
if rank == 0:
project = args.tracker_project_name or "fastvideo"
wandb.init(project=project, config=args)
@@ -493,7 +527,7 @@ def main(args):
"phases": num_phases,
})
progress_bar.update(1)
if rank <= 0:
if rank == 0:
wandb.log(
{
"train_loss": loss,
+4 -4
View File
@@ -23,7 +23,7 @@ from tqdm.auto import tqdm
from fastvideo.dataset.latent_datasets import (LatentDataset, latent_collate_function)
from fastvideo.distill.discriminator import Discriminator
from fastvideo.distill.solver import EulerSolver, extract_into_tensor
from fastvideo.models.mochi_hf.mochi_latents_utils import normalize_dit_input
from fastvideo.utils.latents_utils import normalize_dit_input
from fastvideo.models.mochi_hf.pipeline_mochi import linear_quadratic_schedule
from fastvideo.utils.checkpoint import (resume_lora_optimizer, resume_training_generator_discriminator, save_checkpoint,
save_lora_checkpoint)
@@ -296,7 +296,7 @@ def main(args):
noise_random_generator = None
# Handle the repository creation
if rank <= 0 and args.output_dir is not None:
if rank == 0 and args.output_dir is not None:
os.makedirs(args.output_dir, exist_ok=True)
# For mixed precision training we cast all non-trainable weights to half-precision
@@ -462,7 +462,7 @@ def main(args):
len(train_dataloader) / args.gradient_accumulation_steps * args.sp_size / args.train_sp_batch_size)
args.num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch)
if rank <= 0:
if rank == 0:
project = args.tracker_project_name or "fastvideo"
wandb.init(project=project, config=args)
@@ -559,7 +559,7 @@ def main(args):
"step_time": f"{step_time:.2f}s",
})
progress_bar.update(1)
if rank <= 0:
if rank == 0:
wandb.log(
{
"generator_loss": generator_loss,
+486
View File
@@ -0,0 +1,486 @@
# Copyright 2025 The Wan Team and The HuggingFace Team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import math
from typing import Any, Dict, Optional, Tuple, Union
import torch
import torch.nn as nn
import torch.nn.functional as F
from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.loaders import FromOriginalModelMixin, PeftAdapterMixin
from diffusers.utils import USE_PEFT_BACKEND, logging, scale_lora_layers, unscale_lora_layers
from diffusers.models.attention import FeedForward
from diffusers.models.attention_processor import Attention
from diffusers.models.cache_utils import CacheMixin
from diffusers.models.embeddings import PixArtAlphaTextProjection, TimestepEmbedding, Timesteps, get_1d_rotary_pos_embed
from diffusers.models.modeling_outputs import Transformer2DModelOutput
from diffusers.models.modeling_utils import ModelMixin
from diffusers.models.normalization import FP32LayerNorm
from fastvideo.models.flash_attn_no_pad import flash_attn_no_pad
from fastvideo.utils.communications import all_gather, all_to_all_4D
from fastvideo.utils.parallel_states import get_sequence_parallel_state, nccl_info
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
class WanAttnProcessor2_0:
def __init__(self):
if not hasattr(F, "scaled_dot_product_attention"):
raise ImportError("WanAttnProcessor2_0 requires PyTorch 2.0. To use it, please upgrade PyTorch to 2.0.")
def __call__(
self,
attn: Attention,
hidden_states: torch.Tensor,
encoder_hidden_states: Optional[torch.Tensor] = None,
attention_mask: Optional[torch.Tensor] = None,
rotary_emb: Optional[torch.Tensor] = None,
) -> torch.Tensor:
encoder_hidden_states_img = None
if attn.add_k_proj is not None:
# 512 is the context length of the text encoder, hardcoded for now
image_context_length = encoder_hidden_states.shape[1] - 512
encoder_hidden_states_img = encoder_hidden_states[:, :image_context_length]
encoder_hidden_states = encoder_hidden_states[:, image_context_length:]
if encoder_hidden_states is None:
encoder_hidden_states = hidden_states
query = attn.to_q(hidden_states)
key = attn.to_k(encoder_hidden_states)
value = attn.to_v(encoder_hidden_states)
if attn.norm_q is not None:
query = attn.norm_q(query)
if attn.norm_k is not None:
key = attn.norm_k(key)
query = query.unflatten(2, (attn.heads, -1)).transpose(1, 2)
key = key.unflatten(2, (attn.heads, -1)).transpose(1, 2)
value = value.unflatten(2, (attn.heads, -1)).transpose(1, 2)
if rotary_emb is not None:
def apply_rotary_emb(hidden_states: torch.Tensor, freqs: torch.Tensor):
x_rotated = torch.view_as_complex(hidden_states.to(torch.float64).unflatten(3, (-1, 2)))
x_out = torch.view_as_real(x_rotated * freqs).flatten(3, 4)
return x_out.type_as(hidden_states)
query = apply_rotary_emb(query, rotary_emb)
key = apply_rotary_emb(key, rotary_emb)
# I2V task
hidden_states_img = None
if encoder_hidden_states_img is not None:
key_img = attn.add_k_proj(encoder_hidden_states_img)
key_img = attn.norm_added_k(key_img)
value_img = attn.add_v_proj(encoder_hidden_states_img)
key_img = key_img.unflatten(2, (attn.heads, -1)).transpose(1, 2)
value_img = value_img.unflatten(2, (attn.heads, -1)).transpose(1, 2)
hidden_states_img = F.scaled_dot_product_attention(
query, key_img, value_img, attn_mask=None, dropout_p=0.0, is_causal=False
)
hidden_states_img = hidden_states_img.transpose(1, 2).flatten(2, 3)
hidden_states_img = hidden_states_img.type_as(query)
hidden_states = F.scaled_dot_product_attention(
query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False
)
hidden_states = hidden_states.transpose(1, 2).flatten(2, 3)
hidden_states = hidden_states.type_as(query)
if hidden_states_img is not None:
hidden_states = hidden_states + hidden_states_img
hidden_states = attn.to_out[0](hidden_states)
hidden_states = attn.to_out[1](hidden_states)
return hidden_states
class WanImageEmbedding(torch.nn.Module):
def __init__(self, in_features: int, out_features: int, pos_embed_seq_len=None):
super().__init__()
self.norm1 = FP32LayerNorm(in_features)
self.ff = FeedForward(in_features, out_features, mult=1, activation_fn="gelu")
self.norm2 = FP32LayerNorm(out_features)
if pos_embed_seq_len is not None:
self.pos_embed = nn.Parameter(torch.zeros(1, pos_embed_seq_len, in_features))
else:
self.pos_embed = None
def forward(self, encoder_hidden_states_image: torch.Tensor) -> torch.Tensor:
if self.pos_embed is not None:
batch_size, seq_len, embed_dim = encoder_hidden_states_image.shape
encoder_hidden_states_image = encoder_hidden_states_image.view(-1, 2 * seq_len, embed_dim)
encoder_hidden_states_image = encoder_hidden_states_image + self.pos_embed
hidden_states = self.norm1(encoder_hidden_states_image)
hidden_states = self.ff(hidden_states)
hidden_states = self.norm2(hidden_states)
return hidden_states
class WanTimeTextImageEmbedding(nn.Module):
def __init__(
self,
dim: int,
time_freq_dim: int,
time_proj_dim: int,
text_embed_dim: int,
image_embed_dim: Optional[int] = None,
pos_embed_seq_len: Optional[int] = None,
):
super().__init__()
self.timesteps_proj = Timesteps(num_channels=time_freq_dim, flip_sin_to_cos=True, downscale_freq_shift=0)
self.time_embedder = TimestepEmbedding(in_channels=time_freq_dim, time_embed_dim=dim)
self.act_fn = nn.SiLU()
self.time_proj = nn.Linear(dim, time_proj_dim)
self.text_embedder = PixArtAlphaTextProjection(text_embed_dim, dim, act_fn="gelu_tanh")
self.image_embedder = None
if image_embed_dim is not None:
self.image_embedder = WanImageEmbedding(image_embed_dim, dim, pos_embed_seq_len=pos_embed_seq_len)
def forward(
self,
timestep: torch.Tensor,
encoder_hidden_states: torch.Tensor,
encoder_hidden_states_image: Optional[torch.Tensor] = None,
):
timestep = self.timesteps_proj(timestep)
time_embedder_dtype = next(iter(self.time_embedder.parameters())).dtype
if timestep.dtype != time_embedder_dtype and time_embedder_dtype != torch.int8:
timestep = timestep.to(time_embedder_dtype)
temb = self.time_embedder(timestep).type_as(encoder_hidden_states)
timestep_proj = self.time_proj(self.act_fn(temb))
encoder_hidden_states = self.text_embedder(encoder_hidden_states)
if encoder_hidden_states_image is not None:
encoder_hidden_states_image = self.image_embedder(encoder_hidden_states_image)
return temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image
class WanRotaryPosEmbed(nn.Module):
def __init__(
self, attention_head_dim: int, patch_size: Tuple[int, int, int], max_seq_len: int, theta: float = 10000.0
):
super().__init__()
self.attention_head_dim = attention_head_dim
self.patch_size = patch_size
self.max_seq_len = max_seq_len
h_dim = w_dim = 2 * (attention_head_dim // 6)
t_dim = attention_head_dim - h_dim - w_dim
freqs = []
for dim in [t_dim, h_dim, w_dim]:
freq = get_1d_rotary_pos_embed(
dim, max_seq_len, theta, use_real=False, repeat_interleave_real=False, freqs_dtype=torch.float64
)
freqs.append(freq)
self.freqs = torch.cat(freqs, dim=1)
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
batch_size, num_channels, num_frames, height, width = hidden_states.shape
p_t, p_h, p_w = self.patch_size
ppf, pph, ppw = num_frames // p_t, height // p_h, width // p_w
freqs = self.freqs.to(hidden_states.device)
freqs = freqs.split_with_sizes(
[
self.attention_head_dim // 2 - 2 * (self.attention_head_dim // 6),
self.attention_head_dim // 6,
self.attention_head_dim // 6,
],
dim=1,
)
freqs_f = freqs[0][:ppf].view(ppf, 1, 1, -1).expand(ppf, pph, ppw, -1)
freqs_h = freqs[1][:pph].view(1, pph, 1, -1).expand(ppf, pph, ppw, -1)
freqs_w = freqs[2][:ppw].view(1, 1, ppw, -1).expand(ppf, pph, ppw, -1)
freqs = torch.cat([freqs_f, freqs_h, freqs_w], dim=-1).reshape(1, 1, ppf * pph * ppw, -1)
return freqs
class WanTransformerBlock(nn.Module):
def __init__(
self,
dim: int,
ffn_dim: int,
num_heads: int,
qk_norm: str = "rms_norm_across_heads",
cross_attn_norm: bool = False,
eps: float = 1e-6,
added_kv_proj_dim: Optional[int] = None,
):
super().__init__()
# 1. Self-attention
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
self.attn1 = Attention(
query_dim=dim,
heads=num_heads,
kv_heads=num_heads,
dim_head=dim // num_heads,
qk_norm=qk_norm,
eps=eps,
bias=True,
cross_attention_dim=None,
out_bias=True,
processor=WanAttnProcessor2_0(),
)
# 2. Cross-attention
self.attn2 = Attention(
query_dim=dim,
heads=num_heads,
kv_heads=num_heads,
dim_head=dim // num_heads,
qk_norm=qk_norm,
eps=eps,
bias=True,
cross_attention_dim=None,
out_bias=True,
added_kv_proj_dim=added_kv_proj_dim,
added_proj_bias=True,
processor=WanAttnProcessor2_0(),
)
self.norm2 = FP32LayerNorm(dim, eps, elementwise_affine=True) if cross_attn_norm else nn.Identity()
# 3. Feed-forward
self.ffn = FeedForward(dim, inner_dim=ffn_dim, activation_fn="gelu-approximate")
self.norm3 = FP32LayerNorm(dim, eps, elementwise_affine=False)
self.scale_shift_table = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor,
temb: torch.Tensor,
rotary_emb: torch.Tensor,
) -> torch.Tensor:
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = (
self.scale_shift_table + temb.float()
).chunk(6, dim=1)
# 1. Self-attention
norm_hidden_states = (self.norm1(hidden_states.float()) * (1 + scale_msa) + shift_msa).type_as(hidden_states)
attn_output = self.attn1(hidden_states=norm_hidden_states, rotary_emb=rotary_emb)
hidden_states = (hidden_states.float() + attn_output * gate_msa).type_as(hidden_states)
# 2. Cross-attention
norm_hidden_states = self.norm2(hidden_states.float()).type_as(hidden_states)
attn_output = self.attn2(hidden_states=norm_hidden_states, encoder_hidden_states=encoder_hidden_states)
hidden_states = hidden_states + attn_output
# 3. Feed-forward
norm_hidden_states = (self.norm3(hidden_states.float()) * (1 + c_scale_msa) + c_shift_msa).type_as(
hidden_states
)
ff_output = self.ffn(norm_hidden_states)
hidden_states = (hidden_states.float() + ff_output.float() * c_gate_msa).type_as(hidden_states)
return hidden_states
class WanTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin, CacheMixin):
r"""
A Transformer model for video-like data used in the Wan model.
Args:
patch_size (`Tuple[int]`, defaults to `(1, 2, 2)`):
3D patch dimensions for video embedding (t_patch, h_patch, w_patch).
num_attention_heads (`int`, defaults to `40`):
Fixed length for text embeddings.
attention_head_dim (`int`, defaults to `128`):
The number of channels in each head.
in_channels (`int`, defaults to `16`):
The number of channels in the input.
out_channels (`int`, defaults to `16`):
The number of channels in the output.
text_dim (`int`, defaults to `512`):
Input dimension for text embeddings.
freq_dim (`int`, defaults to `256`):
Dimension for sinusoidal time embeddings.
ffn_dim (`int`, defaults to `13824`):
Intermediate dimension in feed-forward network.
num_layers (`int`, defaults to `40`):
The number of layers of transformer blocks to use.
window_size (`Tuple[int]`, defaults to `(-1, -1)`):
Window size for local attention (-1 indicates global attention).
cross_attn_norm (`bool`, defaults to `True`):
Enable cross-attention normalization.
qk_norm (`bool`, defaults to `True`):
Enable query/key normalization.
eps (`float`, defaults to `1e-6`):
Epsilon value for normalization layers.
add_img_emb (`bool`, defaults to `False`):
Whether to use img_emb.
added_kv_proj_dim (`int`, *optional*, defaults to `None`):
The number of channels to use for the added key and value projections. If `None`, no projection is used.
"""
_supports_gradient_checkpointing = True
_skip_layerwise_casting_patterns = ["patch_embedding", "condition_embedder", "norm"]
_no_split_modules = ["WanTransformerBlock"]
_keep_in_fp32_modules = ["time_embedder", "scale_shift_table", "norm1", "norm2", "norm3"]
_keys_to_ignore_on_load_unexpected = ["norm_added_q"]
@register_to_config
def __init__(
self,
patch_size: Tuple[int] = (1, 2, 2),
num_attention_heads: int = 40,
attention_head_dim: int = 128,
in_channels: int = 16,
out_channels: int = 16,
text_dim: int = 4096,
freq_dim: int = 256,
ffn_dim: int = 13824,
num_layers: int = 40,
cross_attn_norm: bool = True,
qk_norm: Optional[str] = "rms_norm_across_heads",
eps: float = 1e-6,
image_dim: Optional[int] = None,
added_kv_proj_dim: Optional[int] = None,
rope_max_seq_len: int = 1024,
pos_embed_seq_len: Optional[int] = None,
) -> None:
super().__init__()
inner_dim = num_attention_heads * attention_head_dim
out_channels = out_channels or in_channels
# 1. Patch & position embedding
self.rope = WanRotaryPosEmbed(attention_head_dim, patch_size, rope_max_seq_len)
self.patch_embedding = nn.Conv3d(in_channels, inner_dim, kernel_size=patch_size, stride=patch_size)
# 2. Condition embeddings
# image_embedding_dim=1280 for I2V model
self.condition_embedder = WanTimeTextImageEmbedding(
dim=inner_dim,
time_freq_dim=freq_dim,
time_proj_dim=inner_dim * 6,
text_embed_dim=text_dim,
image_embed_dim=image_dim,
pos_embed_seq_len=pos_embed_seq_len,
)
# 3. Transformer blocks
self.blocks = nn.ModuleList(
[
WanTransformerBlock(
inner_dim, ffn_dim, num_attention_heads, qk_norm, cross_attn_norm, eps, added_kv_proj_dim
)
for _ in range(num_layers)
]
)
# 4. Output norm & projection
self.norm_out = FP32LayerNorm(inner_dim, eps, elementwise_affine=False)
self.proj_out = nn.Linear(inner_dim, out_channels * math.prod(patch_size))
self.scale_shift_table = nn.Parameter(torch.randn(1, 2, inner_dim) / inner_dim**0.5)
self.gradient_checkpointing = False
def forward(
self,
hidden_states: torch.Tensor,
timestep: torch.LongTensor,
encoder_hidden_states: torch.Tensor,
encoder_hidden_states_image: Optional[torch.Tensor] = None,
return_dict: bool = True,
attention_kwargs: Optional[Dict[str, Any]] = None,
) -> Union[torch.Tensor, Dict[str, torch.Tensor]]:
if attention_kwargs is not None:
attention_kwargs = attention_kwargs.copy()
lora_scale = attention_kwargs.pop("scale", 1.0)
else:
lora_scale = 1.0
if USE_PEFT_BACKEND:
# weight the lora layers by setting `lora_scale` for each PEFT layer
scale_lora_layers(self, lora_scale)
else:
if attention_kwargs is not None and attention_kwargs.get("scale", None) is not None:
logger.warning(
"Passing `scale` via `attention_kwargs` when not using the PEFT backend is ineffective."
)
batch_size, num_channels, num_frames, height, width = hidden_states.shape
p_t, p_h, p_w = self.config.patch_size
post_patch_num_frames = num_frames // p_t
post_patch_height = height // p_h
post_patch_width = width // p_w
rotary_emb = self.rope(hidden_states)
hidden_states = self.patch_embedding(hidden_states)
hidden_states = hidden_states.flatten(2).transpose(1, 2)
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
timestep, encoder_hidden_states, encoder_hidden_states_image
)
timestep_proj = timestep_proj.unflatten(1, (6, -1))
if encoder_hidden_states_image is not None:
encoder_hidden_states = torch.concat([encoder_hidden_states_image, encoder_hidden_states], dim=1)
# 4. Transformer blocks
if torch.is_grad_enabled() and self.gradient_checkpointing:
for block in self.blocks:
hidden_states = self._gradient_checkpointing_func(
block, hidden_states, encoder_hidden_states, timestep_proj, rotary_emb
)
else:
for block in self.blocks:
hidden_states = block(hidden_states, encoder_hidden_states, timestep_proj, rotary_emb)
# 5. Output norm, projection & unpatchify
shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2, dim=1)
# Move the shift and scale tensors to the same device as hidden_states.
# When using multi-GPU inference via accelerate these will be on the
# first device rather than the last device, which hidden_states ends up
# on.
shift = shift.to(hidden_states.device)
scale = scale.to(hidden_states.device)
hidden_states = (self.norm_out(hidden_states.float()) * (1 + scale) + shift).type_as(hidden_states)
hidden_states = self.proj_out(hidden_states)
hidden_states = hidden_states.reshape(
batch_size, post_patch_num_frames, post_patch_height, post_patch_width, p_t, p_h, p_w, -1
)
hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6)
output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
if USE_PEFT_BACKEND:
# remove `lora_scale` from each PEFT layer
unscale_lora_layers(self, lora_scale)
if not return_dict:
return (output,)
return Transformer2DModelOutput(sample=output)
+609
View File
@@ -0,0 +1,609 @@
# Copyright 2025 The Wan Team and The HuggingFace Team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import html
from typing import Any, Callable, Dict, List, Optional, Union
import regex as re
import torch
from transformers import AutoTokenizer, UMT5EncoderModel
from diffusers.callbacks import MultiPipelineCallbacks, PipelineCallback
from diffusers.loaders import WanLoraLoaderMixin
from diffusers.models import AutoencoderKLWan, WanTransformer3DModel
from diffusers.schedulers import FlowMatchEulerDiscreteScheduler
from diffusers.utils import is_ftfy_available, is_torch_xla_available, logging, replace_example_docstring
from diffusers.utils.torch_utils import randn_tensor
from diffusers.video_processor import VideoProcessor
from diffusers.pipelines.pipeline_utils import DiffusionPipeline
from diffusers.pipelines.wan.pipeline_output import WanPipelineOutput
from einops import rearrange
from transformers import UMT5EncoderModel, T5TokenizerFast
from fastvideo.models.mochi_hf.modeling_wan import WanTransformer3DModel
from fastvideo.utils.communications import all_gather
from fastvideo.utils.parallel_states import get_sequence_parallel_state, nccl_info
if is_torch_xla_available():
import torch_xla.core.xla_model as xm
XLA_AVAILABLE = True
else:
XLA_AVAILABLE = False
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
if is_ftfy_available():
import ftfy
EXAMPLE_DOC_STRING = """
Examples:
```python
>>> import torch
>>> from diffusers.utils import export_to_video
>>> from diffusers import AutoencoderKLWan, WanPipeline
>>> from diffusers.schedulers.scheduling_unipc_multistep import UniPCMultistepScheduler
>>> # Available models: Wan-AI/Wan2.1-T2V-14B-Diffusers, Wan-AI/Wan2.1-T2V-1.3B-Diffusers
>>> model_id = "Wan-AI/Wan2.1-T2V-14B-Diffusers"
>>> vae = AutoencoderKLWan.from_pretrained(model_id, subfolder="vae", torch_dtype=torch.float32)
>>> pipe = WanPipeline.from_pretrained(model_id, vae=vae, torch_dtype=torch.bfloat16)
>>> flow_shift = 5.0 # 5.0 for 720P, 3.0 for 480P
>>> pipe.scheduler = UniPCMultistepScheduler.from_config(pipe.scheduler.config, flow_shift=flow_shift)
>>> pipe.to("cuda")
>>> prompt = "A cat and a dog baking a cake together in a kitchen. The cat is carefully measuring flour, while the dog is stirring the batter with a wooden spoon. The kitchen is cozy, with sunlight streaming through the window."
>>> negative_prompt = "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"
>>> output = pipe(
... prompt=prompt,
... negative_prompt=negative_prompt,
... height=720,
... width=1280,
... num_frames=81,
... guidance_scale=5.0,
... ).frames[0]
>>> export_to_video(output, "output.mp4", fps=16)
```
"""
def basic_clean(text):
text = ftfy.fix_text(text)
text = html.unescape(html.unescape(text))
return text.strip()
def whitespace_clean(text):
text = re.sub(r"\s+", " ", text)
text = text.strip()
return text
def prompt_clean(text):
text = whitespace_clean(basic_clean(text))
return text
class WanPipeline(DiffusionPipeline, WanLoraLoaderMixin):
r"""
Pipeline for text-to-video generation using Wan.
This model inherits from [`DiffusionPipeline`]. Check the superclass documentation for the generic methods
implemented for all pipelines (downloading, saving, running on a particular device, etc.).
Args:
tokenizer ([`T5Tokenizer`]):
Tokenizer from [T5](https://huggingface.co/docs/transformers/en/model_doc/t5#transformers.T5Tokenizer),
specifically the [google/umt5-xxl](https://huggingface.co/google/umt5-xxl) variant.
text_encoder ([`T5EncoderModel`]):
[T5](https://huggingface.co/docs/transformers/en/model_doc/t5#transformers.T5EncoderModel), specifically
the [google/umt5-xxl](https://huggingface.co/google/umt5-xxl) variant.
transformer ([`WanTransformer3DModel`]):
Conditional Transformer to denoise the input latents.
scheduler ([`UniPCMultistepScheduler`]):
A scheduler to be used in combination with `transformer` to denoise the encoded image latents.
vae ([`AutoencoderKLWan`]):
Variational Auto-Encoder (VAE) Model to encode and decode videos to and from latent representations.
"""
model_cpu_offload_seq = "text_encoder->transformer->vae"
_callback_tensor_inputs = ["latents", "prompt_embeds", "negative_prompt_embeds"]
def __init__(
self,
tokenizer: AutoTokenizer,
text_encoder: UMT5EncoderModel,
transformer: WanTransformer3DModel,
vae: AutoencoderKLWan,
scheduler: FlowMatchEulerDiscreteScheduler,
):
super().__init__()
self.register_modules(
vae=vae,
text_encoder=text_encoder,
tokenizer=tokenizer,
transformer=transformer,
scheduler=scheduler,
)
self.vae_scale_factor_temporal = 2 ** sum(self.vae.temperal_downsample) if getattr(self, "vae", None) else 4
self.vae_scale_factor_spatial = 2 ** len(self.vae.temperal_downsample) if getattr(self, "vae", None) else 8
self.video_processor = VideoProcessor(vae_scale_factor=self.vae_scale_factor_spatial)
def _get_t5_prompt_embeds(
self,
prompt: Union[str, List[str]] = None,
num_videos_per_prompt: int = 1,
max_sequence_length: int = 226,
device: Optional[torch.device] = None,
dtype: Optional[torch.dtype] = None,
):
device = device or self._execution_device
dtype = dtype or self.text_encoder.dtype
prompt = [prompt] if isinstance(prompt, str) else prompt
prompt = [prompt_clean(u) for u in prompt]
batch_size = len(prompt)
text_inputs = self.tokenizer(
prompt,
padding="max_length",
max_length=max_sequence_length,
truncation=True,
add_special_tokens=True,
return_attention_mask=True,
return_tensors="pt",
)
text_input_ids, mask = text_inputs.input_ids, text_inputs.attention_mask
seq_lens = mask.gt(0).sum(dim=1).long()
prompt_embeds = self.text_encoder(text_input_ids.to(device), mask.to(device)).last_hidden_state
prompt_embeds = prompt_embeds.to(dtype=dtype, device=device)
prompt_embeds = [u[:v] for u, v in zip(prompt_embeds, seq_lens)]
prompt_embeds = torch.stack(
[torch.cat([u, u.new_zeros(max_sequence_length - u.size(0), u.size(1))]) for u in prompt_embeds], dim=0
)
# duplicate text embeddings for each generation per prompt, using mps friendly method
_, seq_len, _ = prompt_embeds.shape
prompt_embeds = prompt_embeds.repeat(1, num_videos_per_prompt, 1)
prompt_embeds = prompt_embeds.view(batch_size * num_videos_per_prompt, seq_len, -1)
return prompt_embeds
def encode_prompt(
self,
prompt: Union[str, List[str]],
negative_prompt: Optional[Union[str, List[str]]] = None,
do_classifier_free_guidance: bool = True,
num_videos_per_prompt: int = 1,
prompt_embeds: Optional[torch.Tensor] = None,
negative_prompt_embeds: Optional[torch.Tensor] = None,
max_sequence_length: int = 226,
device: Optional[torch.device] = None,
dtype: Optional[torch.dtype] = None,
):
r"""
Encodes the prompt into text encoder hidden states.
Args:
prompt (`str` or `List[str]`, *optional*):
prompt to be encoded
negative_prompt (`str` or `List[str]`, *optional*):
The prompt or prompts not to guide the image generation. If not defined, one has to pass
`negative_prompt_embeds` instead. Ignored when not using guidance (i.e., ignored if `guidance_scale` is
less than `1`).
do_classifier_free_guidance (`bool`, *optional*, defaults to `True`):
Whether to use classifier free guidance or not.
num_videos_per_prompt (`int`, *optional*, defaults to 1):
Number of videos that should be generated per prompt. torch device to place the resulting embeddings on
prompt_embeds (`torch.Tensor`, *optional*):
Pre-generated text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. If not
provided, text embeddings will be generated from `prompt` input argument.
negative_prompt_embeds (`torch.Tensor`, *optional*):
Pre-generated negative text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt
weighting. If not provided, negative_prompt_embeds will be generated from `negative_prompt` input
argument.
device: (`torch.device`, *optional*):
torch device
dtype: (`torch.dtype`, *optional*):
torch dtype
"""
device = device or self._execution_device
prompt = [prompt] if isinstance(prompt, str) else prompt
if prompt is not None:
batch_size = len(prompt)
else:
batch_size = prompt_embeds.shape[0]
if prompt_embeds is None:
prompt_embeds = self._get_t5_prompt_embeds(
prompt=prompt,
num_videos_per_prompt=num_videos_per_prompt,
max_sequence_length=max_sequence_length,
device=device,
dtype=dtype,
)
if do_classifier_free_guidance and negative_prompt_embeds is None:
negative_prompt = negative_prompt or ""
negative_prompt = batch_size * [negative_prompt] if isinstance(negative_prompt, str) else negative_prompt
if prompt is not None and type(prompt) is not type(negative_prompt):
raise TypeError(
f"`negative_prompt` should be the same type to `prompt`, but got {type(negative_prompt)} !="
f" {type(prompt)}."
)
elif batch_size != len(negative_prompt):
raise ValueError(
f"`negative_prompt`: {negative_prompt} has batch size {len(negative_prompt)}, but `prompt`:"
f" {prompt} has batch size {batch_size}. Please make sure that passed `negative_prompt` matches"
" the batch size of `prompt`."
)
negative_prompt_embeds = self._get_t5_prompt_embeds(
prompt=negative_prompt,
num_videos_per_prompt=num_videos_per_prompt,
max_sequence_length=max_sequence_length,
device=device,
dtype=dtype,
)
return prompt_embeds, negative_prompt_embeds
def check_inputs(
self,
prompt,
negative_prompt,
height,
width,
prompt_embeds=None,
negative_prompt_embeds=None,
callback_on_step_end_tensor_inputs=None,
):
if height % 16 != 0 or width % 16 != 0:
raise ValueError(f"`height` and `width` have to be divisible by 16 but are {height} and {width}.")
if callback_on_step_end_tensor_inputs is not None and not all(
k in self._callback_tensor_inputs for k in callback_on_step_end_tensor_inputs
):
raise ValueError(
f"`callback_on_step_end_tensor_inputs` has to be in {self._callback_tensor_inputs}, but found {[k for k in callback_on_step_end_tensor_inputs if k not in self._callback_tensor_inputs]}"
)
if prompt is not None and prompt_embeds is not None:
raise ValueError(
f"Cannot forward both `prompt`: {prompt} and `prompt_embeds`: {prompt_embeds}. Please make sure to"
" only forward one of the two."
)
elif negative_prompt is not None and negative_prompt_embeds is not None:
raise ValueError(
f"Cannot forward both `negative_prompt`: {negative_prompt} and `negative_prompt_embeds`: {negative_prompt_embeds}. Please make sure to"
" only forward one of the two."
)
elif prompt is None and prompt_embeds is None:
raise ValueError(
"Provide either `prompt` or `prompt_embeds`. Cannot leave both `prompt` and `prompt_embeds` undefined."
)
elif prompt is not None and (not isinstance(prompt, str) and not isinstance(prompt, list)):
raise ValueError(f"`prompt` has to be of type `str` or `list` but is {type(prompt)}")
elif negative_prompt is not None and (
not isinstance(negative_prompt, str) and not isinstance(negative_prompt, list)
):
raise ValueError(f"`negative_prompt` has to be of type `str` or `list` but is {type(negative_prompt)}")
def prepare_latents(
self,
batch_size: int,
num_channels_latents: int = 16,
height: int = 480,
width: int = 832,
num_frames: int = 81,
dtype: Optional[torch.dtype] = None,
device: Optional[torch.device] = None,
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
latents: Optional[torch.Tensor] = None,
) -> torch.Tensor:
if latents is not None:
return latents.to(device=device, dtype=dtype)
num_latent_frames = (num_frames - 1) // self.vae_scale_factor_temporal + 1
shape = (
batch_size,
num_channels_latents,
num_latent_frames,
int(height) // self.vae_scale_factor_spatial,
int(width) // self.vae_scale_factor_spatial,
)
if isinstance(generator, list) and len(generator) != batch_size:
raise ValueError(
f"You have passed a list of generators of length {len(generator)}, but requested an effective batch"
f" size of {batch_size}. Make sure the batch size matches the length of the generators."
)
latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
return latents
@property
def guidance_scale(self):
return self._guidance_scale
@property
def do_classifier_free_guidance(self):
return self._guidance_scale > 1.0
@property
def num_timesteps(self):
return self._num_timesteps
@property
def current_timestep(self):
return self._current_timestep
@property
def interrupt(self):
return self._interrupt
@property
def attention_kwargs(self):
return self._attention_kwargs
@torch.no_grad()
@replace_example_docstring(EXAMPLE_DOC_STRING)
def __call__(
self,
prompt: Union[str, List[str]] = None,
negative_prompt: Union[str, List[str]] = None,
height: int = 480,
width: int = 832,
num_frames: int = 81,
num_inference_steps: int = 50,
guidance_scale: float = 5.0,
num_videos_per_prompt: Optional[int] = 1,
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
latents: Optional[torch.Tensor] = None,
prompt_embeds: Optional[torch.Tensor] = None,
negative_prompt_embeds: Optional[torch.Tensor] = None,
output_type: Optional[str] = "np",
return_dict: bool = True,
attention_kwargs: Optional[Dict[str, Any]] = None,
callback_on_step_end: Optional[
Union[Callable[[int, int, Dict], None], PipelineCallback, MultiPipelineCallbacks]
] = None,
callback_on_step_end_tensor_inputs: List[str] = ["latents"],
max_sequence_length: int = 512,
):
r"""
The call function to the pipeline for generation.
Args:
prompt (`str` or `List[str]`, *optional*):
The prompt or prompts to guide the image generation. If not defined, one has to pass `prompt_embeds`.
instead.
height (`int`, defaults to `480`):
The height in pixels of the generated image.
width (`int`, defaults to `832`):
The width in pixels of the generated image.
num_frames (`int`, defaults to `81`):
The number of frames in the generated video.
num_inference_steps (`int`, defaults to `50`):
The number of denoising steps. More denoising steps usually lead to a higher quality image at the
expense of slower inference.
guidance_scale (`float`, defaults to `5.0`):
Guidance scale as defined in [Classifier-Free Diffusion
Guidance](https://huggingface.co/papers/2207.12598). `guidance_scale` is defined as `w` of equation 2.
of [Imagen Paper](https://huggingface.co/papers/2205.11487). Guidance scale is enabled by setting
`guidance_scale > 1`. Higher guidance scale encourages to generate images that are closely linked to
the text `prompt`, usually at the expense of lower image quality.
num_videos_per_prompt (`int`, *optional*, defaults to 1):
The number of images to generate per prompt.
generator (`torch.Generator` or `List[torch.Generator]`, *optional*):
A [`torch.Generator`](https://pytorch.org/docs/stable/generated/torch.Generator.html) to make
generation deterministic.
latents (`torch.Tensor`, *optional*):
Pre-generated noisy latents sampled from a Gaussian distribution, to be used as inputs for image
generation. Can be used to tweak the same generation with different prompts. If not provided, a latents
tensor is generated by sampling using the supplied random `generator`.
prompt_embeds (`torch.Tensor`, *optional*):
Pre-generated text embeddings. Can be used to easily tweak text inputs (prompt weighting). If not
provided, text embeddings are generated from the `prompt` input argument.
output_type (`str`, *optional*, defaults to `"np"`):
The output format of the generated image. Choose between `PIL.Image` or `np.array`.
return_dict (`bool`, *optional*, defaults to `True`):
Whether or not to return a [`WanPipelineOutput`] instead of a plain tuple.
attention_kwargs (`dict`, *optional*):
A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under
`self.processor` in
[diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py).
callback_on_step_end (`Callable`, `PipelineCallback`, `MultiPipelineCallbacks`, *optional*):
A function or a subclass of `PipelineCallback` or `MultiPipelineCallbacks` that is called at the end of
each denoising step during the inference. with the following arguments: `callback_on_step_end(self:
DiffusionPipeline, step: int, timestep: int, callback_kwargs: Dict)`. `callback_kwargs` will include a
list of all tensors as specified by `callback_on_step_end_tensor_inputs`.
callback_on_step_end_tensor_inputs (`List`, *optional*):
The list of tensor inputs for the `callback_on_step_end` function. The tensors specified in the list
will be passed as `callback_kwargs` argument. You will only be able to include variables listed in the
`._callback_tensor_inputs` attribute of your pipeline class.
autocast_dtype (`torch.dtype`, *optional*, defaults to `torch.bfloat16`):
The dtype to use for the torch.amp.autocast.
Examples:
Returns:
[`~WanPipelineOutput`] or `tuple`:
If `return_dict` is `True`, [`WanPipelineOutput`] is returned, otherwise a `tuple` is returned where
the first element is a list with the generated images and the second element is a list of `bool`s
indicating whether the corresponding generated image contains "not-safe-for-work" (nsfw) content.
"""
if isinstance(callback_on_step_end, (PipelineCallback, MultiPipelineCallbacks)):
callback_on_step_end_tensor_inputs = callback_on_step_end.tensor_inputs
# 1. Check inputs. Raise error if not correct
self.check_inputs(
prompt,
negative_prompt,
height,
width,
prompt_embeds,
negative_prompt_embeds,
callback_on_step_end_tensor_inputs,
)
if num_frames % self.vae_scale_factor_temporal != 1:
logger.warning(
f"`num_frames - 1` has to be divisible by {self.vae_scale_factor_temporal}. Rounding to the nearest number."
)
num_frames = num_frames // self.vae_scale_factor_temporal * self.vae_scale_factor_temporal + 1
num_frames = max(num_frames, 1)
self._guidance_scale = guidance_scale
self._attention_kwargs = attention_kwargs
self._current_timestep = None
self._interrupt = False
device = self._execution_device
# 2. Define call parameters
if prompt is not None and isinstance(prompt, str):
batch_size = 1
elif prompt is not None and isinstance(prompt, list):
batch_size = len(prompt)
else:
batch_size = prompt_embeds.shape[0]
# 3. Encode input prompt
prompt_embeds, negative_prompt_embeds = self.encode_prompt(
prompt=prompt,
negative_prompt=negative_prompt,
do_classifier_free_guidance=self.do_classifier_free_guidance,
num_videos_per_prompt=num_videos_per_prompt,
prompt_embeds=prompt_embeds,
negative_prompt_embeds=negative_prompt_embeds,
max_sequence_length=max_sequence_length,
device=device,
)
transformer_dtype = self.transformer.dtype
prompt_embeds = prompt_embeds.to(transformer_dtype)
if negative_prompt_embeds is not None:
negative_prompt_embeds = negative_prompt_embeds.to(transformer_dtype)
# 4. Prepare timesteps
self.scheduler.set_timesteps(num_inference_steps, device=device)
timesteps = self.scheduler.timesteps
# 5. Prepare latent variables
num_channels_latents = self.transformer.config.in_channels
latents = self.prepare_latents(
batch_size * num_videos_per_prompt,
num_channels_latents,
height,
width,
num_frames,
torch.float32,
device,
generator,
latents,
)
world_size, rank = nccl_info.sp_size, nccl_info.rank_within_group
if get_sequence_parallel_state():
latents = rearrange(latents, "b t (n s) h w -> b t n s h w", n=world_size).contiguous()
latents = latents[:, :, rank, :, :, :]
# 6. Denoising loop
num_warmup_steps = len(timesteps) - num_inference_steps * self.scheduler.order
self._num_timesteps = len(timesteps)
self._progress_bar_config = {"disable": nccl_info.rank_within_group != 0}
with self.progress_bar(total=num_inference_steps) as progress_bar:
for i, t in enumerate(timesteps):
if self.interrupt:
continue
self._current_timestep = t
latent_model_input = latents.to(transformer_dtype)
timestep = t.expand(latents.shape[0])
noise_pred = self.transformer(
hidden_states=latent_model_input,
timestep=timestep,
encoder_hidden_states=prompt_embeds,
attention_kwargs=attention_kwargs,
return_dict=False,
)[0]
if self.do_classifier_free_guidance:
noise_uncond = self.transformer(
hidden_states=latent_model_input,
timestep=timestep,
encoder_hidden_states=negative_prompt_embeds,
attention_kwargs=attention_kwargs,
return_dict=False,
)[0]
noise_pred = noise_uncond + guidance_scale * (noise_pred - noise_uncond)
# compute the previous noisy sample x_t -> x_t-1
latents = self.scheduler.step(noise_pred, t, latents, return_dict=False)[0]
if callback_on_step_end is not None:
callback_kwargs = {}
for k in callback_on_step_end_tensor_inputs:
callback_kwargs[k] = locals()[k]
callback_outputs = callback_on_step_end(self, i, t, callback_kwargs)
latents = callback_outputs.pop("latents", latents)
prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds)
negative_prompt_embeds = callback_outputs.pop("negative_prompt_embeds", negative_prompt_embeds)
# call the callback, if provided
if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
progress_bar.update()
if XLA_AVAILABLE:
xm.mark_step()
if get_sequence_parallel_state():
latents = all_gather(latents, dim=2)
self._current_timestep = None
if not output_type == "latent":
latents = latents.to(self.vae.dtype)
latents_mean = (
torch.tensor(self.vae.config.latents_mean)
.view(1, self.vae.config.z_dim, 1, 1, 1)
.to(latents.device, latents.dtype)
)
latents_std = 1.0 / torch.tensor(self.vae.config.latents_std).view(1, self.vae.config.z_dim, 1, 1, 1).to(
latents.device, latents.dtype
)
latents = latents / latents_std + latents_mean
video = self.vae.decode(latents, return_dict=False)[0]
video = self.video_processor.postprocess_video(video, output_type=output_type)
else:
video = latents
# Offload all models
self.maybe_free_model_hooks()
if not return_dict:
return (video,)
return WanPipelineOutput(frames=video)
+2 -2
View File
@@ -86,7 +86,7 @@ def inference(args):
num_inference_steps=args.num_inference_steps,
generator=generator,
).frames
if nccl_info.global_rank <= 0:
if nccl_info.global_rank == 0:
os.makedirs(args.output_path, exist_ok=True)
suffix = prompt.split(".")[0]
export_to_video(
@@ -107,7 +107,7 @@ def inference(args):
generator=generator,
).frames
if nccl_info.global_rank <= 0:
if nccl_info.global_rank == 0:
export_to_video(videos[0], args.output_path + ".mp4", fps=24)
+2 -2
View File
@@ -94,7 +94,7 @@ def main(args):
guidance_scale=args.guidance_scale,
generator=generator,
).frames
if nccl_info.global_rank <= 0:
if nccl_info.global_rank == 0:
os.makedirs(args.output_path, exist_ok=True)
suffix = prompt.split(".")[0]
export_to_video(
@@ -116,7 +116,7 @@ def main(args):
generator=generator,
).frames
if nccl_info.global_rank <= 0:
if nccl_info.global_rank == 0:
export_to_video(videos[0], args.output_path + ".mp4", fps=30)
+4 -4
View File
@@ -20,7 +20,7 @@ from torch.utils.data.distributed import DistributedSampler
from tqdm.auto import tqdm
from fastvideo.dataset.latent_datasets import (LatentDataset, latent_collate_function)
from fastvideo.models.mochi_hf.mochi_latents_utils import normalize_dit_input
from fastvideo.utils.latents_utils import normalize_dit_input
from fastvideo.models.mochi_hf.pipeline_mochi import MochiPipeline
from fastvideo.models.hunyuan_hf.pipeline_hunyuan import HunyuanVideoPipeline
@@ -185,7 +185,7 @@ def main(args):
noise_random_generator = None
# Handle the repository creation
if rank <= 0 and args.output_dir is not None:
if rank == 0 and args.output_dir is not None:
os.makedirs(args.output_dir, exist_ok=True)
# For mixed precision training we cast all non-trainable weights to half-precision
@@ -316,7 +316,7 @@ def main(args):
len(train_dataloader) / args.gradient_accumulation_steps * args.sp_size / args.train_sp_batch_size)
args.num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch)
if rank <= 0:
if rank == 0:
project = args.tracker_project_name or "fastvideo"
wandb.init(project=project, config=args)
@@ -393,7 +393,7 @@ def main(args):
"grad_norm": grad_norm,
})
progress_bar.update(1)
if rank <= 0:
if rank == 0:
wandb.log(
{
"train_loss": loss,
+6 -6
View File
@@ -32,7 +32,7 @@ def save_checkpoint_optimizer(model, optimizer, rank, output_dir, step, discrimi
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:
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)
@@ -60,7 +60,7 @@ def save_checkpoint(transformer, rank, output_dir, step):
):
cpu_state = transformer.state_dict()
# todo move to get_state_dict
if rank <= 0:
if rank == 0:
save_dir = os.path.join(output_dir, f"checkpoint-{step}")
os.makedirs(save_dir, exist_ok=True)
# save using safetensors
@@ -98,7 +98,7 @@ def save_checkpoint_generator_discriminator(
hf_weight_dir = os.path.join(save_dir, "hf_weights")
os.makedirs(hf_weight_dir, exist_ok=True)
# save using safetensors
if rank <= 0:
if rank == 0:
config_dict = dict(model.config)
config_path = os.path.join(hf_weight_dir, "config.json")
# save dict as json
@@ -139,7 +139,7 @@ def save_checkpoint_generator_discriminator(
optim_state = FSDP.optim_state_dict(discriminator, discriminator_optimizer)
model_state = discriminator.state_dict()
state_dict = {"optimizer": optim_state, "model": model_state}
if rank <= 0:
if rank == 0:
discriminator_fsdp_state_fil = os.path.join(discriminator_fsdp_state_dir, "discriminator_state.pt")
torch.save(state_dict, discriminator_fsdp_state_fil)
@@ -178,7 +178,7 @@ def load_full_state_model(model, optimizer, checkpoint_file, rank):
):
discriminator_state = torch.load(checkpoint_file)
model_state = discriminator_state["model"]
if rank <= 0:
if rank == 0:
optim_state = discriminator_state["optimizer"]
else:
optim_state = None
@@ -241,7 +241,7 @@ def save_lora_checkpoint(transformer, optimizer, rank, output_dir, step, pipelin
optimizer,
)
if rank <= 0:
if rank == 0:
save_dir = os.path.join(output_dir, f"lora-checkpoint-{step}")
os.makedirs(save_dir, exist_ok=True)
+88
View File
@@ -0,0 +1,88 @@
import torch
mochi_latents_mean = torch.tensor([
-0.06730895953510081,
-0.038011381506090416,
-0.07477820912866141,
-0.05565264470995561,
0.012767231469026969,
-0.04703542746246419,
0.043896967884726704,
-0.09346305707025976,
-0.09918314763016893,
-0.008729793427399178,
-0.011931556316503654,
-0.0321993391887285,
]).view(1, 12, 1, 1, 1)
mochi_latents_std = torch.tensor([
0.9263795028493863,
0.9248894543193766,
0.9393059390890617,
0.959253732819592,
0.8244560132752793,
0.917259975397747,
0.9294154431013696,
1.3720942357788521,
0.881393668867029,
0.9168315692124348,
0.9185249279345552,
0.9274757570805041,
]).view(1, 12, 1, 1, 1)
mochi_scaling_factor = 1.0
wan_latents_mean = torch.tensor([
-0.7571,
-0.7089,
-0.9113,
0.1075,
-0.1745,
0.9653,
-0.1517,
1.5508,
0.4134,
-0.0715,
0.5517,
-0.3632,
-0.1922,
-0.9497,
0.2503,
-0.2921,
]).view(1, 16, 1, 1, 1)
wan_latents_std = torch.tensor([
2.8184,
1.4541,
2.3275,
2.6558,
1.2196,
1.7708,
2.6052,
2.0743,
3.2687,
2.1526,
2.8652,
1.5579,
1.6382,
1.1253,
2.8251,
1.916,
]).view(1, 16, 1, 1, 1)
def normalize_dit_input(model_type, latents):
if model_type == "mochi":
latents_mean = mochi_latents_mean.to(latents.device, latents.dtype)
latents_std = mochi_latents_std.to(latents.device, latents.dtype)
latents = (latents - latents_mean) / latents_std
return latents
elif model_type == "hunyuan_hf":
return latents * 0.476986
elif model_type == "hunyuan":
return latents * 0.476986
elif model_type == "wan":
latents_mean = wan_latents_mean.to(latents.device, latents.dtype)
latents_std = wan_latents_std.to(latents.device, latents.dtype)
latents = (latents - latents_mean) / latents_std
return latents
else:
raise NotImplementedError(f"model_type {model_type} not supported")
+69 -2
View File
@@ -3,9 +3,9 @@ from pathlib import Path
import torch
import torch.nn.functional as F
from diffusers import AutoencoderKLHunyuanVideo, AutoencoderKLMochi
from diffusers import AutoencoderKLHunyuanVideo, AutoencoderKLMochi, AutoencoderKLWan
from torch import nn
from transformers import AutoTokenizer, T5EncoderModel
from transformers import AutoTokenizer, T5EncoderModel, UMT5EncoderModel
from fastvideo.models.hunyuan.modules.models import (HYVideoDiffusionTransformer, MMDoubleStreamBlock,
MMSingleStreamBlock)
@@ -14,6 +14,7 @@ from fastvideo.models.hunyuan.vae.autoencoder_kl_causal_3d import AutoencoderKLC
from fastvideo.models.hunyuan_hf.modeling_hunyuan import (HunyuanVideoSingleTransformerBlock,
HunyuanVideoTransformer3DModel, HunyuanVideoTransformerBlock)
from fastvideo.models.mochi_hf.modeling_mochi import MochiTransformer3DModel, MochiTransformerBlock
from fastvideo.models.wan_hf.modeling_wan import WanTransformer3DModel, WanTransformerBlock
from fastvideo.utils.logging_ import main_print
hunyuan_config = {
@@ -200,6 +201,48 @@ class MochiTextEncoderWrapper(nn.Module):
return prompt_embeds, prompt_attention_mask
class WanTextEncoderWrapper(nn.Module):
def __init__(self, pretrained_model_name_or_path, device):
super().__init__()
self.text_encoder = UMT5EncoderModel.from_pretrained(os.path.join(pretrained_model_name_or_path,
"text_encoder")).to(device)
self.tokenizer = AutoTokenizer.from_pretrained(os.path.join(pretrained_model_name_or_path, "tokenizer"))
self.max_sequence_length = 256
def encode_prompt(self, prompt):
device = self.text_encoder.device
dtype = self.text_encoder.dtype
prompt = [prompt] if isinstance(prompt, str) else prompt
batch_size = len(prompt)
text_inputs = self.tokenizer(
prompt,
padding="max_length",
max_length=self.max_sequence_length,
truncation=True,
add_special_tokens=True,
return_tensors="pt",
)
text_input_ids = text_inputs.input_ids
prompt_attention_mask = text_inputs.attention_mask
prompt_attention_mask = prompt_attention_mask.bool().to(device)
untruncated_ids = self.tokenizer(prompt, padding="longest", return_tensors="pt").input_ids
if untruncated_ids.shape[-1] >= text_input_ids.shape[-1] and not torch.equal(text_input_ids, untruncated_ids):
removed_text = self.tokenizer.batch_decode(untruncated_ids[:, self.max_sequence_length - 1:-1])
main_print(f"Truncated text input: {prompt} to: {removed_text} for model input.")
prompt_embeds = self.text_encoder(text_input_ids.to(device), attention_mask=prompt_attention_mask)[0]
prompt_embeds = prompt_embeds.to(dtype=dtype, device=device)
# 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_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"
@@ -240,6 +283,20 @@ def load_transformer(
torch_dtype=master_weight_type,
# torch_dtype=torch.bfloat16 if args.use_lora else torch.float32,
)
elif model_type == "wan":
if dit_model_name_or_path:
transformer = WanTransformer3DModel.from_pretrained(
dit_model_name_or_path,
torch_dtype=master_weight_type,
# torch_dtype=torch.bfloat16 if args.use_lora else torch.float32,
)
else:
transformer = WanTransformer3DModel.from_pretrained(
pretrained_model_name_or_path,
subfolder="transformer",
torch_dtype=master_weight_type,
# torch_dtype=torch.bfloat16 if args.use_lora else torch.float32,
)
elif model_type == "hunyuan_hf":
if dit_model_name_or_path:
transformer = HunyuanVideoTransformer3DModel.from_pretrained(
@@ -283,6 +340,12 @@ def load_vae(model_type, pretrained_model_name_or_path):
torch_dtype=weight_dtype).to("cuda")
autocast_type = torch.bfloat16
fps = 24
elif model_type == "wan":
vae = AutoencoderKLWan.from_pretrained(pretrained_model_name_or_path,
subfolder="vae",
torch_dtype=weight_dtype).to("cuda")
autocast_type = torch.bfloat16
fps = 24
elif model_type == "hunyuan":
vae_precision = torch.float32
vae_path = os.path.join(pretrained_model_name_or_path, "hunyuan-video-t2v-720p/vae")
@@ -311,6 +374,8 @@ def load_vae(model_type, pretrained_model_name_or_path):
def load_text_encoder(model_type, pretrained_model_name_or_path, device):
if model_type == "mochi":
text_encoder = MochiTextEncoderWrapper(pretrained_model_name_or_path, device)
elif model_type == "wan":
text_encoder = WanTextEncoderWrapper(pretrained_model_name_or_path, device)
elif model_type == "hunyuan" or "hunyuan_hf":
text_encoder = HunyuanTextEncoderWrapper(pretrained_model_name_or_path, device)
else:
@@ -322,6 +387,8 @@ def get_no_split_modules(transformer):
# if of type MochiTransformer3DModel
if isinstance(transformer, MochiTransformer3DModel):
return (MochiTransformerBlock, )
elif isinstance(transformer, WanTransformer3DModel):
return (WanTransformerBlock, )
elif isinstance(transformer, HunyuanVideoTransformer3DModel):
return (HunyuanVideoSingleTransformerBlock, HunyuanVideoTransformerBlock)
elif isinstance(transformer, HYVideoDiffusionTransformer):
+23 -11
View File
@@ -129,13 +129,22 @@ 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):
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]
if model_type == "wan":
pred_kwargs = {
"hidden_states": latent_model_input,
"encoder_hidden_states": prompt_embeds,
"timestep":timestep,
"return_dict":False,
}
else:
pred_kwargs = {
"hidden_states": latent_model_input,
"encoder_hidden_states": prompt_embeds,
"timestep":timestep,
"encoder_attention_mask":prompt_attention_mask,
"return_dict":False,
}
noise_pred = transformer(**pred_kwargs)[0]
# Mochi CFG + Sampling runs in FP32
noise_pred = noise_pred.to(torch.float32)
@@ -166,10 +175,12 @@ def sample_validation_video(
# denormalize with the mean and std if available and not None
has_latents_mean = (hasattr(vae.config, "latents_mean") and vae.config.latents_mean is not None)
has_latents_std = (hasattr(vae.config, "latents_std") and vae.config.latents_std is not None)
if model_type == "wan":
vae.config.scaling_factor = 1
if has_latents_mean and has_latents_std:
latents_mean = (torch.tensor(vae.config.latents_mean).view(1, 12, 1, 1,
latents_mean = (torch.tensor(vae.config.latents_mean).view(1, num_channels_latents, 1, 1,
1).to(latents.device, latents.dtype))
latents_std = (torch.tensor(vae.config.latents_std).view(1, 12, 1, 1, 1).to(latents.device, latents.dtype))
latents_std = (torch.tensor(vae.config.latents_std).view(1, num_channels_latents, 1, 1, 1).to(latents.device, latents.dtype))
latents = latents * latents_std / vae.config.scaling_factor + latents_mean
else:
latents = latents / vae.config.scaling_factor
@@ -202,14 +213,15 @@ def log_validation(
vae_spatial_scale_factor = 8
vae_temporal_scale_factor = 6
num_channels_latents = 12
elif args.model_type == "hunyuan" or "hunyuan_hf":
elif args.model_type == "hunyuan" or "hunyuan_hf" or "wan":
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.enable_tiling()
if args.model_type != "wan":
vae.enable_tiling()
if scheduler_type == "euler":
scheduler = FlowMatchEulerDiscreteScheduler(shift=shift)
else:
+419
View File
@@ -0,0 +1,419 @@
import json
import os
from collections import defaultdict
from typing import Any, Dict, List, Optional, Tuple
import numpy as np
def configure_sta(mode: str = 'STA_searching',
layer_num: int = 40,
time_step_num: int = 50,
head_num: int = 40,
**kwargs) -> List[List[List[Any]]]:
"""
Configure Sliding Tile Attention (STA) parameters based on the specified mode.
Parameters:
----------
mode : str
The STA mode to use. Options are:
- 'STA_searching': Generate a set of mask candidates for initial search
- 'STA_tuning': Select best mask strategy based on previously saved results
- 'STA_inference': Load and use a previously tuned mask strategy
layer_num: int, number of layers
time_step_num: int, number of timesteps
head_num: int, number of heads
**kwargs : dict
Mode-specific parameters:
For 'STA_searching':
- mask_candidates: list of str, optional, mask candidates to use
- mask_selected: list of int, optional, indices of selected masks
For 'STA_tuning':
- mask_search_files_path: str, required, path to mask search results
- mask_candidates: list of str, optional, mask candidates to use
- mask_selected: list of int, optional, indices of selected masks
- skip_time_steps: int, optional, number of time steps to use full attention (default 12)
- save_dir: str, optional, directory to save mask strategy (default "mask_candidates")
For 'STA_inference':
- load_path: str, optional, path to load mask strategy (default "mask_candidates/mask_strategy.json")
"""
valid_modes = [
'STA_searching', 'STA_tuning', 'STA_inference', 'STA_tuning_cfg'
]
if mode not in valid_modes:
raise ValueError(f"Mode must be one of {valid_modes}, got {mode}")
if mode == 'STA_searching':
# Get parameters with defaults
mask_candidates: Optional[List[str]] = kwargs.get('mask_candidates')
if mask_candidates is None:
raise ValueError(
"mask_candidates is required for STA_searching mode")
mask_selected: List[int] = kwargs.get('mask_selected',
list(range(len(mask_candidates))))
# Parse selected masks
selected_masks: List[List[int]] = []
for index in mask_selected:
mask = mask_candidates[index]
masks_list = [int(x) for x in mask.split(',')]
selected_masks.append(masks_list)
# Create 3D mask structure with fixed dimensions (t=50, l=60)
masks_3d: List[List[List[List[int]]]] = []
for i in range(time_step_num): # Fixed t dimension = 50
row = []
for j in range(layer_num): # Fixed l dimension = 60
row.append(selected_masks) # Add all masks at each position
masks_3d.append(row)
return masks_3d
elif mode == 'STA_tuning':
# Get required parameters
mask_search_files_path: Optional[str] = kwargs.get(
'mask_search_files_path')
if not mask_search_files_path:
raise ValueError(
"mask_search_files_path is required for STA_tuning mode")
# Get optional parameters with defaults
mask_candidates_tuning: Optional[List[str]] = kwargs.get(
'mask_candidates')
if mask_candidates_tuning is None:
raise ValueError("mask_candidates is required for STA_tuning mode")
mask_selected_tuning: List[int] = kwargs.get(
'mask_selected', list(range(len(mask_candidates_tuning))))
skip_time_steps_tuning: Optional[int] = kwargs.get('skip_time_steps')
save_dir_tuning: Optional[str] = kwargs.get('save_dir',
"mask_candidates")
# Parse selected masks
selected_masks_tuning: List[List[int]] = []
for index in mask_selected_tuning:
mask = mask_candidates_tuning[index]
masks_list = [int(x) for x in mask.split(',')]
selected_masks_tuning.append(masks_list)
# Read JSON results
results = read_specific_json_files(mask_search_files_path)
averaged_results = average_head_losses(results, selected_masks_tuning)
# Add full attention mask for specific cases
full_attention_mask_tuning: Optional[List[int]] = kwargs.get(
'full_attention_mask')
if full_attention_mask_tuning is not None:
selected_masks_tuning.append(full_attention_mask_tuning)
# Select best mask strategy
timesteps_tuning: int = kwargs.get('timesteps', time_step_num)
if skip_time_steps_tuning is None:
skip_time_steps_tuning = 12
mask_strategy, sparsity, strategy_counts = select_best_mask_strategy(
averaged_results, selected_masks_tuning, skip_time_steps_tuning,
timesteps_tuning, head_num)
# Save mask strategy
if save_dir_tuning is not None:
os.makedirs(save_dir_tuning, exist_ok=True)
file_path = os.path.join(
save_dir_tuning,
f'mask_strategy_s{skip_time_steps_tuning}.json')
with open(file_path, 'w') as f:
json.dump(mask_strategy, f, indent=4)
print(f"Successfully saved mask_strategy to {file_path}")
# Print sparsity and strategy counts for information
print(f"Overall sparsity: {sparsity:.4f}")
print("\nStrategy usage counts:")
total_heads = time_step_num * layer_num * head_num # Fixed dimensions
for strategy, count in strategy_counts.items():
print(
f"Strategy {strategy}: {count} heads ({count/total_heads*100:.2f}%)"
)
# Convert dictionary to 3D list with fixed dimensions
mask_strategy_3d = dict_to_3d_list(mask_strategy,
t_max=time_step_num,
l_max=layer_num,
h_max=head_num)
return mask_strategy_3d
elif mode == 'STA_tuning_cfg':
# Get required parameters for both positive and negative paths
mask_search_files_path_pos: Optional[str] = kwargs.get(
'mask_search_files_path_pos')
mask_search_files_path_neg: Optional[str] = kwargs.get(
'mask_search_files_path_neg')
save_dir_cfg: Optional[str] = kwargs.get('save_dir')
if not mask_search_files_path_pos or not mask_search_files_path_neg or not save_dir_cfg:
raise ValueError(
"mask_search_files_path_pos, mask_search_files_path_neg, and save_dir are required for STA_tuning_cfg mode"
)
# Get optional parameters with defaults
mask_candidates_cfg: Optional[List[str]] = kwargs.get('mask_candidates')
if mask_candidates_cfg is None:
raise ValueError(
"mask_candidates is required for STA_tuning_cfg mode")
mask_selected_cfg: List[int] = kwargs.get(
'mask_selected', list(range(len(mask_candidates_cfg))))
skip_time_steps_cfg: Optional[int] = kwargs.get('skip_time_steps')
# Parse selected masks
selected_masks_cfg: List[List[int]] = []
for index in mask_selected_cfg:
mask = mask_candidates_cfg[index]
masks_list = [int(x) for x in mask.split(',')]
selected_masks_cfg.append(masks_list)
# Read JSON results for both positive and negative paths
pos_results = read_specific_json_files(mask_search_files_path_pos)
neg_results = read_specific_json_files(mask_search_files_path_neg)
# Combine positive and negative results into one list
combined_results = pos_results + neg_results
# Average the combined results
averaged_results = average_head_losses(combined_results,
selected_masks_cfg)
# Add full attention mask for specific cases
full_attention_mask_cfg: Optional[List[int]] = kwargs.get(
'full_attention_mask')
if full_attention_mask_cfg is not None:
selected_masks_cfg.append(full_attention_mask_cfg)
timesteps_cfg: int = kwargs.get('timesteps', time_step_num)
if skip_time_steps_cfg is None:
skip_time_steps_cfg = 12
# Select best mask strategy using combined results
mask_strategy, sparsity, strategy_counts = select_best_mask_strategy(
averaged_results, selected_masks_cfg, skip_time_steps_cfg,
timesteps_cfg, head_num)
# Save mask strategy
os.makedirs(save_dir_cfg, exist_ok=True)
file_path = os.path.join(save_dir_cfg,
f'mask_strategy_s{skip_time_steps_cfg}.json')
with open(file_path, 'w') as f:
json.dump(mask_strategy, f, indent=4)
print(f"Successfully saved mask_strategy to {file_path}")
# Print sparsity and strategy counts for information
print(f"Overall sparsity: {sparsity:.4f}")
print("\nStrategy usage counts:")
total_heads = time_step_num * layer_num * head_num # Fixed dimensions
for strategy, count in strategy_counts.items():
print(
f"Strategy {strategy}: {count} heads ({count/total_heads*100:.2f}%)"
)
# Convert dictionary to 3D list with fixed dimensions
mask_strategy_3d = dict_to_3d_list(mask_strategy,
t_max=time_step_num,
l_max=layer_num,
h_max=head_num)
return mask_strategy_3d
else: # STA_inference
# Get parameters with defaults
load_path: Optional[str] = kwargs.get(
'load_path', "mask_candidates/mask_strategy.json")
if load_path is None:
raise ValueError("load_path is required for STA_inference mode")
# Load previously saved mask strategy
with open(load_path) as f:
mask_strategy = json.load(f)
# Convert dictionary to 3D list with fixed dimensions
mask_strategy_3d = dict_to_3d_list(mask_strategy,
t_max=time_step_num,
l_max=layer_num,
h_max=head_num)
return mask_strategy_3d
# Helper functions
def read_specific_json_files(folder_path: str) -> List[Dict[str, Any]]:
"""Read and parse JSON files containing mask search results."""
json_contents: List[Dict[str, Any]] = []
# List files only in the current directory (no walk)
files = os.listdir(folder_path)
# Filter files
matching_files = [f for f in files if 'mask' in f and f.endswith('.json')]
print(f"Found {len(matching_files)} matching files: {matching_files}")
for file_name in matching_files:
file_path = os.path.join(folder_path, file_name)
with open(file_path) as file:
data = json.load(file)
json_contents.append(data)
return json_contents
def average_head_losses(
results: List[Dict[str, Any]],
selected_masks: List[List[int]]) -> Dict[str, Dict[str, np.ndarray]]:
"""Average losses across all prompts for each mask strategy."""
# Initialize a dictionary to store the averaged results
averaged_losses: Dict[str, Dict[str, np.ndarray]] = {}
loss_type = 'L2_loss'
# Get all loss types (e.g., 'L2_loss')
averaged_losses[loss_type] = {}
for mask in selected_masks:
mask_str = str(mask)
data_shape = np.array(results[0][loss_type][mask_str]).shape
accumulated_data = np.zeros(data_shape)
# Sum across all prompts
for prompt_result in results:
accumulated_data += np.array(prompt_result[loss_type][mask_str])
# Average by dividing by number of prompts
averaged_data = accumulated_data / len(results)
averaged_losses[loss_type][mask_str] = averaged_data
return averaged_losses
def select_best_mask_strategy(
averaged_results: Dict[str, Dict[str, np.ndarray]],
selected_masks: List[List[int]],
skip_time_steps: int = 12,
timesteps: int = 50,
head_num: int = 40
) -> Tuple[Dict[str, List[int]], float, Dict[str, int]]:
"""Select the best mask strategy for each head based on loss minimization."""
best_mask_strategy: Dict[str, List[int]] = {}
loss_type = 'L2_loss'
# Get the shape of time steps and layers
layers = len(averaged_results[loss_type][str(selected_masks[0])][0])
# Counter for sparsity calculation
total_tokens = 0 # total number of masked tokens
total_length = 0 # total sequence length
strategy_counts: Dict[str, int] = {
str(strategy): 0
for strategy in selected_masks
}
full_attn_strategy = selected_masks[-1] # Last strategy is full attention
print(f"Strategy {full_attn_strategy}, skip first {skip_time_steps} steps ")
for t in range(timesteps):
for layer_idx in range(layers):
for h in range(head_num):
if t < skip_time_steps: # First steps use full attention
strategy = full_attn_strategy
else:
# Get losses for this head across all strategies
head_losses = []
for strategy in selected_masks[:
-1]: # Exclude full attention
head_losses.append(averaged_results[loss_type][str(
strategy)][t][layer_idx][h])
# Find which strategy gives minimum loss
best_strategy_idx = np.argmin(head_losses)
strategy = selected_masks[best_strategy_idx]
best_mask_strategy[f'{t}_{layer_idx}_{h}'] = strategy
# Calculate sparsity
nums = strategy # strategy is already a list of numbers
total_tokens += nums[0] * nums[1] * nums[
2] # masked tokens for chosen strategy
total_length += full_attn_strategy[0] * full_attn_strategy[
1] * full_attn_strategy[2]
# Count strategy usage
strategy_counts[str(strategy)] += 1
overall_sparsity = 1 - total_tokens / total_length
return best_mask_strategy, overall_sparsity, strategy_counts
def dict_to_3d_list(mask_strategy: Optional[Dict[str, List[int]]],
t_max: int = 50,
l_max: int = 60,
h_max: int = 24) -> List[List[List[Optional[List[int]]]]]:
result: List[List[List[Optional[List[int]]]]] = [[[
None for _ in range(h_max)
] for _ in range(l_max)] for _ in range(t_max)]
if mask_strategy is None:
return result
for key, value in mask_strategy.items():
t, layer_idx, h = map(int, key.split('_'))
result[t][layer_idx][h] = value
return result
def save_mask_search_results(
mask_search_final_result: List[Dict[str, List[float]]],
prompt: str,
mask_strategies: List[str],
output_dir: str = 'output/mask_search_result/') -> Optional[str]:
if not mask_search_final_result:
print("No mask search results to save")
return None
# Create result dictionary with defaultdict for nested lists
mask_search_dict: Dict[str, Dict[str, List[List[float]]]] = {
"L2_loss": defaultdict(list),
"L1_loss": defaultdict(list)
}
mask_selected = list(range(len(mask_strategies)))
selected_masks: List[List[int]] = []
for index in mask_selected:
mask = mask_strategies[index]
masks_list = [int(x) for x in mask.split(',')]
selected_masks.append(masks_list)
# Process each mask strategy
for i, mask_strategy in enumerate(selected_masks):
mask_strategy_str = str(mask_strategy)
# Process L2 loss
step_results: List[List[float]] = []
for step_data in mask_search_final_result:
if isinstance(step_data, dict) and "L2_loss" in step_data:
layer_losses = [float(loss) for loss in step_data["L2_loss"]]
step_results.append(layer_losses)
mask_search_dict["L2_loss"][mask_strategy_str] = step_results
step_results = []
for step_data in mask_search_final_result:
if isinstance(step_data, dict) and "L1_loss" in step_data:
layer_losses = [float(loss) for loss in step_data["L1_loss"]]
step_results.append(layer_losses)
mask_search_dict["L1_loss"][mask_strategy_str] = step_results
# Create the output directory if it doesn't exist
os.makedirs(output_dir, exist_ok=True)
# Create a filename based on the first 20 characters of the prompt
filename = prompt[:50].replace(" ", "_")
filepath = os.path.join(output_dir, f'mask_search_{filename}.json')
# Save the results to a JSON file
with open(filepath, 'w') as f:
json.dump(mask_search_dict, f, indent=4)
print(f"Successfully saved mask research results to {filepath}")
return filepath
@@ -1,6 +1,6 @@
import json
from dataclasses import dataclass
from typing import List, Optional, Type
from typing import Any, Dict, List, Optional, Type
import torch
from einops import rearrange
@@ -13,6 +13,7 @@ from fastvideo.v1.attention.backends.abstract import (AttentionBackend,
AttentionMetadataBuilder)
from fastvideo.v1.distributed import get_sp_group
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.forward_context import ForwardContext, get_forward_context
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
@@ -20,7 +21,9 @@ logger = init_logger(__name__)
# TODO(will-refactor): move this to a utils file
def dict_to_3d_list(mask_strategy) -> List[List[List[Optional[torch.Tensor]]]]:
def dict_to_3d_list(
mask_strategy: Dict[str,
Any]) -> List[List[List[Optional[torch.Tensor]]]]:
indices = [tuple(map(int, key.split('_'))) for key in mask_strategy]
max_timesteps_idx = max(
@@ -42,14 +45,14 @@ def dict_to_3d_list(mask_strategy) -> List[List[List[Optional[torch.Tensor]]]]:
class RangeDict(dict):
def __getitem__(self, item):
def __getitem__(self, item: int) -> str:
for key in self.keys():
if isinstance(key, tuple):
low, high = key
if low <= item <= high:
return super().__getitem__(key)
return str(super().__getitem__(key))
elif key == item:
return super().__getitem__(key)
return str(super().__getitem__(key))
raise KeyError(f"seq_len {item} not supported for STA")
@@ -82,6 +85,8 @@ class SlidingTileAttentionBackend(AttentionBackend):
@dataclass
class SlidingTileAttentionMetadata(AttentionMetadata):
current_timestep: int
STA_param: List[List[
Any]] # each timestep with one metadata, shape [num_layers, num_heads]
class SlidingTileAttentionMetadataBuilder(AttentionMetadataBuilder):
@@ -98,8 +103,12 @@ class SlidingTileAttentionMetadataBuilder(AttentionMetadataBuilder):
forward_batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> SlidingTileAttentionMetadata:
return SlidingTileAttentionMetadata(current_timestep=current_timestep, )
param = forward_batch.STA_param
if param is None:
return SlidingTileAttentionMetadata(
current_timestep=current_timestep, STA_param=[])
return SlidingTileAttentionMetadata(current_timestep=current_timestep,
STA_param=param[current_timestep])
class SlidingTileAttentionImpl(AttentionImpl):
@@ -120,12 +129,12 @@ class SlidingTileAttentionImpl(AttentionImpl):
if config_file is None:
raise ValueError("FASTVIDEO_ATTENTION_CONFIG is not set")
# TODO(kevin): get mask strategy for different STA modes
with open(config_file) as f:
mask_strategy = json.load(f)
mask_strategy = dict_to_3d_list(mask_strategy)
self.mask_strategy = dict_to_3d_list(mask_strategy)
self.prefix = prefix
self.mask_strategy = mask_strategy
sp_group = get_sp_group()
self.sp_size = sp_group.world_size
# STA config
@@ -205,16 +214,24 @@ class SlidingTileAttentionImpl(AttentionImpl):
v: torch.Tensor,
attn_metadata: SlidingTileAttentionMetadata,
) -> torch.Tensor:
assert self.mask_strategy is not None, "mask_strategy cannot be None for SlidingTileAttention"
assert self.mask_strategy[
0] is not None, "mask_strategy[0] cannot be None for SlidingTileAttention"
if self.mask_strategy is None:
raise ValueError(
"mask_strategy cannot be None for SlidingTileAttention")
if self.mask_strategy[0] is None:
raise ValueError(
"mask_strategy[0] cannot be None for SlidingTileAttention")
timestep = attn_metadata.current_timestep
forward_context: ForwardContext = get_forward_context()
forward_batch = forward_context.forward_batch
if forward_batch is None:
raise ValueError("forward_batch cannot be None")
# pattern:'.double_blocks.0.attn.impl' or '.single_blocks.0.attn.impl'
layer_idx = int(self.prefix.split('.')[-3])
# TODO: remove hardcode
if attn_metadata.STA_param is None or len(
attn_metadata.STA_param) <= layer_idx:
raise ValueError("Invalid STA_param")
STA_param = attn_metadata.STA_param[layer_idx]
text_length = q.shape[1] - self.img_seq_length
has_text = text_length > 0
@@ -227,15 +244,62 @@ class SlidingTileAttentionImpl(AttentionImpl):
sp_group = get_sp_group()
current_rank = sp_group.rank_in_group
start_head = current_rank * head_num
windows = [
self.mask_strategy[timestep][layer_idx][head_idx + start_head]
for head_idx in range(head_num)
]
# if has_text is False:
# from IPython import embed
# embed()
hidden_states = sliding_tile_attention(
query, key, value, windows, text_length, has_text,
self.img_latent_shape_str).transpose(1, 2)
# searching or tuning mode
if len(STA_param) < head_num * sp_group.world_size:
sparse_attn_hidden_states_all = []
full_mask_window = STA_param[-1]
for window_size in STA_param[:-1]:
sparse_hidden_states = sliding_tile_attention(
query, key, value, [window_size] * head_num, text_length,
has_text, self.img_latent_shape_str).transpose(1, 2)
sparse_attn_hidden_states_all.append(sparse_hidden_states)
hidden_states = sliding_tile_attention(
query, key, value, [full_mask_window] * head_num, text_length,
has_text, self.img_latent_shape_str).transpose(1, 2)
attn_L2_loss = []
attn_L1_loss = []
# average loss across all heads
for sparse_attn_hidden_states in sparse_attn_hidden_states_all:
# L2 loss
attn_L2_loss_ = torch.mean((sparse_attn_hidden_states.float() -
hidden_states.float())**2,
dim=[0, 1, 3]).cpu().numpy()
attn_L2_loss_ = [round(float(x), 6) for x in attn_L2_loss_]
attn_L2_loss.append(attn_L2_loss_)
# L1 loss
attn_L1_loss_ = torch.mean(
torch.abs(sparse_attn_hidden_states.float() -
hidden_states.float()),
dim=[0, 1, 3]).cpu().numpy()
attn_L1_loss_ = [round(float(x), 6) for x in attn_L1_loss_]
attn_L1_loss.append(attn_L1_loss_)
layer_loss_save = {"L2_loss": attn_L2_loss, "L1_loss": attn_L1_loss}
if forward_batch.is_cfg_negative:
if forward_batch.mask_search_final_result_neg is not None:
forward_batch.mask_search_final_result_neg[timestep].append(
layer_loss_save)
else:
if forward_batch.mask_search_final_result_pos is not None:
forward_batch.mask_search_final_result_pos[timestep].append(
layer_loss_save)
else:
# windows = [
# self.mask_strategy[timestep][layer_idx][head_idx + start_head]
# for head_idx in range(head_num)
# ]
windows = [
STA_param[head_idx + start_head] for head_idx in range(head_num)
]
# if has_text is False:
# from IPython import embed
# embed()
hidden_states = sliding_tile_attention(
query, key, value, windows, text_length, has_text,
self.img_latent_shape_str).transpose(1, 2)
return hidden_states
+3 -2
View File
@@ -13,6 +13,7 @@ from fastvideo.v1.distributed.parallel_state import (
get_sequence_model_parallel_rank, get_sequence_model_parallel_world_size)
from fastvideo.v1.forward_context import ForwardContext, get_forward_context
from fastvideo.v1.platforms import _Backend
from fastvideo.v1.utils import get_compute_dtype
class DistributedAttention(nn.Module):
@@ -38,7 +39,7 @@ class DistributedAttention(nn.Module):
if num_kv_heads is None:
num_kv_heads = num_heads
dtype = torch.get_default_dtype()
dtype = get_compute_dtype()
attn_backend = get_attn_backend(
head_size,
dtype,
@@ -155,7 +156,7 @@ class LocalAttention(nn.Module):
if num_kv_heads is None:
num_kv_heads = num_heads
dtype = torch.get_default_dtype()
dtype = get_compute_dtype()
attn_backend = get_attn_backend(
head_size,
dtype,
+3 -1
View File
@@ -1,5 +1,5 @@
from dataclasses import dataclass, field
from typing import Any, Optional, Tuple
from typing import Any, List, Optional, Tuple
from fastvideo.v1.configs.models.base import ArchConfig, ModelConfig
from fastvideo.v1.layers.quantization import QuantizationConfig
@@ -11,6 +11,7 @@ class DiTArchConfig(ArchConfig):
_fsdp_shard_conditions: list = field(default_factory=list)
_compile_conditions: list = field(default_factory=list)
_param_names_mapping: dict = field(default_factory=dict)
_lora_param_names_mapping: dict = field(default_factory=dict)
_supported_attention_backends: Tuple[_Backend,
...] = (_Backend.SLIDING_TILE_ATTN,
_Backend.SAGE_ATTN,
@@ -20,6 +21,7 @@ class DiTArchConfig(ArchConfig):
hidden_size: int = 0
num_attention_heads: int = 0
num_channels_latents: int = 0
exclude_lora_layers: List[str] = field(default_factory=list)
def __post_init__(self) -> None:
if not self._compile_conditions:
@@ -1,5 +1,5 @@
from dataclasses import dataclass, field
from typing import Optional, Tuple
from typing import List, Optional, Tuple
import torch
@@ -163,6 +163,8 @@ class HunyuanVideoArchConfig(DiTArchConfig):
pooled_projection_dim: int = 768
rope_theta: int = 256
qk_norm: str = "rms_norm"
exclude_lora_layers: List[str] = field(
default_factory=lambda: ["img_in", "txt_in", "time_in", "vector_in"])
def __post_init__(self):
super().__post_init__()
@@ -51,6 +51,7 @@ class StepVideoArchConfig(DiTArchConfig):
default_factory=lambda: [6144, 1024])
attention_type: Optional[str] = "torch"
use_additional_conditions: Optional[bool] = False
exclude_lora_layers: List[str] = field(default_factory=lambda: [])
def __post_init__(self):
self.hidden_size = self.num_attention_heads * self.attention_head_dim
+19 -1
View File
@@ -1,5 +1,5 @@
from dataclasses import dataclass, field
from typing import Optional, Tuple
from typing import List, Optional, Tuple
from fastvideo.v1.configs.models.dits.base import DiTArchConfig, DiTConfig
@@ -51,6 +51,23 @@ class WanVideoArchConfig(DiTArchConfig):
r"blocks\.(\d+)\.norm2\.(.*)$":
r"blocks.\1.self_attn_residual_norm.norm.\2",
})
# Some LoRA adapters use the original official layer names instead of hf layer names,
# so apply this before the param_names_mapping
_lora_param_names_mapping: dict = field(
default_factory=lambda: {
r"^blocks\.(\d+)\.self_attn\.q\.(.*)$": r"blocks.\1.attn1.to_q.\2",
r"^blocks\.(\d+)\.self_attn\.k\.(.*)$": r"blocks.\1.attn1.to_k.\2",
r"^blocks\.(\d+)\.self_attn\.v\.(.*)$": r"blocks.\1.attn1.to_v.\2",
r"^blocks\.(\d+)\.self_attn\.o\.(.*)$":
r"blocks.\1.attn1.to_out.0.\2",
r"^blocks\.(\d+)\.cross_attn\.q\.(.*)$": r"blocks.\1.attn2.to_q.\2",
r"^blocks\.(\d+)\.cross_attn\.k\.(.*)$": r"blocks.\1.attn2.to_k.\2",
r"^blocks\.(\d+)\.cross_attn\.v\.(.*)$": r"blocks.\1.attn2.to_v.\2",
r"^blocks\.(\d+)\.cross_attn\.o\.(.*)$":
r"blocks.\1.attn2.to_out.0.\2",
r"^blocks\.(\d+)\.ffn\.0\.(.*)$": r"blocks.\1.ffn.fc_in.\2",
r"^blocks\.(\d+)\.ffn\.2\.(.*)$": r"blocks.\1.ffn.fc_out.\2",
})
patch_size: Tuple[int, int, int] = (1, 2, 2)
text_len = 512
@@ -68,6 +85,7 @@ class WanVideoArchConfig(DiTArchConfig):
image_dim: Optional[int] = None
added_kv_proj_dim: Optional[int] = None
rope_max_seq_len: int = 1024
exclude_lora_layers: List[str] = field(default_factory=lambda: ["embedder"])
def __post_init__(self):
super().__post_init__()
+2 -1
View File
@@ -27,7 +27,6 @@ class PipelineConfig:
# Video generation parameters
embedded_cfg_scale: float = 6.0
flow_shift: Optional[float] = None
use_cpu_offload: bool = False
disable_autocast: bool = False
# Model configuration
@@ -55,6 +54,8 @@ class PipelineConfig:
# STA (Spatial-Temporal Attention) parameters
mask_strategy_file_path: Optional[str] = None
STA_mode: str = "STA_inference"
skip_time_steps: int = 15
# Compilation
enable_torch_compile: bool = False
@@ -68,9 +68,6 @@ class HunyuanConfig(PipelineConfig):
embedded_cfg_scale: int = 6
flow_shift: int = 7
# Video parameters
use_cpu_offload: bool = True
# Text encoding stage
text_encoder_configs: Tuple[EncoderConfig, ...] = field(
default_factory=lambda: (LlamaConfig(), CLIPTextConfig()))
@@ -18,9 +18,6 @@ class StepVideoT2VConfig(PipelineConfig):
vae_tiling: bool = False
vae_sp: bool = False
# Video parameters
use_cpu_offload: bool = True
# Denoising stage
flow_shift: int = 13
timesteps_scale: bool = False
-3
View File
@@ -37,9 +37,6 @@ class WanT2V480PConfig(PipelineConfig):
vae_tiling: bool = False
vae_sp: bool = False
# Video parameters
use_cpu_offload: bool = True
# Denoising stage
flow_shift: int = 3
+40 -2
View File
@@ -9,7 +9,45 @@ frameworks that can handle parquet or lance file.
import pyarrow as pa
pyarrow_schema = pa.schema([
pyarrow_schema_i2v = pa.schema([
pa.field("id", pa.string()),
# --- Image/Video VAE latents ---
# Tensors are stored as raw bytes with shape and dtype info for loading
pa.field("vae_latent_bytes", pa.binary()),
# e.g., [C, T, H, W] or [C, H, W]
pa.field("vae_latent_shape", pa.list_(pa.int64())),
# e.g., 'float32'
pa.field("vae_latent_dtype", pa.string()),
# --- Text encoder output tensor ---
# Tensors are stored as raw bytes with shape and dtype info for loading
pa.field("text_embedding_bytes", pa.binary()),
# e.g., [SeqLen, Dim]
pa.field("text_embedding_shape", pa.list_(pa.int64())),
# e.g., 'bfloat16' or 'float32'
pa.field("text_embedding_dtype", pa.string()),
pa.field("text_attention_mask_bytes", pa.binary()),
# e.g., [SeqLen]
pa.field("text_attention_mask_shape", pa.list_(pa.int64())),
# e.g., 'bool' or 'int8'
pa.field("text_attention_mask_dtype", pa.string()),
#I2V
pa.field("clip_feature_bytes", pa.binary()),
pa.field("clip_feature_shape", pa.list_(pa.int64())),
pa.field("clip_feature_dtype", pa.string()),
# --- Metadata ---
pa.field("file_name", pa.string()),
pa.field("caption", pa.string()),
pa.field("media_type", pa.string()), # 'image' or 'video'
pa.field("width", pa.int64()),
pa.field("height", pa.int64()),
# -- Video-specific (can be null/default for images) ---
# Number of frames processed (e.g., 1 for image, N for video)
pa.field("num_frames", pa.int64()),
pa.field("duration_sec", pa.float64()),
pa.field("fps", pa.float64()),
])
pyarrow_schema_t2v = pa.schema([
pa.field("id", pa.string()),
# --- Image/Video VAE latents ---
# Tensors are stored as raw bytes with shape and dtype info for loading
@@ -41,4 +79,4 @@ pyarrow_schema = pa.schema([
pa.field("num_frames", pa.int64()),
pa.field("duration_sec", pa.float64()),
pa.field("fps", pa.float64()),
])
])
+136
View File
@@ -0,0 +1,136 @@
import argparse
import json
import os
import time
from multiprocessing import Pool, cpu_count
from pathlib import Path
import torchvision
from tqdm import tqdm
def get_video_info(video_path):
"""Get video information using torchvision."""
# Read video tensor (T, C, H, W)
video_tensor, _, info = torchvision.io.read_video(str(video_path),
output_format="TCHW",
pts_unit="sec")
num_frames = video_tensor.shape[0]
height = video_tensor.shape[2]
width = video_tensor.shape[3]
fps = info.get("video_fps", 0)
duration = num_frames / fps if fps > 0 else 0
# Extract name
_, _, videos_dir, video_name = str(video_path).split("/")
return {
"path": str(video_name),
"resolution": {
"width": width,
"height": height
},
"size": os.path.getsize(video_path),
"fps": fps,
"duration": duration,
"num_frames": num_frames
}
def prepare_dataset_json(folder_path,
output_name="videos2caption.json",
num_workers=None) -> None:
"""Prepare dataset information from a folder containing videos and prompt.txt."""
folder_path = Path(folder_path)
# Read prompt file
prompt_file = folder_path / "prompt.txt"
if not prompt_file.exists():
raise FileNotFoundError(f"prompt.txt not found in {folder_path}")
with open(prompt_file) as f:
prompts = [line.strip() for line in f.readlines() if line.strip()]
# Read videos file
videos_file = folder_path / "videos.txt"
if not videos_file.exists():
raise FileNotFoundError(f"videos.txt not found in {folder_path}")
with open(videos_file) as f:
video_paths = [line.strip() for line in f.readlines() if line.strip()]
if len(prompts) != len(video_paths):
raise ValueError(
f"Number of prompts ({len(prompts)}) does not match number of videos ({len(video_paths)})"
)
# Prepare arguments for multiprocessing
process_args = [folder_path / video_path for video_path in video_paths]
# Determine number of workers
if num_workers is None:
num_workers = max(1, cpu_count() - 1) # Leave one CPU free
# Process videos in parallel
start_time = time.time()
with Pool(num_workers) as pool:
results = list(
tqdm(pool.imap(get_video_info, process_args),
total=len(process_args),
desc="Processing videos",
unit="video"))
# Combine results with prompts
dataset_info = []
for result, prompt in zip(results, prompts):
result["cap"] = [prompt]
dataset_info.append(result)
# Calculate total processing time
total_time = time.time() - start_time
total_videos = len(dataset_info)
avg_time_per_video = total_time / total_videos if total_videos > 0 else 0
print("\nProcessing completed:")
print(f"Total videos processed: {total_videos}")
print(f"Total time: {total_time:.2f} seconds")
print(f"Average time per video: {avg_time_per_video:.2f} seconds")
# Save to JSON file
output_file = folder_path / output_name
with open(output_file, 'w') as f:
json.dump(dataset_info, f, indent=2)
# Create merge.txt
merge_file = folder_path / "merge.txt"
with open(merge_file, 'w') as f:
f.write(f"{folder_path}/videos,{output_file}\n")
print(f"Dataset information saved to {output_file}")
print(f"Merge file created at {merge_file}")
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(
description='Prepare video dataset information in JSON format')
parser.add_argument(
'--folder',
type=str,
required=True,
help='Path to the folder containing videos and prompt.txt')
parser.add_argument(
'--output',
type=str,
default='videos2caption.json',
help='Name of the output JSON file (default: videos2caption.json)')
parser.add_argument('--workers',
type=int,
default=32,
help='Number of worker processes (default: 16)')
return parser.parse_args()
if __name__ == "__main__":
args = parse_args()
prepare_dataset_json(args.folder, args.output, args.workers)
-20
View File
@@ -107,23 +107,3 @@ def latent_collate_function(batch):
prompt_attention_masks = torch.stack(prompt_attention_masks, dim=0)
latents = torch.stack(latent_list, dim=0)
return latents, prompt_embeds, latent_attn_mask, prompt_attention_masks
if __name__ == "__main__":
dataset = LatentDataset("data/Mochi-Synthetic-Data/merge.txt",
num_latent_t=28,
cfg_rate=0.0)
dataloader = torch.utils.data.DataLoader(dataset,
batch_size=2,
shuffle=False,
collate_fn=latent_collate_function)
for latent, prompt_embed, latent_attn_mask, prompt_attention_mask in dataloader:
print(
latent.shape,
prompt_embed.shape,
latent_attn_mask.shape,
prompt_attention_mask.shape,
)
import pdb
pdb.set_trace()
+136 -35
View File
@@ -15,7 +15,8 @@ from torch import distributed as dist
from torch.utils.data import Dataset
from torchdata.stateful_dataloader import StatefulDataLoader
from fastvideo.v1.distributed import (get_sequence_model_parallel_rank,
from fastvideo.v1.distributed import (get_dp_group,
get_sequence_model_parallel_rank,
get_sp_group)
from fastvideo.v1.logger import init_logger
@@ -32,18 +33,31 @@ class ParquetVideoTextDataset(Dataset):
world_size: int = 1,
cfg_rate: float = 0.0,
num_latent_t: int = 2,
seed: int = 0):
seed: int = 0,
validation: bool = False):
super().__init__()
self.path = str(path)
self.batch_size = batch_size
self.rank = rank
self.local_rank = get_sequence_model_parallel_rank()
self.sp_world_size = world_size
self.sp_group = get_sp_group()
self.dp_group = get_dp_group()
self.dp_world_size = self.dp_group.world_size
self.sp_world_size = self.sp_group.world_size
self.world_size = int(os.getenv("WORLD_SIZE", 1))
self.cfg_rate = cfg_rate
self.num_latent_t = num_latent_t
self.local_indices = None
self.plan_output_dir = os.path.join(self.path, "data_plan.json")
self.validation = validation
# Negative prompt caching
self.neg_metadata = None
self.cached_neg_prompt: Dict[str, Any] | None = None
self.plan_output_dir = os.path.join(
self.path,
f"data_plan_{self.world_size}_{self.sp_world_size}_{self.dp_world_size}.json"
)
ranks = get_sp_group().ranks
group_ranks: List[List] = [[] for _ in range(self.world_size)]
@@ -54,39 +68,126 @@ class ParquetVideoTextDataset(Dataset):
# This will be useful when resume training
if os.path.exists(self.plan_output_dir):
print(f"Using existing plan from {self.plan_output_dir}")
return
else:
print(f"Creating new plan for {self.plan_output_dir}")
# Find all parquet files recursively, and record num_rows for each file
print(f"Scanning for parquet files in {self.path}")
metadatas = []
for root, _, files in os.walk(self.path):
for file in sorted(files):
if file.endswith('.parquet'):
file_path = os.path.join(root, file)
num_rows = pq.ParquetFile(
file_path).metadata.num_rows
for row_idx in range(num_rows):
metadatas.append((file_path, row_idx))
# Find all parquet files recursively, and record num_rows for each file
print(f"Scanning for parquet files in {self.path}")
metadatas = []
for root, _, files in os.walk(self.path):
for file in sorted(files):
if file.endswith('.parquet'):
file_path = os.path.join(root, file)
num_rows = pq.ParquetFile(file_path).metadata.num_rows
for row_idx in range(num_rows):
metadatas.append((file_path, row_idx))
# the negative prompt is always the first row in the first
# parquet file
if validation:
self.neg_metadata = metadatas[0]
metadatas = metadatas[1:]
# Generate the plan that distribute rows among workers
random.seed(seed)
random.shuffle(metadatas)
# Generate the plan that distribute rows among workers
random.seed(seed)
random.shuffle(metadatas)
# Get all sp groups
# e.g. if num_gpus = 4, sp_size = 2
# group_ranks = [(0, 1), (2, 3)]
# We will assign the same batches of data to ranks in the same sp group, and we'll assign different batches to ranks in different sp groups
# e.g. plan = {0: [row 1, row 4], 1: [row 1, row 4], 2: [row 2, row 3], 3: [row 2, row 3]}
group_ranks_list: List[Any] = list(
set(tuple(r) for r in group_ranks))
num_sp_groups = len(group_ranks_list)
plan = defaultdict(list)
for idx, metadata in enumerate(metadatas):
sp_group_idx = idx % num_sp_groups
for global_rank in group_ranks_list[sp_group_idx]:
plan[global_rank].append(metadata)
# Get all sp groups
# e.g. if num_gpus = 4, sp_size = 2
# group_ranks = [(0, 1), (2, 3)]
# We will assign the same batches of data to ranks in the same sp group, and we'll assign different batches to ranks in different sp groups
# e.g. plan = {0: [row 1, row 4], 1: [row 1, row 4], 2: [row 2, row 3], 3: [row 2, row 3]}
group_ranks_list: List[Any] = list(
set(tuple(r) for r in group_ranks))
num_sp_groups = len(group_ranks_list)
plan = defaultdict(list)
for idx, metadata in enumerate(metadatas):
sp_group_idx = idx % num_sp_groups
for global_rank in group_ranks_list[sp_group_idx]:
plan[global_rank].append(metadata)
with open(self.plan_output_dir, "w") as f:
json.dump(plan, f)
if validation:
assert self.neg_metadata is not None
plan["negative_prompt"] = [self.neg_metadata]
with open(self.plan_output_dir, "w") as f:
json.dump(plan, f)
else:
pass
dist.barrier()
if validation:
with open(self.plan_output_dir) as f:
plan = json.load(f)
self.neg_metadata = plan["negative_prompt"][0]
def _load_and_cache_negative_prompt(self) -> None:
"""Load and cache the negative prompt. Only rank 0 in each SP group should call this."""
if not self.validation or self.neg_metadata is None:
return
if self.cached_neg_prompt is not None:
return
# Only rank 0 in each SP group should read the negative prompt
try:
file_path, row_idx = self.neg_metadata
parquet_file = pq.ParquetFile(file_path)
# Since negative prompt is always the first row (row_idx = 0),
# it's always in the first row group
row_group_index = 0
local_index = row_idx # This will be 0 for the negative prompt
row_group = parquet_file.read_row_group(row_group_index).to_pydict()
row_dict = {k: v[local_index] for k, v in row_group.items()}
del row_group
# Process the negative prompt row
self.cached_neg_prompt = self._process_row(row_dict)
except Exception as e:
logger.error("Failed to load negative prompt: %s", e)
self.cached_neg_prompt = None
def get_validation_negative_prompt(
self
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, Dict[str, Any]]:
"""
Get the negative prompt for validation.
This method ensures the negative prompt is loaded and cached properly.
Returns the processed negative prompt data (latents, embeddings, masks, info).
"""
if not self.validation:
raise ValueError(
"get_validation_negative_prompt() can only be called in validation mode"
)
# Load and cache if needed (only rank 0 in SP group will actually load)
if self.cached_neg_prompt is None:
self._load_and_cache_negative_prompt()
if self.cached_neg_prompt is None:
raise RuntimeError(
f"Rank {self.rank} (SP rank {self.local_rank}): Could not retrieve negative prompt data"
)
# Extract the components
lat, emb, mask, info = (self.cached_neg_prompt["latents"],
self.cached_neg_prompt["embeddings"],
self.cached_neg_prompt["masks"],
self.cached_neg_prompt["info"])
# Apply the same processing as in __getitem__
if lat.numel() == 0: # Validation parquet
return lat, emb, mask, info
else:
lat = lat[:, -self.num_latent_t:]
if self.sp_world_size > 1:
lat = rearrange(lat,
"t (n s) h w -> t n s h w",
n=self.sp_world_size).contiguous()
lat = lat[:, self.local_rank, :, :, :]
return lat, emb, mask, info
def __len__(self):
if self.local_indices is None:
@@ -118,9 +219,9 @@ class ParquetVideoTextDataset(Dataset):
cumulative = 0
for i in range(parquet_file.num_row_groups):
num_rows = parquet_file.metadata.row_group(i).num_rows
if cumulative + num_rows > idx:
if cumulative + num_rows > row_idx:
row_group_index = i
local_index = idx - cumulative
local_index = row_idx - cumulative
break
cumulative += num_rows
+3 -1
View File
@@ -138,6 +138,7 @@ class T2V_dataset(Dataset):
video_path = dataset_prog.cap_list[idx]["path"]
assert os.path.exists(video_path), f"file {video_path} do not exist!"
frame_indices = dataset_prog.cap_list[idx]["sample_frame_index"]
torchvision_video, _, metadata = torchvision.io.read_video(
video_path, output_format="TCHW")
video = torchvision_video[frame_indices]
@@ -270,7 +271,8 @@ class T2V_dataset(Dataset):
cnt_resolution_mismatch += 1
continue
# import ipdb;ipdb.set_trace()
# if path == 'finetrainers/3dgs-dissolve/videos/1.mp4':
# from IPython import embed; embed()
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 * (
+8 -2
View File
@@ -2,8 +2,10 @@
from fastvideo.v1.distributed.communication_op import *
from fastvideo.v1.distributed.parallel_state import (
cleanup_dist_env_and_memory, get_sequence_model_parallel_rank,
get_sequence_model_parallel_world_size, get_tensor_model_parallel_rank,
cleanup_dist_env_and_memory, get_data_parallel_rank,
get_data_parallel_world_size, get_dp_group,
get_sequence_model_parallel_rank, get_sequence_model_parallel_world_size,
get_sp_group, get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size, get_world_group,
init_distributed_environment, initialize_model_parallel,
model_parallel_is_initialized)
@@ -12,11 +14,15 @@ from fastvideo.v1.distributed.utils import *
__all__ = [
"init_distributed_environment",
"initialize_model_parallel",
"get_data_parallel_world_size",
"get_data_parallel_rank",
"get_sequence_model_parallel_rank",
"get_sequence_model_parallel_world_size",
"get_tensor_model_parallel_rank",
"get_tensor_model_parallel_world_size",
"cleanup_dist_env_and_memory",
"get_world_group",
"get_dp_group",
"get_sp_group",
"model_parallel_is_initialized",
]
+48 -6
View File
@@ -704,7 +704,7 @@ def init_world_group(ranks: List[int], local_rank: int,
group_ranks=[ranks],
local_rank=local_rank,
torch_distributed_backend=backend,
use_device_communicator=False,
use_device_communicator=True,
group_name="world",
)
@@ -747,10 +747,10 @@ def set_custom_all_reduce(enable: bool):
def init_distributed_environment(
world_size: int = -1,
rank: int = -1,
world_size: int = 1,
rank: int = 0,
distributed_init_method: str = "env://",
local_rank: int = -1,
local_rank: int = 0,
backend: str = "nccl",
):
logger.debug(
@@ -794,9 +794,18 @@ def get_sp_group() -> GroupCoordinator:
return _SP
_DP: Optional[GroupCoordinator] = None
def get_dp_group() -> GroupCoordinator:
assert _DP is not None, ("data parallel group is not initialized")
return _DP
def initialize_model_parallel(
tensor_model_parallel_size: int = 1,
sequence_model_parallel_size: int = 1,
data_parallel_size: int = 1,
backend: Optional[str] = None,
) -> None:
"""
@@ -852,6 +861,22 @@ def initialize_model_parallel(
backend,
group_name="sp")
# Build the data parallel groups.
num_data_parallel_groups: int = (world_size // data_parallel_size)
global _DP
assert _DP is None, ("data parallel group is already initialized")
group_ranks = []
for i in range(num_data_parallel_groups):
ranks = list(range(i * data_parallel_size,
(i + 1) * data_parallel_size))
group_ranks.append(ranks)
_DP = init_model_parallel_group(group_ranks,
get_world_group().local_rank,
backend,
group_name="dp")
def get_sequence_model_parallel_world_size() -> int:
"""Return world size for the sequence model parallel group."""
@@ -863,9 +888,20 @@ def get_sequence_model_parallel_rank() -> int:
return get_sp_group().rank_in_group
def get_data_parallel_world_size() -> int:
"""Return world size for the data parallel group."""
return get_dp_group().world_size
def get_data_parallel_rank() -> int:
"""Return my rank for the data parallel group."""
return get_dp_group().rank_in_group
def ensure_model_parallel_initialized(
tensor_model_parallel_size: int,
sequence_model_parallel_size: int,
data_parallel_size: int,
backend: Optional[str] = None,
) -> None:
"""Helper to initialize model parallel groups if they are not initialized,
@@ -876,7 +912,8 @@ def ensure_model_parallel_initialized(
get_world_group().device_group)
if not model_parallel_is_initialized():
initialize_model_parallel(tensor_model_parallel_size,
sequence_model_parallel_size, backend)
sequence_model_parallel_size,
data_parallel_size, backend)
return
assert (
@@ -895,7 +932,7 @@ def ensure_model_parallel_initialized(
def model_parallel_is_initialized() -> bool:
"""Check if tensor, sequence parallel groups are initialized."""
return _TP is not None and _SP is not None
return _TP is not None and _SP is not None and _DP is not None
_TP_STATE_PATCHED = False
@@ -948,6 +985,11 @@ def destroy_model_parallel() -> None:
_SP.destroy()
_SP = None
global _DP
if _DP:
_DP.destroy()
_DP = None
def destroy_distributed_environment() -> None:
global _WORLD
+8 -6
View File
@@ -44,8 +44,8 @@ class VideoGenerator:
Initialize the video generator.
Args:
pipeline: The pipeline to use for inference
fastvideo_args: The inference arguments
executor_class: The executor class to use for inference
"""
self.fastvideo_args = fastvideo_args
self.executor = executor_class(fastvideo_args)
@@ -118,7 +118,6 @@ class VideoGenerator:
# initialize_distributed_and_parallelism(fastvideo_args)
executor_class = Executor.get_class(fastvideo_args)
return cls(
fastvideo_args=fastvideo_args,
executor_class=executor_class,
@@ -276,10 +275,10 @@ class VideoGenerator:
# Save video if requested
if batch.save_video:
save_path = batch.output_path
if save_path:
os.makedirs(os.path.dirname(save_path), exist_ok=True)
video_path = os.path.join(save_path, f"{prompt[:100]}.mp4")
output_path = batch.output_path
if output_path:
os.makedirs(output_path, exist_ok=True)
video_path = os.path.join(output_path, f"{prompt[:100]}.mp4")
imageio.mimsave(video_path, frames, fps=batch.fps, format="mp4")
logger.info("Saved video to %s", video_path)
else:
@@ -295,6 +294,9 @@ class VideoGenerator:
"generation_time": gen_time
}
def set_lora_adapter(self, lora_nickname: str, lora_path: str) -> None:
self.executor.set_lora_adapter(lora_nickname, lora_path)
def shutdown(self):
"""
Shutdown the video generator.
+78 -4
View File
@@ -44,6 +44,8 @@ class FastVideoArgs:
num_gpus: int = 1
tp_size: Optional[int] = None
sp_size: Optional[int] = None
dp_size: int = 1
dp_shards: Optional[int] = None
dist_timeout: Optional[int] = None # timeout for torch.distributed
# Video generation parameters
@@ -55,6 +57,8 @@ class FastVideoArgs:
# DiT configuration
dit_config: DiTConfig = field(default_factory=DiTConfig)
precision: str = "bf16"
use_cpu_offload: bool = True
use_fsdp_inference: bool = True
# VAE configuration
vae_precision: str = "fp16"
@@ -70,7 +74,7 @@ class FastVideoArgs:
# Text encoder configuration
DEFAULT_TEXT_ENCODER_PRECISIONS = (
"fp16",
"fp16",
# "fp16",
)
text_encoder_precisions: Tuple[str, ...] = field(
default_factory=lambda: FastVideoArgs.DEFAULT_TEXT_ENCODER_PRECISIONS)
@@ -82,10 +86,19 @@ class FastVideoArgs:
default_factory=lambda: (postprocess_text, ))
# STA (Spatial-Temporal Attention) parameters
STA_mode: str = "STA_inference"
skip_time_steps: int = 15
# LoRA parameters
lora_path: Optional[str] = None
lora_nickname: Optional[
str] = "default" # for swapping adapters in the pipeline
lora_target_names: Optional[List[
str]] = None # can restrict list of layers to adapt, e.g. ["q_proj"]
# STA parameters
mask_strategy_file_path: Optional[str] = None
enable_torch_compile: bool = False
use_cpu_offload: bool = False
disable_autocast: bool = False
# StepVideo specific parameters
@@ -179,6 +192,20 @@ class FastVideoArgs:
default=FastVideoArgs.sp_size,
help="The sequence parallelism size.",
)
parser.add_argument(
"--data-parallel-size",
"--dp-size",
type=int,
default=FastVideoArgs.dp_size,
help="The data parallelism size.",
)
parser.add_argument(
"--data-parallel-shards",
"--dp-shards",
type=int,
default=FastVideoArgs.dp_shards,
help="The data parallelism shards.",
)
parser.add_argument(
"--dist-timeout",
type=int,
@@ -254,6 +281,21 @@ class FastVideoArgs:
)
# STA (Spatial-Temporal Attention) parameters
parser.add_argument(
"--STA-mode",
type=str,
default=FastVideoArgs.STA_mode,
choices=[
"STA_inference", "STA_searching", "STA_tuning", "STA_tuning_cfg"
],
help="STA mode",
)
parser.add_argument(
"--skip-time-steps",
type=int,
default=FastVideoArgs.skip_time_steps,
help="Number of time steps to warmup (full attention) for STA",
)
parser.add_argument(
"--mask-strategy-file-path",
type=str,
@@ -269,8 +311,16 @@ class FastVideoArgs:
parser.add_argument(
"--use-cpu-offload",
action=StoreBoolean,
help="Use CPU offload for the model load",
help=
"Use CPU offload for model inference. Enable if run out of memory with FSDP.",
)
parser.add_argument(
"--use-fsdp-inference",
action=StoreBoolean,
help=
"Use FSDP for inference by sharding the model weights. Latency is very low due to prefetch--enable if run out of memory.",
)
parser.add_argument(
"--disable-autocast",
action=StoreBoolean,
@@ -332,6 +382,10 @@ class FastVideoArgs:
kwargs[attr] = args.tensor_parallel_size
elif attr == 'sp_size' and hasattr(args, 'sequence_parallel_size'):
kwargs[attr] = args.sequence_parallel_size
elif attr == 'dp_size' and hasattr(args, 'data_parallel_size'):
kwargs[attr] = args.data_parallel_size
elif attr == 'dp_shards' and hasattr(args, 'data_parallel_shards'):
kwargs[attr] = args.data_parallel_shards
elif attr == 'flow_shift' and hasattr(args, 'shift'):
kwargs[attr] = args.shift
# Use getattr with default value from the dataclass for potentially missing attributes
@@ -343,10 +397,20 @@ class FastVideoArgs:
def check_fastvideo_args(self) -> None:
"""Validate inference arguments for consistency"""
if not self.inference_mode:
assert self.dp_size is not None, "dp_size must be set for training"
assert self.dp_shards is not None, "dp_shards must be set for training"
assert self.sp_size is not None, "sp_size must be set for training"
if self.tp_size is None:
self.tp_size = self.num_gpus
if self.sp_size is None:
self.sp_size = self.num_gpus
if self.dp_shards is None:
self.dp_shards = self.num_gpus
assert self.sp_size <= self.num_gpus and self.num_gpus % self.sp_size == 0, "num_gpus must >= and be divisible by sp_size"
assert self.dp_size <= self.num_gpus and self.num_gpus % self.dp_size == 0, "num_gpus must >= and be divisible by dp_size"
assert self.dp_shards <= self.num_gpus and self.num_gpus % self.dp_shards == 0, "num_gpus must >= and be divisible by dp_shards"
if self.num_gpus < max(self.tp_size, self.sp_size):
self.num_gpus = max(self.tp_size, self.sp_size)
@@ -472,7 +536,7 @@ class TrainingArgs(FastVideoArgs):
validation_steps: float = 0.0
log_validation: bool = False
tracker_project_name: str = ""
# seed: int
seed: Optional[int] = None
# output
output_dir: str = ""
@@ -520,6 +584,9 @@ class TrainingArgs(FastVideoArgs):
# master_weight_type
master_weight_type: str = ""
# For fast checking in LoRA pipeline
training_mode: bool = True
@classmethod
def from_cli_args(cls, args: argparse.Namespace) -> "TrainingArgs":
# Get all fields from the dataclass
@@ -535,6 +602,10 @@ class TrainingArgs(FastVideoArgs):
kwargs[attr] = args.sequence_parallel_size
elif attr == 'flow_shift' and hasattr(args, 'shift'):
kwargs[attr] = args.shift
elif attr == 'dp_size' and hasattr(args, 'data_parallel_size'):
kwargs[attr] = args.data_parallel_size
elif attr == 'dp_shards' and hasattr(args, 'data_parallel_shards'):
kwargs[attr] = args.data_parallel_shards
# Use getattr with default value from the dataclass for potentially missing attributes
else:
default_value = getattr(cls, attr, None)
@@ -630,6 +701,9 @@ class TrainingArgs(FastVideoArgs):
parser.add_argument("--tracker-project-name",
type=str,
help="Project name for tracking")
parser.add_argument("--seed",
type=int,
help="Seed for deterministic training")
# Output configuration
parser.add_argument("--output-dir",
-207
View File
@@ -1,207 +0,0 @@
# type: ignore
# SPDX-License-Identifier: Apache-2.0
"""
Inference module for diffusion models.
This module provides classes and functions for running inference with diffusion models.
"""
import time
from typing import Any, Dict
import torch
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines import (ComposedPipelineBase, ForwardBatch,
build_pipeline)
# TODO(will): remove, check if this is hunyuan specific
from fastvideo.v1.utils import align_to
logger = init_logger(__name__)
class InferenceEngine:
"""
Engine for running inference with diffusion models.
"""
def __init__(
self,
pipeline: ComposedPipelineBase,
fastvideo_args: FastVideoArgs,
):
"""
Initialize the inference engine.
Args:
pipeline: The pipeline to use for inference.
fastvideo_args: The inference arguments.
default_negative_prompt: The default negative prompt to use.
"""
self.pipeline = pipeline
self.fastvideo_args = fastvideo_args
@classmethod
def create_engine(
cls,
fastvideo_args: FastVideoArgs,
) -> "InferenceEngine":
"""
Create an inference engine with the specified arguments.
Args:
fastvideo_args: The inference arguments.
model_loader_cls: The model loader class to use. If None, it will be
determined from the model type.
pipeline_type: The type of pipeline to create. If None, it will be
determined from the model type.
Returns:
The created inference engine.
Raises:
ValueError: If the model type is not recognized or if the pipeline type
is not recognized.
"""
logger.info("Building pipeline...")
# TODO(will): I don't really like this api.
# it should be something closer to pipeline_cls.from_pretrained(...)
# this way for training we can just do pipeline_cls.from_pretrained(
# checkpoint_path) and have it handle everything.
# TODO(Peiyuan): Then maybe we should only pass in model path and device, not the entire inference args?
pipeline = build_pipeline(fastvideo_args)
logger.info("Pipeline Ready")
# Create the inference engine
return cls(pipeline, fastvideo_args)
def run(
self,
prompt: str,
fastvideo_args: FastVideoArgs,
) -> Dict[str, Any]:
"""
Run inference with the pipeline.
Args:
prompt: The prompt to use for generation.
negative_prompt: The negative prompt to use. If None, the default will be used.
seed: The random seed to use. If None, a random seed will be used.
**kwargs: Additional arguments to pass to the pipeline.
Returns:
A dictionary containing the generated videos and metadata.
"""
out_dict: Dict[str, Any] = dict()
num_videos_per_prompt = fastvideo_args.num_videos
seed = fastvideo_args.seed
height = fastvideo_args.height
width = fastvideo_args.width
video_length = fastvideo_args.num_frames
negative_prompt = fastvideo_args.neg_prompt
infer_steps = fastvideo_args.num_inference_steps
guidance_scale = fastvideo_args.guidance_scale
flow_shift = fastvideo_args.flow_shift
embedded_guidance_scale = fastvideo_args.embedded_cfg_scale
image_path = fastvideo_args.image_path
# ========================================================================
# Arguments: target_width, target_height, target_video_length
# ========================================================================
if width <= 0 or height <= 0 or video_length <= 0:
raise ValueError(
f"`height` and `width` and `video_length` must be positive integers, got height={height}, width={width}, video_length={video_length}"
)
if (video_length - 1) % 4 != 0:
raise ValueError(
f"`video_length-1` must be a multiple of 4, got {video_length}")
target_height = align_to(height, 16)
target_width = align_to(width, 16)
target_video_length = video_length
out_dict["size"] = (target_height, target_width, target_video_length)
# ========================================================================
# Arguments: prompt, new_prompt, negative_prompt
# ========================================================================
if not isinstance(prompt, str):
raise TypeError(
f"`prompt` must be a string, but got {type(prompt)}")
prompt = prompt.strip()
# negative prompt
if negative_prompt is not None:
negative_prompt = negative_prompt.strip()
# TODO(PY): move to hunyuan stage
latents_size = [(video_length - 1) // 4 + 1, height // 8, width // 8]
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
# ========================================================================
# Print infer args
# ========================================================================
debug_str = f"""
height: {target_height}
width: {target_width}
video_length: {target_video_length}
prompt: {prompt}
neg_prompt: {negative_prompt}
seed: {seed}
infer_steps: {infer_steps}
num_videos_per_prompt: {num_videos_per_prompt}
guidance_scale: {guidance_scale}
n_tokens: {n_tokens}
flow_shift: {flow_shift}
embedded_guidance_scale: {embedded_guidance_scale}"""
logger.info(debug_str)
# return
# sp_group = get_sp_group()
# local_rank = sp_group.rank
device = torch.device(fastvideo_args.device_str)
batch = ForwardBatch(
image_path=image_path,
prompt=prompt,
negative_prompt=negative_prompt,
num_videos_per_prompt=num_videos_per_prompt,
height=fastvideo_args.height,
width=fastvideo_args.width,
num_frames=fastvideo_args.num_frames,
num_inference_steps=fastvideo_args.num_inference_steps,
guidance_scale=fastvideo_args.guidance_scale,
# generator=generator,
eta=0.0,
n_tokens=n_tokens,
data_type="video" if fastvideo_args.num_frames > 1 else "image",
device=device,
extra={}, # Any additional parameters
)
print('===============================================')
print(batch)
print('===============================================')
print('===============================================')
print(fastvideo_args)
# ========================================================================
# Pipeline inference
# ========================================================================
start_time = time.perf_counter()
samples = self.pipeline.forward(
batch=batch,
fastvideo_args=fastvideo_args,
).output
# TODO(will): fix and move to hunyuan stage
# out_dict["seeds"] = batch.seeds
out_dict["samples"] = samples
out_dict["prompts"] = prompt
gen_time = time.perf_counter() - start_time
logger.info("Success, time: %s", gen_time)
return out_dict
+295
View File
@@ -0,0 +1,295 @@
# Code adapted from SGLang https://github.com/sgl-project/sglang/blob/main/python/sglang/srt/lora/layers.py
from typing import Dict, List, Tuple, Type, Union
import torch
from torch import nn
from torch.distributed.tensor import DTensor, distribute_tensor
from fastvideo.v1.distributed import (get_tensor_model_parallel_rank,
split_tensor_along_last_dim,
tensor_model_parallel_all_gather,
tensor_model_parallel_all_reduce)
from fastvideo.v1.layers.linear import (ColumnParallelLinear, LinearBase,
MergedColumnParallelLinear,
QKVParallelLinear, ReplicatedLinear,
RowParallelLinear)
from fastvideo.v1.layers.vocab_parallel_embedding import VocabParallelEmbedding
class BaseLayerWithLoRA(nn.Module):
def __init__(
self,
base_layer: nn.Module,
):
super().__init__()
self.base_layer: nn.Module = base_layer
self.lora_A: torch.Tensor = None
self.lora_B: torch.Tensor = None
self.merged: bool = False
self.weight = base_layer.weight
self.cpu_weight = base_layer.weight.to("cpu")
self.unmerge_count = 0
# indicates adapter weights don't contain this layer
# (which shouldn't normally happen, but we want to separate it from the case of erroneous merging)
self.disable_lora: bool = False
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.base_layer.forward(x)
def slice_lora_a_weights(self, A: torch.Tensor) -> torch.Tensor:
return A
def slice_lora_b_weights(self, B: torch.Tensor) -> torch.Tensor:
return B
def set_lora_weights(self,
A: torch.Tensor,
B: torch.Tensor,
training_mode: bool = False) -> None:
self.lora_A = A # share storage with weights in the pipeline
self.lora_B = B
self.disable_lora = False
if not training_mode:
self.merge_lora_weights()
@torch.no_grad()
def merge_lora_weights(self) -> None:
if self.disable_lora:
return
if self.merged:
raise ValueError(
"LoRA weights already merged. Please unmerge them first.")
assert self.lora_A is not None and self.lora_B is not None, "LoRA weights not set. Please set them first."
if isinstance(self.base_layer.weight, DTensor):
mesh = self.base_layer.weight.data.device_mesh
placements = self.base_layer.weight.data.placements
current_device = self.base_layer.weight.data.device
data = self.base_layer.weight.data.to(
f"cuda:{torch.cuda.current_device()}").full_tensor()
data += (self.slice_lora_b_weights(self.lora_B)
@ self.slice_lora_a_weights(self.lora_A)).to(data)
self.base_layer.weight.data = distribute_tensor(
data, mesh, placements=placements).to(current_device)
else:
current_device = self.base_layer.weight.data.device
data = self.base_layer.weight.data.to(
f"cuda:{torch.cuda.current_device()}")
data += \
(self.slice_lora_b_weights(self.lora_B) @ self.slice_lora_a_weights(self.lora_A)).to(data)
self.base_layer.weight.data = data.to(current_device)
self.merged = True
@torch.no_grad()
def unmerge_lora_weights(self) -> None:
if self.disable_lora:
return
if not self.merged:
raise ValueError(
"LoRA weights not merged. Please merge them first before unmerging."
)
self.unmerge_count += 1
# Avoid precision loss
if self.unmerge_count % 3 == 0:
self.base_layer.weight.data = self.cpu_weight.data.to(
self.base_layer.weight)
if isinstance(self.base_layer.weight, DTensor):
mesh = self.base_layer.weight.data.device_mesh
placement = self.base_layer.weight.data.placements
device = self.base_layer.weight.data.device
data = self.base_layer.weight.data.to(
f"cuda:{torch.cuda.current_device()}").full_tensor()
data -= self.slice_lora_b_weights(
self.lora_B) @ self.slice_lora_a_weights(self.lora_A)
self.base_layer.weight.data = distribute_tensor(
data, mesh, placements=placement).to(device)
else:
self.base_layer.weight.data -= \
self.slice_lora_b_weights(self.lora_B) @\
self.slice_lora_a_weights(self.lora_A)
self.merged = False
class VocabParallelEmbeddingWithLoRA(BaseLayerWithLoRA):
"""
Vocab parallel embedding layer with support for LoRA (Low-Rank Adaptation).
Note: The current version does not yet implement the LoRA functionality.
This class behaves exactly the same as the base VocabParallelEmbedding.
Future versions will integrate LoRA functionality to support efficient parameter fine-tuning.
"""
def __init__(
self,
base_layer: VocabParallelEmbedding,
) -> None:
super().__init__(base_layer)
def forward(self, input_: torch.Tensor) -> torch.Tensor:
raise NotImplementedError(
"We don't support VocabParallelEmbeddingWithLoRA yet.")
class ColumnParallelLinearWithLoRA(BaseLayerWithLoRA):
def __init__(
self,
base_layer: ColumnParallelLinear,
) -> None:
super().__init__(base_layer)
def forward(self, input_: torch.Tensor) -> torch.Tensor:
# duplicate the logic in ColumnParallelLinear
bias = self.base_layer.bias if not self.base_layer.skip_bias_add else None
output_parallel = self.base_layer.quant_method.apply(
self.base_layer, input_, bias)
if self.base_layer.gather_output:
output = tensor_model_parallel_all_gather(output_parallel)
else:
output = output_parallel
output_bias = self.base_layer.bias if self.base_layer.skip_bias_add else None
return output, output_bias
def slice_lora_a_weights(self, A: torch.Tensor) -> torch.Tensor:
return A
def slice_lora_b_weights(self, B: torch.Tensor) -> torch.Tensor:
tp_rank = get_tensor_model_parallel_rank()
shard_size = self.base_layer.output_partition_sizes[0]
start_idx = tp_rank * shard_size
end_idx = (tp_rank + 1) * shard_size
B = B[start_idx:end_idx, :]
return B
class MergedColumnParallelLinearWithLoRA(ColumnParallelLinearWithLoRA):
def __init__(
self,
base_layer: MergedColumnParallelLinear,
) -> None:
super().__init__(base_layer)
def slice_lora_a_weights(self, A: torch.Tensor) -> torch.Tensor:
return A.to(self.base_layer.weight)
def slice_lora_b_weights(self, B: torch.Tensor) -> torch.Tensor:
tp_rank = get_tensor_model_parallel_rank()
# Since the outputs for both gate and up are identical, we use a random one.
shard_size = self.base_layer.output_partition_sizes[0]
start_idx = tp_rank * shard_size
end_idx = (tp_rank + 1) * shard_size
return B[:, start_idx:end_idx, :]
class QKVParallelLinearWithLoRA(ColumnParallelLinearWithLoRA):
def __init__(
self,
base_layer: QKVParallelLinear,
) -> None:
super().__init__(base_layer)
def slice_lora_a_weights(self, A: torch.Tensor) -> torch.Tensor:
return A
def slice_lora_b_weights(
self, B: List[torch.Tensor]) -> Tuple[torch.Tensor, torch.Tensor]:
tp_rank = get_tensor_model_parallel_rank()
B_q, B_kv = B
base_layer = self.base_layer
q_proj_shard_size = base_layer.q_proj_shard_size
kv_proj_shard_size = base_layer.kv_proj_shard_size
num_kv_head_replicas = base_layer.num_kv_head_replicas
q_start_idx = q_proj_shard_size * tp_rank
q_end_idx = q_start_idx + q_proj_shard_size
kv_shard_id = tp_rank // num_kv_head_replicas
kv_start_idx = kv_proj_shard_size * kv_shard_id
kv_end_idx = kv_start_idx + kv_proj_shard_size
return B_q[q_start_idx:q_end_idx, :], B_kv[:,
kv_start_idx:kv_end_idx, :]
class RowParallelLinearWithLoRA(BaseLayerWithLoRA):
def __init__(
self,
base_layer: RowParallelLinear,
) -> None:
super().__init__(base_layer)
def forward(self, input_: torch.Tensor):
# duplicate the logic in RowParallelLinear
if self.base_layer.input_is_parallel:
input_parallel = input_
else:
tp_rank = get_tensor_model_parallel_rank()
splitted_input = split_tensor_along_last_dim(
input_, num_partitions=self.base_layer.tp_size)
input_parallel = splitted_input[tp_rank].contiguous()
output_parallel = self.base_layer.quant_method.apply(
self.base_layer, input_parallel)
if self.set_lora:
output_parallel = self.apply_lora(output_parallel, input_parallel)
if self.base_layer.reduce_results and self.base_layer.tp_size > 1:
output_ = tensor_model_parallel_all_reduce(output_parallel)
else:
output_ = output_parallel
if not self.base_layer.skip_bias_add:
output = (output_ + self.base_layer.bias
if self.base_layer.bias is not None else output_)
output_bias = None
else:
output = output_
output_bias = self.base_layer.bias
return output, output_bias
def slice_lora_a_weights(self, A: torch.Tensor) -> torch.Tensor:
tp_rank = get_tensor_model_parallel_rank()
shard_size = self.base_layer.input_size_per_partition
start_idx = tp_rank * shard_size
end_idx = (tp_rank + 1) * shard_size
A = A[:, start_idx:end_idx].contiguous()
return A
def slice_lora_b_weights(self, B: torch.Tensor) -> torch.Tensor:
return B
def get_lora_layer(layer: nn.Module) -> Union[BaseLayerWithLoRA, None]:
supported_layer_types: Dict[Type[LinearBase], Type[BaseLayerWithLoRA]] = {
# the order matters
# VocabParallelEmbedding: VocabParallelEmbeddingWithLoRA,
QKVParallelLinear: QKVParallelLinearWithLoRA,
MergedColumnParallelLinear: MergedColumnParallelLinearWithLoRA,
ColumnParallelLinear: ColumnParallelLinearWithLoRA,
RowParallelLinear: RowParallelLinearWithLoRA,
ReplicatedLinear: BaseLayerWithLoRA,
}
for src_layer_type, lora_layer_type in supported_layer_types.items():
if isinstance(layer, src_layer_type): # pylint: disable=unidiomatic-typecheck
ret = lora_layer_type(layer)
return ret
return None
# source: https://github.com/vllm-project/vllm/blob/93b38bea5dd03e1b140ca997dfaadef86f8f1855/vllm/lora/utils.py#L9
def replace_submodule(model: nn.Module, module_name: str,
new_module: nn.Module) -> nn.Module:
"""Replace a submodule in a model with a new module."""
parent = model.get_submodule(".".join(module_name.split(".")[:-1]))
target_name = module_name.split(".")[-1]
setattr(parent, target_name, new_module)
return new_module
+1
View File
@@ -78,6 +78,7 @@ class CachableDiT(BaseDiT):
# These are required class attributes that should be overridden by concrete implementations
_fsdp_shard_conditions = []
_param_names_mapping = {}
_lora_param_names_mapping: dict = {}
# Ensure these instance attributes are properly defined in subclasses
hidden_size: int
num_attention_heads: int
+1
View File
@@ -441,6 +441,7 @@ class HunyuanVideoTransformer3DModel(CachableDiT):
_supported_attention_backends = HunyuanVideoConfig(
)._supported_attention_backends
_param_names_mapping = HunyuanVideoConfig()._param_names_mapping
_lora_param_names_mapping = HunyuanVideoConfig()._lora_param_names_mapping
def __init__(self, config: HunyuanVideoConfig, hf_config: dict[str, Any]):
super().__init__(config=config, hf_config=hf_config)
+1
View File
@@ -459,6 +459,7 @@ class StepVideoModel(BaseDiT):
# lambda n, m: "pos_embed" in n # If needed for the patch embedding.
]
_param_names_mapping = StepVideoConfig()._param_names_mapping
_lora_param_names_mapping = StepVideoConfig()._lora_param_names_mapping
_supported_attention_backends = StepVideoConfig(
)._supported_attention_backends
+2
View File
@@ -319,6 +319,7 @@ class WanTransformerBlock(nn.Module):
query = query.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
value = value.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
# Apply rotary embeddings
cos, sin = freqs_cis
query, key = _apply_rotary_emb(query, cos, sin,
@@ -359,6 +360,7 @@ class WanTransformer3DModel(CachableDiT):
_supported_attention_backends = WanVideoConfig(
)._supported_attention_backends
_param_names_mapping = WanVideoConfig()._param_names_mapping
_lora_param_names_mapping = WanVideoConfig()._lora_param_names_mapping
def __init__(self, config: WanVideoConfig, hf_config: dict[str,
Any]) -> None:
+15 -5
View File
@@ -15,10 +15,10 @@ from safetensors.torch import load_file as safetensors_load_file
from transformers import AutoImageProcessor, AutoTokenizer
from transformers.utils import SAFE_WEIGHTS_INDEX_NAME
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.hf_transformer_utils import get_diffusers_config
from fastvideo.v1.models.loader.fsdp_load import load_fsdp_model
from fastvideo.v1.models.loader.fsdp_load import maybe_load_fsdp_model
from fastvideo.v1.models.loader.utils import set_default_torch_dtype
from fastvideo.v1.models.loader.weight_utils import (
filter_duplicate_safetensors_files, filter_files_not_needed_for_inference,
@@ -391,12 +391,19 @@ class TransformerLoader(ComponentLoader):
len(safetensors_list), model_path)
# initialize_sequence_parallel_group(fastvideo_args.sp_size)
default_dtype = PRECISION_TO_TYPE[fastvideo_args.precision]
if fastvideo_args.training_mode:
assert isinstance(
fastvideo_args, TrainingArgs
), "fastvideo_args must be a TrainingArgs object when training_mode is True"
default_dtype = PRECISION_TO_TYPE[fastvideo_args.master_weight_type]
else:
default_dtype = PRECISION_TO_TYPE[fastvideo_args.precision]
# Load the model using FSDP loader
logger.info("Loading model from %s, default_dtype: %s", cls_name,
default_dtype)
model = load_fsdp_model(
assert fastvideo_args.dp_shards is not None
model = maybe_load_fsdp_model(
model_cls=model_cls,
init_params={
"config": dit_config,
@@ -404,13 +411,16 @@ class TransformerLoader(ComponentLoader):
},
weight_dir_list=safetensors_list,
device=fastvideo_args.device,
data_parallel_size=fastvideo_args.dp_size,
data_parallel_shards=fastvideo_args.dp_shards,
cpu_offload=fastvideo_args.use_cpu_offload,
fsdp_inference=fastvideo_args.use_fsdp_inference,
default_dtype=default_dtype,
# TODO(will): make these configurable
param_dtype=torch.bfloat16,
reduce_dtype=torch.float32,
output_dtype=None,
)
training_mode=fastvideo_args.training_mode)
if fastvideo_args.enable_torch_compile:
logger.info("Torch Compile enabled for DiT")
for n, m in reversed(list(model.named_modules())):
+56 -69
View File
@@ -5,24 +5,23 @@
# Copyright 2025 The FastVideo Authors.
import contextlib
import re
from collections import defaultdict
from itertools import chain
from typing import (Any, Callable, DefaultDict, Dict, Generator, Hashable, List,
Optional, Tuple, Type)
from typing import (Any, Callable, DefaultDict, Dict, Generator, List, Optional,
Tuple, Type, Union)
import torch
from torch import nn
from torch.distributed import DeviceMesh, init_device_mesh
from torch.distributed._tensor import distribute_tensor
from torch.distributed.fsdp import (CPUOffloadPolicy, MixedPrecisionPolicy,
fully_shard)
from torch.distributed.fsdp import (CPUOffloadPolicy, FSDPModule,
MixedPrecisionPolicy, fully_shard)
from torch.nn.modules.module import _IncompatibleKeys
from fastvideo.v1.distributed.parallel_state import (
get_sequence_model_parallel_world_size)
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.loader.utils import get_param_names_mapping
from fastvideo.v1.models.loader.weight_utils import safetensors_weights_iterator
from fastvideo.v1.utils import set_mixed_precision_policy
logger = init_logger(__name__)
@@ -55,77 +54,60 @@ def set_default_dtype(dtype: torch.dtype) -> Generator[None, None, None]:
torch.set_default_dtype(old_dtype)
def get_param_names_mapping(
mapping_dict: Dict[str, str]) -> Callable[[str], tuple[str, Any, Any]]:
"""
Creates a mapping function that transforms parameter names using regex patterns.
Args:
mapping_dict (Dict[str, str]): Dictionary mapping regex patterns to replacement patterns
param_name (str): The parameter name to be transformed
Returns:
Callable[[str], str]: A function that maps parameter names from source to target format
"""
def mapping_fn(name: str) -> tuple[str, Any, Any]:
# Try to match and transform the name using the regex patterns in mapping_dict
for pattern, replacement in mapping_dict.items():
match = re.match(pattern, name)
if match:
merge_index = None
total_splitted_params = None
if isinstance(replacement, tuple):
merge_index = replacement[1]
total_splitted_params = replacement[2]
replacement = replacement[0]
name = re.sub(pattern, replacement, name)
return name, merge_index, total_splitted_params
# If no pattern matches, return the original name
return name, None, None
return mapping_fn
# TODO(PY): add compile option
def load_fsdp_model(
def maybe_load_fsdp_model(
model_cls: Type[nn.Module],
init_params: Dict[str, Any],
weight_dir_list: List[str],
device: torch.device,
data_parallel_size: int,
data_parallel_shards: int,
default_dtype: torch.dtype,
param_dtype: torch.dtype,
reduce_dtype: torch.dtype,
cpu_offload: bool = False,
fsdp_inference: bool = False,
output_dtype: Optional[torch.dtype] = None,
training_mode: bool = True,
) -> torch.nn.Module:
"""
Load the model with FSDP if is training, else load the model without FSDP.
"""
# NOTE(will): cast_forward_inputs=True shouldn't be needed as we are
# manually casting the inputs to the model
mp_policy = MixedPrecisionPolicy(param_dtype,
reduce_dtype,
output_dtype,
cast_forward_inputs=True)
cast_forward_inputs=False)
set_mixed_precision_policy(master_dtype=default_dtype,
param_dtype=param_dtype,
reduce_dtype=reduce_dtype,
output_dtype=output_dtype)
with set_default_dtype(default_dtype), torch.device("meta"):
model = model_cls(**init_params)
dp_size = data_parallel_size if fsdp_inference or training_mode else 1
device_mesh = init_device_mesh(
"cuda",
mesh_shape=(get_sequence_model_parallel_world_size(), ),
mesh_dim_names=("dp", ),
# (Replicate(), Shard(dim=0))
mesh_shape=(dp_size, data_parallel_shards),
mesh_dim_names=("dp", "sp"),
)
shard_model(model,
cpu_offload=cpu_offload,
reshard_after_forward=True,
mp_policy=mp_policy,
dp_mesh=device_mesh["dp"])
mesh=device_mesh)
weight_iterator = safetensors_weights_iterator(weight_dir_list)
param_names_mapping_fn = get_param_names_mapping(model._param_names_mapping)
load_fsdp_model_from_full_model_state_dict(
load_model_from_full_model_state_dict(
model,
weight_iterator,
device,
param_dtype,
strict=True,
cpu_offload=cpu_offload,
param_names_mapping=param_names_mapping_fn,
@@ -134,8 +116,9 @@ def load_fsdp_model(
if p.is_meta:
raise RuntimeError(
f"Unexpected param or buffer {n} on meta device.")
for p in model.parameters():
p.requires_grad = False
if isinstance(p, torch.nn.Parameter):
p.requires_grad = False
return model
@@ -146,6 +129,7 @@ def shard_model(
reshard_after_forward: bool = True,
mp_policy: Optional[MixedPrecisionPolicy] = None,
dp_mesh: Optional[DeviceMesh] = None,
mesh: Optional[DeviceMesh] = None,
) -> None:
"""
Utility to shard a model with FSDP using the PyTorch Distributed fully_shard API.
@@ -172,7 +156,7 @@ def shard_model(
"""
fsdp_kwargs = {
"reshard_after_forward": reshard_after_forward,
"mesh": dp_mesh,
"mesh": mesh,
"mp_policy": mp_policy,
}
if cpu_offload:
@@ -201,25 +185,28 @@ def shard_model(
# TODO(PY): device mesh for cfg parallel
def load_fsdp_model_from_full_model_state_dict(
model: torch.nn.Module,
def load_model_from_full_model_state_dict(
model: Union[FSDPModule, torch.nn.Module],
full_sd_iterator: Generator[Tuple[str, torch.Tensor], None, None],
device: torch.device,
param_dtype: torch.dtype,
strict: bool = False,
cpu_offload: bool = False,
param_names_mapping: Optional[Callable[[str], tuple[str, Any, Any]]] = None,
training_mode: bool = True,
) -> _IncompatibleKeys:
"""
Converting full state dict into a sharded state dict
and loading it into FSDP model
and loading it into FSDP model (if training) or normal huggingface model
Args:
model (FSDPModule): Model to generate fully qualified names for cpu_state_dict
model (Union[FSDPModule, torch.nn.Module]): Model to generate fully qualified names for cpu_state_dict
full_sd_iterator (Generator): an iterator yielding (param_name, tensor) pairs
device (torch.device): device used to move full state dict tensors
param_dtype (torch.dtype): dtype used to move full state dict tensors
strict (bool): flag to check if to load the model in strict mode
cpu_offload (bool): flag to check if offload to CPU is enabled
cpu_offload (bool): flag to check if FSDP offload is enabled
param_names_mapping (Optional[Callable[[str], str]]): a function that maps full param name to sharded param name
training_mode (bool): apply FSDP only for training
Returns:
``NamedTuple`` with ``missing_keys`` and ``unexpected_keys`` fields:
* **missing_keys** is a list of str containing the missing keys
@@ -228,19 +215,18 @@ def load_fsdp_model_from_full_model_state_dict(
Raises:
NotImplementedError: If got FSDP with more than 1D.
"""
meta_sharded_sd = model.state_dict()
meta_sd = model.state_dict()
sharded_sd = {}
to_merge_params: DefaultDict[Hashable, Dict[Any, Any]] = defaultdict(dict)
to_merge_params: DefaultDict[str, Dict[Any, Any]] = defaultdict(dict)
for source_param_name, full_tensor in full_sd_iterator:
assert param_names_mapping is not None
target_param_name, merge_index, num_params_to_merge = param_names_mapping(
source_param_name)
if merge_index is not None:
to_merge_params[target_param_name][merge_index] = full_tensor
if len(to_merge_params[target_param_name]) == num_params_to_merge:
# cat at dim=1 according to the merge_index order
# cat at output dim according to the merge_index order
sorted_tensors = [
to_merge_params[target_param_name][i]
for i in range(num_params_to_merge)
@@ -250,24 +236,25 @@ def load_fsdp_model_from_full_model_state_dict(
else:
continue
sharded_meta_param = meta_sharded_sd.get(target_param_name)
if sharded_meta_param is None:
meta_sharded_param = meta_sd.get(target_param_name)
if meta_sharded_param is None:
raise ValueError(
f"Parameter {source_param_name}-->{target_param_name} not found in meta sharded state dict"
)
full_tensor = full_tensor.to(sharded_meta_param.dtype).to(device)
if not hasattr(sharded_meta_param, "device_mesh"):
if not hasattr(meta_sharded_param, "device_mesh"):
full_tensor = full_tensor.to(device=device, dtype=param_dtype)
# In cases where parts of the model aren't sharded, some parameters will be plain tensors
sharded_tensor = full_tensor
else:
full_tensor = full_tensor.to(device=device, dtype=param_dtype)
sharded_tensor = distribute_tensor(
full_tensor,
sharded_meta_param.device_mesh,
sharded_meta_param.placements,
meta_sharded_param.device_mesh,
meta_sharded_param.placements,
)
if cpu_offload:
sharded_tensor = sharded_tensor.cpu()
if cpu_offload:
sharded_tensor = sharded_tensor.cpu()
sharded_sd[target_param_name] = nn.Parameter(sharded_tensor)
# choose `assign=True` since we cannot call `copy_` on meta tensor
return model.load_state_dict(sharded_sd, strict=strict, assign=True)
+36
View File
@@ -1,6 +1,8 @@
# SPDX-License-Identifier: Apache-2.0
"""Utilities for selecting and loading models."""
import contextlib
import re
from typing import Any, Callable, Dict
import torch
@@ -16,3 +18,37 @@ def set_default_torch_dtype(dtype: torch.dtype):
torch.set_default_dtype(dtype)
yield
torch.set_default_dtype(old_dtype)
def get_param_names_mapping(
mapping_dict: Dict[str, str]) -> Callable[[str], tuple[str, Any, Any]]:
"""
Creates a mapping function that transforms parameter names using regex patterns.
Args:
mapping_dict (Dict[str, str]): Dictionary mapping regex patterns to replacement patterns
param_name (str): The parameter name to be transformed
Returns:
Callable[[str], str]: A function that maps parameter names from source to target format
"""
def mapping_fn(name: str) -> tuple[str, Any, Any]:
# Try to match and transform the name using the regex patterns in mapping_dict
for pattern, replacement in mapping_dict.items():
match = re.match(pattern, name)
if match:
merge_index = None
total_splitted_params = None
if isinstance(replacement, tuple):
merge_index = replacement[1]
total_splitted_params = replacement[2]
replacement = replacement[0]
name = re.sub(pattern, replacement, name)
return name, merge_index, total_splitted_params
# If no pattern matches, return the original name
return name, None, None
return mapping_fn
@@ -0,0 +1,817 @@
# Copied from https://github.com/huggingface/diffusers/blob/v0.31.0/src/diffusers/schedulers/scheduling_unipc_multistep.py
# Convert unipc for flow matching
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
import math
from typing import Any, List, Optional, Tuple, Union
import numpy as np
import torch
from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.schedulers.scheduling_utils import (KarrasDiffusionSchedulers,
SchedulerMixin,
SchedulerOutput)
from diffusers.utils import deprecate, is_scipy_available
from fastvideo.v1.models.schedulers.base import BaseScheduler
if is_scipy_available():
pass
class FlowUniPCMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
"""
`UniPCMultistepScheduler` is a training-free framework designed for the fast sampling of diffusion models.
This model inherits from [`SchedulerMixin`] and [`ConfigMixin`]. Check the superclass documentation for the generic
methods the library implements for all schedulers such as loading and saving.
Args:
num_train_timesteps (`int`, defaults to 1000):
The number of diffusion steps to train the model.
solver_order (`int`, default `2`):
The UniPC order which can be any positive integer. The effective order of accuracy is `solver_order + 1`
due to the UniC. It is recommended to use `solver_order=2` for guided sampling, and `solver_order=3` for
unconditional sampling.
prediction_type (`str`, defaults to "flow_prediction"):
Prediction type of the scheduler function; must be `flow_prediction` for this scheduler, which predicts
the flow of the diffusion process.
thresholding (`bool`, defaults to `False`):
Whether to use the "dynamic thresholding" method. This is unsuitable for latent-space diffusion models such
as Stable Diffusion.
dynamic_thresholding_ratio (`float`, defaults to 0.995):
The ratio for the dynamic thresholding method. Valid only when `thresholding=True`.
sample_max_value (`float`, defaults to 1.0):
The threshold value for dynamic thresholding. Valid only when `thresholding=True` and `predict_x0=True`.
predict_x0 (`bool`, defaults to `True`):
Whether to use the updating algorithm on the predicted x0.
solver_type (`str`, default `bh2`):
Solver type for UniPC. It is recommended to use `bh1` for unconditional sampling when steps < 10, and `bh2`
otherwise.
lower_order_final (`bool`, default `True`):
Whether to use lower-order solvers in the final steps. Only valid for < 15 inference steps. This can
stabilize the sampling of DPMSolver for steps < 15, especially for steps <= 10.
disable_corrector (`list`, default `[]`):
Decides which step to disable the corrector to mitigate the misalignment between `epsilon_theta(x_t, c)`
and `epsilon_theta(x_t^c, c)` which can influence convergence for a large guidance scale. Corrector is
usually disabled during the first few steps.
solver_p (`SchedulerMixin`, default `None`):
Any other scheduler that if specified, the algorithm becomes `solver_p + UniC`.
use_karras_sigmas (`bool`, *optional*, defaults to `False`):
Whether to use Karras sigmas for step sizes in the noise schedule during the sampling process. If `True`,
the sigmas are determined according to a sequence of noise levels {σi}.
use_exponential_sigmas (`bool`, *optional*, defaults to `False`):
Whether to use exponential sigmas for step sizes in the noise schedule during the sampling process.
timestep_spacing (`str`, defaults to `"linspace"`):
The way the timesteps should be scaled. Refer to Table 2 of the [Common Diffusion Noise Schedules and
Sample Steps are Flawed](https://huggingface.co/papers/2305.08891) for more information.
steps_offset (`int`, defaults to 0):
An offset added to the inference steps, as required by some model families.
final_sigmas_type (`str`, defaults to `"zero"`):
The final `sigma` value for the noise schedule during the sampling process. If `"sigma_min"`, the final
sigma is the same as the last sigma in the training schedule. If `zero`, the final sigma is set to 0.
"""
_compatibles = [e.name for e in KarrasDiffusionSchedulers]
order = 1
@register_to_config
def __init__(
self,
num_train_timesteps: int = 1000,
solver_order: int = 2,
prediction_type: str = "flow_prediction",
shift: Optional[float] = 1.0,
use_dynamic_shifting=False,
thresholding: bool = False,
dynamic_thresholding_ratio: float = 0.995,
sample_max_value: float = 1.0,
predict_x0: bool = True,
solver_type: str = "bh2",
lower_order_final: bool = True,
disable_corrector: Tuple = (),
solver_p: SchedulerMixin = None,
timestep_spacing: str = "linspace",
steps_offset: int = 0,
final_sigmas_type: Optional[str] = "zero", # "zero", "sigma_min"
**kwargs):
if solver_type not in ["bh1", "bh2"]:
if solver_type in ["midpoint", "heun", "logrho"]:
self.register_to_config(solver_type="bh2")
else:
raise NotImplementedError(
f"{solver_type} is not implemented for {self.__class__}")
self.predict_x0 = predict_x0
# setable values
self.num_inference_steps: Optional[int] = None
alphas = np.linspace(1, 1 / num_train_timesteps,
num_train_timesteps)[::-1].copy()
sigmas = 1.0 - alphas
sigmas = torch.from_numpy(sigmas).to(dtype=torch.float32)
if not use_dynamic_shifting:
# when use_dynamic_shifting is True, we apply the timestep shifting on the fly based on the image resolution
assert shift is not None
sigmas = shift * sigmas / (1 +
(shift - 1) * sigmas) # pyright: ignore
self.sigmas = sigmas
self.timesteps = sigmas * num_train_timesteps
self.model_outputs = [None] * solver_order
self.timestep_list: List[Optional[Any]] = [None] * solver_order
self.lower_order_nums = 0
self.disable_corrector = list(disable_corrector)
self.solver_p = solver_p
self.last_sample = None
self._step_index: Optional[int] = None
self._begin_index: Optional[int] = None
self.sigmas = self.sigmas.to(
"cpu") # to avoid too much CPU/GPU communication
self.sigma_min = self.sigmas[-1].item()
self.sigma_max = self.sigmas[0].item()
@property
def step_index(self):
"""
The index counter for current timestep. It will increase 1 after each scheduler step.
"""
return self._step_index
@property
def begin_index(self):
"""
The index for the first timestep. It should be set from pipeline with `set_begin_index` method.
"""
return self._begin_index
def set_shift(self, shift: float) -> None:
self.config.shift = shift
# Copied from diffusers.schedulers.scheduling_dpmsolver_multistep.DPMSolverMultistepScheduler.set_begin_index
def set_begin_index(self, begin_index: int = 0):
"""
Sets the begin index for the scheduler. This function should be run from pipeline before the inference.
Args:
begin_index (`int`):
The begin index for the scheduler.
"""
self._begin_index = begin_index
# Modified from diffusers.schedulers.scheduling_flow_match_euler_discrete.FlowMatchEulerDiscreteScheduler.set_timesteps
def set_timesteps(
self,
num_inference_steps: Union[int, None] = None,
device: Union[str, torch.device] = None,
sigmas: Optional[List[float]] = None,
mu: Optional[Union[float, None]] = None,
shift: Optional[Union[float, None]] = None,
):
"""
Sets the discrete timesteps used for the diffusion chain (to be run before inference).
Args:
num_inference_steps (`int`):
Total number of the spacing of the time steps.
device (`str` or `torch.device`, *optional*):
The device to which the timesteps should be moved to. If `None`, the timesteps are not moved.
"""
if self.config.use_dynamic_shifting and mu is None:
raise ValueError(
" you have to pass a value for `mu` when `use_dynamic_shifting` is set to be `True`"
)
if sigmas is None:
assert num_inference_steps is not None
sigmas = np.linspace(self.sigma_max, self.sigma_min,
num_inference_steps +
1).copy()[:-1] # pyright: ignore
if self.config.use_dynamic_shifting:
assert mu is not None
sigmas = self.time_shift(mu, 1.0, sigmas) # pyright: ignore
else:
if shift is None:
shift = self.config.shift
assert isinstance(sigmas, np.ndarray)
sigmas = shift * sigmas / (1 +
(shift - 1) * sigmas) # pyright: ignore
if self.config.final_sigmas_type == "sigma_min":
sigma_last = ((1 - self.alphas_cumprod[0]) /
self.alphas_cumprod[0])**0.5
elif self.config.final_sigmas_type == "zero":
sigma_last = 0
else:
raise ValueError(
f"`final_sigmas_type` must be one of 'zero', or 'sigma_min', but got {self.config.final_sigmas_type}"
)
timesteps = sigmas * self.config.num_train_timesteps
sigmas = np.concatenate([sigmas, [sigma_last]
]).astype(np.float32) # pyright: ignore
self.sigmas = torch.from_numpy(sigmas)
self.timesteps = torch.from_numpy(timesteps).to(device=device,
dtype=torch.int64)
self.num_inference_steps = len(timesteps)
self.model_outputs = [
None,
] * self.config.solver_order
self.lower_order_nums = 0
self.last_sample = None
if self.solver_p:
self.solver_p.set_timesteps(self.num_inference_steps, device=device)
# add an index counter for schedulers that allow duplicated timesteps
self._step_index = None
self._begin_index = None
self.sigmas = self.sigmas.to(
"cpu") # to avoid too much CPU/GPU communication
# Copied from diffusers.schedulers.scheduling_ddpm.DDPMScheduler._threshold_sample
def _threshold_sample(self, sample: torch.Tensor) -> torch.Tensor:
"""
"Dynamic thresholding: At each sampling step we set s to a certain percentile absolute pixel value in xt0 (the
prediction of x_0 at timestep t), and if s > 1, then we threshold xt0 to the range [-s, s] and then divide by
s. Dynamic thresholding pushes saturated pixels (those near -1 and 1) inwards, thereby actively preventing
pixels from saturation at each step. We find that dynamic thresholding results in significantly better
photorealism as well as better image-text alignment, especially when using very large guidance weights."
https://arxiv.org/abs/2205.11487
"""
dtype = sample.dtype
batch_size, channels, *remaining_dims = sample.shape
if dtype not in (torch.float32, torch.float64):
sample = sample.float(
) # upcast for quantile calculation, and clamp not implemented for cpu half
# Flatten sample for doing quantile calculation along each image
sample = sample.reshape(batch_size, channels * np.prod(remaining_dims))
abs_sample = sample.abs() # "a certain percentile absolute pixel value"
s = torch.quantile(abs_sample,
self.config.dynamic_thresholding_ratio,
dim=1)
s = torch.clamp(
s, min=1, max=self.config.sample_max_value
) # When clamped to min=1, equivalent to standard clipping to [-1, 1]
s = s.unsqueeze(
1) # (batch_size, 1) because clamp will broadcast along dim=0
sample = torch.clamp(
sample, -s, s
) / s # "we threshold xt0 to the range [-s, s] and then divide by s"
sample = sample.reshape(batch_size, channels, *remaining_dims)
sample = sample.to(dtype)
return sample
# Copied from diffusers.schedulers.scheduling_flow_match_euler_discrete.FlowMatchEulerDiscreteScheduler._sigma_to_t
def _sigma_to_t(self, sigma):
return sigma * self.config.num_train_timesteps
def _sigma_to_alpha_sigma_t(self, sigma) -> Tuple[Any, Any]:
return 1 - sigma, sigma
# Copied from diffusers.schedulers.scheduling_flow_match_euler_discrete.set_timesteps
def time_shift(self, mu: float, sigma: float, t: torch.Tensor):
return math.exp(mu) / (math.exp(mu) + (1 / t - 1)**sigma)
def convert_model_output(
self,
model_output: torch.Tensor,
*args,
sample: torch.Tensor = None,
**kwargs,
) -> torch.Tensor:
r"""
Convert the model output to the corresponding type the UniPC algorithm needs.
Args:
model_output (`torch.Tensor`):
The direct output from the learned diffusion model.
timestep (`int`):
The current discrete timestep in the diffusion chain.
sample (`torch.Tensor`):
A current instance of a sample created by the diffusion process.
Returns:
`torch.Tensor`:
The converted model output.
"""
timestep = args[0] if len(args) > 0 else kwargs.pop("timestep", None)
if sample is None:
if len(args) > 1:
sample = args[1]
else:
raise ValueError(
"missing `sample` as a required keyword argument")
if timestep is not None:
deprecate(
"timesteps",
"1.0.0",
"Passing `timesteps` is deprecated and has no effect as model output conversion is now handled via an internal counter `self.step_index`",
)
sigma = self.sigmas[self.step_index]
alpha_t, sigma_t = self._sigma_to_alpha_sigma_t(sigma)
if self.predict_x0:
if self.config.prediction_type == "flow_prediction":
sigma_t = self.sigmas[self.step_index]
x0_pred = sample - sigma_t * model_output
else:
raise ValueError(
f"prediction_type given as {self.config.prediction_type} must be one of `epsilon`, `sample`,"
" `v_prediction` or `flow_prediction` for the UniPCMultistepScheduler."
)
if self.config.thresholding:
x0_pred = self._threshold_sample(x0_pred)
return x0_pred
else:
if self.config.prediction_type == "flow_prediction":
sigma_t = self.sigmas[self.step_index]
epsilon = sample - (1 - sigma_t) * model_output
else:
raise ValueError(
f"prediction_type given as {self.config.prediction_type} must be one of `epsilon`, `sample`,"
" `v_prediction` or `flow_prediction` for the UniPCMultistepScheduler."
)
if self.config.thresholding:
sigma_t = self.sigmas[self.step_index]
x0_pred = sample - sigma_t * model_output
x0_pred = self._threshold_sample(x0_pred)
epsilon = model_output + x0_pred
return epsilon
def multistep_uni_p_bh_update(
self,
model_output: torch.Tensor,
*args,
sample: torch.Tensor = None,
order: Optional[int] = None, # pyright: ignore
**kwargs,
) -> torch.Tensor:
"""
One step for the UniP (B(h) version). Alternatively, `self.solver_p` is used if is specified.
Args:
model_output (`torch.Tensor`):
The direct output from the learned diffusion model at the current timestep.
prev_timestep (`int`):
The previous discrete timestep in the diffusion chain.
sample (`torch.Tensor`):
A current instance of a sample created by the diffusion process.
order (`int`):
The order of UniP at this timestep (corresponds to the *p* in UniPC-p).
Returns:
`torch.Tensor`:
The sample tensor at the previous timestep.
"""
prev_timestep = args[0] if len(args) > 0 else kwargs.pop(
"prev_timestep", None)
if sample is None:
if len(args) > 1:
sample = args[1]
else:
raise ValueError(
" missing `sample` as a required keyword argument")
if order is None:
if len(args) > 2:
order = args[2]
else:
raise ValueError(
" missing `order` as a required keyword argument")
if prev_timestep is not None:
deprecate(
"prev_timestep",
"1.0.0",
"Passing `prev_timestep` is deprecated and has no effect as model output conversion is now handled via an internal counter `self.step_index`",
)
model_output_list = self.model_outputs
s0 = self.timestep_list[-1]
m0 = model_output_list[-1]
x = sample
if self.solver_p:
x_t = self.solver_p.step(model_output, s0, x).prev_sample
return x_t
sigma_t, sigma_s0 = self.sigmas[self.step_index + 1], self.sigmas[
self.step_index] # pyright: ignore
alpha_t, sigma_t = self._sigma_to_alpha_sigma_t(sigma_t)
alpha_s0, sigma_s0 = self._sigma_to_alpha_sigma_t(sigma_s0)
lambda_t = torch.log(alpha_t) - torch.log(sigma_t)
lambda_s0 = torch.log(alpha_s0) - torch.log(sigma_s0)
h = lambda_t - lambda_s0
device = sample.device
rks = []
D1s: Optional[List[Any]] = []
for i in range(1, order):
si = self.step_index - i # pyright: ignore
mi = model_output_list[-(i + 1)]
alpha_si, sigma_si = self._sigma_to_alpha_sigma_t(self.sigmas[si])
lambda_si = torch.log(alpha_si) - torch.log(sigma_si)
rk = (lambda_si - lambda_s0) / h
rks.append(rk)
assert mi is not None
D1s.append((mi - m0) / rk) # pyright: ignore
rks.append(1.0)
rks = torch.tensor(rks, device=device)
R = []
b = []
hh = -h if self.predict_x0 else h
h_phi_1 = torch.expm1(hh) # h\phi_1(h) = e^h - 1
h_phi_k = h_phi_1 / hh - 1
factorial_i = 1
if self.config.solver_type == "bh1":
B_h = hh
elif self.config.solver_type == "bh2":
B_h = torch.expm1(hh)
else:
raise NotImplementedError()
for i in range(1, order + 1):
R.append(torch.pow(rks, i - 1))
b.append(h_phi_k * factorial_i / B_h)
factorial_i *= i + 1
h_phi_k = h_phi_k / hh - 1 / factorial_i
R = torch.stack(R)
b = torch.tensor(b, device=device)
if D1s is not None and len(D1s) > 0:
D1s = torch.stack(D1s, dim=1) # (B, K)
# for order 2, we use a simplified version
if order == 2:
rhos_p = torch.tensor([0.5], dtype=x.dtype, device=device)
else:
assert isinstance(R, torch.Tensor)
rhos_p = torch.linalg.solve(R[:-1, :-1],
b[:-1]).to(device).to(x.dtype)
else:
D1s = None
if self.predict_x0:
x_t_ = sigma_t / sigma_s0 * x - alpha_t * h_phi_1 * m0
if D1s is not None:
pred_res = torch.einsum("k,bkc...->bc...", rhos_p,
D1s) # pyright: ignore
else:
pred_res = 0
x_t = x_t_ - alpha_t * B_h * pred_res
else:
x_t_ = alpha_t / alpha_s0 * x - sigma_t * h_phi_1 * m0
if D1s is not None:
pred_res = torch.einsum("k,bkc...->bc...", rhos_p,
D1s) # pyright: ignore
else:
pred_res = 0
x_t = x_t_ - sigma_t * B_h * pred_res
x_t = x_t.to(x.dtype)
return x_t
def multistep_uni_c_bh_update(
self,
this_model_output: torch.Tensor,
*args,
last_sample: torch.Tensor = None,
this_sample: torch.Tensor = None,
order: Optional[int] = None, # pyright: ignore
**kwargs,
) -> torch.Tensor:
"""
One step for the UniC (B(h) version).
Args:
this_model_output (`torch.Tensor`):
The model outputs at `x_t`.
this_timestep (`int`):
The current timestep `t`.
last_sample (`torch.Tensor`):
The generated sample before the last predictor `x_{t-1}`.
this_sample (`torch.Tensor`):
The generated sample after the last predictor `x_{t}`.
order (`int`):
The `p` of UniC-p at this step. The effective order of accuracy should be `order + 1`.
Returns:
`torch.Tensor`:
The corrected sample tensor at the current timestep.
"""
this_timestep = args[0] if len(args) > 0 else kwargs.pop(
"this_timestep", None)
if last_sample is None:
if len(args) > 1:
last_sample = args[1]
else:
raise ValueError(
" missing`last_sample` as a required keyword argument")
if this_sample is None:
if len(args) > 2:
this_sample = args[2]
else:
raise ValueError(
" missing`this_sample` as a required keyword argument")
if order is None:
if len(args) > 3:
order = args[3]
else:
raise ValueError(
" missing`order` as a required keyword argument")
if this_timestep is not None:
deprecate(
"this_timestep",
"1.0.0",
"Passing `this_timestep` is deprecated and has no effect as model output conversion is now handled via an internal counter `self.step_index`",
)
model_output_list = self.model_outputs
m0 = model_output_list[-1]
x = last_sample
x_t = this_sample
model_t = this_model_output
sigma_t, sigma_s0 = self.sigmas[self.step_index], self.sigmas[
self.step_index - 1] # pyright: ignore
alpha_t, sigma_t = self._sigma_to_alpha_sigma_t(sigma_t)
alpha_s0, sigma_s0 = self._sigma_to_alpha_sigma_t(sigma_s0)
lambda_t = torch.log(alpha_t) - torch.log(sigma_t)
lambda_s0 = torch.log(alpha_s0) - torch.log(sigma_s0)
h = lambda_t - lambda_s0
device = this_sample.device
rks = []
D1s: Optional[List[Any]] = []
for i in range(1, order):
si = self.step_index - (i + 1) # pyright: ignore
mi = model_output_list[-(i + 1)]
alpha_si, sigma_si = self._sigma_to_alpha_sigma_t(self.sigmas[si])
lambda_si = torch.log(alpha_si) - torch.log(sigma_si)
rk = (lambda_si - lambda_s0) / h
rks.append(rk)
assert mi is not None
D1s.append((mi - m0) / rk) # pyright: ignore
rks.append(1.0)
rks = torch.tensor(rks, device=device)
R = []
b = []
hh = -h if self.predict_x0 else h
h_phi_1 = torch.expm1(hh) # h\phi_1(h) = e^h - 1
h_phi_k = h_phi_1 / hh - 1
factorial_i = 1
if self.config.solver_type == "bh1":
B_h = hh
elif self.config.solver_type == "bh2":
B_h = torch.expm1(hh)
else:
raise NotImplementedError()
for i in range(1, order + 1):
R.append(torch.pow(rks, i - 1))
b.append(h_phi_k * factorial_i / B_h)
factorial_i *= i + 1
h_phi_k = h_phi_k / hh - 1 / factorial_i
R = torch.stack(R)
b = torch.tensor(b, device=device)
if D1s is not None and len(D1s) > 0:
D1s = torch.stack(D1s, dim=1)
else:
D1s = None
# for order 1, we use a simplified version
if order == 1:
rhos_c = torch.tensor([0.5], dtype=x.dtype, device=device)
else:
rhos_c = torch.linalg.solve(R, b).to(device).to(x.dtype)
if self.predict_x0:
x_t_ = sigma_t / sigma_s0 * x - alpha_t * h_phi_1 * m0
if D1s is not None:
corr_res = torch.einsum("k,bkc...->bc...", rhos_c[:-1], D1s)
else:
corr_res = 0
D1_t = model_t - m0
x_t = x_t_ - alpha_t * B_h * (corr_res + rhos_c[-1] * D1_t)
else:
x_t_ = alpha_t / alpha_s0 * x - sigma_t * h_phi_1 * m0
if D1s is not None:
corr_res = torch.einsum("k,bkc...->bc...", rhos_c[:-1], D1s)
else:
corr_res = 0
D1_t = model_t - m0
x_t = x_t_ - sigma_t * B_h * (corr_res + rhos_c[-1] * D1_t)
x_t = x_t.to(x.dtype)
return x_t
def index_for_timestep(self, timestep, schedule_timesteps=None) -> int:
if schedule_timesteps is None:
schedule_timesteps = self.timesteps
indices = (schedule_timesteps == timestep).nonzero()
# The sigma index that is taken for the **very** first `step`
# is always the second index (or the last index if there is only 1)
# This way we can ensure we don't accidentally skip a sigma in
# case we start in the middle of the denoising schedule (e.g. for image-to-image)
pos = 1 if len(indices) > 1 else 0
step_index: int = indices[pos].item()
return step_index
# Copied from diffusers.schedulers.scheduling_dpmsolver_multistep.DPMSolverMultistepScheduler._init_step_index
def _init_step_index(self, timestep) -> None:
"""
Initialize the step_index counter for the scheduler.
"""
if self.begin_index is None:
if isinstance(timestep, torch.Tensor):
timestep = timestep.to(self.timesteps.device)
self._step_index = self.index_for_timestep(timestep)
else:
self._step_index = self._begin_index
def step(self,
model_output: torch.Tensor,
timestep: Union[int, torch.Tensor],
sample: torch.Tensor,
return_dict: bool = True,
generator=None) -> Union[SchedulerOutput, Tuple]:
"""
Predict the sample from the previous timestep by reversing the SDE. This function propagates the sample with
the multistep UniPC.
Args:
model_output (`torch.Tensor`):
The direct output from learned diffusion model.
timestep (`int`):
The current discrete timestep in the diffusion chain.
sample (`torch.Tensor`):
A current instance of a sample created by the diffusion process.
return_dict (`bool`):
Whether or not to return a [`~schedulers.scheduling_utils.SchedulerOutput`] or `tuple`.
Returns:
[`~schedulers.scheduling_utils.SchedulerOutput`] or `tuple`:
If return_dict is `True`, [`~schedulers.scheduling_utils.SchedulerOutput`] is returned, otherwise a
tuple is returned where the first element is the sample tensor.
"""
if self.num_inference_steps is None:
raise ValueError(
"Number of inference steps is 'None', you need to run 'set_timesteps' after creating the scheduler"
)
if self.step_index is None:
self._init_step_index(timestep)
use_corrector = (
self.step_index > 0
and self.step_index - 1 not in self.disable_corrector
and self.last_sample is not None # pyright: ignore
)
model_output_convert = self.convert_model_output(model_output,
sample=sample)
if use_corrector:
sample = self.multistep_uni_c_bh_update(
this_model_output=model_output_convert,
last_sample=self.last_sample,
this_sample=sample,
order=self.this_order,
)
for i in range(self.config.solver_order - 1):
self.model_outputs[i] = self.model_outputs[i + 1]
self.timestep_list[i] = self.timestep_list[i + 1]
self.model_outputs[-1] = model_output_convert
self.timestep_list[-1] = timestep # pyright: ignore
if self.config.lower_order_final:
this_order = min(self.config.solver_order,
len(self.timesteps) -
self.step_index) # pyright: ignore
else:
this_order = self.config.solver_order
self.this_order: int = min(this_order, self.lower_order_nums +
1) # warmup for multistep
assert self.this_order > 0
self.last_sample = sample
prev_sample = self.multistep_uni_p_bh_update(
model_output=
model_output, # pass the original non-converted model output, in case solver-p is used
sample=sample,
order=self.this_order,
)
if self.lower_order_nums < self.config.solver_order:
self.lower_order_nums += 1
# upon completion increase step index by one
assert self._step_index is not None
self._step_index += 1 # pyright: ignore
if not return_dict:
return (prev_sample, )
return SchedulerOutput(prev_sample=prev_sample)
def scale_model_input(self, sample: torch.Tensor, *args,
**kwargs) -> torch.Tensor:
"""
Ensures interchangeability with schedulers that need to scale the denoising model input depending on the
current timestep.
Args:
sample (`torch.Tensor`):
The input sample.
Returns:
`torch.Tensor`:
A scaled input sample.
"""
return sample
# Copied from diffusers.schedulers.scheduling_dpmsolver_multistep.DPMSolverMultistepScheduler.add_noise
def add_noise(
self,
original_samples: torch.Tensor,
noise: torch.Tensor,
timesteps: torch.IntTensor,
) -> torch.Tensor:
# Make sure sigmas and timesteps have the same device and dtype as original_samples
sigmas = self.sigmas.to(device=original_samples.device,
dtype=original_samples.dtype)
if original_samples.device.type == "mps" and torch.is_floating_point(
timesteps):
# mps does not support float64
schedule_timesteps = self.timesteps.to(original_samples.device,
dtype=torch.float32)
timesteps = timesteps.to(original_samples.device,
dtype=torch.float32)
else:
schedule_timesteps = self.timesteps.to(original_samples.device)
timesteps = timesteps.to(original_samples.device)
# begin_index is None when the scheduler is used for training or pipeline does not implement set_begin_index
if self.begin_index is None:
step_indices = [
self.index_for_timestep(t, schedule_timesteps)
for t in timesteps
]
elif self.step_index is not None:
# add_noise is called after first denoising step (for inpainting)
step_indices = [self.step_index] * timesteps.shape[0]
else:
# add noise is called before first denoising step to create initial latent(img2img)
step_indices = [self.begin_index] * timesteps.shape[0]
sigma = sigmas[step_indices].flatten()
while len(sigma.shape) < len(original_samples.shape):
sigma = sigma.unsqueeze(-1)
alpha_t, sigma_t = self._sigma_to_alpha_sigma_t(sigma)
noisy_samples = alpha_t * original_samples + sigma_t * noise
return noisy_samples
def __len__(self):
return self.config.num_train_timesteps
+11 -2
View File
@@ -5,9 +5,12 @@ Diffusion pipelines for fastvideo.v1.
This package contains diffusion pipelines for generating videos and images.
"""
from typing import cast
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.v1.pipelines.lora_pipeline import LoRAPipeline
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.pipeline_registry import PipelineRegistry
from fastvideo.v1.utils import (maybe_download_model,
@@ -16,7 +19,12 @@ from fastvideo.v1.utils import (maybe_download_model,
logger = init_logger(__name__)
def build_pipeline(fastvideo_args: FastVideoArgs) -> ComposedPipelineBase:
class PipelineWithLoRA(LoRAPipeline, ComposedPipelineBase):
"""Type for a pipeline that has both ComposedPipelineBase and LoRAPipeline functionality."""
pass
def build_pipeline(fastvideo_args: FastVideoArgs) -> PipelineWithLoRA:
"""
Only works with valid hf diffusers configs. (model_index.json)
We want to build a pipeline based on the inference args mode_path:
@@ -45,7 +53,7 @@ def build_pipeline(fastvideo_args: FastVideoArgs) -> ComposedPipelineBase:
logger.info("Pipeline instantiated")
# pipeline is now initialized and ready to use
return pipeline
return cast(PipelineWithLoRA, pipeline)
__all__ = [
@@ -54,4 +62,5 @@ __all__ = [
"ComposedPipelineBase",
"PipelineRegistry",
"ForwardBatch",
"LoRAPipeline",
]
@@ -42,6 +42,7 @@ class ComposedPipelineBase(ABC):
_required_config_modules: List[str] = []
training_args: Optional[TrainingArgs] = None
fastvideo_args: Optional[FastVideoArgs] = None
modules: Dict[str, torch.nn.Module] = {}
# TODO(will): args should support both inference args and training args
def __init__(self,
@@ -160,6 +161,7 @@ class ComposedPipelineBase(ABC):
for key, value in config_args.items():
setattr(fastvideo_args, key, value)
fastvideo_args.num_gpus = int(os.environ.get("WORLD_SIZE", 1))
fastvideo_args.use_cpu_offload = False
# make sure we are in training mode
fastvideo_args.inference_mode = False
@@ -167,8 +169,9 @@ class ComposedPipelineBase(ABC):
# model is loaded with the correct precision. Subsequently we will
# use FSDP2's MixedPrecisionPolicy to set the precision for the
# fwd, bwd, and other operations' precision.
fastvideo_args.precision = fastvideo_args.master_weight_type
assert fastvideo_args.precision == 'fp32', 'only fp32 is supported for training'
# fastvideo_args.precision = fastvideo_args.master_weight_type
assert fastvideo_args.master_weight_type == 'fp32', 'only fp32 is supported for training'
# assert fastvideo_args.precision == 'fp32', 'only fp32 is supported for training'
fastvideo_args.check_fastvideo_args()
@@ -187,7 +190,7 @@ class ComposedPipelineBase(ABC):
if local_rank == -1 or world_size == -1 or rank == -1:
raise ValueError(
"Local rank, world size, and rank must be set. Use torchrun to launch the script."
"Local rank, world size, and rank must be set. Use torchrun to launch the script or pass rank to the worker process."
)
torch.cuda.set_device(local_rank)
@@ -198,7 +201,8 @@ class ComposedPipelineBase(ABC):
assert fastvideo_args.sp_size is not None, "sp_size must be set"
initialize_model_parallel(
tensor_model_parallel_size=fastvideo_args.tp_size,
sequence_model_parallel_size=fastvideo_args.sp_size)
sequence_model_parallel_size=fastvideo_args.sp_size,
data_parallel_size=fastvideo_args.dp_size)
device = torch.device(f"cuda:{local_rank}")
fastvideo_args.device = device
@@ -246,7 +250,7 @@ class ComposedPipelineBase(ABC):
@abstractmethod
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""
Create the pipeline stages.
Create the inference pipeline stages.
"""
raise NotImplementedError
+149
View File
@@ -0,0 +1,149 @@
import logging
from collections import defaultdict
from typing import Any, DefaultDict, Dict, Hashable, List, Optional
import torch
import torch.distributed as dist
from safetensors.torch import load_file
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.layers.lora.linear import (BaseLayerWithLoRA, get_lora_layer,
replace_submodule)
from fastvideo.v1.models.loader.utils import get_param_names_mapping
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.v1.utils import maybe_download_lora
logger = logging.getLogger(__name__)
class LoRAPipeline(ComposedPipelineBase):
"""
Pipeline that supports injecting LoRA adapters into the diffusion transformer.
TODO: support training.
"""
lora_adapters: Dict[str, Dict[str, torch.Tensor]] = defaultdict(
dict) # state dicts of loaded lora adapters
cur_adapter_name: str = ""
lora_layers: Dict[str, BaseLayerWithLoRA] = {}
fastvideo_args: FastVideoArgs
exclude_lora_layers: List[str] = []
device: torch.device = torch.device(f"cuda:{torch.cuda.current_device()}")
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
self.exclude_lora_layers = self.modules[
"transformer"].config.arch_config.exclude_lora_layers
self.convert_to_lora_layers()
if self.fastvideo_args.lora_path is not None:
self.set_lora_adapter(
self.fastvideo_args.lora_nickname, # type: ignore
self.fastvideo_args.lora_path)
def is_target_layer(self, module_name: str) -> bool:
if self.fastvideo_args.lora_target_names is None:
return True
return any(target_name in module_name
for target_name in self.fastvideo_args.lora_target_names)
def convert_to_lora_layers(self) -> None:
"""
Converts the transformer to a LoRA transformer.
"""
for name, layer in self.modules["transformer"].named_modules():
if not self.is_target_layer(name):
continue
excluded = False
for exclude_layer in self.exclude_lora_layers:
if exclude_layer in name:
excluded = True
break
if excluded:
continue
layer = get_lora_layer(layer)
if layer is not None:
self.lora_layers[name] = layer
replace_submodule(self.modules["transformer"], name, layer)
def set_lora_adapter(self,
lora_nickname: str,
lora_path: Optional[str] = None): # type: ignore
"""
Loads a LoRA adapter into the pipeline and applies it to the transformer.
Args:
lora_nickname: The "nick name" of the adapter when referenced in the pipeline.
lora_path: The path to the adapter, either a local path or a Hugging Face repo id.
"""
if lora_nickname not in self.lora_adapters and lora_path is None:
raise ValueError(
f"Adapter {lora_nickname} not found in the pipeline. Please provide lora_path to load it."
)
adapter_updated = False
rank = dist.get_rank()
if lora_path is not None:
lora_local_path = maybe_download_lora(lora_path)
lora_state_dict = load_file(lora_local_path)
# Map the hf layer names to our custom layer names
param_names_mapping_fn = get_param_names_mapping(
self.modules["transformer"]._param_names_mapping)
lora_param_names_mapping_fn = get_param_names_mapping(
self.modules["transformer"]._lora_param_names_mapping)
to_merge_params: DefaultDict[Hashable,
Dict[Any, Any]] = defaultdict(dict)
for name, weight in lora_state_dict.items():
name = ".".join(
name.split(".")
[1:-1]) # remove the transformer prefix and .weight suffix
name, _, _ = lora_param_names_mapping_fn(name)
target_name, merge_index, num_params_to_merge = param_names_mapping_fn(
name)
# for (in_dim, r) @ (r, out_dim), we only merge (r, out_dim * n) where n is the number of linear layers to fuse
# see param mapping in HunyuanVideoArchConfig
if merge_index is not None and "lora_B" in name:
to_merge_params[target_name][merge_index] = weight
if len(to_merge_params[target_name]) == num_params_to_merge:
# cat at output dim according to the merge_index order
sorted_tensors = [
to_merge_params[target_name][i]
for i in range(num_params_to_merge)
]
weight = torch.cat(sorted_tensors, dim=1)
del to_merge_params[target_name]
else:
continue
self.lora_adapters[lora_nickname][target_name] = weight.to(
self.device)
adapter_updated = True
logger.info("Rank %d: loaded LoRA adapter %s", rank, lora_path)
if not adapter_updated and lora_nickname == self.cur_adapter_name:
return
# Merge the new adapter
adapted_count = 0
for name, layer in self.lora_layers.items():
lora_A_name = name + ".lora_A"
lora_B_name = name + ".lora_B"
if lora_A_name in self.lora_adapters[lora_nickname]\
and lora_B_name in self.lora_adapters[lora_nickname]:
if layer.merged:
layer.unmerge_lora_weights()
layer.set_lora_weights(
self.lora_adapters[lora_nickname][lora_A_name],
self.lora_adapters[lora_nickname][lora_B_name],
training_mode=self.fastvideo_args.training_mode)
adapted_count += 1
else:
if rank == 0:
logger.warning(
"LoRA adapter %s does not contain the weights for layer %s. LoRA will not be applied to it.",
lora_path, name)
layer.disable_lora = True
logger.info("Rank %d: LoRA adapter %s applied to %d layers", rank,
lora_path, adapted_count)
self.cur_adapter_name = lora_nickname
+12 -1
View File
@@ -7,7 +7,8 @@ This module defines the dataclasses used to pass state between pipeline componen
in a functional manner, reducing the need for explicit parameter passing.
"""
from dataclasses import dataclass, field
import pprint
from dataclasses import asdict, dataclass, field
from typing import Any, Dict, List, Optional, Union
import torch
@@ -114,10 +115,20 @@ class ForwardBatch:
enable_teacache: bool = False
teacache_params: Optional[TeaCacheParams | WanTeaCacheParams] = None
# STA parameters
STA_param: Optional[List] = None
is_cfg_negative: bool = False
mask_search_final_result_pos: Optional[List[List]] = None
mask_search_final_result_neg: Optional[List[List]] = None
def __post_init__(self):
"""Initialize dependent fields after dataclass initialization."""
# Set do_classifier_free_guidance based on guidance scale and negative prompt
if self.guidance_scale > 1.0:
self.do_classifier_free_guidance = True
if self.negative_prompt_embeds is None:
self.negative_prompt_embeds = []
def __str__(self):
return pprint.pformat(asdict(self), indent=2, width=120)
+3 -2
View File
@@ -6,10 +6,11 @@ import importlib
import pkgutil
from dataclasses import dataclass, field
from functools import lru_cache
from typing import AbstractSet, Dict, Optional, Tuple, Type
from typing import AbstractSet, Dict, Optional, Tuple, Type, Union
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.v1.pipelines.lora_pipeline import LoRAPipeline
logger = init_logger(__name__)
@@ -33,7 +34,7 @@ class _PipelineRegistry:
def resolve_pipeline_cls(
self,
architecture: str,
) -> Tuple[Type[ComposedPipelineBase], str]:
) -> Tuple[Union[Type[ComposedPipelineBase], Type[LoRAPipeline]], str]:
if not architecture:
logger.warning("No pipeline architecture is specified")
@@ -0,0 +1,92 @@
# SPDX-License-Identifier: Apache-2.0
"""
I2V Data Preprocessing pipeline implementation.
This module contains an implementation of the I2V Data Preprocessing pipeline
using the modular pipeline architecture.
"""
from typing import Any, Dict, List, Optional
import numpy as np
import torch
from PIL import Image
from fastvideo.v1.dataset.dataloader.schema import pyarrow_schema_i2v
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.pipelines.preprocess_pipeline_base import (
BasePreprocessPipeline)
class PreprocessPipeline_I2V(BasePreprocessPipeline):
"""I2V preprocessing pipeline implementation."""
_required_config_modules = [
"text_encoder", "tokenizer", "vae", "image_encoder", "image_processor"
]
def get_schema_fields(self) -> List[str]:
"""Get the schema fields for I2V pipeline."""
return [f.name for f in pyarrow_schema_i2v]
def get_extra_features(self, valid_data: Dict[str, Any],
fastvideo_args: FastVideoArgs) -> Dict[str, Any]:
"""Get CLIP features from the first frame of each video."""
first_frame = valid_data["pixel_values"][:, :, 0, :, :].permute(
0, 2, 3, 1) # (B, C, T, H, W) -> (B, H, W, C)
processed_images = []
for frame in first_frame:
frame_pil = Image.fromarray(frame.cpu().numpy().astype(np.uint8))
processed_img = self.get_module("image_processor")(
images=frame_pil, return_tensors="pt")
processed_images.append(processed_img)
# Get CLIP features
pixel_values = torch.cat(
[img['pixel_values'] for img in processed_images],
dim=0).to(fastvideo_args.device)
with torch.no_grad():
image_inputs = {'pixel_values': pixel_values}
with set_forward_context(current_timestep=0, attn_metadata=None):
clip_features = self.get_module("image_encoder")(**image_inputs)
clip_features = clip_features.last_hidden_state
return {"clip_feature": clip_features}
def create_record(
self,
video_name: str,
vae_latent: np.ndarray,
text_embedding: np.ndarray,
text_attention_mask: np.ndarray,
valid_data: Optional[Dict[str, Any]],
idx: int,
extra_features: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
"""Create a record for the Parquet dataset with CLIP features."""
record = super().create_record(video_name=video_name,
vae_latent=vae_latent,
text_embedding=text_embedding,
text_attention_mask=text_attention_mask,
valid_data=valid_data,
idx=idx,
extra_features=extra_features)
if extra_features and "clip_feature" in extra_features:
clip_feature = extra_features["clip_feature"]
record.update({
"clip_feature_bytes": clip_feature.tobytes(),
"clip_feature_shape": list(clip_feature.shape),
"clip_feature_dtype": str(clip_feature.dtype),
})
else:
record.update({
"clip_feature_bytes": b"",
"clip_feature_shape": [],
"clip_feature_dtype": "",
})
return record
EntryClass = PreprocessPipeline_I2V
@@ -0,0 +1,23 @@
# SPDX-License-Identifier: Apache-2.0
"""
T2V Data Preprocessing pipeline implementation.
This module contains an implementation of the T2V Data Preprocessing pipeline
using the modular pipeline architecture.
"""
from fastvideo.v1.dataset.dataloader.schema import pyarrow_schema_t2v
from fastvideo.v1.pipelines.preprocess_pipeline_base import (
BasePreprocessPipeline)
class PreprocessPipeline_T2V(BasePreprocessPipeline):
"""T2V preprocessing pipeline implementation."""
_required_config_modules = ["text_encoder", "tokenizer", "vae"]
def get_schema_fields(self):
"""Get the schema fields for T2V pipeline."""
return [f.name for f in pyarrow_schema_t2v]
EntryClass = PreprocessPipeline_T2V
@@ -1,15 +1,8 @@
# SPDX-License-Identifier: Apache-2.0
"""
T2V Data Preprocessing pipeline implementation.
This module contains an implementation of the T2V Data Preprocessing pipeline
using the modular pipeline architecture.
"""
import gc
import multiprocessing
import os
from concurrent.futures import ProcessPoolExecutor
from typing import Any, Dict
from typing import Any, Dict, List, Optional
import numpy as np
import pyarrow as pa
@@ -19,26 +12,22 @@ from torch.utils.data import DataLoader
from torch.utils.data.distributed import DistributedSampler
from tqdm import tqdm
from fastvideo.v1.configs.sample import SamplingParam
from fastvideo.v1.dataset import getdataset
from fastvideo.v1.dataset.dataloader.schema import pyarrow_schema
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages import TextEncodingStage
# TODO(will): move PRECISION_TO_TYPE to better place
logger = init_logger(__name__)
class PreprocessPipeline(ComposedPipelineBase):
_required_config_modules = ["text_encoder", "tokenizer", "vae"]
class BasePreprocessPipeline(ComposedPipelineBase):
"""Base class for preprocessing pipelines that handles common functionality."""
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
self.add_stage(stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
@@ -58,6 +47,53 @@ class PreprocessPipeline(ComposedPipelineBase):
self.preprocess_validation_text(fastvideo_args, args)
self.preprocess_video_and_text(fastvideo_args, args)
def get_extra_features(self, valid_data: Dict[str, Any],
fastvideo_args: FastVideoArgs) -> Dict[str, Any]:
"""Get additional features specific to the pipeline type. Override in subclasses."""
return {}
def get_schema_fields(self) -> List[str]:
"""Get the schema fields for the pipeline type. Override in subclasses."""
raise NotImplementedError
def create_record(
self,
video_name: str,
vae_latent: np.ndarray,
text_embedding: np.ndarray,
text_attention_mask: np.ndarray,
valid_data: Optional[Dict[str, Any]],
idx: int,
extra_features: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
"""Create a record for the Parquet dataset."""
record = {
"id": video_name,
"vae_latent_bytes": vae_latent.tobytes(),
"vae_latent_shape": list(vae_latent.shape),
"vae_latent_dtype": str(vae_latent.dtype),
"text_embedding_bytes": text_embedding.tobytes(),
"text_embedding_shape": list(text_embedding.shape),
"text_embedding_dtype": str(text_embedding.dtype),
"text_attention_mask_bytes": text_attention_mask.tobytes(),
"text_attention_mask_shape": list(text_attention_mask.shape),
"text_attention_mask_dtype": str(text_attention_mask.dtype),
"file_name": video_name,
"caption": valid_data["text"][idx] if valid_data else "",
"media_type": "video",
"width":
valid_data["pixel_values"][idx].shape[-2] if valid_data else 0,
"height":
valid_data["pixel_values"][idx].shape[-1] if valid_data else 0,
"num_frames":
vae_latent.shape[1] if len(vae_latent.shape) > 1 else 0,
"duration_sec":
float(valid_data["duration"][idx]) if valid_data else 0.0,
"fps": float(valid_data["fps"][idx]) if valid_data else 0.0,
}
if extra_features:
record.update(extra_features)
return record
def preprocess_video_and_text(self, fastvideo_args: FastVideoArgs, args):
os.makedirs(args.output_dir, exist_ok=True)
# Create directory for combined data
@@ -94,6 +130,7 @@ class PreprocessPipeline(ComposedPipelineBase):
desc="Processing videos",
unit="batch",
disable=local_rank != 0)
for batch_idx, data in enumerate(pbar):
if data is None:
continue
@@ -127,8 +164,11 @@ class PreprocessPipeline(ComposedPipelineBase):
valid_data["pixel_values"].to(
fastvideo_args.device)).mean
batch_captions = valid_data["text"]
# Get extra features if needed
extra_features = self.get_extra_features(
valid_data, fastvideo_args)
batch_captions = valid_data["text"]
batch = ForwardBatch(
data_type="video",
prompt=batch_captions,
@@ -170,7 +210,6 @@ class PreprocessPipeline(ComposedPipelineBase):
# Get the corresponding latent and info using video name
latent = latents[idx].cpu()
video_name = os.path.basename(video_path).split(".")[0]
height, width = valid_data["pixel_values"][idx].shape[-2:]
# Convert tensors to numpy arrays
vae_latent = latent.cpu().numpy()
@@ -178,28 +217,25 @@ class PreprocessPipeline(ComposedPipelineBase):
text_attention_mask = prompt_attention_mask[idx].cpu().numpy(
).astype(np.uint8)
# Get extra features for this sample if needed
sample_extra_features = {}
if extra_features:
for key, value in extra_features.items():
if isinstance(value, torch.Tensor):
sample_extra_features[key] = value[idx].cpu().numpy(
)
else:
sample_extra_features[key] = value[idx]
# Create record for Parquet dataset
record = {
"id": video_name,
"vae_latent_bytes": vae_latent.tobytes(),
"vae_latent_shape": list(vae_latent.shape),
"vae_latent_dtype": str(vae_latent.dtype),
"text_embedding_bytes": text_embedding.tobytes(),
"text_embedding_shape": list(text_embedding.shape),
"text_embedding_dtype": str(text_embedding.dtype),
"text_attention_mask_bytes": text_attention_mask.tobytes(),
"text_attention_mask_shape":
list(text_attention_mask.shape),
"text_attention_mask_dtype": str(text_attention_mask.dtype),
"file_name": video_name,
"caption": valid_data["text"][idx],
"media_type": "video",
"width": width,
"height": height,
"num_frames": latents[idx].shape[1],
"duration_sec": float(valid_data["duration"][idx]),
"fps": float(valid_data["fps"][idx]),
}
record = self.create_record(
video_name=video_name,
vae_latent=vae_latent,
text_embedding=text_embedding,
text_attention_mask=text_attention_mask,
valid_data=valid_data,
idx=idx,
extra_features=sample_extra_features)
batch_data.append(record)
if batch_data:
@@ -208,57 +244,30 @@ class PreprocessPipeline(ComposedPipelineBase):
desc="Writing to Parquet dataset",
unit="batch")
# Convert batch data to PyArrow arrays
arrays = [
pa.array([record["id"] for record in batch_data]),
pa.array(
[record["vae_latent_bytes"] for record in batch_data],
type=pa.binary()),
pa.array(
[record["vae_latent_shape"] for record in batch_data],
type=pa.list_(pa.int32())),
pa.array(
[record["vae_latent_dtype"] for record in batch_data]),
pa.array([
record["text_embedding_bytes"] for record in batch_data
],
type=pa.binary()),
pa.array([
record["text_embedding_shape"] for record in batch_data
],
type=pa.list_(pa.int32())),
pa.array([
record["text_embedding_dtype"] for record in batch_data
]),
pa.array([
record["text_attention_mask_bytes"]
for record in batch_data
],
type=pa.binary()),
pa.array([
record["text_attention_mask_shape"]
for record in batch_data
],
type=pa.list_(pa.int32())),
pa.array([
record["text_attention_mask_dtype"]
for record in batch_data
]),
pa.array([record["file_name"] for record in batch_data]),
pa.array([record["caption"] for record in batch_data]),
pa.array([record["media_type"] for record in batch_data]),
pa.array([record["width"] for record in batch_data],
type=pa.int32()),
pa.array([record["height"] for record in batch_data],
type=pa.int32()),
pa.array([record["num_frames"] for record in batch_data],
type=pa.int32()),
pa.array([record["duration_sec"] for record in batch_data],
type=pa.float32()),
pa.array([record["fps"] for record in batch_data],
type=pa.float32()),
]
table = pa.Table.from_arrays(
arrays, names=[f.name for f in pyarrow_schema])
arrays = []
for field in self.get_schema_fields():
if field.endswith('_bytes'):
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.binary()))
elif field.endswith('_shape'):
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.list_(pa.int32())))
elif field in ['width', 'height', 'num_frames']:
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.int32()))
elif field in ['duration_sec', 'fps']:
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.float32()))
else:
arrays.append(
pa.array([record[field] for record in batch_data]))
table = pa.Table.from_arrays(arrays,
names=self.get_schema_fields())
write_pbar.update(1)
write_pbar.close()
@@ -270,83 +279,17 @@ class PreprocessPipeline(ComposedPipelineBase):
logger.info("Collected batch with %s samples", len(table))
if num_processed_samples >= args.flush_frequency:
assert hasattr(self, 'all_tables') and self.all_tables
print(f"Combining {len(self.all_tables)} batches...")
combined_table = pa.concat_tables(self.all_tables)
assert len(combined_table) == num_processed_samples
print(f"Total samples collected: {len(combined_table)}")
# Calculate total number of chunks needed, discarding remainder
total_chunks = max(
num_processed_samples // args.samples_per_file, 1)
print(
f"Fixed samples per parquet file: {args.samples_per_file}")
print(f"Total number of parquet files: {total_chunks}")
print(
f"Total samples to be processed: {total_chunks * args.samples_per_file} (discarding {num_processed_samples % args.samples_per_file} samples)"
)
# Split work among processes
num_workers = int(min(multiprocessing.cpu_count(),
total_chunks))
chunks_per_worker = (total_chunks + num_workers -
1) // num_workers
print(
f"Using {num_workers} workers to process {total_chunks} chunks"
)
logger.info("Chunks per worker: %s", chunks_per_worker)
# Prepare work ranges
work_ranges = []
for i in range(num_workers):
start_idx = i * chunks_per_worker
end_idx = min((i + 1) * chunks_per_worker, total_chunks)
if start_idx < total_chunks:
work_ranges.append(
(start_idx, end_idx, combined_table, i,
combined_parquet_dir, args.samples_per_file))
total_written = 0
failed_ranges = []
with ProcessPoolExecutor(max_workers=num_workers) as executor:
futures = {
executor.submit(self.process_chunk_range, work_range):
work_range
for work_range in work_ranges
}
for future in tqdm(futures, desc="Processing chunks"):
try:
written = future.result()
total_written += written
logger.info("Processed chunk with %s samples",
written)
except Exception as e:
work_range = futures[future]
failed_ranges.append(work_range)
logger.error("Failed to process range %s-%s: %s",
work_range[0], work_range[1], str(e))
# Retry failed ranges sequentially
if failed_ranges:
logger.warning("Retrying %s failed ranges sequentially",
len(failed_ranges))
for work_range in failed_ranges:
try:
total_written += self.process_chunk_range(
work_range)
except Exception as e:
logger.error(
"Failed to process range %s-%s after retry: %s",
work_range[0], work_range[1], str(e))
logger.info("Total samples written: %s", total_written)
self._flush_tables(num_processed_samples, args,
combined_parquet_dir)
num_processed_samples = 0
self.all_tables = []
def preprocess_validation_text(self, fastvideo_args: FastVideoArgs, args):
"""Process validation text prompts and save them to parquet files.
This base implementation handles the common validation text processing logic.
Subclasses can override this method to add pipeline-specific features.
"""
# Create Parquet dataset directory for validation
validation_parquet_dir = os.path.join(args.output_dir,
"validation_parquet_dataset")
@@ -358,7 +301,10 @@ class PreprocessPipeline(ComposedPipelineBase):
# Prepare batch data for Parquet dataset
batch_data = []
sampling_param = SamplingParam.from_pretrained(
fastvideo_args.model_path)
if sampling_param.negative_prompt:
prompts = [sampling_param.negative_prompt] + prompts
# Add progress bar for validation text preprocessing
pbar = tqdm(enumerate(prompts),
desc="Processing validation prompts",
@@ -392,26 +338,14 @@ class PreprocessPipeline(ComposedPipelineBase):
text_embedding.shape, text_attention_mask.shape)
# Create record for Parquet dataset
record = {
"id": file_name,
"vae_latent_bytes": b"", # Not available for validation
"vae_latent_shape": [],
"vae_latent_dtype": "",
"text_embedding_bytes": text_embedding.tobytes(),
"text_embedding_shape": list(text_embedding.shape),
"text_embedding_dtype": str(text_embedding.dtype),
"text_attention_mask_bytes": text_attention_mask.tobytes(),
"text_attention_mask_shape": list(text_attention_mask.shape),
"text_attention_mask_dtype": str(text_attention_mask.dtype),
"file_name": file_name,
"caption": prompt,
"media_type": "video",
"width": 0, # Not available for validation
"height": 0, # Not available for validation
"num_frames": 0, # Not available for validation
"duration_sec": 0.0, # Not available for validation
"fps": 0.0, # Not available for validation
}
record = self.create_record(video_name=file_name,
vae_latent=np.array([],
dtype=np.float32),
text_embedding=text_embedding,
text_attention_mask=text_attention_mask,
valid_data=None,
idx=0,
extra_features=None)
batch_data.append(record)
logger.info("Saved validation sample: %s", file_name)
@@ -422,48 +356,29 @@ class PreprocessPipeline(ComposedPipelineBase):
desc="Writing to Parquet dataset",
unit="batch")
# Convert batch data to PyArrow arrays
arrays = [
pa.array([record["id"] for record in batch_data]),
pa.array([record["vae_latent_bytes"] for record in batch_data],
type=pa.binary()),
pa.array([record["vae_latent_shape"] for record in batch_data],
type=pa.list_(pa.int32())),
pa.array([record["vae_latent_dtype"] for record in batch_data]),
pa.array(
[record["text_embedding_bytes"] for record in batch_data],
type=pa.binary()),
pa.array(
[record["text_embedding_shape"] for record in batch_data],
type=pa.list_(pa.int32())),
pa.array(
[record["text_embedding_dtype"] for record in batch_data]),
pa.array([
record["text_attention_mask_bytes"] for record in batch_data
],
type=pa.binary()),
pa.array([
record["text_attention_mask_shape"] for record in batch_data
],
type=pa.list_(pa.int32())),
pa.array([
record["text_attention_mask_dtype"] for record in batch_data
]),
pa.array([record["file_name"] for record in batch_data]),
pa.array([record["caption"] for record in batch_data]),
pa.array([record["media_type"] for record in batch_data]),
pa.array([record["width"] for record in batch_data],
type=pa.int32()),
pa.array([record["height"] for record in batch_data],
type=pa.int32()),
pa.array([record["num_frames"] for record in batch_data],
type=pa.int32()),
pa.array([record["duration_sec"] for record in batch_data],
type=pa.float32()),
pa.array([record["fps"] for record in batch_data],
type=pa.float32()),
]
table = pa.Table.from_arrays(arrays,
names=[f.name for f in pyarrow_schema])
arrays = []
for field in self.get_schema_fields():
if field.endswith('_bytes'):
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.binary()))
elif field.endswith('_shape'):
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.list_(pa.int32())))
elif field in ['width', 'height', 'num_frames']:
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.int32()))
elif field in ['duration_sec', 'fps']:
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.float32()))
else:
arrays.append(
pa.array([record[field] for record in batch_data]))
table = pa.Table.from_arrays(arrays, names=self.get_schema_fields())
write_pbar.update(1)
write_pbar.close()
@@ -504,6 +419,74 @@ class PreprocessPipeline(ComposedPipelineBase):
del table
gc.collect() # Force garbage collection
def _flush_tables(self, num_processed_samples: int, args,
combined_parquet_dir: str):
"""Flush collected tables to disk."""
assert hasattr(self, 'all_tables') and self.all_tables
print(f"Combining {len(self.all_tables)} batches...")
combined_table = pa.concat_tables(self.all_tables)
assert len(combined_table) == num_processed_samples
print(f"Total samples collected: {len(combined_table)}")
# Calculate total number of chunks needed, discarding remainder
total_chunks = max(num_processed_samples // args.samples_per_file, 1)
print(f"Fixed samples per parquet file: {args.samples_per_file}")
print(f"Total number of parquet files: {total_chunks}")
print(
f"Total samples to be processed: {total_chunks * args.samples_per_file} (discarding {num_processed_samples % args.samples_per_file} samples)"
)
# Split work among processes
num_workers = int(min(multiprocessing.cpu_count(), total_chunks))
chunks_per_worker = (total_chunks + num_workers - 1) // num_workers
print(f"Using {num_workers} workers to process {total_chunks} chunks")
logger.info("Chunks per worker: %s", chunks_per_worker)
# Prepare work ranges
work_ranges = []
for i in range(num_workers):
start_idx = i * chunks_per_worker
end_idx = min((i + 1) * chunks_per_worker, total_chunks)
if start_idx < total_chunks:
work_ranges.append(
(start_idx, end_idx, combined_table, i,
combined_parquet_dir, args.samples_per_file))
total_written = 0
failed_ranges = []
with ProcessPoolExecutor(max_workers=num_workers) as executor:
futures = {
executor.submit(self.process_chunk_range, work_range):
work_range
for work_range in work_ranges
}
for future in tqdm(futures, desc="Processing chunks"):
try:
written = future.result()
total_written += written
logger.info("Processed chunk with %s samples", written)
except Exception as e:
work_range = futures[future]
failed_ranges.append(work_range)
logger.error("Failed to process range %s-%s: %s",
work_range[0], work_range[1], str(e))
# Retry failed ranges sequentially
if failed_ranges:
logger.warning("Retrying %s failed ranges sequentially",
len(failed_ranges))
for work_range in failed_ranges:
try:
total_written += self.process_chunk_range(work_range)
except Exception as e:
logger.error(
"Failed to process range %s-%s after retry: %s",
work_range[0], work_range[1], str(e))
logger.info("Total samples written: %s", total_written)
@staticmethod
def process_chunk_range(args: Any) -> int:
start_idx, end_idx, table, worker_id, output_dir, samples_per_file = args
@@ -554,6 +537,3 @@ class PreprocessPipeline(ComposedPipelineBase):
logger.error("Error processing chunks %s-%s for worker %s: %s",
start_idx, end_idx, worker_id, str(e))
raise
EntryClass = PreprocessPipeline
@@ -36,6 +36,7 @@ class ConditioningStage(PipelineStage):
Returns:
The batch with applied conditioning.
"""
# TODO!!
if not batch.do_classifier_free_guidance:
return batch
else:
+160 -9
View File
@@ -47,6 +47,17 @@ class DenoisingStage(PipelineStage):
super().__init__()
self.transformer = transformer
self.scheduler = scheduler
# when used for validation, transformer is None as it is taking from the
# training loop
if transformer is not None:
attn_head_size = self.transformer.hidden_size // self.transformer.num_attention_heads
self.attn_backend = get_attn_backend(
head_size=attn_head_size,
dtype=torch.float16, # TODO(will): hack
supported_attention_backends=(_Backend.SLIDING_TILE_ATTN,
_Backend.FLASH_ATTN,
_Backend.TORCH_SDPA) # hack
)
def forward(
self,
@@ -151,6 +162,10 @@ class DenoisingStage(PipelineStage):
},
)
# Prepare STA parameters
if st_attn_available and self.attn_backend == SlidingTileAttentionBackend:
self.prepare_sta_param(batch, fastvideo_args)
# Get latents and embeddings
latents = batch.latents
prompt_embeds = batch.prompt_embeds
@@ -192,15 +207,6 @@ class DenoisingStage(PipelineStage):
dtype=target_dtype,
enabled=autocast_enabled):
# TODO(will-refactor): all of this should be in the stage's init
attn_head_size = self.transformer.hidden_size // self.transformer.num_attention_heads
self.attn_backend = get_attn_backend(
head_size=attn_head_size,
dtype=torch.float16, # TODO(will): hack
supported_attention_backends=(
_Backend.SLIDING_TILE_ATTN, _Backend.FLASH_ATTN,
_Backend.TORCH_SDPA) # hack
)
if st_attn_available and self.attn_backend == SlidingTileAttentionBackend:
self.attn_metadata_builder_cls = self.attn_backend.get_builder_cls(
)
@@ -222,6 +228,7 @@ class DenoisingStage(PipelineStage):
# support torch dynamo compilation. They pass in
# attn_metadata, vllm_config, and num_tokens. We can pass in
# fastvideo_args or training_args, and attn_metadata.
batch.is_cfg_negative = False
with set_forward_context(
current_timestep=i,
attn_metadata=attn_metadata,
@@ -240,6 +247,7 @@ class DenoisingStage(PipelineStage):
# Apply guidance
if batch.do_classifier_free_guidance:
batch.is_cfg_negative = True
with set_forward_context(
current_timestep=i,
attn_metadata=attn_metadata,
@@ -289,6 +297,10 @@ class DenoisingStage(PipelineStage):
# Update batch with final latents
batch.latents = latents
# Save STA mask search results if needed
if st_attn_available and self.attn_backend == SlidingTileAttentionBackend and fastvideo_args.STA_mode == 'STA_searching':
self.save_sta_search_results(batch)
if fastvideo_args.use_cpu_offload:
self.transformer.to('cpu')
torch.cuda.empty_cache()
@@ -360,3 +372,142 @@ class DenoisingStage(PipelineStage):
noise_cfg = (guidance_rescale * noise_pred_rescaled +
(1 - guidance_rescale) * noise_cfg)
return noise_cfg
def prepare_sta_param(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs):
"""
Prepare Sliding Tile Attention (STA) parameters and settings.
Args:
batch: The current batch information.
fastvideo_args: The inference arguments.
"""
# TODO(kevin): STA mask search, currently only support Wan2.1 with 69x768x1280
from fastvideo.v1.STA_configuration import configure_sta
STA_mode = fastvideo_args.STA_mode
skip_time_steps = fastvideo_args.skip_time_steps
if batch.timesteps is None:
raise ValueError("Timesteps must be provided")
timesteps_num = batch.timesteps.shape[0]
logger.info("STA_mode: %s", STA_mode)
if (batch.num_frames, batch.height,
batch.width) != (69, 768, 1280) and STA_mode != "STA_inference":
raise NotImplementedError(
"STA mask search/tuning is not supported for this resolution")
if STA_mode == "STA_searching" or STA_mode == "STA_tuning" or STA_mode == "STA_tuning_cfg":
size = (batch.width, batch.height)
if size == (1280, 768):
# TODO: make it configurable
sparse_mask_candidates_searching = [
"3, 1, 10", "1, 5, 7", "3, 3, 3", "1, 6, 5", "1, 3, 10",
"3, 6, 1"
]
sparse_mask_candidates_tuning = [
"3, 1, 10", "1, 5, 7", "3, 3, 3", "1, 6, 5", "1, 3, 10",
"3, 6, 1"
]
full_mask = ["3,6,10"]
else:
raise NotImplementedError(
"STA mask search is not supported for this resolution")
layer_num = self.transformer.config.num_layers
# specific for HunyuanVideo
if hasattr(self.transformer.config, "num_single_layers"):
layer_num += self.transformer.config.num_single_layers
head_num = self.transformer.config.num_attention_heads
if STA_mode == "STA_searching":
STA_param = configure_sta(
mode='STA_searching',
layer_num=layer_num,
head_num=head_num,
time_step_num=timesteps_num,
mask_candidates=sparse_mask_candidates_searching +
full_mask, # last is full mask; Can add more sparse masks while keep last one as full mask
)
elif STA_mode == 'STA_tuning':
STA_param = configure_sta(
mode='STA_tuning',
layer_num=layer_num,
head_num=head_num,
time_step_num=timesteps_num,
mask_search_files_path=
f'output/mask_search_result_pos_{size[0]}x{size[1]}/',
mask_candidates=sparse_mask_candidates_tuning,
full_attention_mask=[int(x) for x in full_mask[0].split(',')],
skip_time_steps=
skip_time_steps, # Use full attention for first 12 steps
save_dir=
f'output/mask_search_strategy_{size[0]}x{size[1]}/', # Custom save directory
timesteps=timesteps_num)
elif STA_mode == 'STA_tuning_cfg':
STA_param = configure_sta(
mode='STA_tuning_cfg',
layer_num=layer_num,
head_num=head_num,
time_step_num=timesteps_num,
mask_search_files_path_pos=
f'output/mask_search_result_pos_{size[0]}x{size[1]}/',
mask_search_files_path_neg=
f'output/mask_search_result_neg_{size[0]}x{size[1]}/',
mask_candidates=sparse_mask_candidates_tuning,
full_attention_mask=[int(x) for x in full_mask[0].split(',')],
skip_time_steps=skip_time_steps,
save_dir=f'output/mask_search_strategy_{size[0]}x{size[1]}/',
timesteps=timesteps_num)
elif STA_mode == 'STA_inference':
import fastvideo.v1.envs as envs
config_file = envs.FASTVIDEO_ATTENTION_CONFIG
if config_file is None:
raise ValueError("FASTVIDEO_ATTENTION_CONFIG is not set")
STA_param = configure_sta(mode='STA_inference',
layer_num=layer_num,
head_num=head_num,
time_step_num=timesteps_num,
load_path=config_file)
batch.STA_param = STA_param
batch.mask_search_final_result_pos = [[] for _ in range(timesteps_num)]
batch.mask_search_final_result_neg = [[] for _ in range(timesteps_num)]
def save_sta_search_results(self, batch: ForwardBatch):
"""
Save the STA mask search results.
Args:
batch: The current batch information.
"""
size = (batch.width, batch.height)
if size == (1280, 768):
# TODO: make it configurable
sparse_mask_candidates_searching = [
"3, 1, 10", "1, 5, 7", "3, 3, 3", "1, 6, 5", "1, 3, 10",
"3, 6, 1"
]
else:
raise NotImplementedError(
"STA mask search is not supported for this resolution")
from fastvideo.v1.STA_configuration import save_mask_search_results
if batch.mask_search_final_result_pos is not None and batch.prompt is not None:
save_mask_search_results(
[
dict(layer_data)
for layer_data in batch.mask_search_final_result_pos
],
prompt=str(batch.prompt),
mask_strategies=sparse_mask_candidates_searching,
output_dir=f'output/mask_search_result_pos_{size[0]}x{size[1]}/'
)
if batch.mask_search_final_result_neg is not None and batch.prompt is not None:
save_mask_search_results(
[
dict(layer_data)
for layer_data in batch.mask_search_final_result_neg
],
prompt=str(batch.prompt),
mask_strategies=sparse_mask_candidates_searching,
output_dir=f'output/mask_search_result_neg_{size[0]}x{size[1]}/'
)
@@ -20,6 +20,7 @@ from fastvideo.v1.models.encoders.bert import HunyuanClip # type: ignore
from fastvideo.v1.models.encoders.stepllm import STEP1TextEncoder
from fastvideo.v1.models.loader.component_loader import PipelineComponentLoader
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.v1.pipelines.lora_pipeline import LoRAPipeline
from fastvideo.v1.pipelines.stages import (DecodingStage, DenoisingStage,
InputValidationStage,
LatentPreparationStage,
@@ -29,7 +30,7 @@ from fastvideo.v1.pipelines.stages import (DecodingStage, DenoisingStage,
logger = init_logger(__name__)
class StepVideoPipeline(ComposedPipelineBase):
class StepVideoPipeline(LoRAPipeline, ComposedPipelineBase):
_required_config_modules = ["transformer", "scheduler", "vae"]
@@ -9,6 +9,7 @@ using the modular pipeline architecture.
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.v1.pipelines.lora_pipeline import LoRAPipeline
# isort: off
from fastvideo.v1.pipelines.stages import (
@@ -16,17 +17,23 @@ from fastvideo.v1.pipelines.stages import (
EncodingStage, InputValidationStage, LatentPreparationStage,
TextEncodingStage, TimestepPreparationStage)
# isort: on
from fastvideo.v1.models.schedulers.scheduling_flow_unipc_multistep import (
FlowUniPCMultistepScheduler)
logger = init_logger(__name__)
class WanImageToVideoPipeline(ComposedPipelineBase):
class WanImageToVideoPipeline(LoRAPipeline, ComposedPipelineBase):
_required_config_modules = [
"text_encoder", "tokenizer", "vae", "transformer", "scheduler", \
"image_encoder", "image_processor"
]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
shift=fastvideo_args.flow_shift)
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
+18 -4
View File
@@ -8,7 +8,9 @@ using the modular pipeline architecture.
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.v1.models.schedulers.scheduling_flow_unipc_multistep import (
FlowUniPCMultistepScheduler)
from fastvideo.v1.pipelines import ComposedPipelineBase, LoRAPipeline
from fastvideo.v1.pipelines.stages import (ConditioningStage, DecodingStage,
DenoisingStage, InputValidationStage,
LatentPreparationStage,
@@ -18,13 +20,21 @@ from fastvideo.v1.pipelines.stages import (ConditioningStage, DecodingStage,
logger = init_logger(__name__)
class WanPipeline(ComposedPipelineBase):
class WanPipeline(LoRAPipeline, ComposedPipelineBase):
"""
Wan video diffusion pipeline with LoRA support.
"""
_required_config_modules = [
"text_encoder", "tokenizer", "vae", "transformer", "scheduler"
]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
# We use UniPCMScheduler from Wan2.1 official repo, not the one in diffusers.
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
shift=fastvideo_args.flow_shift)
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
"""Set up pipeline stages with proper dependency injection."""
self.add_stage(stage_name="input_validation_stage",
@@ -63,7 +73,11 @@ class WanValidationPipeline(ComposedPipelineBase):
"""
_required_config_modules = ["vae", "scheduler"]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
shift=fastvideo_args.flow_shift)
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
"""Set up pipeline stages with proper dependency injection."""
self.add_stage(stage_name="timestep_preparation_stage",
stage=TimestepPreparationStage(
@@ -1,77 +0,0 @@
# type: ignore
# SPDX-License-Identifier: Apache-2.0
import os
import sys
import imageio
import numpy as np
import torch
import torchvision
from einops import rearrange
from fastvideo.v1.distributed import (init_distributed_environment,
initialize_model_parallel)
from fastvideo.v1.fastvideo_args import FastVideoArgs, prepare_fastvideo_args
# Fix the import path
from fastvideo.v1.inference_engine import InferenceEngine
def initialize_distributed_and_parallelism(fastvideo_args: FastVideoArgs):
local_rank = int(os.environ.get("LOCAL_RANK", 0))
rank = int(os.environ.get("RANK", 0))
world_size = int(os.environ.get("WORLD_SIZE", 1))
torch.cuda.set_device(local_rank)
init_distributed_environment(world_size=world_size,
rank=rank,
local_rank=local_rank)
device_str = f"cuda:{local_rank}"
fastvideo_args.device_str = device_str
fastvideo_args.device = torch.device(device_str)
assert fastvideo_args.sp_size is not None
assert fastvideo_args.tp_size is not None
initialize_model_parallel(
sequence_model_parallel_size=fastvideo_args.sp_size,
tensor_model_parallel_size=fastvideo_args.tp_size,
)
def main(fastvideo_args: FastVideoArgs):
initialize_distributed_and_parallelism(fastvideo_args)
engine = InferenceEngine.create_engine(fastvideo_args, )
if fastvideo_args.prompt_path is not None:
with open(fastvideo_args.prompt_path) as f:
prompts = [line.strip() for line in f.readlines()]
else:
if fastvideo_args.prompt is None:
raise ValueError("prompt or prompt_path is required")
prompts = [fastvideo_args.prompt]
# Process each prompt
for prompt in prompts:
outputs = engine.run(
prompt=prompt,
fastvideo_args=fastvideo_args,
)
# Process outputs
videos = rearrange(outputs["samples"], "b c t h w -> t b c h w")
frames = []
for x in videos:
x = torchvision.utils.make_grid(x, nrow=6)
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
frames.append((x * 255).numpy().astype(np.uint8))
# Save video
os.makedirs(os.path.dirname(fastvideo_args.output_path), exist_ok=True)
imageio.mimsave(os.path.join(fastvideo_args.output_path,
f"{prompt[:100]}.mp4"),
frames,
fps=fastvideo_args.fps)
if __name__ == "__main__":
fastvideo_args = prepare_fastvideo_args(sys.argv[1:])
main(fastvideo_args)
@@ -1,20 +0,0 @@
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
import os
def main():
print(os.environ["RANK"])
print(os.environ["WORLD_SIZE"])
print(os.environ["LOCAL_RANK"])
print(os.environ["MASTER_ADDR"])
print(os.environ["MASTER_PORT"])
print(os.environ["RANK"])
print(os.environ["WORLD_SIZE"])
print(os.environ["LOCAL_RANK"])
print(os.environ["MASTER_ADDR"])
if __name__ == "__main__":
main()
@@ -257,6 +257,7 @@ def test_inference_similarity(prompt, ATTENTION_BACKEND, model_id):
"flow_shift": BASE_PARAMS["flow_shift"],
"sp_size": BASE_PARAMS["sp_size"],
"tp_size": BASE_PARAMS["tp_size"],
"use_cpu_offload": True,
}
if BASE_PARAMS.get("vae_sp"):
init_kwargs["vae_sp"] = True
@@ -65,6 +65,7 @@ def test_hunyuanvideo_distributed():
precision=precision_str)
args.device = torch.device(f"cuda:{LOCAL_RANK}")
args.dit_config = HunyuanVideoConfig()
args.check_fastvideo_args()
loader = TransformerLoader()
model = loader.load(TRANSFORMER_PATH, "", args)
@@ -37,6 +37,7 @@ def test_wan_transformer():
precision=precision_str)
args.device = device
args.dit_config = WanVideoConfig()
args.check_fastvideo_args()
loader = TransformerLoader()
model2 = loader.load(TRANSFORMER_PATH, "", args).to(device, dtype=precision)
+4
View File
@@ -0,0 +1,4 @@
from .training_pipeline import TrainingPipeline
from .wan_training_pipeline import WanTrainingPipeline
__all__ = ["TrainingPipeline", "WanTrainingPipeline"]
@@ -0,0 +1,107 @@
import random
from typing import Any, Dict, Optional
import numpy as np
import torch
import torch.distributed.checkpoint.stateful
from torch.distributed.checkpoint.state_dict import (StateDictOptions,
get_model_state_dict,
get_optimizer_state_dict,
set_model_state_dict,
set_optimizer_state_dict)
class ModelWrapper(torch.distributed.checkpoint.stateful.Stateful):
def __init__(self, model: torch.nn.Module) -> None:
self.model = model
def state_dict(self) -> Dict[str, Any]:
return get_model_state_dict(self.model) # type: ignore[no-any-return]
def load_state_dict(self, state_dict: Dict[str, Any]) -> None:
set_model_state_dict(
self.model,
model_state_dict=state_dict,
options=StateDictOptions(strict=False),
)
class OptimizerWrapper(torch.distributed.checkpoint.stateful.Stateful):
def __init__(self, model: torch.nn.Module,
optimizer: torch.optim.Optimizer) -> None:
self.model = model
self.optimizer = optimizer
def state_dict(self) -> Dict[str, Any]:
return get_optimizer_state_dict( # type: ignore[no-any-return]
self.model,
self.optimizer,
options=StateDictOptions(flatten_optimizer_state_dict=True),
)
def load_state_dict(self, state_dict: Dict[str, Any]) -> None:
set_optimizer_state_dict(
self.model,
self.optimizer,
optim_state_dict=state_dict,
options=StateDictOptions(flatten_optimizer_state_dict=True),
)
class SchedulerWrapper(torch.distributed.checkpoint.stateful.Stateful):
def __init__(self, scheduler) -> None:
self.scheduler = scheduler
def state_dict(self) -> Dict[str, Any]:
return {"scheduler": self.scheduler.state_dict()}
def load_state_dict(self, state_dict: Dict[str, Any]) -> None:
self.scheduler.load_state_dict(state_dict["scheduler"])
class RandomStateWrapper(torch.distributed.checkpoint.stateful.Stateful):
def __init__(self,
noise_generator: Optional[torch.Generator] = None) -> None:
self.noise_generator = noise_generator
def state_dict(self) -> Dict[str, Any]:
state = {
"torch_rng_state": torch.get_rng_state(),
"numpy_rng_state": np.random.get_state(),
"python_rng_state": random.getstate(),
}
if torch.cuda.is_available():
state["cuda_rng_state"] = torch.cuda.get_rng_state()
if torch.cuda.device_count() > 1:
state["cuda_rng_state_all"] = torch.cuda.get_rng_state_all()
if self.noise_generator is not None:
state["noise_generator_state"] = self.noise_generator.get_state()
return state
def load_state_dict(self, state_dict: Dict[str, Any]) -> None:
if "torch_rng_state" in state_dict:
torch.set_rng_state(state_dict["torch_rng_state"])
if "numpy_rng_state" in state_dict:
np.random.set_state(state_dict["numpy_rng_state"])
if "python_rng_state" in state_dict:
random.setstate(state_dict["python_rng_state"])
# Restore CUDA random state
if torch.cuda.is_available():
if "cuda_rng_state" in state_dict:
torch.cuda.set_rng_state(state_dict["cuda_rng_state"])
if "cuda_rng_state_all" in state_dict:
torch.cuda.set_rng_state_all(state_dict["cuda_rng_state_all"])
# Restore noise generator state
if "noise_generator_state" in state_dict and self.noise_generator is not None:
self.noise_generator.set_state(state_dict["noise_generator_state"])
+113 -55
View File
@@ -1,11 +1,14 @@
import gc
import math
import os
import traceback
from abc import ABC, abstractmethod
from typing import Any, Dict, Iterator
import imageio
import numpy as np
import torch
import torch.distributed as dist
import torchvision
from diffusers.optimization import get_scheduler
from einops import rearrange
@@ -13,7 +16,7 @@ from torchdata.stateful_dataloader import StatefulDataLoader
from fastvideo.v1.configs.sample import SamplingParam
from fastvideo.v1.dataset.parquet_datasets import ParquetVideoTextDataset
from fastvideo.v1.distributed import get_sp_group
from fastvideo.v1.distributed import get_sp_group, get_world_group
from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.logger import init_logger
@@ -37,6 +40,10 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
"""
_required_config_modules = ["scheduler", "transformer"]
validation_pipeline: ComposedPipelineBase
train_dataloader: StatefulDataLoader
train_loader_iter: Iterator[tuple[torch.Tensor, torch.Tensor, torch.Tensor,
Dict[str, Any]]]
current_epoch: int = 0
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
raise RuntimeError(
@@ -45,10 +52,13 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
def initialize_training_pipeline(self, training_args: TrainingArgs):
logger.info("Initializing training pipeline...")
self.device = training_args.device
world_group = get_world_group()
self.world_size = world_group.world_size
self.rank = world_group.rank
self.sp_group = get_sp_group()
self.world_size = self.sp_group.world_size
self.rank = self.sp_group.rank
self.local_rank = self.sp_group.local_rank
self.rank_in_sp_group = self.sp_group.rank_in_group
self.sp_world_size = self.sp_group.world_size
self.local_rank = world_group.local_rank
self.transformer = self.get_module("transformer")
assert self.transformer is not None
@@ -97,11 +107,26 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
prefetch_factor=2,
shuffle=False,
pin_memory=True,
pin_memory_device=f"cuda:{torch.cuda.current_device()}",
drop_last=True)
self.noise_scheduler = noise_scheduler
if self.rank <= 0:
assert training_args.gradient_accumulation_steps is not None
assert training_args.sp_size is not None
assert training_args.train_sp_batch_size is not None
assert training_args.max_train_steps is not None
self.num_update_steps_per_epoch = math.ceil(
len(self.train_dataloader) /
training_args.gradient_accumulation_steps * training_args.sp_size /
training_args.train_sp_batch_size)
self.num_train_epochs = math.ceil(training_args.max_train_steps /
self.num_update_steps_per_epoch)
# TODO(will): is there a cleaner way to track epochs?
self.current_epoch = 0
if self.rank == 0:
project = training_args.tracker_project_name or "fastvideo"
wandb.init(project=project, config=training_args)
@@ -122,7 +147,8 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
raise NotImplementedError(
"Training pipeline must implement this method")
def log_validation(self, transformer, training_args, global_step) -> None:
@torch.no_grad()
def _log_validation(self, transformer, training_args, global_step) -> None:
assert training_args is not None
training_args.inference_mode = True
training_args.use_cpu_offload = False
@@ -131,36 +157,48 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
if self.validation_pipeline is None:
raise ValueError("Validation pipeline is not set")
logger.info("Starting validation")
# Create sampling parameters if not provided
sampling_param = SamplingParam.from_pretrained(training_args.model_path)
# Set deterministic seed for validation
validation_seed = training_args.seed if training_args.seed is not None else 42
torch.manual_seed(validation_seed)
torch.cuda.manual_seed_all(validation_seed)
logger.info("Using validation seed: %s", validation_seed)
# Prepare validation prompts
logger.info('fastvideo_args.validation_prompt_dir: %s',
training_args.validation_prompt_dir)
validation_dataset = ParquetVideoTextDataset(
training_args.validation_prompt_dir,
batch_size=1,
rank=0,
world_size=1,
cfg_rate=0,
num_latent_t=training_args.num_latent_t)
rank=self.rank,
world_size=self.world_size,
cfg_rate=training_args.cfg,
num_latent_t=training_args.num_latent_t,
validation=True)
if sampling_param.negative_prompt:
_, negative_prompt_embeds, negative_prompt_attention_mask, _ = validation_dataset.get_validation_negative_prompt(
)
validation_dataloader = StatefulDataLoader(
validation_dataset,
batch_size=1,
num_workers=1, # Reduce number of workers to avoid memory issues
num_workers=5, # Reduce number of workers to avoid memory issues
prefetch_factor=2,
shuffle=False,
pin_memory=True,
pin_memory_device=f"cuda:{torch.cuda.current_device()}",
drop_last=False)
transformer.requires_grad_(False)
for p in transformer.parameters():
p.requires_grad = False
transformer.eval()
# Add the transformer to the validation pipeline
self.validation_pipeline.add_module("transformer", transformer)
# TODO(Peiyuan): those logic should be inside add_module
self.validation_pipeline.latent_preparation_stage.transformer = transformer # type: ignore[attr-defined]
self.validation_pipeline.denoising_stage.transformer = transformer # type: ignore[attr-defined]
@@ -168,9 +206,8 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
videos = []
captions = []
for _, embeddings, masks, infos in validation_dataloader:
logger.info("infos: %s", infos)
caption = infos['caption']
captions.append(caption)
captions.extend(caption)
prompt_embeds = embeddings.to(training_args.device)
prompt_attention_mask = masks.to(training_args.device)
@@ -183,38 +220,47 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
temporal_compression_factor = training_args.vae_config.arch_config.temporal_compression_ratio
num_frames = (training_args.num_latent_t -
1) * temporal_compression_factor + 1
logger.info(
"validation num_frames: %s, temporal_compression_factor: %s, num_latent_t: %s",
num_frames, temporal_compression_factor,
training_args.num_latent_t)
# Prepare batch for validation
batch = ForwardBatch(
data_type="video",
latents=None,
# seed=sampling_param.seed,
seed=validation_seed, # Use deterministic seed
generator=torch.Generator(
device="cpu").manual_seed(validation_seed),
prompt_embeds=[prompt_embeds],
prompt_attention_mask=[prompt_attention_mask],
negative_prompt_embeds=[negative_prompt_embeds],
negative_attention_mask=[negative_prompt_attention_mask],
# make sure we use the same height, width, and num_frames as the training pipeline
height=training_args.num_height,
width=training_args.num_width,
num_frames=num_frames,
# TODO(will): validation_sampling_steps and
# validation_guidance_scale are actually passed in as a list of
# values, like "10,20,30". The validation should be run for each
# combination of values.
# num_inference_steps=fastvideo_args.validation_sampling_steps,
num_inference_steps=sampling_param.num_inference_steps,
# guidance_scale=fastvideo_args.validation_guidance_scale,
guidance_scale=1,
guidance_scale=sampling_param.guidance_scale,
n_tokens=n_tokens,
do_classifier_free_guidance=False,
eta=0.0,
extra={},
)
# Run validation inference
with torch.inference_mode(), torch.autocast("cuda",
dtype=torch.bfloat16):
with torch.no_grad(), torch.autocast("cuda", dtype=torch.bfloat16):
output_batch = self.validation_pipeline.forward(
batch, training_args)
samples = output_batch.output
# Re-enable gradients for training
transformer.requires_grad_(True)
transformer.train()
if self.rank_in_sp_group != 0:
continue
# Process outputs
video = rearrange(samples, "b c t h w -> t b c h w")
frames = []
@@ -225,34 +271,46 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
videos.append(frames)
# Log validation results
rank = int(os.environ.get("RANK", 0))
world_group = get_world_group()
num_sp_groups = world_group.world_size // self.sp_group.world_size
if rank == 0:
video_filenames = []
video_captions = []
for i, video in enumerate(videos):
caption = captions[i]
os.makedirs(training_args.output_dir, exist_ok=True)
filename = os.path.join(
training_args.output_dir,
f"validation_step_{global_step}_video_{i}.mp4")
imageio.mimsave(filename, video, fps=sampling_param.fps)
video_filenames.append(filename)
video_captions.append(
caption) # Store the caption for each video
# Only sp_group leaders (rank_in_sp_group == 0) need to send their
# results to global rank 0
if self.rank_in_sp_group == 0:
if self.rank == 0:
# Global rank 0 collects results from all sp_group leaders
all_videos = videos # Start with own results
all_captions = captions
logs = {
"validation_videos": [
wandb.Video(filename,
caption=caption) for filename, caption in zip(
video_filenames, video_captions)
]
}
wandb.log(logs, step=global_step)
# Receive from other sp_group leaders
for sp_group_idx in range(1, num_sp_groups):
src_rank = sp_group_idx * self.sp_world_size # Global rank of other sp_group leaders
recv_videos = world_group.recv_object(src=src_rank)
recv_captions = world_group.recv_object(src=src_rank)
all_videos.extend(recv_videos)
all_captions.extend(recv_captions)
# Re-enable gradients for training
transformer.requires_grad_(True)
transformer.train()
video_filenames = []
for i, (video,
caption) in enumerate(zip(all_videos, all_captions)):
os.makedirs(training_args.output_dir, exist_ok=True)
filename = os.path.join(
training_args.output_dir,
f"validation_step_{global_step}_video_{i}.mp4")
imageio.mimsave(filename, video, fps=sampling_param.fps)
video_filenames.append(filename)
logs = {
"validation_videos": [
wandb.Video(filename, caption=caption) for filename,
caption in zip(video_filenames, all_captions)
]
}
wandb.log(logs, step=global_step)
else:
# Other sp_group leaders send their results to global rank 0
world_group.send_object(videos, dst=0)
world_group.send_object(captions, dst=0)
gc.collect()
torch.cuda.empty_cache()
@@ -339,7 +397,7 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
absolute_errors: list[float] = []
param_count = 0
rank = int(os.environ.get("RANK", 0))
rank = dist.get_rank()
sp_group = get_sp_group()
for name, param in transformer.named_parameters():
sp_group.barrier()
@@ -381,7 +439,7 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
with torch.no_grad():
# only have a single rank modify the parameter
# because we are using FSDP
if rank <= 0:
if rank == 0:
flat_param[check_idx] = orig_value + delta
loss = compute_loss()
if delta > 0:
@@ -399,7 +457,7 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
abs(numerical_grad), 1e-3)
absolute_errors.append(abs_error)
if self.rank <= 0:
if self.rank == 0:
logger.info(
"%s[%s]: analytical=%.5f, numerical=%.5f, abs_error=%.2e, rel_error=%.2f%%",
name, check_idx, analytical_grad, numerical_grad,
@@ -408,7 +466,7 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
# param_count += 1
# Compute and log statistics
if rank <= 0 and absolute_errors:
if rank == 0 and absolute_errors:
min_err, max_err, mean_err = min(absolute_errors), max(
absolute_errors
), sum(absolute_errors) / len(absolute_errors)
+159 -20
View File
@@ -1,21 +1,57 @@
import json
import math
import os
from typing import List, Optional, Tuple, Union
import time
from typing import Any, Dict, List, Optional, Tuple, Union
import torch
import torch.distributed as dist
from torch.distributed.fsdp import FullStateDictConfig
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
from torch.distributed.fsdp import StateDictType
import torch.distributed.checkpoint as dcp
import torch.distributed.checkpoint.stateful
from safetensors.torch import save_file
from fastvideo.v1.logger import init_logger
from fastvideo.v1.training.checkpointing_utils import (ModelWrapper,
OptimizerWrapper,
RandomStateWrapper,
SchedulerWrapper)
logger = init_logger(__name__)
_HAS_ERRORED_CLIP_GRAD_NORM_WHILE_HANDLING_FAILING_DTENSOR_CASES = False
def gather_state_dict_on_cpu_rank0(
model,
device: Optional[torch.device] = None,
) -> Dict[str, Any]:
rank = dist.get_rank()
cpu_state_dict = {}
sharded_sd = model.state_dict()
for param_name, param in sharded_sd.items():
if hasattr(param, "_local_tensor"):
# DTensor case
if param.is_cpu:
# Gather directly on CPU
param = param.full_tensor()
else:
if device is not None:
param = param.to(device)
param = param.full_tensor()
else:
# Regular tensor case
if param.is_cpu:
pass
else:
if device is not None:
param = param.to(device)
if rank == 0:
cpu_state_dict[param_name] = param.cpu()
return cpu_state_dict
def compute_density_for_timestep_sampling(
weighting_scheme: str,
batch_size: int,
@@ -66,24 +102,67 @@ def get_sigmas(noise_scheduler,
return sigma
def save_checkpoint(transformer, rank, output_dir, step) -> None:
# Configure FSDP to save full state dict
FSDP.set_state_dict_type(
transformer,
state_dict_type=StateDictType.FULL_STATE_DICT,
state_dict_config=FullStateDictConfig(offload_to_cpu=True,
rank0_only=True),
)
def save_checkpoint(transformer,
rank,
output_dir,
step,
optimizer=None,
dataloader=None,
scheduler=None,
noise_generator=None) -> None:
"""
Save checkpoint following finetrainer's distributed checkpoint approach.
Saves both distributed checkpoint and consolidated model weights.
"""
save_dir = os.path.join(output_dir, f"checkpoint-{step}")
os.makedirs(save_dir, exist_ok=True)
# Now get the state dict
cpu_state = transformer.state_dict()
states = {
"model": ModelWrapper(transformer),
"random_state": RandomStateWrapper(noise_generator),
}
# Save it (only on rank 0 since we used rank0_only=True)
if rank <= 0:
save_dir = os.path.join(output_dir, f"checkpoint-{step}")
os.makedirs(save_dir, exist_ok=True)
weight_path = os.path.join(save_dir, "diffusion_pytorch_model.pt")
torch.save(cpu_state, weight_path)
if optimizer is not None:
states["optimizer"] = OptimizerWrapper(transformer, optimizer)
if dataloader is not None:
states["dataloader"] = dataloader
if scheduler is not None:
states["scheduler"] = SchedulerWrapper(scheduler)
dcp_dir = os.path.join(save_dir, "distributed_checkpoint")
logger.info("rank: %s, saving distributed checkpoint to %s",
rank,
dcp_dir,
local_main_process_only=False)
begin_time = time.perf_counter()
dcp.save(states, checkpoint_id=dcp_dir)
end_time = time.perf_counter()
logger.info("rank: %s, distributed checkpoint saved in %.2f seconds",
rank,
end_time - begin_time,
local_main_process_only=False)
cpu_state = gather_state_dict_on_cpu_rank0(transformer, device=None)
if rank == 0:
# Save model weights (consolidated)
weight_path = os.path.join(save_dir,
"diffusion_pytorch_model.safetensors")
logger.info("rank: %s, saving consolidated checkpoint to %s",
rank,
weight_path,
local_main_process_only=False)
save_file(cpu_state, weight_path)
logger.info("rank: %s, consolidated checkpoint saved to %s",
rank,
weight_path,
local_main_process_only=False)
# Save model config
config_dict = transformer.hf_config
if "dtype" in config_dict:
del config_dict["dtype"] # TODO
@@ -94,6 +173,66 @@ def save_checkpoint(transformer, rank, output_dir, step) -> None:
logger.info("--> checkpoint saved at step %s to %s", step, weight_path)
def load_checkpoint(transformer,
rank,
checkpoint_path,
optimizer=None,
dataloader=None,
scheduler=None,
noise_generator=None) -> int:
"""
Load checkpoint following finetrainer's distributed checkpoint approach.
Returns the step number from which training should resume.
"""
if not os.path.exists(checkpoint_path):
logger.warning("Checkpoint path %s does not exist", checkpoint_path)
return 0
# Extract step number from checkpoint path
step = int(os.path.basename(checkpoint_path).split('-')[-1])
if rank == 0:
logger.info("Loading checkpoint from step %s", step)
dcp_dir = os.path.join(checkpoint_path, "distributed_checkpoint")
if not os.path.exists(dcp_dir):
logger.warning("Distributed checkpoint directory %s does not exist",
dcp_dir)
return 0
states = {
"model": ModelWrapper(transformer),
"random_state": RandomStateWrapper(noise_generator),
}
if optimizer is not None:
states["optimizer"] = OptimizerWrapper(transformer, optimizer)
if dataloader is not None:
states["dataloader"] = dataloader
if scheduler is not None:
states["scheduler"] = SchedulerWrapper(scheduler)
logger.info("rank: %s, loading distributed checkpoint from %s",
rank,
dcp_dir,
local_main_process_only=False)
begin_time = time.perf_counter()
dcp.load(states, checkpoint_id=dcp_dir)
end_time = time.perf_counter()
logger.info("rank: %s, distributed checkpoint loaded in %.2f seconds",
rank,
end_time - begin_time,
local_main_process_only=False)
logger.info("--> checkpoint loaded from step %s", step)
return step
def normalize_dit_input(model_type, latents, args=None) -> torch.Tensor:
if model_type == "hunyuan_hf" or model_type == "hunyuan":
return latents * 0.476986
+82 -32
View File
@@ -1,23 +1,28 @@
import random
import sys
import time
from collections import deque
from copy import deepcopy
import numpy as np
import torch
from diffusers import FlowMatchEulerDiscreteScheduler
from tqdm.auto import tqdm
from fastvideo.v1.distributed import cleanup_dist_env_and_memory, get_sp_group
from fastvideo.v1.distributed import (cleanup_dist_env_and_memory, get_sp_group,
get_world_group)
from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.schedulers.scheduling_flow_unipc_multistep import (
FlowUniPCMultistepScheduler)
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.wan.wan_pipeline import WanValidationPipeline
from fastvideo.v1.training.training_pipeline import TrainingPipeline
from fastvideo.v1.training.training_utils import (
clip_grad_norm_while_handling_failing_dtensor_cases,
compute_density_for_timestep_sampling, get_sigmas, normalize_dit_input,
save_checkpoint)
compute_density_for_timestep_sampling, get_sigmas, load_checkpoint,
normalize_dit_input, save_checkpoint)
import wandb # isort: skip
@@ -33,6 +38,10 @@ class WanTrainingPipeline(TrainingPipeline):
"""
_required_config_modules = ["scheduler", "transformer"]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
shift=fastvideo_args.flow_shift)
def create_training_stages(self, training_args: TrainingArgs):
"""
May be used in future refactors.
@@ -46,7 +55,7 @@ class WanTrainingPipeline(TrainingPipeline):
args_copy.inference_mode = True
args_copy.vae_config.load_encoder = False
validation_pipeline = WanValidationPipeline.from_pretrained(
args.model_path, args=None, inference_mode=True)
training_args.model_path, args=None, inference_mode=True)
self.validation_pipeline = validation_pipeline
@@ -74,13 +83,25 @@ class WanTrainingPipeline(TrainingPipeline):
total_loss = 0.0
optimizer.zero_grad()
for _ in range(gradient_accumulation_steps):
(
latents,
encoder_hidden_states,
encoder_attention_mask,
infos,
) = next(loader_iter)
# Get next batch, handling epoch boundaries gracefully
batch = next(self.train_loader_iter, None) # type: ignore
if batch is None:
self.current_epoch += 1
logger.info("Starting epoch %s", self.current_epoch)
# Reset iterator for next epoch
self.train_loader_iter = iter(self.train_dataloader)
# Get first batch of new epoch
batch = next(self.train_loader_iter)
latents, encoder_hidden_states, encoder_attention_mask, infos = batch
# logger.info("rank: %s, caption: %s",
# self.rank,
# infos['caption'],
# local_main_process_only=False)
# TODO(will): don't hardcode bfloat16
latents = latents.to(self.training_args.device,
dtype=torch.bfloat16)
encoder_hidden_states = encoder_hidden_states.to(
@@ -126,7 +147,7 @@ class WanTrainingPipeline(TrainingPipeline):
dtype=torch.bfloat16)
with set_forward_context(current_timestep=timesteps,
attn_metadata=None):
model_pred = transformer(**input_kwargs)[0]
model_pred = transformer(**input_kwargs)
if precondition_outputs:
model_pred = noisy_model_input - model_pred * sigmas
@@ -138,8 +159,10 @@ class WanTrainingPipeline(TrainingPipeline):
loss.backward()
avg_loss = loss.detach().clone()
sp_group = get_sp_group()
sp_group.all_reduce(avg_loss, op=torch.distributed.ReduceOp.AVG)
# logger.info(f"rank: {self.rank}, avg_loss: {avg_loss.item()}",
# local_main_process_only=False)
world_group = get_world_group()
world_group.all_reduce(avg_loss, op=torch.distributed.ReduceOp.AVG)
total_loss += avg_loss.item()
# TODO(will): perhaps move this into transformer api so that we can do
@@ -166,7 +189,18 @@ class WanTrainingPipeline(TrainingPipeline):
fastvideo_args: FastVideoArgs,
):
assert self.training_args is not None
noise_random_generator = None
# Set random seeds for deterministic training
seed = self.training_args.seed if self.training_args.seed is not None else 42
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
torch.cuda.manual_seed_all(seed)
noise_random_generator = torch.Generator(device="cpu").manual_seed(seed)
logger.info("Initialized random seeds with seed: %s", seed)
noise_scheduler = FlowMatchEulerDiscreteScheduler()
@@ -178,10 +212,11 @@ class WanTrainingPipeline(TrainingPipeline):
self.training_args.sp_size *
self.training_args.train_sp_batch_size)
logger.info("***** Running training *****")
# logger.info(f" Num examples = {len(train_dataset)}")
# logger.info(f" Dataloader size = {len(train_dataloader)}")
# logger.info(f" Num Epochs = {args.num_train_epochs}")
logger.info(" Resume training from step %s", self.init_steps)
logger.info(" Num examples = %s", len(self.train_dataset))
logger.info(" Dataloader size = %s", len(self.train_dataloader))
logger.info(" Num Epochs = %s", self.num_train_epochs)
logger.info(" Resume training from step %s",
self.init_steps) # type: ignore
logger.info(" Instantaneous batch size per device = %s",
self.training_args.train_batch_size)
logger.info(
@@ -200,11 +235,21 @@ class WanTrainingPipeline(TrainingPipeline):
logger.info(" Master weight dtype: %s",
self.transformer.parameters().__next__().dtype)
# Potentially load in the weights and states from a previous save
if self.training_args.resume_from_checkpoint:
assert NotImplementedError(
"resume_from_checkpoint is not supported now.")
# TODO
logger.info("Loading checkpoint from %s",
self.training_args.resume_from_checkpoint)
resumed_step = load_checkpoint(
self.transformer, self.rank,
self.training_args.resume_from_checkpoint, self.optimizer,
self.train_dataloader, self.lr_scheduler,
noise_random_generator)
if resumed_step > 0:
self.init_steps = resumed_step
logger.info("Successfully resumed from step %s", resumed_step)
else:
logger.warning(
"Failed to load checkpoint, starting from step 0")
self.init_steps = 0
progress_bar = tqdm(
range(0, self.training_args.max_train_steps),
@@ -214,7 +259,7 @@ class WanTrainingPipeline(TrainingPipeline):
disable=self.local_rank > 0,
)
loader_iter = iter(self.train_dataloader)
self.train_loader_iter = iter(self.train_dataloader)
step_times: deque[float] = deque(maxlen=100)
@@ -225,8 +270,9 @@ class WanTrainingPipeline(TrainingPipeline):
gpu_memory_usage = torch.cuda.memory_allocated() / 1024**2
logger.info("GPU memory usage before train_one_step: %s MB",
gpu_memory_usage)
for step in range(self.init_steps + 1, args.max_train_steps + 1):
self._log_validation(self.transformer, self.training_args, 1)
for step in range(self.init_steps + 1,
self.training_args.max_train_steps + 1):
start_time = time.perf_counter()
loss, grad_norm = self.train_one_step(
@@ -235,7 +281,7 @@ class WanTrainingPipeline(TrainingPipeline):
"wan",
self.optimizer,
self.lr_scheduler,
loader_iter,
self.train_loader_iter,
noise_scheduler,
noise_random_generator,
self.training_args.gradient_accumulation_steps,
@@ -258,7 +304,8 @@ class WanTrainingPipeline(TrainingPipeline):
# Manual gradient checking - only at first step
if step == 1 and ENABLE_GRADIENT_CHECK:
logger.info("Performing gradient check at step %s", step)
self.setup_gradient_check(args, loader_iter, noise_scheduler,
self.setup_gradient_check(args, self.train_loader_iter,
noise_scheduler,
noise_random_generator)
progress_bar.set_postfix({
@@ -267,7 +314,7 @@ class WanTrainingPipeline(TrainingPipeline):
"grad_norm": grad_norm,
})
progress_bar.update(1)
if self.rank <= 0:
if self.rank == 0:
wandb.log(
{
"train_loss": loss,
@@ -279,17 +326,20 @@ class WanTrainingPipeline(TrainingPipeline):
step=step,
)
if step % self.training_args.checkpointing_steps == 0:
# Your existing checkpoint saving code
save_checkpoint(self.transformer, self.rank,
self.training_args.output_dir, step)
self.training_args.output_dir, step,
self.optimizer, self.train_dataloader,
self.lr_scheduler, noise_random_generator)
self.transformer.train()
self.sp_group.barrier()
if self.training_args.log_validation and step % self.training_args.validation_steps == 0:
self.log_validation(self.transformer, self.training_args, step)
self._log_validation(self.transformer, self.training_args, step)
save_checkpoint(self.transformer, self.rank,
self.training_args.output_dir,
self.training_args.max_train_steps)
self.training_args.max_train_steps, self.optimizer,
self.train_dataloader, self.lr_scheduler,
noise_random_generator)
if get_sp_group():
cleanup_dist_env_and_memory()
+97 -10
View File
@@ -10,10 +10,12 @@ import json
import math
import os
import signal
import socket
import sys
import tempfile
import threading
import traceback
from dataclasses import fields, is_dataclass
from dataclasses import dataclass, fields, is_dataclass
from functools import partial, wraps
from typing import (Any, Callable, Dict, List, Optional, Tuple, Type, TypeVar,
Union, cast)
@@ -22,7 +24,10 @@ import cloudpickle
import filelock
import torch
import yaml
from diffusers.loaders.lora_base import (
_best_guess_weight_name) # watch out for potetential removal from diffusers
from huggingface_hub import snapshot_download
from remote_pdb import RemotePdb
import fastvideo.v1.envs as envs
from fastvideo.v1.logger import init_logger
@@ -448,14 +453,14 @@ def import_pynvml():
return pynvml
def maybe_download_model(model_path: str,
def maybe_download_model(model_name_or_path: str,
local_dir: Optional[str] = None,
download: bool = True) -> str:
"""
Check if the model path is a Hugging Face Hub model ID and download it if needed.
Args:
model_path: Local path or Hugging Face Hub model ID
model_name_or_path: Local path or Hugging Face Hub model ID
local_dir: Local directory to save the model
download: Whether to download the model from Hugging Face Hub
@@ -464,27 +469,47 @@ def maybe_download_model(model_path: str,
"""
# If the path exists locally, return it
if os.path.exists(model_path):
logger.info("Model already exists locally at %s", model_path)
return model_path
if os.path.exists(model_name_or_path):
logger.info("Model already exists locally at %s", model_name_or_path)
return model_name_or_path
# Otherwise, assume it's a HF Hub model ID and try to download it
try:
logger.info("Downloading model snapshot from HF Hub for %s...",
model_path)
with get_lock(model_path):
model_name_or_path)
with get_lock(model_name_or_path):
local_path = snapshot_download(
repo_id=model_path,
repo_id=model_name_or_path,
ignore_patterns=["*.onnx", "*.msgpack"],
local_dir=local_dir)
logger.info("Downloaded model to %s", local_path)
return str(local_path)
except Exception as e:
raise ValueError(
f"Could not find model at {model_path} and failed to download from HF Hub: {e}"
f"Could not find model at {model_name_or_path} and failed to download from HF Hub: {e}"
) from e
def maybe_download_lora(model_name_or_path: str,
local_dir: Optional[str] = None,
download: bool = True) -> str:
"""
Check if the model path is a Hugging Face Hub model ID and download it if needed.
Args:
model_name_or_path: Local path or Hugging Face Hub model ID
local_dir: Local directory to save the model
download: Whether to download the model from Hugging Face Hub
Returns:
Local path to the model
"""
local_path = maybe_download_model(model_name_or_path, local_dir, download)
weight_name = _best_guess_weight_name(model_name_or_path,
file_extension=".safetensors")
return os.path.join(local_path, weight_name)
def verify_model_config_and_directory(model_path: str) -> Dict[str, Any]:
"""
Verify that the model directory contains a valid diffusers configuration.
@@ -645,3 +670,65 @@ class TypeBasedDispatcher:
if isinstance(obj, ty):
return fn(obj)
raise ValueError(f"Invalid object: {obj}")
# For non-torch.distributed debugging
def remote_breakpoint() -> None:
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
s.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
s.bind(("localhost", 0)) # Let the OS pick an ephemeral port.
port = s.getsockname()[1]
RemotePdb(host="localhost", port=port).set_trace()
@dataclass
class MixedPrecisionState:
master_dtype: Optional[torch.dtype] = None
param_dtype: Optional[torch.dtype] = None
reduce_dtype: Optional[torch.dtype] = None
output_dtype: Optional[torch.dtype] = None
compute_dtype: Optional[torch.dtype] = None
# Thread-local storage for mixed precision state
_mixed_precision_state = threading.local()
def get_mixed_precision_state() -> MixedPrecisionState:
"""Get the current mixed precision state."""
if not hasattr(_mixed_precision_state, 'state'):
raise ValueError("Mixed precision state not set")
return cast(MixedPrecisionState, _mixed_precision_state.state)
def set_mixed_precision_policy(master_dtype: torch.dtype,
param_dtype: torch.dtype,
reduce_dtype: torch.dtype,
output_dtype: Optional[torch.dtype] = None):
"""Set mixed precision policy globally.
Args:
param_dtype: Parameter dtype used for training
reduce_dtype: Reduction dtype used for gradients
output_dtype: Optional output dtype
"""
state = MixedPrecisionState(
master_dtype=master_dtype,
param_dtype=param_dtype,
reduce_dtype=reduce_dtype,
output_dtype=output_dtype,
)
_mixed_precision_state.state = state
def get_compute_dtype() -> torch.dtype:
"""Get the current compute dtype from mixed precision policy.
Returns:
torch.dtype: The compute dtype to use, defaults to get_default_dtype() if no policy set
"""
if not hasattr(_mixed_precision_state, 'state'):
return torch.get_default_dtype()
else:
state = get_mixed_precision_state()
return state.param_dtype
+7
View File
@@ -47,6 +47,13 @@ class Executor(ABC):
})
return cast(ForwardBatch, outputs[0]["output_batch"])
@abstractmethod
def set_lora_adapter(self, lora_nickname: str, lora_path: str) -> None:
"""
Set the LoRA adapter for the workers.
"""
raise NotImplementedError
@abstractmethod
def collective_rpc(self,
method: Union[str, Callable[..., _R]],
+10 -5
View File
@@ -5,6 +5,7 @@ import multiprocessing as mp
import os
import signal
import sys
from multiprocessing.connection import Connection
from typing import Any, Dict, Optional, TextIO, cast
import psutil
@@ -29,14 +30,14 @@ RESET = '\033[0;0m'
class Worker:
def __init__(self, fastvideo_args: FastVideoArgs, local_rank: int,
rank: int, pipe):
rank: int, pipe: Connection, master_port: int):
self.fastvideo_args = fastvideo_args
self.local_rank = local_rank
self.rank = rank
# TODO(will): don't hardcode this
self.distributed_init_method = "env://"
self.pipe = pipe
self.master_port = master_port
self.init_device()
# Init request dispatcher
@@ -76,7 +77,7 @@ class Worker:
f"Unsupported device: {self.fastvideo_args.device_str}")
os.environ["MASTER_ADDR"] = "localhost"
os.environ["MASTER_PORT"] = "29503"
os.environ["MASTER_PORT"] = str(self.master_port)
os.environ["LOCAL_RANK"] = str(self.local_rank)
os.environ["RANK"] = str(self.rank)
@@ -92,6 +93,9 @@ class Worker:
output_batch = self.pipeline.forward(forward_batch, self.fastvideo_args)
return cast(ForwardBatch, output_batch)
def set_lora_adapter(self, lora_nickname: str, lora_path: str) -> None:
self.pipeline.set_lora_adapter(lora_nickname, lora_path)
def shutdown(self) -> Dict[str, Any]:
"""Gracefully shut down the worker process"""
logger.info("Worker %d shutting down...",
@@ -191,7 +195,7 @@ def init_worker_distributed_environment(
def run_worker_process(fastvideo_args: FastVideoArgs, local_rank: int,
rank: int, pipe):
rank: int, pipe: Connection, master_port: int):
# Add process-specific prefix to stdout and stderr
process_name = mp.current_process().name
pid = os.getpid()
@@ -206,8 +210,9 @@ def run_worker_process(fastvideo_args: FastVideoArgs, local_rank: int,
logger.info("Worker %d initializing...",
rank,
local_main_process_only=False)
try:
worker = Worker(fastvideo_args, local_rank, rank, pipe)
worker = Worker(fastvideo_args, local_rank, rank, pipe, master_port)
logger.info("Worker %d sending ready", rank)
pipe.send({
"status": "ready",
+19 -1
View File
@@ -3,6 +3,7 @@ import contextlib
import multiprocessing as mp
import os
import signal
import socket
import time
from multiprocessing.process import BaseProcess
from typing import Any, Callable, List, Optional, Union, cast
@@ -28,6 +29,15 @@ class MultiprocExecutor(Executor):
self.workers: List[BaseProcess] = []
self.worker_pipes = []
self.master_port = None
for port in range(29503, 65535):
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
if s.connect_ex(('localhost', port)) != 0:
self.master_port = port
break
if self.master_port is None:
raise ValueError("No unused port found to use as master port")
# Create pipes and start workers
for rank in range(self.world_size):
@@ -39,7 +49,8 @@ class MultiprocExecutor(Executor):
kwargs=dict(fastvideo_args=self.fastvideo_args,
local_rank=rank,
rank=rank,
pipe=worker_pipe))
pipe=worker_pipe,
master_port=self.master_port))
worker.start()
self.workers.append(worker)
@@ -62,6 +73,13 @@ class MultiprocExecutor(Executor):
})
return cast(ForwardBatch, responses[0]["output_batch"])
def set_lora_adapter(self, lora_nickname: str, lora_path: str) -> None:
self.collective_rpc("set_lora_adapter",
kwargs={
"lora_nickname": lora_nickname,
"lora_path": lora_path
})
def collective_rpc(self,
method: Union[str, Callable],
timeout: Optional[float] = None,
+8 -7
View File
@@ -37,21 +37,22 @@ dependencies = [
"gradio>=5.22.0", "moviepy==1.0.3", "flask",
"flask_restful", "aiohttp", "huggingface_hub", "cloudpickle",
# System & Monitoring Tools
"gpustat", "watch",
"gpustat", "watch", "remote-pdb",
# Kernel & Packaging
"wheel",
# Training Dependencies
"torchdata",
"pyarrow",
"datasets",
"av",
]
[project.optional-dependencies]
# flash-attn: pip install flash-attn==2.7.4.post1 --no-cache-dir --no-build-isolation
train = [
"torchdata",
"pyarrow",
"datasets",
]
lint = [
"pre-commit==4.0.1",
@@ -63,7 +64,7 @@ test = [
"pytest",
]
dev = [ "fastvideo[lint]", "fastvideo[test]", "fastvideo[train]", ]
dev = [ "fastvideo[lint]", "fastvideo[test]", ]
[project.scripts]
fastvideo = "fastvideo.v1.entrypoints.cli.main:main"
+39
View File
@@ -0,0 +1,39 @@
#https://huggingface.co/Wan-AI/Wan2.1-T2V-1.3B-Diffusers/tree/main
DATA_DIR=./data
torchrun --nnodes 1 --nproc_per_node 8\
fastvideo/distill.py\
--seed 42\
--pretrained_model_name_or_path $DATA_DIR/wan\
--model_type "wan" \
--cache_dir "$DATA_DIR/.cache"\
--data_json_path "$DATA_DIR/Image-Vid-Finetune-Wan/videos2caption.json"\
--validation_prompt_dir "$DATA_DIR/Image-Vid-Finetune-Wan/validation"\
--gradient_checkpointing\
--train_batch_size=1\
--num_latent_t 32 \
--sp_size 1 \
--train_sp_batch_size 1\
--dataloader_num_workers 4\
--gradient_accumulation_steps=1\
--max_train_steps=320\
--learning_rate=1e-6\
--mixed_precision="bf16"\
--checkpointing_steps=64\
--validation_steps 64\
--validation_sampling_steps "8" \
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--cfg 5.0\
--log_validation\
--output_dir="$DATA_DIR/outputs/wan"\
--tracker_project_name Wan_PCM \
--num_height 480 \
--num_width 832 \
--num_frames 81 \
--shift 3 \
--validation_guidance_scale "1.0" \
--num_euler_timesteps 50 \
--multi_phased_distill_schedule "4000-1" \
--not_apply_cfg_solver
+5 -1
View File
@@ -8,6 +8,8 @@ NUM_GPUS=1
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
CHECKPOINT_PATH="$DATA_DIR/outputs/wan_finetune/checkpoint-5"
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
torchrun --nnodes 1 --nproc_per_node $NUM_GPUS\
fastvideo/v1/training/wan_training_pipeline.py\
@@ -47,4 +49,6 @@ torchrun --nnodes 1 --nproc_per_node $NUM_GPUS\
--weight_decay 0.01 \
--not_apply_cfg_solver \
--master_weight_type "fp32" \
--max_grad_norm 1.0
--max_grad_norm 1.0 \
# --resume_from_checkpoint "$CHECKPOINT_PATH"
+11 -13
View File
@@ -6,19 +6,17 @@ export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
# Note that the tp_size and sp_size should be the same and equal to the number
# of GPUs. They are used for different parallel groups. sp_size is used for
# dit model and tp_size is used for encoder models.
torchrun --nnodes=1 --nproc_per_node=$num_gpus --master_port 29503 \
fastvideo/v1/sample/v1_fastvideo_inference.py \
--sp_size $num_gpus \
--tp_size $num_gpus \
fastvideo generate \
--model-path $MODEL_BASE \
--sp-size $num_gpus \
--tp-size $num_gpus \
--height 720 \
--width 1280 \
--num_frames 125 \
--num_inference_steps 6 \
--guidance_scale 1 \
--embedded_cfg_scale 6 \
--flow_shift 17 \
--prompt_path ./assets/prompt.txt \
--num-frames 125 \
--num-inference-steps 6 \
--guidance-scale 1 \
--embedded-cfg-scale 6 \
--flow-shift 17 \
--prompt "A beautiful woman in a red dress walking down a street" \
--seed 1024 \
--output_path outputs_video/ \
--model_path $MODEL_BASE \
--vae-sp
--output-path outputs_video/
+11 -13
View File
@@ -7,19 +7,17 @@ export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
# Note that the tp_size and sp_size should be the same and equal to the number
# of GPUs. They are used for different parallel groups. sp_size is used for
# dit model and tp_size is used for encoder models.
torchrun --nnodes=1 --nproc_per_node=$num_gpus --master_port 29503 \
fastvideo/v1/sample/v1_fastvideo_inference.py \
--sp_size $num_gpus \
--tp_size $num_gpus \
fastvideo generate \
--model-path $MODEL_BASE \
--sp-size $num_gpus \
--tp-size $num_gpus \
--height 720 \
--width 1280 \
--num_frames 125 \
--num_inference_steps 50 \
--guidance_scale 1 \
--embedded_cfg_scale 6 \
--flow_shift 7 \
--prompt_path ./assets/prompt.txt \
--num-frames 125 \
--num-inference-steps 50 \
--guidance-scale 1 \
--embedded-cfg-scale 6 \
--flow-shift 7 \
--prompt "A beautiful woman in a red dress walking down a street" \
--seed 1024 \
--output_path outputs_video/ \
--model_path $MODEL_BASE \
--vae-sp
--output-path outputs_video/
+11 -13
View File
@@ -8,19 +8,17 @@ export MODEL_BASE=hunyuanvideo-community/HunyuanVideo
# Note that the tp_size and sp_size should be the same and equal to the number
# of GPUs. They are used for different parallel groups. sp_size is used for
# dit model and tp_size is used for encoder models.
torchrun --nnodes=1 --nproc_per_node=$num_gpus --master_port 29503 \
fastvideo/v1/sample/v1_fastvideo_inference.py \
--sp_size ${num_gpus} \
--tp_size ${num_gpus} \
fastvideo generate \
--model-path $MODEL_BASE \
--sp-size ${num_gpus} \
--tp-size ${num_gpus} \
--height 768 \
--width 1280 \
--num_frames 117 \
--num_inference_steps 50 \
--guidance_scale 1 \
--embedded_cfg_scale 6 \
--flow_shift 7 \
--prompt_path ./assets/prompt.txt \
--num-frames 117 \
--num-inference-steps 50 \
--guidance-scale 1 \
--embedded-cfg-scale 6 \
--flow-shift 7 \
--prompt "A beautiful woman in a red dress walking down a street" \
--seed 1024 \
--output_path outputs_video/ \
--model_path $MODEL_BASE \
--vae-sp
--output-path outputs_video/
+12 -13
View File
@@ -5,19 +5,18 @@ export FASTVIDEO_ATTENTION_BACKEND=
num_gpus=2
url='127.0.0.1'
model_dir=data/stepvideo-t2v
torchrun --nnodes=1 --nproc_per_node=$num_gpus --master_port 29503 \
fastvideo/v1/sample/v1_fastvideo_inference.py \
--sp_size ${num_gpus} \
--tp_size ${num_gpus} \
fastvideo generate \
--model-path $model_dir \
--sp-size ${num_gpus} \
--tp-size ${num_gpus} \
--height 256 \
--width 256 \
--num_frames 29 \
--num_inference_steps 50 \
--embedded_cfg_scale 9.0 \
--guidance_scale 9.0 \
--prompt_path ./assets/prompt.txt \
--num-frames 29 \
--num-inference-steps 50 \
--embedded-cfg-scale 9.0 \
--guidance-scale 9.0 \
--prompt "A beautiful woman in a red dress walking down a street" \
--seed 1024 \
--output_path outputs_stepvideo/ \
--model_path $model_dir \
--flow_shift 13.0 \
--vae_precision bf16
--output-path outputs_stepvideo/ \
--flow-shift 13.0 \
--vae-precision bf16
+10 -14
View File
@@ -7,21 +7,17 @@ export MODEL_BASE=Wan-AI/Wan2.1-T2V-1.3B-Diffusers
# Note that the tp_size and sp_size should be the same and equal to the number
# of GPUs. They are used for different parallel groups. sp_size is used for
# dit model and tp_size is used for encoder models.
torchrun --nnodes=1 --nproc_per_node=$num_gpus --master_port 29503 \
fastvideo/v1/sample/v1_fastvideo_inference.py \
--sp_size $num_gpus \
--tp_size $num_gpus \
fastvideo generate \
--model-path $MODEL_BASE \
--sp-size $num_gpus \
--tp-size $num_gpus \
--height 480 \
--width 832 \
--num_frames 77 \
--num_inference_steps 50 \
--num-frames 77 \
--num-inference-steps 50 \
--fps 16 \
--guidance_scale 3.0 \
--prompt_path ./assets/prompt.txt \
--neg_prompt "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards" \
--guidance-scale 3.0 \
--prompt "A beautiful woman in a red dress walking down a street" \
--negative-prompt "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards" \
--seed 1024 \
--output_path outputs_video/ \
--model_path $MODEL_BASE \
--vae-sp \
--text-encoder-precision "fp32" \
--use-cpu-offload
--output-path outputs_video/
+10 -14
View File
@@ -7,21 +7,17 @@ export MODEL_BASE=Wan-AI/Wan2.1-T2V-14B-Diffusers
# Note that the tp_size and sp_size should be the same and equal to the number
# of GPUs. They are used for different parallel groups. sp_size is used for
# dit model and tp_size is used for encoder models.
torchrun --nnodes=1 --nproc_per_node=$num_gpus --master_port 29503 \
fastvideo/v1/sample/v1_fastvideo_inference.py \
--sp_size $num_gpus \
--tp_size $num_gpus \
fastvideo generate \
--model-path $MODEL_BASE \
--sp-size $num_gpus \
--tp-size $num_gpus \
--height 768 \
--width 1280 \
--num_frames 69 \
--num_inference_steps 50 \
--num-frames 69 \
--num-inference-steps 50 \
--fps 16 \
--guidance_scale 5.0 \
--prompt_path ./assets/prompt.txt \
--neg_prompt "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards" \
--guidance-scale 5.0 \
--prompt "A beautiful woman in a red dress walking down a street" \
--negative-prompt "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards" \
--seed 12345 \
--output_path outputs_video/ \
--model_path $MODEL_BASE \
--vae-sp \
--text-encoder-precision "fp32" \
--use-cpu-offload
--output-path outputs_video/
+11 -15
View File
@@ -7,23 +7,19 @@ export MODEL_BASE=Wan-AI/Wan2.1-I2V-14B-480P-Diffusers
# Note that the tp_size and sp_size should be the same and equal to the number
# of GPUs. They are used for different parallel groups. sp_size is used for
# dit model and tp_size is used for encoder models.
torchrun --nnodes=1 --nproc_per_node=$num_gpus --master_port 29503 \
fastvideo/v1/sample/v1_fastvideo_inference.py \
--sp_size $num_gpus \
--tp_size $num_gpus \
fastvideo generate \
--model-path $MODEL_BASE \
--sp-size $num_gpus \
--tp-size $num_gpus \
--height 480 \
--width 832 \
--num_frames 77 \
--num_inference_steps 40 \
--num-frames 77 \
--num-inference-steps 40 \
--fps 16 \
--flow_shift 3.0 \
--guidance_scale 5.0 \
--image_path "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg" \
--flow-shift 3.0 \
--guidance-scale 5.0 \
--image-path "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg" \
--prompt "An astronaut hatching from an egg, on the surface of the moon, the darkness and depth of space realised in the background. High quality, ultrarealistic detail and breath-taking movie-like camera shot." \
--neg_prompt "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards" \
--negative-prompt "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards" \
--seed 1024 \
--output_path outputs_i2v/ \
--model_path $MODEL_BASE \
--vae-sp \
--text-encoder-precision "fp32" \
--use-cpu-offload
--output-path outputs_i2v/
+20 -10
View File
@@ -1,23 +1,33 @@
# export WANDB_MODE="offline"
GPU_NUM=1 # 2,4,8
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
MODEL_PATH="data/wan-diffusers"
MODEL_TYPE="wan"
DATA_MERGE_PATH="your/path/to/Mixkit-Src/merge.txt"
OUTPUT_DIR="your/path"
DATA_MERGE_PATH="data/Image-Vid-Finetune-Src/merge.txt"
OUTPUT_DIR="data/Image-Vid-Finetune-Wan"
VALIDATION_PATH="assets/prompt.txt"
torchrun --nproc_per_node=$GPU_NUM \
fastvideo/data_preprocess/preprocess.py \
fastvideo/data_preprocess/preprocess_vae_latents.py \
--model_path $MODEL_PATH \
--data_merge_path $DATA_MERGE_PATH \
--preprocess_video_batch_size=4 \
--train_batch_size=1 \
--max_height=480 \
--max_width=832 \
--max_width=848 \
--num_frames=81 \
--dataloader_num_workers 1 \
--output_dir=$OUTPUT_DIR \
--model_type $MODEL_TYPE \
--train_fps 16 \
--validation_prompt_txt $VALIDATION_PATH \
--samples_per_file 108 \
--flush_frequency 108
--train_fps 24
torchrun --nproc_per_node=$GPU_NUM \
fastvideo/data_preprocess/preprocess_text_embeddings.py \
--model_type $MODEL_TYPE \
--model_path $MODEL_PATH \
--output_dir=$OUTPUT_DIR
torchrun --nproc_per_node=1 \
fastvideo/data_preprocess/preprocess_validation_text_embeddings.py \
--model_type $MODEL_TYPE \
--model_path $MODEL_PATH \
--output_dir=$OUTPUT_DIR \
--validation_prompt_txt $VALIDATION_PATH
@@ -0,0 +1,24 @@
# export WANDB_MODE="offline"
GPU_NUM=1 # 2,4,8
MODEL_PATH="Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
MODEL_TYPE="wan"
DATA_MERGE_PATH="finetrainers/crush-smol/merge.txt"
OUTPUT_DIR="crush-smol_preprocess"
VALIDATION_PATH="assets/prompt.txt"
torchrun --nproc_per_node=$GPU_NUM \
fastvideo/data_preprocess/preprocess.py \
--model_path $MODEL_PATH \
--data_merge_path $DATA_MERGE_PATH \
--preprocess_video_batch_size 8 \
--max_height 480 \
--max_width 832 \
--num_frames 77 \
--dataloader_num_workers 0 \
--output_dir=$OUTPUT_DIR \
--model_type $MODEL_TYPE \
--train_fps 16 \
--validation_prompt_txt $VALIDATION_PATH \
--samples_per_file 16 \
--flush_frequency 32 \
--preprocess_task "i2v"
@@ -0,0 +1,24 @@
# export WANDB_MODE="offline"
GPU_NUM=1 # 2,4,8
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
MODEL_TYPE="wan"
DATA_MERGE_PATH="data/crush-smol/merge.txt"
OUTPUT_DIR="data/crush-smol/latents"
VALIDATION_PATH="assets/prompt.txt"
torchrun --nproc_per_node=$GPU_NUM \
fastvideo/data_preprocess/preprocess.py \
--model_path $MODEL_PATH \
--data_merge_path $DATA_MERGE_PATH \
--preprocess_video_batch_size 8 \
--max_height 480 \
--max_width 832 \
--num_frames 81 \
--dataloader_num_workers 0 \
--output_dir=$OUTPUT_DIR \
--model_type $MODEL_TYPE \
--train_fps 16 \
--validation_prompt_txt $VALIDATION_PATH \
--samples_per_file 32 \
--flush_frequency 128 \
--preprocess_task "t2v"