From 37bd73564702fa47ffe4576c87da6f82a50cb60e Mon Sep 17 00:00:00 2001 From: Bubbliiiing <47347516+bubbliiiing@users.noreply.github.com> Date: Tue, 9 Sep 2025 10:20:23 +0800 Subject: [PATCH] Update VACE (#306) --- comfyui/cogvideox_fun/nodes.py | 12 +- comfyui/wan2_1/nodes.py | 12 +- comfyui/wan2_1_fun/nodes.py | 12 +- comfyui/wan2_2/nodes.py | 12 +- config/wan2.2/wan_civitai_5b.yaml | 2 +- config/wan2.2/wan_civitai_s2v.yaml | 44 ++ examples/wan2.1_vace/predict_i2v.py | 311 ++++++++ examples/wan2.1_vace/predict_s2v.py | 311 ++++++++ examples/wan2.1_vace/predict_v2v_control.py | 311 ++++++++ examples/wan2.2/predict_s2v.py | 23 +- examples/wan2.2_fun/predict_t2v.py | 2 +- videox_fun/models/__init__.py | 3 +- videox_fun/models/cache_utils.py | 5 +- videox_fun/models/wan_transformer3d_s2v.py | 19 +- videox_fun/models/wan_transformer3d_vace.py | 354 +++++++++ videox_fun/pipeline/__init__.py | 7 +- videox_fun/pipeline/pipeline_wan2_2_s2v.py | 73 +- videox_fun/pipeline/pipeline_wan_vace.py | 785 ++++++++++++++++++++ 18 files changed, 2201 insertions(+), 97 deletions(-) create mode 100644 config/wan2.2/wan_civitai_s2v.yaml create mode 100644 examples/wan2.1_vace/predict_i2v.py create mode 100644 examples/wan2.1_vace/predict_s2v.py create mode 100644 examples/wan2.1_vace/predict_v2v_control.py create mode 100644 videox_fun/models/wan_transformer3d_vace.py create mode 100644 videox_fun/pipeline/pipeline_wan_vace.py diff --git a/comfyui/cogvideox_fun/nodes.py b/comfyui/cogvideox_fun/nodes.py index 963bf22..1d1dd5d 100755 --- a/comfyui/cogvideox_fun/nodes.py +++ b/comfyui/cogvideox_fun/nodes.py @@ -104,10 +104,16 @@ class LoadCogVideoXFunModel: if os.path.exists(candidate_path): model_name = candidate_path break - + try: + if os.path.exists(eas_cache_dir): + list_dirs = os.listdir(eas_cache_dir) + else: + list_dirs = [] + except: + list_dirs = [] # If model_name is still None, check eas_cache_dir for each possible folder if model_name is None and os.path.exists(eas_cache_dir): - for folder in possible_folders: + for folder in possible_folders + list_dirs: candidate_path = os.path.join(eas_cache_dir, folder, model) if os.path.exists(candidate_path): model_name = candidate_path @@ -115,7 +121,7 @@ class LoadCogVideoXFunModel: # If model_name is still None, prompt the user to download the model if model_name is None: - print(f"Please download cogvideoxfun model to one of the following directories:") + print(f"Please download videoxfun model to one of the following directories:") for folder in possible_folders: print(f"- {os.path.join(folder_paths.models_dir, folder)}") if os.path.exists(eas_cache_dir): diff --git a/comfyui/wan2_1/nodes.py b/comfyui/wan2_1/nodes.py index fda736a..1cd91d3 100755 --- a/comfyui/wan2_1/nodes.py +++ b/comfyui/wan2_1/nodes.py @@ -112,10 +112,16 @@ class LoadWanModel: if os.path.exists(candidate_path): model_name = candidate_path break - + try: + if os.path.exists(eas_cache_dir): + list_dirs = os.listdir(eas_cache_dir) + else: + list_dirs = [] + except: + list_dirs = [] # If model_name is still None, check eas_cache_dir for each possible folder if model_name is None and os.path.exists(eas_cache_dir): - for folder in possible_folders: + for folder in possible_folders + list_dirs: candidate_path = os.path.join(eas_cache_dir, folder, model) if os.path.exists(candidate_path): model_name = candidate_path @@ -123,7 +129,7 @@ class LoadWanModel: # If model_name is still None, prompt the user to download the model if model_name is None: - print(f"Please download cogvideoxfun model to one of the following directories:") + print(f"Please download videoxfun model to one of the following directories:") for folder in possible_folders: print(f"- {os.path.join(folder_paths.models_dir, folder)}") if os.path.exists(eas_cache_dir): diff --git a/comfyui/wan2_1_fun/nodes.py b/comfyui/wan2_1_fun/nodes.py index 7374e3d..9653ac8 100755 --- a/comfyui/wan2_1_fun/nodes.py +++ b/comfyui/wan2_1_fun/nodes.py @@ -118,10 +118,16 @@ class LoadWanFunModel: if os.path.exists(candidate_path): model_name = candidate_path break - + try: + if os.path.exists(eas_cache_dir): + list_dirs = os.listdir(eas_cache_dir) + else: + list_dirs = [] + except: + list_dirs = [] # If model_name is still None, check eas_cache_dir for each possible folder if model_name is None and os.path.exists(eas_cache_dir): - for folder in possible_folders: + for folder in possible_folders + list_dirs: candidate_path = os.path.join(eas_cache_dir, folder, model) if os.path.exists(candidate_path): model_name = candidate_path @@ -129,7 +135,7 @@ class LoadWanFunModel: # If model_name is still None, prompt the user to download the model if model_name is None: - print(f"Please download cogvideoxfun model to one of the following directories:") + print(f"Please download videoxfun model to one of the following directories:") for folder in possible_folders: print(f"- {os.path.join(folder_paths.models_dir, folder)}") if os.path.exists(eas_cache_dir): diff --git a/comfyui/wan2_2/nodes.py b/comfyui/wan2_2/nodes.py index 291e744..4dc4135 100755 --- a/comfyui/wan2_2/nodes.py +++ b/comfyui/wan2_2/nodes.py @@ -115,10 +115,16 @@ class LoadWan2_2Model: if os.path.exists(candidate_path): model_name = candidate_path break - + try: + if os.path.exists(eas_cache_dir): + list_dirs = os.listdir(eas_cache_dir) + else: + list_dirs = [] + except: + list_dirs = [] # If model_name is still None, check eas_cache_dir for each possible folder if model_name is None and os.path.exists(eas_cache_dir): - for folder in possible_folders: + for folder in possible_folders + list_dirs: candidate_path = os.path.join(eas_cache_dir, folder, model) if os.path.exists(candidate_path): model_name = candidate_path @@ -126,7 +132,7 @@ class LoadWan2_2Model: # If model_name is still None, prompt the user to download the model if model_name is None: - print(f"Please download cogvideoxfun model to one of the following directories:") + print(f"Please download videoxfun model to one of the following directories:") for folder in possible_folders: print(f"- {os.path.join(folder_paths.models_dir, folder)}") if os.path.exists(eas_cache_dir): diff --git a/config/wan2.2/wan_civitai_5b.yaml b/config/wan2.2/wan_civitai_5b.yaml index 137bf5d..f465e28 100644 --- a/config/wan2.2/wan_civitai_5b.yaml +++ b/config/wan2.2/wan_civitai_5b.yaml @@ -30,7 +30,7 @@ text_encoder_kwargs: scheduler_kwargs: scheduler_subpath: null num_train_timesteps: 1000 - shift: 12.0 + shift: 5.0 use_dynamic_shifting: false base_shift: 0.5 max_shift: 1.15 diff --git a/config/wan2.2/wan_civitai_s2v.yaml b/config/wan2.2/wan_civitai_s2v.yaml new file mode 100644 index 0000000..81d9921 --- /dev/null +++ b/config/wan2.2/wan_civitai_s2v.yaml @@ -0,0 +1,44 @@ +format: civitai +pipeline: Wan +transformer_additional_kwargs: + transformer_low_noise_model_subpath: ./ + transformer_combination_type: "single" + dict_mapping: + in_dim: in_channels + dim: hidden_size + +vae_kwargs: + vae_type: "AutoencoderKLWan" + vae_subpath: Wan2.1_VAE.pth + temporal_compression_ratio: 4 + spatial_compression_ratio: 8 + +text_encoder_kwargs: + text_encoder_subpath: models_t5_umt5-xxl-enc-bf16.pth + tokenizer_subpath: google/umt5-xxl + text_length: 512 + vocab: 256384 + dim: 4096 + dim_attn: 4096 + dim_ffn: 10240 + num_heads: 64 + num_layers: 24 + num_buckets: 32 + shared_pos: False + dropout: 0.0 + +audio_encoder_kwargs: + audio_encoder_subpath: wav2vec2-large-xlsr-53-english + +scheduler_kwargs: + scheduler_subpath: null + num_train_timesteps: 1000 + shift: 3.0 + use_dynamic_shifting: false + base_shift: 0.5 + max_shift: 1.15 + base_image_seq_len: 256 + max_image_seq_len: 4096 + +image_encoder_kwargs: + image_encoder_subpath: models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth \ No newline at end of file diff --git a/examples/wan2.1_vace/predict_i2v.py b/examples/wan2.1_vace/predict_i2v.py new file mode 100644 index 0000000..47667b3 --- /dev/null +++ b/examples/wan2.1_vace/predict_i2v.py @@ -0,0 +1,311 @@ +import os +import sys + +import numpy as np +import torch +from diffusers import FlowMatchEulerDiscreteScheduler +from omegaconf import OmegaConf +from PIL import Image +from transformers import AutoTokenizer + +current_file_path = os.path.abspath(__file__) +project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))] +for project_root in project_roots: + sys.path.insert(0, project_root) if project_root not in sys.path else None + +from videox_fun.dist import set_multi_gpus_devices, shard_model +from videox_fun.models import (AutoencoderKLWan, AutoTokenizer, + WanT5EncoderModel, VaceWanTransformer3DModel) +from videox_fun.data.dataset_image_video import process_pose_file +from videox_fun.models.cache_utils import get_teacache_coefficients +from videox_fun.pipeline import WanVacePipeline, WanPipeline +from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, + convert_weight_dtype_wrapper, + replace_parameters_by_name) +from videox_fun.utils.lora_utils import merge_lora, unmerge_lora +from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent, get_image_latent, + get_video_to_video_latent, + save_videos_grid) +from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler +from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler + +# GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, sequential_cpu_offload]. +# model_full_load means that the entire model will be moved to the GPU. +# +# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU, +# and the transformer model has been quantized to float8, which can save more GPU memory. +# +# model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory. +# +# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use, +# and the transformer model has been quantized to float8, which can save more GPU memory. +# +# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use, +# resulting in slower speeds but saving a large amount of GPU memory. +GPU_memory_mode = "sequential_cpu_offload" +# Multi GPUs config +# Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used. +# For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4. +# If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1. +ulysses_degree = 1 +ring_degree = 1 +# Use FSDP to save more GPU memory in multi gpus. +fsdp_dit = False +fsdp_text_encoder = True +# Compile will give a speedup in fixed resolution and need a little GPU memory. +# The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload. +compile_dit = False + +# Support TeaCache. +enable_teacache = True +# Recommended to be set between 0.05 and 0.30. A larger threshold can cache more steps, speeding up the inference process, +# but it may cause slight differences between the generated content and the original content. +# # --------------------------------------------------------------------------------------------------- # +# | Model Name | threshold | Model Name | threshold | +# | Wan2.1-VACE-1.3B | 0.05~0.10 | Wan2.1-VACE-14B | 0.10~0.15 | +# # --------------------------------------------------------------------------------------------------- # +teacache_threshold = 0.05 +# The number of steps to skip TeaCache at the beginning of the inference process, which can +# reduce the impact of TeaCache on generated video quality. +num_skip_start_steps = 5 +# Whether to offload TeaCache tensors to cpu to save a little bit of GPU memory. +teacache_offload = False + +# Skip some cfg steps in inference for acceleration +# Recommended to be set between 0.00 and 0.25 +cfg_skip_ratio = 0 + +# Riflex config +enable_riflex = False +# Index of intrinsic frequency +riflex_k = 6 + +# Config and model path +config_path = "config/wan2.1/wan_civitai.yaml" +# model path +model_name = "models/Diffusion_Transformer/Wan2.1-VACE-1.3B" + +# Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++" +sampler_name = "Flow_Unipc" +# [NOTE]: Noise schedule shift parameter. Affects temporal dynamics. +# Used when the sampler is in "Flow_Unipc", "Flow_DPM++". +shift = 16 + +# Load pretrained model if need +transformer_path = None +vae_path = None +lora_path = None + +# Other params +sample_size = [480, 832] +video_length = 81 +fps = 16 + +# Use torch.float16 if GPU does not support torch.bfloat16 +# ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16 +weight_dtype = torch.bfloat16 +control_video = None +start_image = "asset/1.png" +end_image = None +subject_ref_images = None +vace_context_scale = 1.00 + +# 使用更长的neg prompt如"模糊,突变,变形,失真,画面暗,文本字幕,画面固定,连环画,漫画,线稿,没有主体。",可以增加稳定性 +# 在neg prompt中添加"安静,固定"等词语可以增加动态性。 +prompt = "一只棕色的狗舔了一下它的舌头,坐在舒适房间里的浅色沙发上。在狗的后面,架子上有一幅镶框的画,周围是粉红色的花朵。房间里柔和温暖的灯光营造出舒适的氛围。" +negative_prompt = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走" + +# Using longer neg prompt such as "Blurring, mutation, deformation, distortion, dark and solid, comics, text subtitles, line art." can increase stability +# Adding words such as "quiet, solid" to the neg prompt can increase dynamism. +# prompt = "A young woman with beautiful, clear eyes and blonde hair stands in the forest, wearing a white dress and a crown. Her expression is serene, reminiscent of a movie star, with fair and youthful skin. Her brown long hair flows in the wind. The video quality is very high, with a clear view. High quality, masterpiece, best quality, high resolution, ultra-fine, fantastical." +# negative_prompt = "Twisted body, limb deformities, text captions, comic, static, ugly, error, messy code." +guidance_scale = 5.0 +seed = 43 +num_inference_steps = 40 +lora_weight = 0.55 +save_path = "samples/vace-videos" + +device = set_multi_gpus_devices(ulysses_degree, ring_degree) +config = OmegaConf.load(config_path) + +transformer = VaceWanTransformer3DModel.from_pretrained( + os.path.join(model_name, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + low_cpu_mem_usage=True, + torch_dtype=weight_dtype, +) + +if transformer_path is not None: + print(f"From checkpoint: {transformer_path}") + if transformer_path.endswith("safetensors"): + from safetensors.torch import load_file, safe_open + state_dict = load_file(transformer_path) + else: + state_dict = torch.load(transformer_path, map_location="cpu") + state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict + + m, u = transformer.load_state_dict(state_dict, strict=False) + print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") + +# Get Vae +vae = AutoencoderKLWan.from_pretrained( + os.path.join(model_name, config['vae_kwargs'].get('vae_subpath', 'vae')), + additional_kwargs=OmegaConf.to_container(config['vae_kwargs']), +).to(weight_dtype) + +if vae_path is not None: + print(f"From checkpoint: {vae_path}") + if vae_path.endswith("safetensors"): + from safetensors.torch import load_file, safe_open + state_dict = load_file(vae_path) + else: + state_dict = torch.load(vae_path, map_location="cpu") + state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict + + m, u = vae.load_state_dict(state_dict, strict=False) + print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") + +# Get Tokenizer +tokenizer = AutoTokenizer.from_pretrained( + os.path.join(model_name, config['text_encoder_kwargs'].get('tokenizer_subpath', 'tokenizer')), +) + +# Get Text encoder +text_encoder = WanT5EncoderModel.from_pretrained( + os.path.join(model_name, config['text_encoder_kwargs'].get('text_encoder_subpath', 'text_encoder')), + additional_kwargs=OmegaConf.to_container(config['text_encoder_kwargs']), + low_cpu_mem_usage=True, + torch_dtype=weight_dtype, +) +text_encoder = text_encoder.eval() + +# Get Scheduler +Chosen_Scheduler = scheduler_dict = { + "Flow": FlowMatchEulerDiscreteScheduler, + "Flow_Unipc": FlowUniPCMultistepScheduler, + "Flow_DPM++": FlowDPMSolverMultistepScheduler, +}[sampler_name] +if sampler_name == "Flow_Unipc" or sampler_name == "Flow_DPM++": + config['scheduler_kwargs']['shift'] = 1 +scheduler = Chosen_Scheduler( + **filter_kwargs(Chosen_Scheduler, OmegaConf.to_container(config['scheduler_kwargs'])) +) + +# Get Pipeline +pipeline = WanVacePipeline( + transformer=transformer, + vae=vae, + tokenizer=tokenizer, + text_encoder=text_encoder, + scheduler=scheduler, +) +if ulysses_degree > 1 or ring_degree > 1: + from functools import partial + transformer.enable_multi_gpus_inference() + if fsdp_dit: + shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype) + pipeline.transformer = shard_fn(pipeline.transformer) + print("Add FSDP DIT") + if fsdp_text_encoder: + shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype) + pipeline.text_encoder = shard_fn(pipeline.text_encoder) + print("Add FSDP TEXT ENCODER") + +if compile_dit: + for i in range(len(pipeline.transformer.blocks)): + pipeline.transformer.blocks[i] = torch.compile(pipeline.transformer.blocks[i]) + print("Add Compile") + +if GPU_memory_mode == "sequential_cpu_offload": + replace_parameters_by_name(transformer, ["modulation",], device=device) + transformer.freqs = transformer.freqs.to(device=device) + pipeline.enable_sequential_cpu_offload(device=device) +elif GPU_memory_mode == "model_cpu_offload_and_qfloat8": + convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device) + convert_weight_dtype_wrapper(transformer, weight_dtype) + pipeline.enable_model_cpu_offload(device=device) +elif GPU_memory_mode == "model_cpu_offload": + pipeline.enable_model_cpu_offload(device=device) +elif GPU_memory_mode == "model_full_load_and_qfloat8": + convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device) + convert_weight_dtype_wrapper(transformer, weight_dtype) + pipeline.to(device=device) +else: + pipeline.to(device=device) + +coefficients = get_teacache_coefficients(model_name) if enable_teacache else None +if coefficients is not None: + print(f"Enable TeaCache with threshold {teacache_threshold} and skip the first {num_skip_start_steps} steps.") + pipeline.transformer.enable_teacache( + coefficients, num_inference_steps, teacache_threshold, num_skip_start_steps=num_skip_start_steps, offload=teacache_offload + ) + +if cfg_skip_ratio is not None: + print(f"Enable cfg_skip_ratio {cfg_skip_ratio}.") + pipeline.transformer.enable_cfg_skip(cfg_skip_ratio, num_inference_steps) + +generator = torch.Generator(device=device).manual_seed(seed) + +if lora_path is not None: + pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device) + +with torch.no_grad(): + video_length = int((video_length - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1 + latent_frames = (video_length - 1) // vae.config.temporal_compression_ratio + 1 + + if enable_riflex: + pipeline.transformer.enable_riflex(k = riflex_k, L_test = latent_frames) + + if subject_ref_images is not None: + subject_ref_images = [get_image_latent(_subject_ref_image, sample_size=sample_size, padding=True) for _subject_ref_image in subject_ref_images] + subject_ref_images = torch.cat(subject_ref_images, dim=2) + + inpaint_video, inpaint_video_mask, clip_image = get_image_to_video_latent(start_image, end_image, video_length=video_length, sample_size=sample_size) + + control_video, _, _, _ = get_video_to_video_latent(control_video, video_length=video_length, sample_size=sample_size, fps=fps, ref_image=None) + + sample = pipeline( + prompt, + num_frames = video_length, + negative_prompt = negative_prompt, + height = sample_size[0], + width = sample_size[1], + generator = generator, + guidance_scale = guidance_scale, + num_inference_steps = num_inference_steps, + + video = inpaint_video, + mask_video = inpaint_video_mask, + control_video = control_video, + subject_ref_images = subject_ref_images, + shift = shift, + vace_context_scale = vace_context_scale, + ).videos + +if lora_path is not None: + pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device) + +def save_results(): + if not os.path.exists(save_path): + os.makedirs(save_path, exist_ok=True) + + index = len([path for path in os.listdir(save_path)]) + 1 + prefix = str(index).zfill(8) + if video_length == 1: + video_path = os.path.join(save_path, prefix + ".png") + + image = sample[0, :, 0] + image = image.transpose(0, 1).transpose(1, 2) + image = (image * 255).numpy().astype(np.uint8) + image = Image.fromarray(image) + image.save(video_path) + else: + video_path = os.path.join(save_path, prefix + ".mp4") + save_videos_grid(sample, video_path, fps=fps) + +if ulysses_degree * ring_degree > 1: + import torch.distributed as dist + if dist.get_rank() == 0: + save_results() +else: + save_results() \ No newline at end of file diff --git a/examples/wan2.1_vace/predict_s2v.py b/examples/wan2.1_vace/predict_s2v.py new file mode 100644 index 0000000..e78a034 --- /dev/null +++ b/examples/wan2.1_vace/predict_s2v.py @@ -0,0 +1,311 @@ +import os +import sys + +import numpy as np +import torch +from diffusers import FlowMatchEulerDiscreteScheduler +from omegaconf import OmegaConf +from PIL import Image +from transformers import AutoTokenizer + +current_file_path = os.path.abspath(__file__) +project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))] +for project_root in project_roots: + sys.path.insert(0, project_root) if project_root not in sys.path else None + +from videox_fun.dist import set_multi_gpus_devices, shard_model +from videox_fun.models import (AutoencoderKLWan, AutoTokenizer, + WanT5EncoderModel, VaceWanTransformer3DModel) +from videox_fun.data.dataset_image_video import process_pose_file +from videox_fun.models.cache_utils import get_teacache_coefficients +from videox_fun.pipeline import WanVacePipeline, WanPipeline +from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, + convert_weight_dtype_wrapper, + replace_parameters_by_name) +from videox_fun.utils.lora_utils import merge_lora, unmerge_lora +from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent, get_image_latent, + get_video_to_video_latent, + save_videos_grid) +from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler +from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler + +# GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, sequential_cpu_offload]. +# model_full_load means that the entire model will be moved to the GPU. +# +# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU, +# and the transformer model has been quantized to float8, which can save more GPU memory. +# +# model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory. +# +# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use, +# and the transformer model has been quantized to float8, which can save more GPU memory. +# +# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use, +# resulting in slower speeds but saving a large amount of GPU memory. +GPU_memory_mode = "sequential_cpu_offload" +# Multi GPUs config +# Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used. +# For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4. +# If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1. +ulysses_degree = 1 +ring_degree = 1 +# Use FSDP to save more GPU memory in multi gpus. +fsdp_dit = False +fsdp_text_encoder = True +# Compile will give a speedup in fixed resolution and need a little GPU memory. +# The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload. +compile_dit = False + +# Support TeaCache. +enable_teacache = True +# Recommended to be set between 0.05 and 0.30. A larger threshold can cache more steps, speeding up the inference process, +# but it may cause slight differences between the generated content and the original content. +# # --------------------------------------------------------------------------------------------------- # +# | Model Name | threshold | Model Name | threshold | +# | Wan2.1-VACE-1.3B | 0.05~0.10 | Wan2.1-VACE-14B | 0.10~0.15 | +# # --------------------------------------------------------------------------------------------------- # +teacache_threshold = 0.05 +# The number of steps to skip TeaCache at the beginning of the inference process, which can +# reduce the impact of TeaCache on generated video quality. +num_skip_start_steps = 5 +# Whether to offload TeaCache tensors to cpu to save a little bit of GPU memory. +teacache_offload = False + +# Skip some cfg steps in inference for acceleration +# Recommended to be set between 0.00 and 0.25 +cfg_skip_ratio = 0 + +# Riflex config +enable_riflex = False +# Index of intrinsic frequency +riflex_k = 6 + +# Config and model path +config_path = "config/wan2.1/wan_civitai.yaml" +# model path +model_name = "models/Diffusion_Transformer/Wan2.1-VACE-1.3B" + +# Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++" +sampler_name = "Flow_Unipc" +# [NOTE]: Noise schedule shift parameter. Affects temporal dynamics. +# Used when the sampler is in "Flow_Unipc", "Flow_DPM++". +shift = 16 + +# Load pretrained model if need +transformer_path = None +vae_path = None +lora_path = None + +# Other params +sample_size = [832, 480] +video_length = 81 +fps = 16 +vace_context_scale = 1.00 + +# Use torch.float16 if GPU does not support torch.bfloat16 +# ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16 +weight_dtype = torch.bfloat16 +control_video = None +start_image = None +end_image = None +subject_ref_images = ["asset/ref_1.png", "asset/ref_2.png"] + +# 使用更长的neg prompt如"模糊,突变,变形,失真,画面暗,文本字幕,画面固定,连环画,漫画,线稿,没有主体。",可以增加稳定性 +# 在neg prompt中添加"安静,固定"等词语可以增加动态性。 +prompt = "暖阳漫过草地,扎着双马尾、头戴绿色蝴蝶结、身穿浅绿色连衣裙的小女孩蹲在盛开的雏菊旁。她身旁一只棕白相间的狗狗吐着舌头,毛茸茸尾巴欢快摇晃。小女孩笑着举起黄红配色、带有蓝色按钮的玩具相机,将和狗狗的欢乐瞬间定格。" +negative_prompt = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走" + +# Using longer neg prompt such as "Blurring, mutation, deformation, distortion, dark and solid, comics, text subtitles, line art." can increase stability +# Adding words such as "quiet, solid" to the neg prompt can increase dynamism. +# prompt = "A young woman with beautiful, clear eyes and blonde hair stands in the forest, wearing a white dress and a crown. Her expression is serene, reminiscent of a movie star, with fair and youthful skin. Her brown long hair flows in the wind. The video quality is very high, with a clear view. High quality, masterpiece, best quality, high resolution, ultra-fine, fantastical." +# negative_prompt = "Twisted body, limb deformities, text captions, comic, static, ugly, error, messy code." +guidance_scale = 5.0 +seed = 43 +num_inference_steps = 40 +lora_weight = 0.55 +save_path = "samples/vace-videos" + +device = set_multi_gpus_devices(ulysses_degree, ring_degree) +config = OmegaConf.load(config_path) + +transformer = VaceWanTransformer3DModel.from_pretrained( + os.path.join(model_name, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + low_cpu_mem_usage=True, + torch_dtype=weight_dtype, +) + +if transformer_path is not None: + print(f"From checkpoint: {transformer_path}") + if transformer_path.endswith("safetensors"): + from safetensors.torch import load_file, safe_open + state_dict = load_file(transformer_path) + else: + state_dict = torch.load(transformer_path, map_location="cpu") + state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict + + m, u = transformer.load_state_dict(state_dict, strict=False) + print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") + +# Get Vae +vae = AutoencoderKLWan.from_pretrained( + os.path.join(model_name, config['vae_kwargs'].get('vae_subpath', 'vae')), + additional_kwargs=OmegaConf.to_container(config['vae_kwargs']), +).to(weight_dtype) + +if vae_path is not None: + print(f"From checkpoint: {vae_path}") + if vae_path.endswith("safetensors"): + from safetensors.torch import load_file, safe_open + state_dict = load_file(vae_path) + else: + state_dict = torch.load(vae_path, map_location="cpu") + state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict + + m, u = vae.load_state_dict(state_dict, strict=False) + print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") + +# Get Tokenizer +tokenizer = AutoTokenizer.from_pretrained( + os.path.join(model_name, config['text_encoder_kwargs'].get('tokenizer_subpath', 'tokenizer')), +) + +# Get Text encoder +text_encoder = WanT5EncoderModel.from_pretrained( + os.path.join(model_name, config['text_encoder_kwargs'].get('text_encoder_subpath', 'text_encoder')), + additional_kwargs=OmegaConf.to_container(config['text_encoder_kwargs']), + low_cpu_mem_usage=True, + torch_dtype=weight_dtype, +) +text_encoder = text_encoder.eval() + +# Get Scheduler +Chosen_Scheduler = scheduler_dict = { + "Flow": FlowMatchEulerDiscreteScheduler, + "Flow_Unipc": FlowUniPCMultistepScheduler, + "Flow_DPM++": FlowDPMSolverMultistepScheduler, +}[sampler_name] +if sampler_name == "Flow_Unipc" or sampler_name == "Flow_DPM++": + config['scheduler_kwargs']['shift'] = 1 +scheduler = Chosen_Scheduler( + **filter_kwargs(Chosen_Scheduler, OmegaConf.to_container(config['scheduler_kwargs'])) +) + +# Get Pipeline +pipeline = WanVacePipeline( + transformer=transformer, + vae=vae, + tokenizer=tokenizer, + text_encoder=text_encoder, + scheduler=scheduler, +) +if ulysses_degree > 1 or ring_degree > 1: + from functools import partial + transformer.enable_multi_gpus_inference() + if fsdp_dit: + shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype) + pipeline.transformer = shard_fn(pipeline.transformer) + print("Add FSDP DIT") + if fsdp_text_encoder: + shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype) + pipeline.text_encoder = shard_fn(pipeline.text_encoder) + print("Add FSDP TEXT ENCODER") + +if compile_dit: + for i in range(len(pipeline.transformer.blocks)): + pipeline.transformer.blocks[i] = torch.compile(pipeline.transformer.blocks[i]) + print("Add Compile") + +if GPU_memory_mode == "sequential_cpu_offload": + replace_parameters_by_name(transformer, ["modulation",], device=device) + transformer.freqs = transformer.freqs.to(device=device) + pipeline.enable_sequential_cpu_offload(device=device) +elif GPU_memory_mode == "model_cpu_offload_and_qfloat8": + convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device) + convert_weight_dtype_wrapper(transformer, weight_dtype) + pipeline.enable_model_cpu_offload(device=device) +elif GPU_memory_mode == "model_cpu_offload": + pipeline.enable_model_cpu_offload(device=device) +elif GPU_memory_mode == "model_full_load_and_qfloat8": + convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device) + convert_weight_dtype_wrapper(transformer, weight_dtype) + pipeline.to(device=device) +else: + pipeline.to(device=device) + +coefficients = get_teacache_coefficients(model_name) if enable_teacache else None +if coefficients is not None: + print(f"Enable TeaCache with threshold {teacache_threshold} and skip the first {num_skip_start_steps} steps.") + pipeline.transformer.enable_teacache( + coefficients, num_inference_steps, teacache_threshold, num_skip_start_steps=num_skip_start_steps, offload=teacache_offload + ) + +if cfg_skip_ratio is not None: + print(f"Enable cfg_skip_ratio {cfg_skip_ratio}.") + pipeline.transformer.enable_cfg_skip(cfg_skip_ratio, num_inference_steps) + +generator = torch.Generator(device=device).manual_seed(seed) + +if lora_path is not None: + pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device) + +with torch.no_grad(): + video_length = int((video_length - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1 + latent_frames = (video_length - 1) // vae.config.temporal_compression_ratio + 1 + + if enable_riflex: + pipeline.transformer.enable_riflex(k = riflex_k, L_test = latent_frames) + + if subject_ref_images is not None: + subject_ref_images = [get_image_latent(_subject_ref_image, sample_size=sample_size, padding=True) for _subject_ref_image in subject_ref_images] + subject_ref_images = torch.cat(subject_ref_images, dim=2) + + inpaint_video, inpaint_video_mask, clip_image = get_image_to_video_latent(start_image, end_image, video_length=video_length, sample_size=sample_size) + + control_video, _, _, _ = get_video_to_video_latent(control_video, video_length=video_length, sample_size=sample_size, fps=fps, ref_image=None) + + sample = pipeline( + prompt, + num_frames = video_length, + negative_prompt = negative_prompt, + height = sample_size[0], + width = sample_size[1], + generator = generator, + guidance_scale = guidance_scale, + num_inference_steps = num_inference_steps, + + video = inpaint_video, + mask_video = inpaint_video_mask, + control_video = control_video, + subject_ref_images = subject_ref_images, + shift = shift, + vace_context_scale = vace_context_scale, + ).videos + +if lora_path is not None: + pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device) + +def save_results(): + if not os.path.exists(save_path): + os.makedirs(save_path, exist_ok=True) + + index = len([path for path in os.listdir(save_path)]) + 1 + prefix = str(index).zfill(8) + if video_length == 1: + video_path = os.path.join(save_path, prefix + ".png") + + image = sample[0, :, 0] + image = image.transpose(0, 1).transpose(1, 2) + image = (image * 255).numpy().astype(np.uint8) + image = Image.fromarray(image) + image.save(video_path) + else: + video_path = os.path.join(save_path, prefix + ".mp4") + save_videos_grid(sample, video_path, fps=fps) + +if ulysses_degree * ring_degree > 1: + import torch.distributed as dist + if dist.get_rank() == 0: + save_results() +else: + save_results() \ No newline at end of file diff --git a/examples/wan2.1_vace/predict_v2v_control.py b/examples/wan2.1_vace/predict_v2v_control.py new file mode 100644 index 0000000..3f58b9a --- /dev/null +++ b/examples/wan2.1_vace/predict_v2v_control.py @@ -0,0 +1,311 @@ +import os +import sys + +import numpy as np +import torch +from diffusers import FlowMatchEulerDiscreteScheduler +from omegaconf import OmegaConf +from PIL import Image +from transformers import AutoTokenizer + +current_file_path = os.path.abspath(__file__) +project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))] +for project_root in project_roots: + sys.path.insert(0, project_root) if project_root not in sys.path else None + +from videox_fun.dist import set_multi_gpus_devices, shard_model +from videox_fun.models import (AutoencoderKLWan, AutoTokenizer, + WanT5EncoderModel, VaceWanTransformer3DModel) +from videox_fun.data.dataset_image_video import process_pose_file +from videox_fun.models.cache_utils import get_teacache_coefficients +from videox_fun.pipeline import WanVacePipeline, WanPipeline +from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, + convert_weight_dtype_wrapper, + replace_parameters_by_name) +from videox_fun.utils.lora_utils import merge_lora, unmerge_lora +from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent, get_image_latent, + get_video_to_video_latent, + save_videos_grid) +from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler +from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler + +# GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, sequential_cpu_offload]. +# model_full_load means that the entire model will be moved to the GPU. +# +# model_full_load_and_qfloat8 means that the entire model will be moved to the GPU, +# and the transformer model has been quantized to float8, which can save more GPU memory. +# +# model_cpu_offload means that the entire model will be moved to the CPU after use, which can save some GPU memory. +# +# model_cpu_offload_and_qfloat8 indicates that the entire model will be moved to the CPU after use, +# and the transformer model has been quantized to float8, which can save more GPU memory. +# +# sequential_cpu_offload means that each layer of the model will be moved to the CPU after use, +# resulting in slower speeds but saving a large amount of GPU memory. +GPU_memory_mode = "sequential_cpu_offload" +# Multi GPUs config +# Please ensure that the product of ulysses_degree and ring_degree equals the number of GPUs used. +# For example, if you are using 8 GPUs, you can set ulysses_degree = 2 and ring_degree = 4. +# If you are using 1 GPU, you can set ulysses_degree = 1 and ring_degree = 1. +ulysses_degree = 1 +ring_degree = 1 +# Use FSDP to save more GPU memory in multi gpus. +fsdp_dit = False +fsdp_text_encoder = True +# Compile will give a speedup in fixed resolution and need a little GPU memory. +# The compile_dit is not compatible with the fsdp_dit and sequential_cpu_offload. +compile_dit = False + +# Support TeaCache. +enable_teacache = True +# Recommended to be set between 0.05 and 0.30. A larger threshold can cache more steps, speeding up the inference process, +# but it may cause slight differences between the generated content and the original content. +# # --------------------------------------------------------------------------------------------------- # +# | Model Name | threshold | Model Name | threshold | +# | Wan2.1-VACE-1.3B | 0.05~0.10 | Wan2.1-VACE-14B | 0.10~0.15 | +# # --------------------------------------------------------------------------------------------------- # +teacache_threshold = 0.05 +# The number of steps to skip TeaCache at the beginning of the inference process, which can +# reduce the impact of TeaCache on generated video quality. +num_skip_start_steps = 5 +# Whether to offload TeaCache tensors to cpu to save a little bit of GPU memory. +teacache_offload = False + +# Skip some cfg steps in inference for acceleration +# Recommended to be set between 0.00 and 0.25 +cfg_skip_ratio = 0 + +# Riflex config +enable_riflex = False +# Index of intrinsic frequency +riflex_k = 6 + +# Config and model path +config_path = "config/wan2.1/wan_civitai.yaml" +# model path +model_name = "models/Diffusion_Transformer/Wan2.1-VACE-1.3B" + +# Choose the sampler in "Flow", "Flow_Unipc", "Flow_DPM++" +sampler_name = "Flow_Unipc" +# [NOTE]: Noise schedule shift parameter. Affects temporal dynamics. +# Used when the sampler is in "Flow_Unipc", "Flow_DPM++". +shift = 16 + +# Load pretrained model if need +transformer_path = None +vae_path = None +lora_path = None + +# Other params +sample_size = [832, 480] +video_length = 81 +fps = 16 +vace_context_scale = 1.00 + +# Use torch.float16 if GPU does not support torch.bfloat16 +# ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16 +weight_dtype = torch.bfloat16 +control_video = "asset/pose.mp4" +start_image = None +end_image = None +subject_ref_images = None + +# 使用更长的neg prompt如"模糊,突变,变形,失真,画面暗,文本字幕,画面固定,连环画,漫画,线稿,没有主体。",可以增加稳定性 +# 在neg prompt中添加"安静,固定"等词语可以增加动态性。 +prompt = "一位年轻女子站在阳光明媚的海岸线上,身穿深蓝色背心与清爽的白色衬衫,外搭一条简洁的白色围裙,围裙在轻拂的海风中微微飘动。她拥有一头鲜艳的紫色长发,在风中轻盈舞动,发间系着一个精致的黑色蝴蝶结,与身后柔和的蔚蓝天空形成鲜明对比。她面容清秀,眉目精致,透着一股甜美的青春气息;神情柔和,略带羞涩,目光静静地凝望着远方的地平线,双手自然交叠于身前,仿佛沉浸在思绪之中。在她身后,是辽阔无垠、波光粼粼的大海,阳光洒在海面上,映出温暖的金色光晕。" +negative_prompt = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走" + +# Using longer neg prompt such as "Blurring, mutation, deformation, distortion, dark and solid, comics, text subtitles, line art." can increase stability +# Adding words such as "quiet, solid" to the neg prompt can increase dynamism. +# prompt = "A young woman with beautiful, clear eyes and blonde hair stands in the forest, wearing a white dress and a crown. Her expression is serene, reminiscent of a movie star, with fair and youthful skin. Her brown long hair flows in the wind. The video quality is very high, with a clear view. High quality, masterpiece, best quality, high resolution, ultra-fine, fantastical." +# negative_prompt = "Twisted body, limb deformities, text captions, comic, static, ugly, error, messy code." +guidance_scale = 5.0 +seed = 43 +num_inference_steps = 40 +lora_weight = 0.55 +save_path = "samples/vace-videos" + +device = set_multi_gpus_devices(ulysses_degree, ring_degree) +config = OmegaConf.load(config_path) + +transformer = VaceWanTransformer3DModel.from_pretrained( + os.path.join(model_name, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + low_cpu_mem_usage=True, + torch_dtype=weight_dtype, +) + +if transformer_path is not None: + print(f"From checkpoint: {transformer_path}") + if transformer_path.endswith("safetensors"): + from safetensors.torch import load_file, safe_open + state_dict = load_file(transformer_path) + else: + state_dict = torch.load(transformer_path, map_location="cpu") + state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict + + m, u = transformer.load_state_dict(state_dict, strict=False) + print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") + +# Get Vae +vae = AutoencoderKLWan.from_pretrained( + os.path.join(model_name, config['vae_kwargs'].get('vae_subpath', 'vae')), + additional_kwargs=OmegaConf.to_container(config['vae_kwargs']), +).to(weight_dtype) + +if vae_path is not None: + print(f"From checkpoint: {vae_path}") + if vae_path.endswith("safetensors"): + from safetensors.torch import load_file, safe_open + state_dict = load_file(vae_path) + else: + state_dict = torch.load(vae_path, map_location="cpu") + state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict + + m, u = vae.load_state_dict(state_dict, strict=False) + print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") + +# Get Tokenizer +tokenizer = AutoTokenizer.from_pretrained( + os.path.join(model_name, config['text_encoder_kwargs'].get('tokenizer_subpath', 'tokenizer')), +) + +# Get Text encoder +text_encoder = WanT5EncoderModel.from_pretrained( + os.path.join(model_name, config['text_encoder_kwargs'].get('text_encoder_subpath', 'text_encoder')), + additional_kwargs=OmegaConf.to_container(config['text_encoder_kwargs']), + low_cpu_mem_usage=True, + torch_dtype=weight_dtype, +) +text_encoder = text_encoder.eval() + +# Get Scheduler +Chosen_Scheduler = scheduler_dict = { + "Flow": FlowMatchEulerDiscreteScheduler, + "Flow_Unipc": FlowUniPCMultistepScheduler, + "Flow_DPM++": FlowDPMSolverMultistepScheduler, +}[sampler_name] +if sampler_name == "Flow_Unipc" or sampler_name == "Flow_DPM++": + config['scheduler_kwargs']['shift'] = 1 +scheduler = Chosen_Scheduler( + **filter_kwargs(Chosen_Scheduler, OmegaConf.to_container(config['scheduler_kwargs'])) +) + +# Get Pipeline +pipeline = WanVacePipeline( + transformer=transformer, + vae=vae, + tokenizer=tokenizer, + text_encoder=text_encoder, + scheduler=scheduler, +) +if ulysses_degree > 1 or ring_degree > 1: + from functools import partial + transformer.enable_multi_gpus_inference() + if fsdp_dit: + shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype) + pipeline.transformer = shard_fn(pipeline.transformer) + print("Add FSDP DIT") + if fsdp_text_encoder: + shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype) + pipeline.text_encoder = shard_fn(pipeline.text_encoder) + print("Add FSDP TEXT ENCODER") + +if compile_dit: + for i in range(len(pipeline.transformer.blocks)): + pipeline.transformer.blocks[i] = torch.compile(pipeline.transformer.blocks[i]) + print("Add Compile") + +if GPU_memory_mode == "sequential_cpu_offload": + replace_parameters_by_name(transformer, ["modulation",], device=device) + transformer.freqs = transformer.freqs.to(device=device) + pipeline.enable_sequential_cpu_offload(device=device) +elif GPU_memory_mode == "model_cpu_offload_and_qfloat8": + convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device) + convert_weight_dtype_wrapper(transformer, weight_dtype) + pipeline.enable_model_cpu_offload(device=device) +elif GPU_memory_mode == "model_cpu_offload": + pipeline.enable_model_cpu_offload(device=device) +elif GPU_memory_mode == "model_full_load_and_qfloat8": + convert_model_weight_to_float8(transformer, exclude_module_name=["modulation",], device=device) + convert_weight_dtype_wrapper(transformer, weight_dtype) + pipeline.to(device=device) +else: + pipeline.to(device=device) + +coefficients = get_teacache_coefficients(model_name) if enable_teacache else None +if coefficients is not None: + print(f"Enable TeaCache with threshold {teacache_threshold} and skip the first {num_skip_start_steps} steps.") + pipeline.transformer.enable_teacache( + coefficients, num_inference_steps, teacache_threshold, num_skip_start_steps=num_skip_start_steps, offload=teacache_offload + ) + +if cfg_skip_ratio is not None: + print(f"Enable cfg_skip_ratio {cfg_skip_ratio}.") + pipeline.transformer.enable_cfg_skip(cfg_skip_ratio, num_inference_steps) + +generator = torch.Generator(device=device).manual_seed(seed) + +if lora_path is not None: + pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device) + +with torch.no_grad(): + video_length = int((video_length - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if video_length != 1 else 1 + latent_frames = (video_length - 1) // vae.config.temporal_compression_ratio + 1 + + if enable_riflex: + pipeline.transformer.enable_riflex(k = riflex_k, L_test = latent_frames) + + if subject_ref_images is not None: + subject_ref_images = [get_image_latent(_subject_ref_image, sample_size=sample_size, padding=True) for _subject_ref_image in subject_ref_images] + subject_ref_images = torch.cat(subject_ref_images, dim=2) + + inpaint_video, inpaint_video_mask, clip_image = get_image_to_video_latent(start_image, end_image, video_length=video_length, sample_size=sample_size) + + control_video, _, _, _ = get_video_to_video_latent(control_video, video_length=video_length, sample_size=sample_size, fps=fps, ref_image=None) + + sample = pipeline( + prompt, + num_frames = video_length, + negative_prompt = negative_prompt, + height = sample_size[0], + width = sample_size[1], + generator = generator, + guidance_scale = guidance_scale, + num_inference_steps = num_inference_steps, + + video = inpaint_video, + mask_video = inpaint_video_mask, + control_video = control_video, + subject_ref_images = subject_ref_images, + shift = shift, + vace_context_scale = vace_context_scale, + ).videos + +if lora_path is not None: + pipeline = unmerge_lora(pipeline, lora_path, lora_weight, device=device) + +def save_results(): + if not os.path.exists(save_path): + os.makedirs(save_path, exist_ok=True) + + index = len([path for path in os.listdir(save_path)]) + 1 + prefix = str(index).zfill(8) + if video_length == 1: + video_path = os.path.join(save_path, prefix + ".png") + + image = sample[0, :, 0] + image = image.transpose(0, 1).transpose(1, 2) + image = (image * 255).numpy().astype(np.uint8) + image = Image.fromarray(image) + image.save(video_path) + else: + video_path = os.path.join(save_path, prefix + ".mp4") + save_videos_grid(sample, video_path, fps=fps) + +if ulysses_degree * ring_degree > 1: + import torch.distributed as dist + if dist.get_rank() == 0: + save_results() +else: + save_results() \ No newline at end of file diff --git a/examples/wan2.2/predict_s2v.py b/examples/wan2.2/predict_s2v.py index 9034291..98c0a1d 100644 --- a/examples/wan2.2/predict_s2v.py +++ b/examples/wan2.2/predict_s2v.py @@ -13,17 +13,22 @@ for project_root in project_roots: sys.path.insert(0, project_root) if project_root not in sys.path else None from videox_fun.dist import set_multi_gpus_devices, shard_model -from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8, AutoTokenizer, CLIPModel, WanAudioEncoder, - WanT5EncoderModel, Wan2_2Transformer3DModel_S2V) +from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8, + AutoTokenizer, CLIPModel, + Wan2_2Transformer3DModel_S2V, WanAudioEncoder, + WanT5EncoderModel) from videox_fun.models.cache_utils import get_teacache_coefficients from videox_fun.pipeline import Wan2_2S2VPipeline -from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, replace_parameters_by_name, - convert_weight_dtype_wrapper) -from videox_fun.utils.lora_utils import merge_lora, unmerge_lora -from videox_fun.utils.utils import (filter_kwargs, get_image_to_video_latent, get_image_latent, - save_videos_grid, merge_video_audio) from videox_fun.utils.fm_solvers import FlowDPMSolverMultistepScheduler from videox_fun.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler +from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, + convert_weight_dtype_wrapper, + replace_parameters_by_name) +from videox_fun.utils.lora_utils import merge_lora, unmerge_lora +from videox_fun.utils.utils import (filter_kwargs, get_image_latent, + get_image_to_video_latent, + get_video_to_video_latent, + merge_video_audio, save_videos_grid) # GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, sequential_cpu_offload]. # model_full_load means that the entire model will be moved to the GPU. @@ -106,6 +111,7 @@ fps = 16 # ome graphics cards, such as v100, 2080ti, do not support torch.bfloat16 weight_dtype = torch.bfloat16 # If you want to generate from text, please set the validation_image_start = None and validation_image_end = None +control_video = "asset/pose.mp4" ref_image = "asset/8.png" audio_path = "asset/talk.wav" @@ -312,6 +318,8 @@ with torch.no_grad(): if ref_image is not None: ref_image = get_image_latent(ref_image, sample_size=sample_size) + pose_video, _, _, _ = get_video_to_video_latent(control_video, video_length=video_length, sample_size=sample_size, fps=fps, ref_image=None) + sample = pipeline( prompt, num_frames = video_length, @@ -324,6 +332,7 @@ with torch.no_grad(): boundary = boundary, ref_image = ref_image, + pose_video = pose_video, audio_path = audio_path, shift = shift, fps = fps diff --git a/examples/wan2.2_fun/predict_t2v.py b/examples/wan2.2_fun/predict_t2v.py index 8f17b96..0f7c5bf 100644 --- a/examples/wan2.2_fun/predict_t2v.py +++ b/examples/wan2.2_fun/predict_t2v.py @@ -13,7 +13,7 @@ for project_root in project_roots: sys.path.insert(0, project_root) if project_root not in sys.path else None from videox_fun.dist import set_multi_gpus_devices, shard_model -from videox_fun.models import (AutoencoderKLWan, AutoTokenizer, CLIPModel, +from videox_fun.models import (AutoencoderKLWan, AutoTokenizer, CLIPModel, AutoencoderKLWan3_8, WanT5EncoderModel, Wan2_2Transformer3DModel) from videox_fun.models.cache_utils import get_teacache_coefficients from videox_fun.pipeline import Wan2_2FunInpaintPipeline diff --git a/videox_fun/models/__init__.py b/videox_fun/models/__init__.py index 41503bf..279f18c 100755 --- a/videox_fun/models/__init__.py +++ b/videox_fun/models/__init__.py @@ -1,6 +1,6 @@ import importlib.util -from diffusers import AutoencoderKL +from diffusers import AutoencoderKL from transformers import (AutoTokenizer, CLIPImageProcessor, CLIPTextModel, CLIPTokenizer, CLIPVisionModelWithProjection, T5EncoderModel, T5Tokenizer, T5TokenizerFast) @@ -22,6 +22,7 @@ from .wan_text_encoder import WanT5EncoderModel from .wan_transformer3d import (Wan2_2Transformer3DModel, WanRMSNorm, WanSelfAttention, WanTransformer3DModel) from .wan_transformer3d_s2v import Wan2_2Transformer3DModel_S2V +from .wan_transformer3d_vace import VaceWanTransformer3DModel from .wan_vae import AutoencoderKLWan, AutoencoderKLWan_ from .wan_vae3_8 import AutoencoderKLWan2_2_, AutoencoderKLWan3_8 diff --git a/videox_fun/models/cache_utils.py b/videox_fun/models/cache_utils.py index 2779ba6..bc03419 100755 --- a/videox_fun/models/cache_utils.py +++ b/videox_fun/models/cache_utils.py @@ -2,7 +2,8 @@ import numpy as np import torch def get_teacache_coefficients(model_name): - if "wan2.1-t2v-1.3b" in model_name.lower() or "wan2.1-fun-1.3b" in model_name.lower() or "wan2.1-fun-v1.1-1.3b" in model_name.lower(): + if "wan2.1-t2v-1.3b" in model_name.lower() or "wan2.1-fun-1.3b" in model_name.lower() \ + or "wan2.1-fun-v1.1-1.3b" in model_name.lower() or "wan2.1-vace-1.3b" in model_name.lower(): return [-5.21862437e+04, 9.23041404e+03, -5.28275948e+02, 1.36987616e+01, -4.99875664e-02] elif "wan2.1-t2v-14b" in model_name.lower(): return [-3.03318725e+05, 4.90537029e+04, -2.65530556e+03, 5.87365115e+01, -3.15583525e-01] @@ -10,7 +11,7 @@ def get_teacache_coefficients(model_name): return [2.57151496e+05, -3.54229917e+04, 1.40286849e+03, -1.35890334e+01, 1.32517977e-01] elif "wan2.1-i2v-14b-720p" in model_name.lower() or "wan2.1-fun-14b" in model_name.lower() or "wan2.2-fun" in model_name.lower() \ or "wan2.2-i2v-a14b" in model_name.lower() or "wan2.2-t2v-a14b" in model_name.lower() or "wan2.2-ti2v-5b" in model_name.lower() \ - or "wan2.2-s2v" in model_name.lower() : + or "wan2.2-s2v" in model_name.lower() or "wan2.1-vace-14b" in model_name.lower(): return [8.10705460e+03, 2.13393892e+03, -3.72934672e+02, 1.66203073e+01, -4.17769401e-02] else: print(f"The model {model_name} is not supported by TeaCache.") diff --git a/videox_fun/models/wan_transformer3d_s2v.py b/videox_fun/models/wan_transformer3d_s2v.py index 7bb69f0..14e3853 100644 --- a/videox_fun/models/wan_transformer3d_s2v.py +++ b/videox_fun/models/wan_transformer3d_s2v.py @@ -1,32 +1,27 @@ # Modified from https://github.com/Wan-Video/Wan2.2/blob/main/wan/modules/s2v/model_s2v.py # Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved. -import glob -import json import math -import os import types from copy import deepcopy +from typing import Any, Dict -import numpy as np import torch import torch.cuda.amp as amp import torch.nn as nn -from diffusers.configuration_utils import ConfigMixin, register_to_config -from diffusers.models.modeling_utils import ModelMixin +from diffusers.configuration_utils import register_to_config +from diffusers.utils import is_torch_version from einops import rearrange from ..dist import (get_sequence_parallel_rank, get_sequence_parallel_world_size, get_sp_group, - usp_attn_s2v_forward, xFuserLongContextAttention) -from ..utils import cfg_skip -from .attention_utils import attention, flash_attention -from .cache_utils import TeaCache + usp_attn_s2v_forward) +from .attention_utils import attention from .wan_audio_injector import (AudioInjector_WAN, CausalAudioEncoder, FramePackMotioner, MotionerTransformers, rope_precompute) -from .wan_transformer3d import (Head, WanAttentionBlock, WanLayerNorm, Wan2_2Transformer3DModel, - WanSelfAttention, rope_params, +from .wan_transformer3d import (Wan2_2Transformer3DModel, WanAttentionBlock, + WanLayerNorm, WanSelfAttention, sinusoidal_embedding_1d) diff --git a/videox_fun/models/wan_transformer3d_vace.py b/videox_fun/models/wan_transformer3d_vace.py new file mode 100644 index 0000000..3a562e6 --- /dev/null +++ b/videox_fun/models/wan_transformer3d_vace.py @@ -0,0 +1,354 @@ +# Modified from https://github.com/ali-vilab/VACE/blob/main/vace/models/wan/wan_vace.py +# -*- coding: utf-8 -*- +# Copyright (c) Alibaba, Inc. and its affiliates. +import torch +import torch.cuda.amp as amp +import torch.nn as nn +from diffusers.configuration_utils import register_to_config + +from .wan_transformer3d import (WanAttentionBlock, WanTransformer3DModel, + sinusoidal_embedding_1d) + + +class VaceWanAttentionBlock(WanAttentionBlock): + def __init__( + self, + cross_attn_type, + dim, + ffn_dim, + num_heads, + window_size=(-1, -1), + qk_norm=True, + cross_attn_norm=False, + eps=1e-6, + block_id=0 + ): + super().__init__(cross_attn_type, dim, ffn_dim, num_heads, window_size, qk_norm, cross_attn_norm, eps) + self.block_id = block_id + if block_id == 0: + self.before_proj = nn.Linear(self.dim, self.dim) + nn.init.zeros_(self.before_proj.weight) + nn.init.zeros_(self.before_proj.bias) + self.after_proj = nn.Linear(self.dim, self.dim) + nn.init.zeros_(self.after_proj.weight) + nn.init.zeros_(self.after_proj.bias) + + def forward(self, c, x, **kwargs): + if self.block_id == 0: + c = self.before_proj(c) + x + all_c = [] + else: + all_c = list(torch.unbind(c)) + c = all_c.pop(-1) + + c = super().forward(c, **kwargs) + c_skip = self.after_proj(c) + all_c += [c_skip, c] + c = torch.stack(all_c) + return c + + +class BaseWanAttentionBlock(WanAttentionBlock): + def __init__( + self, + cross_attn_type, + dim, + ffn_dim, + num_heads, + window_size=(-1, -1), + qk_norm=True, + cross_attn_norm=False, + eps=1e-6, + block_id=None + ): + super().__init__(cross_attn_type, dim, ffn_dim, num_heads, window_size, qk_norm, cross_attn_norm, eps) + self.block_id = block_id + + def forward(self, x, hints, context_scale=1.0, **kwargs): + x = super().forward(x, **kwargs) + if self.block_id is not None: + x = x + hints[self.block_id] * context_scale + return x + + +class VaceWanTransformer3DModel(WanTransformer3DModel): + @register_to_config + def __init__(self, + vace_layers=None, + vace_in_dim=None, + model_type='t2v', + patch_size=(1, 2, 2), + text_len=512, + in_dim=16, + dim=2048, + ffn_dim=8192, + freq_dim=256, + text_dim=4096, + out_dim=16, + num_heads=16, + num_layers=32, + window_size=(-1, -1), + qk_norm=True, + cross_attn_norm=True, + eps=1e-6): + model_type = "t2v" # TODO: Hard code for both preview and official versions. + super().__init__(model_type, patch_size, text_len, in_dim, dim, ffn_dim, freq_dim, text_dim, out_dim, + num_heads, num_layers, window_size, qk_norm, cross_attn_norm, eps) + + self.vace_layers = [i for i in range(0, self.num_layers, 2)] if vace_layers is None else vace_layers + self.vace_in_dim = self.in_dim if vace_in_dim is None else vace_in_dim + + assert 0 in self.vace_layers + self.vace_layers_mapping = {i: n for n, i in enumerate(self.vace_layers)} + + # blocks + self.blocks = nn.ModuleList([ + BaseWanAttentionBlock('t2v_cross_attn', self.dim, self.ffn_dim, self.num_heads, self.window_size, self.qk_norm, + self.cross_attn_norm, self.eps, + block_id=self.vace_layers_mapping[i] if i in self.vace_layers else None) + for i in range(self.num_layers) + ]) + + # vace blocks + self.vace_blocks = nn.ModuleList([ + VaceWanAttentionBlock('t2v_cross_attn', self.dim, self.ffn_dim, self.num_heads, self.window_size, self.qk_norm, + self.cross_attn_norm, self.eps, block_id=i) + for i in self.vace_layers + ]) + + # vace patch embeddings + self.vace_patch_embedding = nn.Conv3d( + self.vace_in_dim, self.dim, kernel_size=self.patch_size, stride=self.patch_size + ) + + def forward_vace( + self, + x, + vace_context, + seq_len, + kwargs + ): + # embeddings + c = [self.vace_patch_embedding(u.unsqueeze(0)) for u in vace_context] + c = [u.flatten(2).transpose(1, 2) for u in c] + c = torch.cat([ + torch.cat([u, u.new_zeros(1, seq_len - u.size(1), u.size(2))], + dim=1) for u in c + ]) + # # Context Parallel + # if self.sp_world_size > 1: + # c = torch.chunk(c, self.sp_world_size, dim=1)[self.sp_world_rank] + + # arguments + new_kwargs = dict(x=x) + new_kwargs.update(kwargs) + + for block in self.vace_blocks: + c = block(c, **new_kwargs) + hints = torch.unbind(c)[:-1] + return hints + + def forward( + self, + x, + t, + vace_context, + context, + seq_len, + vace_context_scale=1.0, + clip_fea=None, + y=None, + cond_flag=True + ): + r""" + Forward pass through the diffusion model + + Args: + x (List[Tensor]): + List of input video tensors, each with shape [C_in, F, H, W] + t (Tensor): + Diffusion timesteps tensor of shape [B] + context (List[Tensor]): + List of text embeddings each with shape [L, C] + seq_len (`int`): + Maximum sequence length for positional encoding + clip_fea (Tensor, *optional*): + CLIP image features for image-to-video mode + y (List[Tensor], *optional*): + Conditional video inputs for image-to-video mode, same shape as x + + Returns: + List[Tensor]: + List of denoised video tensors with original input shapes [C_out, F, H / 8, W / 8] + """ + # if self.model_type == 'i2v': + # assert clip_fea is not None and y is not None + # params + device = self.patch_embedding.weight.device + if self.freqs.device != device: + self.freqs = self.freqs.to(device) + + # if y is not None: + # x = [torch.cat([u, v], dim=0) for u, v in zip(x, y)] + + # embeddings + x = [self.patch_embedding(u.unsqueeze(0)) for u in x] + grid_sizes = torch.stack( + [torch.tensor(u.shape[2:], dtype=torch.long) for u in x]) + x = [u.flatten(2).transpose(1, 2) for u in x] + seq_lens = torch.tensor([u.size(1) for u in x], dtype=torch.long) + assert seq_lens.max() <= seq_len + x = torch.cat([ + torch.cat([u, u.new_zeros(1, seq_len - u.size(1), u.size(2))], + dim=1) for u in x + ]) + + # time embeddings + with amp.autocast(dtype=torch.float32): + e = self.time_embedding( + sinusoidal_embedding_1d(self.freq_dim, t).float()) + e0 = self.time_projection(e).unflatten(1, (6, self.dim)) + assert e.dtype == torch.float32 and e0.dtype == torch.float32 + + # context + context_lens = None + context = self.text_embedding( + torch.stack([ + torch.cat( + [u, u.new_zeros(self.text_len - u.size(0), u.size(1))]) + for u in context + ])) + + # arguments + kwargs = dict( + e=e0, + seq_lens=seq_lens, + grid_sizes=grid_sizes, + freqs=self.freqs, + context=context, + context_lens=context_lens) + hints = self.forward_vace(x, vace_context, seq_len, kwargs) + + # Context Parallel + if self.sp_world_size > 1: + x = torch.chunk(x, self.sp_world_size, dim=1)[self.sp_world_rank] + hints = [torch.chunk(u, self.sp_world_size, dim=1)[self.sp_world_rank] for u in hints] + + kwargs['hints'] = hints + kwargs['context_scale'] = vace_context_scale + + # TeaCache + if self.teacache is not None: + if cond_flag: + if t.dim() != 1: + modulated_inp = e0[:, -1, :] + else: + modulated_inp = e0 + skip_flag = self.teacache.cnt < self.teacache.num_skip_start_steps + if skip_flag: + self.should_calc = True + self.teacache.accumulated_rel_l1_distance = 0 + else: + if cond_flag: + rel_l1_distance = self.teacache.compute_rel_l1_distance(self.teacache.previous_modulated_input, modulated_inp) + self.teacache.accumulated_rel_l1_distance += self.teacache.rescale_func(rel_l1_distance) + if self.teacache.accumulated_rel_l1_distance < self.teacache.rel_l1_thresh: + self.should_calc = False + else: + self.should_calc = True + self.teacache.accumulated_rel_l1_distance = 0 + self.teacache.previous_modulated_input = modulated_inp + self.teacache.should_calc = self.should_calc + else: + self.should_calc = self.teacache.should_calc + + # TeaCache + if self.teacache is not None: + if not self.should_calc: + previous_residual = self.teacache.previous_residual_cond if cond_flag else self.teacache.previous_residual_uncond + x = x + previous_residual.to(x.device)[-x.size()[0]:,] + else: + ori_x = x.clone().cpu() if self.teacache.offload else x.clone() + + for block in self.blocks: + if torch.is_grad_enabled() and self.gradient_checkpointing: + + def create_custom_forward(module): + def custom_forward(*inputs): + return module(*inputs) + + return custom_forward + ckpt_kwargs: Dict[str, Any] = {"use_reentrant": False} if is_torch_version(">=", "1.11.0") else {} + x = torch.utils.checkpoint.checkpoint( + create_custom_forward(block), + x, + hints, + vace_context_scale, + e0, + seq_lens, + grid_sizes, + self.freqs, + context, + context_lens, + dtype, + t, + **ckpt_kwargs, + ) + else: + x = block(x, **kwargs) + + if cond_flag: + self.teacache.previous_residual_cond = x.cpu() - ori_x if self.teacache.offload else x - ori_x + else: + self.teacache.previous_residual_uncond = x.cpu() - ori_x if self.teacache.offload else x - ori_x + else: + for block in self.blocks: + if torch.is_grad_enabled() and self.gradient_checkpointing: + + def create_custom_forward(module): + def custom_forward(*inputs): + return module(*inputs) + + return custom_forward + ckpt_kwargs: Dict[str, Any] = {"use_reentrant": False} if is_torch_version(">=", "1.11.0") else {} + x = torch.utils.checkpoint.checkpoint( + create_custom_forward(block), + x, + hints, + vace_context_scale, + e0, + seq_lens, + grid_sizes, + self.freqs, + context, + context_lens, + dtype, + t, + **ckpt_kwargs, + ) + else: + x = block(x, **kwargs) + + # head + if torch.is_grad_enabled() and self.gradient_checkpointing: + def create_custom_forward(module): + def custom_forward(*inputs): + return module(*inputs) + + return custom_forward + ckpt_kwargs: Dict[str, Any] = {"use_reentrant": False} if is_torch_version(">=", "1.11.0") else {} + x = torch.utils.checkpoint.checkpoint(create_custom_forward(self.head), x, e, **ckpt_kwargs) + else: + x = self.head(x, e) + + if self.sp_world_size > 1: + x = self.all_gather(x, dim=1) + + # unpatchify + x = self.unpatchify(x, grid_sizes) + x = torch.stack(x) + if self.teacache is not None and cond_flag: + self.teacache.cnt += 1 + if self.teacache.cnt == self.teacache.num_steps: + self.teacache.reset() + return x \ No newline at end of file diff --git a/videox_fun/pipeline/__init__.py b/videox_fun/pipeline/__init__.py index ec8b7f6..a3d975f 100755 --- a/videox_fun/pipeline/__init__.py +++ b/videox_fun/pipeline/__init__.py @@ -1,17 +1,18 @@ from .pipeline_cogvideox_fun import CogVideoXFunPipeline from .pipeline_cogvideox_fun_control import CogVideoXFunControlPipeline from .pipeline_cogvideox_fun_inpaint import CogVideoXFunInpaintPipeline +from .pipeline_flux import FluxPipeline +from .pipeline_qwenimage import QwenImagePipeline from .pipeline_wan import WanPipeline from .pipeline_wan2_2 import Wan2_2Pipeline from .pipeline_wan2_2_fun_control import Wan2_2FunControlPipeline from .pipeline_wan2_2_fun_inpaint import Wan2_2FunInpaintPipeline +from .pipeline_wan2_2_s2v import Wan2_2S2VPipeline from .pipeline_wan2_2_ti2v import Wan2_2TI2VPipeline from .pipeline_wan_fun_control import WanFunControlPipeline from .pipeline_wan_fun_inpaint import WanFunInpaintPipeline from .pipeline_wan_phantom import WanFunPhantomPipeline -from .pipeline_qwenimage import QwenImagePipeline -from .pipeline_wan2_2_s2v import Wan2_2S2VPipeline -from .pipeline_flux import FluxPipeline +from .pipeline_wan_vace import WanVacePipeline WanFunPipeline = WanPipeline WanI2VPipeline = WanFunInpaintPipeline diff --git a/videox_fun/pipeline/pipeline_wan2_2_s2v.py b/videox_fun/pipeline/pipeline_wan2_2_s2v.py index 5debdfa..64ad75a 100644 --- a/videox_fun/pipeline/pipeline_wan2_2_s2v.py +++ b/videox_fun/pipeline/pipeline_wan2_2_s2v.py @@ -331,67 +331,19 @@ class Wan2_2S2VPipeline(DiffusionPipeline): audio_embed_bucket = audio_embed_bucket.permute(0, 2, 3, 1) return audio_embed_bucket, num_repeat - def read_last_n_frames(self, - video_path, - n_frames, - target_fps=16, - reverse=False): - """ - Read the last `n_frames` from a video at the specified frame rate. - - Parameters: - video_path (str): Path to the video file. - n_frames (int): Number of frames to read. - target_fps (int, optional): Target sampling frame rate. Defaults to 16. - reverse (bool, optional): Whether to read frames in reverse order. - If True, reads the first `n_frames` instead of the last ones. - - Returns: - np.ndarray: A NumPy array of shape [n_frames, H, W, 3], representing the sampled video frames. - """ - vr = VideoReader(video_path) - original_fps = vr.get_avg_fps() - total_frames = len(vr) - - interval = max(1, round(original_fps / target_fps)) - - required_span = (n_frames - 1) * interval - - start_frame = max(0, total_frames - required_span - - 1) if not reverse else 0 - - sampled_indices = [] - for i in range(n_frames): - indice = start_frame + i * interval - if indice >= total_frames: - break - else: - sampled_indices.append(indice) - - return vr.get_batch(sampled_indices).asnumpy() - def encode_pose_latents(self, pose_video, num_repeat, num_frames, size, fps, weight_dtype, device): height, width = size if not pose_video is None: - pose_seq = self.read_last_n_frames(pose_video, n_frames=num_frames * num_repeat, target_fps=fps,reverse=True) + padding_frame_num = num_repeat * num_frames - pose_video.shape[2] + pose_video = torch.cat( + [ + pose_video, + -torch.ones([1, 3, padding_frame_num, height, width]) + ], + dim=2 + ) - resize_opreat = transforms.Resize(min(height, width)) - crop_opreat = transforms.CenterCrop((height, width)) - tensor_trans = transforms.ToTensor() - - cond_tensor = torch.from_numpy(pose_seq) - cond_tensor = cond_tensor.permute(0, 3, 1, 2) / 255.0 * 2 - 1.0 - cond_tensor = crop_opreat(resize_opreat(cond_tensor)).permute( - 1, 0, 2, 3).unsqueeze(0) - - padding_frame_num = num_repeat * num_frames - cond_tensor.shape[2] - cond_tensor = torch.cat([ - cond_tensor, - - torch.ones([1, 3, padding_frame_num, height, width]) - ], - dim=2) - - cond_tensors = torch.chunk(cond_tensor, num_repeat, dim=2) + cond_tensors = torch.chunk(pose_video, num_repeat, dim=2) else: cond_tensors = [-torch.ones([1, 3, num_frames, height, width])] @@ -700,6 +652,11 @@ class Wan2_2S2VPipeline(DiffusionPipeline): motion_latents = self.vae.encode(motion_latents)[0].mode() # Get pose cond input if need + if pose_video is not None: + video_length = pose_video.shape[2] + pose_video = self.image_processor.preprocess(rearrange(pose_video, "b c f h w -> (b f) c h w"), height=height, width=width) + pose_video = pose_video.to(dtype=torch.float32) + pose_video = rearrange(pose_video, "(b f) c h w -> b c f h w", f=video_length) pose_latents = self.encode_pose_latents( pose_video=pose_video, num_repeat=num_repeat, @@ -768,7 +725,7 @@ class Wan2_2S2VPipeline(DiffusionPipeline): with torch.no_grad(): left_idx = r * num_frames right_idx = r * num_frames + num_frames - cond_latents = pose_latents[r] if pose_video else pose_latents[0] * 0 + cond_latents = pose_latents[r] if pose_video is not None else pose_latents[0] * 0 cond_latents = cond_latents.to(dtype=weight_dtype, device=device) audio_input = audio_emb[..., left_idx:right_idx] diff --git a/videox_fun/pipeline/pipeline_wan_vace.py b/videox_fun/pipeline/pipeline_wan_vace.py new file mode 100644 index 0000000..ec8b7f5 --- /dev/null +++ b/videox_fun/pipeline/pipeline_wan_vace.py @@ -0,0 +1,785 @@ +import inspect +import math +from dataclasses import dataclass +from typing import Any, Callable, Dict, List, Optional, Tuple, Union + +import numpy as np +import torch +import torch.nn.functional as F +import torchvision.transforms.functional as TF +from diffusers import FlowMatchEulerDiscreteScheduler +from diffusers.callbacks import MultiPipelineCallbacks, PipelineCallback +from diffusers.image_processor import VaeImageProcessor +from diffusers.models.embeddings import get_1d_rotary_pos_embed +from diffusers.pipelines.pipeline_utils import DiffusionPipeline +from diffusers.schedulers import FlowMatchEulerDiscreteScheduler +from diffusers.utils import BaseOutput, logging, replace_example_docstring +from diffusers.utils.torch_utils import randn_tensor +from diffusers.video_processor import VideoProcessor +from einops import rearrange +from PIL import Image +from transformers import T5Tokenizer + +from ..models import (AutoencoderKLWan, AutoTokenizer, + WanT5EncoderModel, VaceWanTransformer3DModel) +from ..utils.fm_solvers import (FlowDPMSolverMultistepScheduler, + get_sampling_sigmas) +from ..utils.fm_solvers_unipc import FlowUniPCMultistepScheduler + +logger = logging.get_logger(__name__) # pylint: disable=invalid-name + + +EXAMPLE_DOC_STRING = """ + Examples: + ```python + pass + ``` +""" + + +# Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.retrieve_timesteps +def retrieve_timesteps( + scheduler, + num_inference_steps: Optional[int] = None, + device: Optional[Union[str, torch.device]] = None, + timesteps: Optional[List[int]] = None, + sigmas: Optional[List[float]] = None, + **kwargs, +): + """ + Calls the scheduler's `set_timesteps` method and retrieves timesteps from the scheduler after the call. Handles + custom timesteps. Any kwargs will be supplied to `scheduler.set_timesteps`. + + Args: + scheduler (`SchedulerMixin`): + The scheduler to get timesteps from. + num_inference_steps (`int`): + The number of diffusion steps used when generating samples with a pre-trained model. If used, `timesteps` + must be `None`. + device (`str` or `torch.device`, *optional*): + The device to which the timesteps should be moved to. If `None`, the timesteps are not moved. + timesteps (`List[int]`, *optional*): + Custom timesteps used to override the timestep spacing strategy of the scheduler. If `timesteps` is passed, + `num_inference_steps` and `sigmas` must be `None`. + sigmas (`List[float]`, *optional*): + Custom sigmas used to override the timestep spacing strategy of the scheduler. If `sigmas` is passed, + `num_inference_steps` and `timesteps` must be `None`. + + Returns: + `Tuple[torch.Tensor, int]`: A tuple where the first element is the timestep schedule from the scheduler and the + second element is the number of inference steps. + """ + if timesteps is not None and sigmas is not None: + raise ValueError("Only one of `timesteps` or `sigmas` can be passed. Please choose one to set custom values") + if timesteps is not None: + accepts_timesteps = "timesteps" in set(inspect.signature(scheduler.set_timesteps).parameters.keys()) + if not accepts_timesteps: + raise ValueError( + f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom" + f" timestep schedules. Please check whether you are using the correct scheduler." + ) + scheduler.set_timesteps(timesteps=timesteps, device=device, **kwargs) + timesteps = scheduler.timesteps + num_inference_steps = len(timesteps) + elif sigmas is not None: + accept_sigmas = "sigmas" in set(inspect.signature(scheduler.set_timesteps).parameters.keys()) + if not accept_sigmas: + raise ValueError( + f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom" + f" sigmas schedules. Please check whether you are using the correct scheduler." + ) + scheduler.set_timesteps(sigmas=sigmas, device=device, **kwargs) + timesteps = scheduler.timesteps + num_inference_steps = len(timesteps) + else: + scheduler.set_timesteps(num_inference_steps, device=device, **kwargs) + timesteps = scheduler.timesteps + return timesteps, num_inference_steps + + +def resize_mask(mask, latent, process_first_frame_only=True): + latent_size = latent.size() + batch_size, channels, num_frames, height, width = mask.shape + + if process_first_frame_only: + target_size = list(latent_size[2:]) + target_size[0] = 1 + first_frame_resized = F.interpolate( + mask[:, :, 0:1, :, :], + size=target_size, + mode='trilinear', + align_corners=False + ) + + target_size = list(latent_size[2:]) + target_size[0] = target_size[0] - 1 + if target_size[0] != 0: + remaining_frames_resized = F.interpolate( + mask[:, :, 1:, :, :], + size=target_size, + mode='trilinear', + align_corners=False + ) + resized_mask = torch.cat([first_frame_resized, remaining_frames_resized], dim=2) + else: + resized_mask = first_frame_resized + else: + target_size = list(latent_size[2:]) + resized_mask = F.interpolate( + mask, + size=target_size, + mode='trilinear', + align_corners=False + ) + return resized_mask + + +@dataclass +class WanPipelineOutput(BaseOutput): + r""" + Output class for CogVideo pipelines. + + Args: + video (`torch.Tensor`, `np.ndarray`, or List[List[PIL.Image.Image]]): + List of video outputs - It can be a nested list of length `batch_size,` with each sub-list containing + denoised PIL image sequences of length `num_frames.` It can also be a NumPy array or Torch tensor of shape + `(batch_size, num_frames, channels, height, width)`. + """ + + videos: torch.Tensor + + +class WanVacePipeline(DiffusionPipeline): + r""" + Pipeline for text-to-video generation using Wan. + + This model inherits from [`DiffusionPipeline`]. Check the superclass documentation for the generic methods the + library implements for all the pipelines (such as downloading or saving, running on a particular device, etc.) + """ + + _optional_components = [] + model_cpu_offload_seq = "text_encoder->transformer->vae" + + _callback_tensor_inputs = [ + "latents", + "prompt_embeds", + "negative_prompt_embeds", + ] + + def __init__( + self, + tokenizer: AutoTokenizer, + text_encoder: WanT5EncoderModel, + vae: AutoencoderKLWan, + transformer: VaceWanTransformer3DModel, + scheduler: FlowMatchEulerDiscreteScheduler, + ): + super().__init__() + + self.register_modules( + tokenizer=tokenizer, text_encoder=text_encoder, vae=vae, transformer=transformer, scheduler=scheduler + ) + + self.video_processor = VideoProcessor(vae_scale_factor=self.vae.spatial_compression_ratio) + self.image_processor = VaeImageProcessor(vae_scale_factor=self.vae.spatial_compression_ratio) + self.mask_processor = VaeImageProcessor( + vae_scale_factor=self.vae.spatial_compression_ratio, do_normalize=False, do_binarize=True, do_convert_grayscale=True + ) + + def _get_t5_prompt_embeds( + self, + prompt: Union[str, List[str]] = None, + num_videos_per_prompt: int = 1, + max_sequence_length: int = 512, + 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 + batch_size = len(prompt) + + text_inputs = self.tokenizer( + prompt, + padding="max_length", + max_length=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 + 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[:, max_sequence_length - 1 : -1]) + logger.warning( + "The following part of your input was truncated because `max_sequence_length` is set to " + f" {max_sequence_length} tokens: {removed_text}" + ) + + seq_lens = prompt_attention_mask.gt(0).sum(dim=1).long() + prompt_embeds = self.text_encoder(text_input_ids.to(device), attention_mask=prompt_attention_mask.to(device))[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.repeat(1, num_videos_per_prompt, 1) + prompt_embeds = prompt_embeds.view(batch_size * num_videos_per_prompt, seq_len, -1) + + return [u[:v] for u, v in zip(prompt_embeds, seq_lens)] + + 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 = 512, + 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 prepare_latents( + self, batch_size, num_channels_latents, num_frames, height, width, dtype, device, generator, latents=None, num_length_latents=None + ): + 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." + ) + + shape = ( + batch_size, + num_channels_latents, + (num_frames - 1) // self.vae.temporal_compression_ratio + 1 if num_length_latents is None else num_length_latents, + height // self.vae.spatial_compression_ratio, + width // self.vae.spatial_compression_ratio, + ) + + if latents is None: + latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype) + else: + latents = latents.to(device) + + # scale the initial noise by the standard deviation required by the scheduler + if hasattr(self.scheduler, "init_noise_sigma"): + latents = latents * self.scheduler.init_noise_sigma + return latents + + def vace_encode_frames(self, frames, ref_images, masks=None, vae=None): + vae = self.vae if vae is None else vae + weight_dtype = frames.dtype + if ref_images is None: + ref_images = [None] * len(frames) + else: + assert len(frames) == len(ref_images) + + if masks is None: + latents = vae.encode(frames)[0].mode() + else: + masks = [torch.where(m > 0.5, 1.0, 0.0).to(weight_dtype) for m in masks] + inactive = [i * (1 - m) + 0 * m for i, m in zip(frames, masks)] + reactive = [i * m + 0 * (1 - m) for i, m in zip(frames, masks)] + inactive = vae.encode(inactive)[0].mode() + reactive = vae.encode(reactive)[0].mode() + latents = [torch.cat((u, c), dim=0) for u, c in zip(inactive, reactive)] + + cat_latents = [] + for latent, refs in zip(latents, ref_images): + if refs is not None: + if masks is None: + ref_latent = vae.encode(refs)[0].mode() + else: + ref_latent = vae.encode(refs)[0].mode() + ref_latent = [torch.cat((u, torch.zeros_like(u)), dim=0) for u in ref_latent] + assert all([x.shape[1] == 1 for x in ref_latent]) + latent = torch.cat([*ref_latent, latent], dim=1) + cat_latents.append(latent) + return cat_latents + + def vace_encode_masks(self, masks, ref_images=None, vae_stride=[4, 8, 8]): + if ref_images is None: + ref_images = [None] * len(masks) + else: + assert len(masks) == len(ref_images) + + result_masks = [] + for mask, refs in zip(masks, ref_images): + c, depth, height, width = mask.shape + new_depth = int((depth + 3) // vae_stride[0]) + height = 2 * (int(height) // (vae_stride[1] * 2)) + width = 2 * (int(width) // (vae_stride[2] * 2)) + + # reshape + mask = mask[0, :, :, :] + mask = mask.view( + depth, height, vae_stride[1], width, vae_stride[1] + ) # depth, height, 8, width, 8 + mask = mask.permute(2, 4, 0, 1, 3) # 8, 8, depth, height, width + mask = mask.reshape( + vae_stride[1] * vae_stride[2], depth, height, width + ) # 8*8, depth, height, width + + # interpolation + mask = F.interpolate(mask.unsqueeze(0), size=(new_depth, height, width), mode='nearest-exact').squeeze(0) + + if refs is not None: + length = len(refs) + mask_pad = torch.zeros_like(mask[:, :length, :, :]) + mask = torch.cat((mask_pad, mask), dim=1) + result_masks.append(mask) + return result_masks + + def vace_latent(self, z, m): + return [torch.cat([zz, mm], dim=0) for zz, mm in zip(z, m)] + + def prepare_control_latents( + self, control, control_image, batch_size, height, width, dtype, device, generator, do_classifier_free_guidance + ): + # resize the control to latents shape as we concatenate the control to the latents + # we do that before converting to dtype to avoid breaking in case we're using cpu_offload + # and half precision + + if control is not None: + control = control.to(device=device, dtype=dtype) + bs = 1 + new_control = [] + for i in range(0, control.shape[0], bs): + control_bs = control[i : i + bs] + control_bs = self.vae.encode(control_bs)[0] + control_bs = control_bs.mode() + new_control.append(control_bs) + control = torch.cat(new_control, dim = 0) + + if control_image is not None: + control_image = control_image.to(device=device, dtype=dtype) + bs = 1 + new_control_pixel_values = [] + for i in range(0, control_image.shape[0], bs): + control_pixel_values_bs = control_image[i : i + bs] + control_pixel_values_bs = self.vae.encode(control_pixel_values_bs)[0] + control_pixel_values_bs = control_pixel_values_bs.mode() + new_control_pixel_values.append(control_pixel_values_bs) + control_image_latents = torch.cat(new_control_pixel_values, dim = 0) + else: + control_image_latents = None + + return control, control_image_latents + + def decode_latents(self, latents: torch.Tensor) -> torch.Tensor: + frames = self.vae.decode(latents.to(self.vae.dtype)).sample + frames = (frames / 2 + 0.5).clamp(0, 1) + # we always cast to float32 as this does not cause significant overhead and is compatible with bfloa16 + frames = frames.cpu().float().numpy() + return frames + + # Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.StableDiffusionPipeline.prepare_extra_step_kwargs + def prepare_extra_step_kwargs(self, generator, eta): + # prepare extra kwargs for the scheduler step, since not all schedulers have the same signature + # eta (η) is only used with the DDIMScheduler, it will be ignored for other schedulers. + # eta corresponds to η in DDIM paper: https://arxiv.org/abs/2010.02502 + # and should be between [0, 1] + + accepts_eta = "eta" in set(inspect.signature(self.scheduler.step).parameters.keys()) + extra_step_kwargs = {} + if accepts_eta: + extra_step_kwargs["eta"] = eta + + # check if the scheduler accepts generator + accepts_generator = "generator" in set(inspect.signature(self.scheduler.step).parameters.keys()) + if accepts_generator: + extra_step_kwargs["generator"] = generator + return extra_step_kwargs + + # Copied from diffusers.pipelines.latte.pipeline_latte.LattePipeline.check_inputs + def check_inputs( + self, + prompt, + height, + width, + negative_prompt, + callback_on_step_end_tensor_inputs, + prompt_embeds=None, + negative_prompt_embeds=None, + ): + if height % 8 != 0 or width % 8 != 0: + raise ValueError(f"`height` and `width` have to be divisible by 8 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 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)}") + + if prompt is not None and negative_prompt_embeds is not None: + raise ValueError( + f"Cannot forward both `prompt`: {prompt} and `negative_prompt_embeds`:" + f" {negative_prompt_embeds}. Please make sure to only forward one of the two." + ) + + if 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`:" + f" {negative_prompt_embeds}. Please make sure to only forward one of the two." + ) + + if prompt_embeds is not None and negative_prompt_embeds is not None: + if prompt_embeds.shape != negative_prompt_embeds.shape: + raise ValueError( + "`prompt_embeds` and `negative_prompt_embeds` must have the same shape when passed directly, but" + f" got: `prompt_embeds` {prompt_embeds.shape} != `negative_prompt_embeds`" + f" {negative_prompt_embeds.shape}." + ) + + @property + def guidance_scale(self): + return self._guidance_scale + + @property + def num_timesteps(self): + return self._num_timesteps + + @property + def attention_kwargs(self): + return self._attention_kwargs + + @property + def interrupt(self): + return self._interrupt + + @torch.no_grad() + @replace_example_docstring(EXAMPLE_DOC_STRING) + def __call__( + self, + prompt: Optional[Union[str, List[str]]] = None, + negative_prompt: Optional[Union[str, List[str]]] = None, + height: int = 480, + width: int = 720, + video: Union[torch.FloatTensor] = None, + mask_video: Union[torch.FloatTensor] = None, + control_video: Union[torch.FloatTensor] = None, + subject_ref_images: Union[torch.FloatTensor] = None, + num_frames: int = 49, + num_inference_steps: int = 50, + timesteps: Optional[List[int]] = None, + guidance_scale: float = 6, + num_videos_per_prompt: int = 1, + eta: float = 0.0, + generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None, + latents: Optional[torch.FloatTensor] = None, + prompt_embeds: Optional[torch.FloatTensor] = None, + negative_prompt_embeds: Optional[torch.FloatTensor] = None, + output_type: str = "numpy", + return_dict: bool = False, + callback_on_step_end: Optional[ + Union[Callable[[int, int, Dict], None], PipelineCallback, MultiPipelineCallbacks] + ] = None, + attention_kwargs: Optional[Dict[str, Any]] = None, + callback_on_step_end_tensor_inputs: List[str] = ["latents"], + max_sequence_length: int = 512, + comfyui_progressbar: bool = False, + shift: int = 5, + ) -> Union[WanPipelineOutput, Tuple]: + """ + Function invoked when calling the pipeline for generation. + Args: + + Examples: + + Returns: + + """ + + if isinstance(callback_on_step_end, (PipelineCallback, MultiPipelineCallbacks)): + callback_on_step_end_tensor_inputs = callback_on_step_end.tensor_inputs + num_videos_per_prompt = 1 + + # 1. Check inputs. Raise error if not correct + self.check_inputs( + prompt, + height, + width, + negative_prompt, + callback_on_step_end_tensor_inputs, + prompt_embeds, + negative_prompt_embeds, + ) + self._guidance_scale = guidance_scale + self._attention_kwargs = attention_kwargs + self._interrupt = False + + # 2. Default 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] + + device = self._execution_device + weight_dtype = self.text_encoder.dtype + + # here `guidance_scale` is defined analog to the guidance weight `w` of equation (2) + # of the Imagen paper: https://arxiv.org/pdf/2205.11487.pdf . `guidance_scale = 1` + # corresponds to doing no classifier free guidance. + do_classifier_free_guidance = guidance_scale > 1.0 + + # 3. Encode input prompt + prompt_embeds, negative_prompt_embeds = self.encode_prompt( + prompt, + negative_prompt, + 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, + ) + if do_classifier_free_guidance: + in_prompt_embeds = negative_prompt_embeds + prompt_embeds + else: + in_prompt_embeds = prompt_embeds + + # 4. Prepare timesteps + if isinstance(self.scheduler, FlowMatchEulerDiscreteScheduler): + timesteps, num_inference_steps = retrieve_timesteps(self.scheduler, num_inference_steps, device, timesteps, mu=1) + elif isinstance(self.scheduler, FlowUniPCMultistepScheduler): + self.scheduler.set_timesteps(num_inference_steps, device=device, shift=shift) + timesteps = self.scheduler.timesteps + elif isinstance(self.scheduler, FlowDPMSolverMultistepScheduler): + sampling_sigmas = get_sampling_sigmas(num_inference_steps, shift) + timesteps, _ = retrieve_timesteps( + self.scheduler, + device=device, + sigmas=sampling_sigmas) + else: + timesteps, num_inference_steps = retrieve_timesteps(self.scheduler, num_inference_steps, device, timesteps) + self._num_timesteps = len(timesteps) + if comfyui_progressbar: + from comfy.utils import ProgressBar + pbar = ProgressBar(num_inference_steps + 2) + + latent_channels = self.vae.config.latent_channels + + if comfyui_progressbar: + pbar.update(1) + + # Prepare mask latent variables + if mask_video is not None: + bs, _, video_length, height, width = video.size() + mask_condition = self.mask_processor.preprocess(rearrange(mask_video, "b c f h w -> (b f) c h w"), height=height, width=width) + mask_condition = mask_condition.to(dtype=torch.float32) + mask_condition = rearrange(mask_condition, "(b f) c h w -> b c f h w", f=video_length) + mask_condition = torch.tile(mask_condition, [1, 3, 1, 1, 1]).to(dtype=weight_dtype, device=device) + + + if control_video is not None: + video_length = control_video.shape[2] + control_video = self.image_processor.preprocess(rearrange(control_video, "b c f h w -> (b f) c h w"), height=height, width=width) + control_video = control_video.to(dtype=torch.float32) + input_video = rearrange(control_video, "(b f) c h w -> b c f h w", f=video_length) + + input_video = input_video.to(dtype=weight_dtype, device=device) + + elif video is not None: + video_length = video.shape[2] + init_video = self.image_processor.preprocess(rearrange(video, "b c f h w -> (b f) c h w"), height=height, width=width) + init_video = init_video.to(dtype=torch.float32) + init_video = rearrange(init_video, "(b f) c h w -> b c f h w", f=video_length).to(dtype=weight_dtype, device=device) + + input_video = init_video * (mask_condition < 0.5) + input_video = input_video.to(dtype=weight_dtype, device=device) + + if subject_ref_images is not None: + video_length = subject_ref_images.shape[2] + subject_ref_images = self.image_processor.preprocess(rearrange(subject_ref_images, "b c f h w -> (b f) c h w"), height=height, width=width) + subject_ref_images = subject_ref_images.to(dtype=torch.float32) + subject_ref_images = rearrange(subject_ref_images, "(b f) c h w -> b c f h w", f=video_length) + subject_ref_images = subject_ref_images.to(dtype=weight_dtype, device=device) + + bs, c, f, h, w = subject_ref_images.size() + new_subject_ref_images = [] + for i in range(bs): + new_subject_ref_images.append([]) + for j in range(f): + new_subject_ref_images[i].append(subject_ref_images[i, :, j:j+1]) + subject_ref_images = new_subject_ref_images + + vace_latents = self.vace_encode_frames(input_video, subject_ref_images, masks=mask_condition, vae=self.vae) + mask_latents = self.vace_encode_masks(mask_condition, subject_ref_images) + vace_context = self.vace_latent(vace_latents, mask_latents) + + # 5. Prepare latents. + latents = self.prepare_latents( + batch_size * num_videos_per_prompt, + latent_channels, + num_frames, + height, + width, + weight_dtype, + device, + generator, + latents, + num_length_latents=vace_latents[0].size(1) + ) + + # 6. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline + extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta) + + target_shape = (self.vae.latent_channels, vace_latents[0].size(1), vace_latents[0].size(2), vace_latents[0].size(3)) + seq_len = math.ceil((target_shape[2] * target_shape[3]) / (self.transformer.config.patch_size[1] * self.transformer.config.patch_size[2]) * target_shape[1]) + # 7. Denoising loop + num_warmup_steps = max(len(timesteps) - num_inference_steps * self.scheduler.order, 0) + self.transformer.num_inference_steps = num_inference_steps + with self.progress_bar(total=num_inference_steps) as progress_bar: + for i, t in enumerate(timesteps): + self.transformer.current_steps = i + + if self.interrupt: + continue + + latent_model_input = torch.cat([latents] * 2) if do_classifier_free_guidance else latents + if hasattr(self.scheduler, "scale_model_input"): + latent_model_input = self.scheduler.scale_model_input(latent_model_input, t) + + vace_context_input = torch.stack(vace_context * 2) if do_classifier_free_guidance else vace_context + + # broadcast to batch dimension in a way that's compatible with ONNX/Core ML + timestep = t.expand(latent_model_input.shape[0]) + + # predict noise model_output + with torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=device): + noise_pred = self.transformer( + x=latent_model_input, + context=in_prompt_embeds, + t=timestep, + vace_context=vace_context_input, + seq_len=seq_len, + ) + + # perform guidance + if do_classifier_free_guidance: + noise_pred_uncond, noise_pred_text = noise_pred.chunk(2) + noise_pred = noise_pred_uncond + self.guidance_scale * (noise_pred_text - noise_pred_uncond) + + # compute the previous noisy sample x_t -> x_t-1 + latents = self.scheduler.step(noise_pred, t, latents, **extra_step_kwargs, 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) + + if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0): + progress_bar.update() + if comfyui_progressbar: + pbar.update(1) + + if subject_ref_images is not None: + len_subject_ref_images = len(subject_ref_images[0]) + latents = latents[:, :, len_subject_ref_images:, :, :] + + if output_type == "numpy": + video = self.decode_latents(latents) + elif not output_type == "latent": + video = self.decode_latents(latents) + video = self.video_processor.postprocess_video(video=video, output_type=output_type) + else: + video = latents + + # Offload all models + self.maybe_free_model_hooks() + + if not return_dict: + video = torch.from_numpy(video) + + return WanPipelineOutput(videos=video)