From e4c74a12f214e6c40bdb63494c006ca8a6c42806 Mon Sep 17 00:00:00 2001 From: bubbliiiing <3323290568@qq.com> Date: Mon, 9 Jun 2025 08:25:49 +0000 Subject: [PATCH 1/4] Update phantom && add fsdp training && Fix bug in image encoder while using sage attention --- examples/phantom/predict_s2v.py | 303 ++++++++ scripts/wan2.1/README_TRAIN.md | 2 +- scripts/wan2.1/README_TRAIN_LORA.md | 2 +- scripts/wan2.1/train.py | 82 ++- scripts/wan2.1/train_lora.py | 75 +- scripts/wan2.1_fun/README_TRAIN.md | 2 +- scripts/wan2.1_fun/README_TRAIN_CONTROL.md | 4 +- .../wan2.1_fun/README_TRAIN_CONTROL_LORA.md | 4 +- scripts/wan2.1_fun/README_TRAIN_LORA.md | 2 +- scripts/wan2.1_fun/train.py | 82 ++- scripts/wan2.1_fun/train_control.py | 82 ++- scripts/wan2.1_fun/train_control_lora.py | 104 ++- scripts/wan2.1_fun/train_lora.py | 90 ++- videox_fun/dist/__init__.py | 17 +- videox_fun/models/__init__.py | 16 +- videox_fun/models/wan_image_encoder.py | 4 +- videox_fun/models/wan_transformer3d.py | 16 +- videox_fun/pipeline/__init__.py | 4 +- videox_fun/pipeline/pipeline_wan_phantom.py | 695 ++++++++++++++++++ videox_fun/utils/utils.py | 32 +- 20 files changed, 1474 insertions(+), 144 deletions(-) create mode 100644 examples/phantom/predict_s2v.py mode change 100644 => 100755 videox_fun/models/wan_image_encoder.py create mode 100644 videox_fun/pipeline/pipeline_wan_phantom.py diff --git a/examples/phantom/predict_s2v.py b/examples/phantom/predict_s2v.py new file mode 100644 index 0000000..c01e142 --- /dev/null +++ b/examples/phantom/predict_s2v.py @@ -0,0 +1,303 @@ +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, WanTransformer3DModel) +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 WanFunPhantomPipeline +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, + save_videos_grid) +from videox_fun.utils.utils import (filter_kwargs, + 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 choosen 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 | Model Name | threshold | +# | Wan2.1-T2V-1.3B | 0.05~0.10 | Wan2.1-T2V-14B | 0.10~0.15 | Wan2.1-I2V-14B-720P | 0.20~0.30 | +# | Wan2.1-I2V-14B-480P | 0.20~0.25 | Wan2.1-Fun-*-1.3B-* | 0.05~0.10 | Wan2.1-Fun-*-14B-* | 0.20~0.30 | +# # --------------------------------------------------------------------------------------------------- # +teacache_threshold = 0.10 +# 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 +# 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-T2V-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++". +# If you want to generate a 480p video, it is recommended to set the shift value to 3.0. +# If you want to generate a 720p video, it is recommended to set the shift value to 5.0. +shift = 3 + +# Load pretrained model if need +transformer_path = "models/Personalized_Model/Phantom-Wan-1.3B.safetensors" +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 +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 = 6.0 +seed = 43 +num_inference_steps = 50 +lora_weight = 0.55 +save_path = "samples/wan-videos-fun-control" + +device = set_multi_gpus_devices(ulysses_degree, ring_degree) +config = OmegaConf.load(config_path) + +transformer = WanTransformer3DModel.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 +Choosen_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 = Choosen_Scheduler( + **filter_kwargs(Choosen_Scheduler, OmegaConf.to_container(config['scheduler_kwargs'])) +) + +# Get Pipeline +pipeline = WanFunPhantomPipeline( + 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) + + 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, + + subject_ref_images = subject_ref_images, + shift = shift, + ).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/scripts/wan2.1/README_TRAIN.md b/scripts/wan2.1/README_TRAIN.md index 4200a01..08332f1 100755 --- a/scripts/wan2.1/README_TRAIN.md +++ b/scripts/wan2.1/README_TRAIN.md @@ -134,7 +134,7 @@ export NCCL_IB_DISABLE=1 export NCCL_P2P_DISABLE=1 NCCL_DEBUG=INFO -accelerate launch --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag true --use_deepspeed --deepspeed_config_file config/zero_stage3_config.json --deepspeed_multinode_launcherr standard scripts/wan2.1/train.py \ +accelerate launch --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag true --use_deepspeed --deepspeed_config_file config/zero_stage3_config.json --deepspeed_multinode_launcher standard scripts/wan2.1/train.py \ --config_path="config/wan2.1/wan_civitai.yaml" \ --pretrained_model_name_or_path=$MODEL_NAME \ --train_data_dir=$DATASET_NAME \ diff --git a/scripts/wan2.1/README_TRAIN_LORA.md b/scripts/wan2.1/README_TRAIN_LORA.md index d11e449..47a2407 100755 --- a/scripts/wan2.1/README_TRAIN_LORA.md +++ b/scripts/wan2.1/README_TRAIN_LORA.md @@ -124,7 +124,7 @@ export NCCL_IB_DISABLE=1 export NCCL_P2P_DISABLE=1 NCCL_DEBUG=INFO -accelerate launch --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag true --use_deepspeed --deepspeed_config_file config/zero_stage3_config.json --deepspeed_multinode_launcherr standard scripts/wan2.1/train.py \ +accelerate launch --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag true --use_deepspeed --deepspeed_config_file config/zero_stage3_config.json --deepspeed_multinode_launcher standard scripts/wan2.1/train_lora.py \ --config_path="config/wan2.1/wan_civitai.yaml" \ --pretrained_model_name_or_path=$MODEL_NAME \ --train_data_dir=$DATASET_NAME \ diff --git a/scripts/wan2.1/train.py b/scripts/wan2.1/train.py index 2b69084..2b9227f 100755 --- a/scripts/wan2.1/train.py +++ b/scripts/wan2.1/train.py @@ -667,6 +667,9 @@ def parse_args(): parser.add_argument( "--use_deepspeed", action="store_true", help="Whether or not to use deepspeed." ) + parser.add_argument( + "--use_fsdp", action="store_true", help="Whether or not to use fsdp." + ) parser.add_argument( "--low_vram", action="store_true", help="Whether enable low_vram mode." ) @@ -757,12 +760,28 @@ def main(): project_config=accelerator_project_config, ) deepspeed_plugin = accelerator.state.deepspeed_plugin + fsdp_plugin = accelerator.state.fsdp_plugin if deepspeed_plugin is not None: zero_stage = int(deepspeed_plugin.zero_stage) + fsdp_stage = 0 print(f"Using DeepSpeed Zero stage: {zero_stage}") + + args.use_deepspeed = True + elif fsdp_plugin is not None: + from torch.distributed.fsdp import ShardingStrategy + zero_stage = 0 + if fsdp_plugin.sharding_strategy is ShardingStrategy.FULL_SHARD: + fsdp_stage = 3 + elif fsdp_plugin.sharding_strategy is ShardingStrategy.SHARD_GRAD_OP: + fsdp_stage = 2 + print(f"Using FSDP stage: {fsdp_stage}") + + args.use_fsdp = True else: zero_stage = 0 + fsdp_stage = 0 print("DeepSpeed is not enabled.") + if accelerator.is_main_process: writer = SummaryWriter(log_dir=logging_dir) @@ -926,7 +945,42 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): - if not zero_stage == 3: + # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format + if fsdp_stage != 0: + def save_model_hook(models, weights, output_dir): + accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) + if accelerator.is_main_process: + from safetensors.torch import save_file + + safetensor_save_path = os.path.join(output_dir, f"diffusion_pytorch_model.safetensors") + save_file(accelerate_state_dict, safetensor_save_path, metadata={"format": "pt"}) + + with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: + pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) + + def load_model_hook(models, input_dir): + pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") + if os.path.exists(pkl_path): + with open(pkl_path, 'rb') as file: + loaded_number, _ = pickle.load(file) + batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) + print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") + + elif zero_stage == 3: + # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format + def save_model_hook(models, weights, output_dir): + if accelerator.is_main_process: + with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: + pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) + + def load_model_hook(models, input_dir): + pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") + if os.path.exists(pkl_path): + with open(pkl_path, 'rb') as file: + loaded_number, _ = pickle.load(file) + batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) + print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") + else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): if accelerator.is_main_process: @@ -974,20 +1028,6 @@ def main(): loaded_number, _ = pickle.load(file) batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - else: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - if accelerator.is_main_process: - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") accelerator.register_save_state_pre_hook(save_model_hook) accelerator.register_load_state_pre_hook(load_model_hook) @@ -1654,7 +1694,7 @@ def main(): ) # Predict the noise residual - with torch.cuda.amp.autocast(dtype=weight_dtype): + with torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): noise_pred = transformer3d( x=noisy_latents, context=prompt_embeds, @@ -1693,7 +1733,7 @@ def main(): # Backpropagate accelerator.backward(loss) if accelerator.sync_gradients: - if not args.use_deepspeed: + if not args.use_deepspeed and not args.use_fsdp: trainable_params_grads = [p.grad for p in trainable_params if p.grad is not None] trainable_params_total_norm = torch.norm(torch.stack([torch.norm(g.detach(), 2) for g in trainable_params_grads]), 2) max_grad_norm = linear_decay(args.max_grad_norm * args.initial_grad_norm_ratio, args.max_grad_norm, args.abnormal_norm_clip_start, global_step) @@ -1704,14 +1744,14 @@ def main(): else: actual_max_grad_norm = args.max_grad_norm - if not args.use_deepspeed and args.report_model_info and accelerator.is_main_process: + if not args.use_deepspeed and not args.use_fsdp and args.report_model_info and accelerator.is_main_process: if trainable_params_total_norm > 1 and global_step > args.abnormal_norm_clip_start: for name, param in transformer3d.named_parameters(): if param.requires_grad: writer.add_scalar(f'gradients/before_clip_norm/{name}', param.grad.norm(), global_step=global_step) norm_sum = accelerator.clip_grad_norm_(trainable_params, actual_max_grad_norm) - if not args.use_deepspeed and args.report_model_info and accelerator.is_main_process: + if not args.use_deepspeed and not args.use_fsdp and args.report_model_info and accelerator.is_main_process: writer.add_scalar(f'gradients/norm_sum', norm_sum, global_step=global_step) writer.add_scalar(f'gradients/actual_max_grad_norm', actual_max_grad_norm, global_step=global_step) optimizer.step() @@ -1729,7 +1769,7 @@ def main(): train_loss = 0.0 if global_step % args.checkpointing_steps == 0: - if args.use_deepspeed or accelerator.is_main_process: + if args.use_deepspeed or args.use_fsdp or accelerator.is_main_process: # _before_ saving state, check if this save would set us over the `checkpoints_total_limit` if args.checkpoints_total_limit is not None: checkpoints = os.listdir(args.output_dir) @@ -1811,7 +1851,7 @@ def main(): if args.use_ema: ema_transformer3d.copy_to(transformer3d.parameters()) - if args.use_deepspeed or accelerator.is_main_process: + if args.use_deepspeed or args.use_fsdp or accelerator.is_main_process: save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") accelerator.save_state(save_path) logger.info(f"Saved state to {save_path}") diff --git a/scripts/wan2.1/train_lora.py b/scripts/wan2.1/train_lora.py index ba51018..42ca41c 100755 --- a/scripts/wan2.1/train_lora.py +++ b/scripts/wan2.1/train_lora.py @@ -682,6 +682,9 @@ def parse_args(): parser.add_argument( "--use_deepspeed", action="store_true", help="Whether or not to use deepspeed." ) + parser.add_argument( + "--use_fsdp", action="store_true", help="Whether or not to use fsdp." + ) parser.add_argument( "--low_vram", action="store_true", help="Whether enable low_vram mode." ) @@ -756,12 +759,28 @@ def main(): project_config=accelerator_project_config, ) deepspeed_plugin = accelerator.state.deepspeed_plugin + fsdp_plugin = accelerator.state.fsdp_plugin if deepspeed_plugin is not None: zero_stage = int(deepspeed_plugin.zero_stage) + fsdp_stage = 0 print(f"Using DeepSpeed Zero stage: {zero_stage}") + + args.use_deepspeed = True + elif fsdp_plugin is not None: + from torch.distributed.fsdp import ShardingStrategy + zero_stage = 0 + if fsdp_plugin.sharding_strategy is ShardingStrategy.FULL_SHARD: + fsdp_stage = 3 + elif fsdp_plugin.sharding_strategy is ShardingStrategy.SHARD_GRAD_OP: + fsdp_stage = 2 + print(f"Using FSDP stage: {fsdp_stage}") + + args.use_fsdp = True else: zero_stage = 0 + fsdp_stage = 0 print("DeepSpeed is not enabled.") + if accelerator.is_main_process: writer = SummaryWriter(log_dir=logging_dir) @@ -911,15 +930,34 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if zero_stage != 3: + if fsdp_stage != 0: + def save_model_hook(models, weights, output_dir): + accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) + if accelerator.is_main_process: + from safetensors.torch import save_file + + safetensor_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model.safetensors") + network_state_dict = {} + for key in accelerate_state_dict: + if "network" in key: + network_state_dict[key.replace("network.", "")] = accelerate_state_dict[key] + + save_file(network_state_dict, safetensor_save_path, metadata={"format": "pt"}) + + with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: + pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) + + def load_model_hook(models, input_dir): + pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") + if os.path.exists(pkl_path): + with open(pkl_path, 'rb') as file: + loaded_number, _ = pickle.load(file) + batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) + print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") + + elif zero_stage == 3: def save_model_hook(models, weights, output_dir): if accelerator.is_main_process: - safetensor_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model.safetensors") - save_model(safetensor_save_path, accelerator.unwrap_model(models[-1])) - if not args.use_deepspeed: - for _ in range(len(weights)): - weights.pop() - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) @@ -933,6 +971,12 @@ def main(): else: def save_model_hook(models, weights, output_dir): if accelerator.is_main_process: + safetensor_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model.safetensors") + save_model(safetensor_save_path, accelerator.unwrap_model(models[-1])) + if not args.use_deepspeed: + for _ in range(len(weights)): + weights.pop() + with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) @@ -1222,9 +1266,16 @@ def main(): ) # Prepare everything with our `accelerator`. - network, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( - network, optimizer, train_dataloader, lr_scheduler - ) + if fsdp_stage != 0: + transformer3d.network = network + transformer3d = transformer3d.to(weight_dtype) + transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( + transformer3d, optimizer, train_dataloader, lr_scheduler + ) + else: + network, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( + network, optimizer, train_dataloader, lr_scheduler + ) # Move text_encode and vae to gpu and cast to weight_dtype vae.to(accelerator.device, dtype=weight_dtype) @@ -1636,7 +1687,7 @@ def main(): target_shape[1] ) # Predict the noise residual - with torch.cuda.amp.autocast(dtype=weight_dtype): + with torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): noise_pred = transformer3d( x=noisy_latents, context=prompt_embeds, @@ -1688,7 +1739,7 @@ def main(): train_loss = 0.0 if global_step % args.checkpointing_steps == 0: - if args.use_deepspeed or accelerator.is_main_process: + if args.use_deepspeed or args.use_fsdp or accelerator.is_main_process: # _before_ saving state, check if this save would set us over the `checkpoints_total_limit` if args.checkpoints_total_limit is not None: checkpoints = os.listdir(args.output_dir) diff --git a/scripts/wan2.1_fun/README_TRAIN.md b/scripts/wan2.1_fun/README_TRAIN.md index 7af888d..95bb626 100755 --- a/scripts/wan2.1_fun/README_TRAIN.md +++ b/scripts/wan2.1_fun/README_TRAIN.md @@ -132,7 +132,7 @@ export NCCL_IB_DISABLE=1 export NCCL_P2P_DISABLE=1 NCCL_DEBUG=INFO -accelerate launch --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag true --use_deepspeed --deepspeed_config_file config/zero_stage3_config.json --deepspeed_multinode_launcherr standard scripts/wan2.1/train.py \ +accelerate launch --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag true --use_deepspeed --deepspeed_config_file config/zero_stage3_config.json --deepspeed_multinode_launcher standard scripts/wan2.1/train.py \ --config_path="config/wan2.1/wan_civitai.yaml" \ --pretrained_model_name_or_path=$MODEL_NAME \ --train_data_dir=$DATASET_NAME \ diff --git a/scripts/wan2.1_fun/README_TRAIN_CONTROL.md b/scripts/wan2.1_fun/README_TRAIN_CONTROL.md index 1168ca0..aec1202 100755 --- a/scripts/wan2.1_fun/README_TRAIN_CONTROL.md +++ b/scripts/wan2.1_fun/README_TRAIN_CONTROL.md @@ -168,7 +168,7 @@ export NCCL_IB_DISABLE=1 export NCCL_P2P_DISABLE=1 NCCL_DEBUG=INFO -accelerate launch --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag true --use_deepspeed --deepspeed_config_file config/zero_stage3_config.json --deepspeed_multinode_launcherr standard scripts/wan2.1_fun/train_control.py \ +accelerate launch --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag true --use_deepspeed --deepspeed_config_file config/zero_stage3_config.json --deepspeed_multinode_launcher standard scripts/wan2.1_fun/train_control.py \ --config_path="config/wan2.1/wan_civitai.yaml" \ --pretrained_model_name_or_path=$MODEL_NAME \ --train_data_dir=$DATASET_NAME \ @@ -319,7 +319,7 @@ export NCCL_IB_DISABLE=1 export NCCL_P2P_DISABLE=1 NCCL_DEBUG=INFO -accelerate launch --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag true --use_deepspeed --deepspeed_config_file config/zero_stage3_config.json --deepspeed_multinode_launcherr standard scripts/wan2.1_fun/train_control.py \ +accelerate launch --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag true --use_deepspeed --deepspeed_config_file config/zero_stage3_config.json --deepspeed_multinode_launcher standard scripts/wan2.1_fun/train_control.py \ --config_path="config/wan2.1/wan_civitai.yaml" \ --pretrained_model_name_or_path=$MODEL_NAME \ --train_data_dir=$DATASET_NAME \ diff --git a/scripts/wan2.1_fun/README_TRAIN_CONTROL_LORA.md b/scripts/wan2.1_fun/README_TRAIN_CONTROL_LORA.md index 414d9ca..fac39a7 100755 --- a/scripts/wan2.1_fun/README_TRAIN_CONTROL_LORA.md +++ b/scripts/wan2.1_fun/README_TRAIN_CONTROL_LORA.md @@ -162,7 +162,7 @@ export NCCL_IB_DISABLE=1 export NCCL_P2P_DISABLE=1 NCCL_DEBUG=INFO -accelerate launch --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag true --use_deepspeed --deepspeed_config_file config/zero_stage3_config.json --deepspeed_multinode_launcherr standard scripts/wan2.1_fun/train_control.py \ +accelerate launch --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag true --use_deepspeed --deepspeed_config_file config/zero_stage3_config.json --deepspeed_multinode_launcher standard scripts/wan2.1_fun/train_control_lora.py \ --config_path="config/wan2.1/wan_civitai.yaml" \ --pretrained_model_name_or_path=$MODEL_NAME \ --train_data_dir=$DATASET_NAME \ @@ -305,7 +305,7 @@ export NCCL_IB_DISABLE=1 export NCCL_P2P_DISABLE=1 NCCL_DEBUG=INFO -accelerate launch --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag true --use_deepspeed --deepspeed_config_file config/zero_stage3_config.json --deepspeed_multinode_launcherr standard scripts/wan2.1_fun/train_control.py \ +accelerate launch --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag true --use_deepspeed --deepspeed_config_file config/zero_stage3_config.json --deepspeed_multinode_launcher standard scripts/wan2.1_fun/train_control_lora.py \ --config_path="config/wan2.1/wan_civitai.yaml" \ --pretrained_model_name_or_path=$MODEL_NAME \ --train_data_dir=$DATASET_NAME \ diff --git a/scripts/wan2.1_fun/README_TRAIN_LORA.md b/scripts/wan2.1_fun/README_TRAIN_LORA.md index e66c6da..ae16c71 100755 --- a/scripts/wan2.1_fun/README_TRAIN_LORA.md +++ b/scripts/wan2.1_fun/README_TRAIN_LORA.md @@ -126,7 +126,7 @@ export NCCL_IB_DISABLE=1 export NCCL_P2P_DISABLE=1 NCCL_DEBUG=INFO -accelerate launch --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag true --use_deepspeed --deepspeed_config_file config/zero_stage3_config.json --deepspeed_multinode_launcherr standard scripts/wan2.1_fun/train.py \ +accelerate launch --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag true --use_deepspeed --deepspeed_config_file config/zero_stage3_config.json --deepspeed_multinode_launcher standard scripts/wan2.1_fun/train_lora.py \ --config_path="config/wan2.1/wan_civitai.yaml" \ --pretrained_model_name_or_path=$MODEL_NAME \ --train_data_dir=$DATASET_NAME \ diff --git a/scripts/wan2.1_fun/train.py b/scripts/wan2.1_fun/train.py index 019981d..7743942 100755 --- a/scripts/wan2.1_fun/train.py +++ b/scripts/wan2.1_fun/train.py @@ -637,6 +637,9 @@ def parse_args(): parser.add_argument( "--use_deepspeed", action="store_true", help="Whether or not to use deepspeed." ) + parser.add_argument( + "--use_fsdp", action="store_true", help="Whether or not to use fsdp." + ) parser.add_argument( "--low_vram", action="store_true", help="Whether enable low_vram mode." ) @@ -727,12 +730,28 @@ def main(): project_config=accelerator_project_config, ) deepspeed_plugin = accelerator.state.deepspeed_plugin + fsdp_plugin = accelerator.state.fsdp_plugin if deepspeed_plugin is not None: zero_stage = int(deepspeed_plugin.zero_stage) + fsdp_stage = 0 print(f"Using DeepSpeed Zero stage: {zero_stage}") + + args.use_deepspeed = True + elif fsdp_plugin is not None: + from torch.distributed.fsdp import ShardingStrategy + zero_stage = 0 + if fsdp_plugin.sharding_strategy is ShardingStrategy.FULL_SHARD: + fsdp_stage = 3 + elif fsdp_plugin.sharding_strategy is ShardingStrategy.SHARD_GRAD_OP: + fsdp_stage = 2 + print(f"Using FSDP stage: {fsdp_stage}") + + args.use_fsdp = True else: zero_stage = 0 + fsdp_stage = 0 print("DeepSpeed is not enabled.") + if accelerator.is_main_process: writer = SummaryWriter(log_dir=logging_dir) @@ -896,7 +915,42 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): - if not zero_stage == 3: + # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format + if fsdp_stage != 0: + def save_model_hook(models, weights, output_dir): + accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) + if accelerator.is_main_process: + from safetensors.torch import save_file + + safetensor_save_path = os.path.join(output_dir, f"diffusion_pytorch_model.safetensors") + save_file(accelerate_state_dict, safetensor_save_path, metadata={"format": "pt"}) + + with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: + pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) + + def load_model_hook(models, input_dir): + pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") + if os.path.exists(pkl_path): + with open(pkl_path, 'rb') as file: + loaded_number, _ = pickle.load(file) + batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) + print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") + + elif zero_stage == 3: + # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format + def save_model_hook(models, weights, output_dir): + if accelerator.is_main_process: + with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: + pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) + + def load_model_hook(models, input_dir): + pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") + if os.path.exists(pkl_path): + with open(pkl_path, 'rb') as file: + loaded_number, _ = pickle.load(file) + batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) + print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") + else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): if accelerator.is_main_process: @@ -944,20 +998,6 @@ def main(): loaded_number, _ = pickle.load(file) batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - else: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - if accelerator.is_main_process: - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") accelerator.register_save_state_pre_hook(save_model_hook) accelerator.register_load_state_pre_hook(load_model_hook) @@ -1663,7 +1703,7 @@ def main(): ) # Predict the noise residual - with torch.cuda.amp.autocast(dtype=weight_dtype): + with torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): noise_pred = transformer3d( x=noisy_latents, context=prompt_embeds, @@ -1702,7 +1742,7 @@ def main(): # Backpropagate accelerator.backward(loss) if accelerator.sync_gradients: - if not args.use_deepspeed: + if not args.use_deepspeed and not args.use_fsdp: trainable_params_grads = [p.grad for p in trainable_params if p.grad is not None] trainable_params_total_norm = torch.norm(torch.stack([torch.norm(g.detach(), 2) for g in trainable_params_grads]), 2) max_grad_norm = linear_decay(args.max_grad_norm * args.initial_grad_norm_ratio, args.max_grad_norm, args.abnormal_norm_clip_start, global_step) @@ -1713,14 +1753,14 @@ def main(): else: actual_max_grad_norm = args.max_grad_norm - if not args.use_deepspeed and args.report_model_info and accelerator.is_main_process: + if not args.use_deepspeed and not args.use_fsdp and args.report_model_info and accelerator.is_main_process: if trainable_params_total_norm > 1 and global_step > args.abnormal_norm_clip_start: for name, param in transformer3d.named_parameters(): if param.requires_grad: writer.add_scalar(f'gradients/before_clip_norm/{name}', param.grad.norm(), global_step=global_step) norm_sum = accelerator.clip_grad_norm_(trainable_params, actual_max_grad_norm) - if not args.use_deepspeed and args.report_model_info and accelerator.is_main_process: + if not args.use_deepspeed and not args.use_fsdp and args.report_model_info and accelerator.is_main_process: writer.add_scalar(f'gradients/norm_sum', norm_sum, global_step=global_step) writer.add_scalar(f'gradients/actual_max_grad_norm', actual_max_grad_norm, global_step=global_step) optimizer.step() @@ -1738,7 +1778,7 @@ def main(): train_loss = 0.0 if global_step % args.checkpointing_steps == 0: - if args.use_deepspeed or accelerator.is_main_process: + if args.use_deepspeed or args.use_fsdp or accelerator.is_main_process: # _before_ saving state, check if this save would set us over the `checkpoints_total_limit` if args.checkpoints_total_limit is not None: checkpoints = os.listdir(args.output_dir) @@ -1820,7 +1860,7 @@ def main(): if args.use_ema: ema_transformer3d.copy_to(transformer3d.parameters()) - if args.use_deepspeed or accelerator.is_main_process: + if args.use_deepspeed or args.use_fsdp or accelerator.is_main_process: save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") accelerator.save_state(save_path) logger.info(f"Saved state to {save_path}") diff --git a/scripts/wan2.1_fun/train_control.py b/scripts/wan2.1_fun/train_control.py index fb361c6..9886c14 100755 --- a/scripts/wan2.1_fun/train_control.py +++ b/scripts/wan2.1_fun/train_control.py @@ -553,6 +553,9 @@ def parse_args(): parser.add_argument( "--use_deepspeed", action="store_true", help="Whether or not to use deepspeed." ) + parser.add_argument( + "--use_fsdp", action="store_true", help="Whether or not to use fsdp." + ) parser.add_argument( "--low_vram", action="store_true", help="Whether enable low_vram mode." ) @@ -659,12 +662,28 @@ def main(): project_config=accelerator_project_config, ) deepspeed_plugin = accelerator.state.deepspeed_plugin + fsdp_plugin = accelerator.state.fsdp_plugin if deepspeed_plugin is not None: zero_stage = int(deepspeed_plugin.zero_stage) + fsdp_stage = 0 print(f"Using DeepSpeed Zero stage: {zero_stage}") + + args.use_deepspeed = True + elif fsdp_plugin is not None: + from torch.distributed.fsdp import ShardingStrategy + zero_stage = 0 + if fsdp_plugin.sharding_strategy is ShardingStrategy.FULL_SHARD: + fsdp_stage = 3 + elif fsdp_plugin.sharding_strategy is ShardingStrategy.SHARD_GRAD_OP: + fsdp_stage = 2 + print(f"Using FSDP stage: {fsdp_stage}") + + args.use_fsdp = True else: zero_stage = 0 + fsdp_stage = 0 print("DeepSpeed is not enabled.") + if accelerator.is_main_process: writer = SummaryWriter(log_dir=logging_dir) @@ -826,7 +845,42 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): - if not zero_stage == 3: + # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format + if fsdp_stage != 0: + def save_model_hook(models, weights, output_dir): + accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) + if accelerator.is_main_process: + from safetensors.torch import save_file + + safetensor_save_path = os.path.join(output_dir, f"diffusion_pytorch_model.safetensors") + save_file(accelerate_state_dict, safetensor_save_path, metadata={"format": "pt"}) + + with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: + pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) + + def load_model_hook(models, input_dir): + pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") + if os.path.exists(pkl_path): + with open(pkl_path, 'rb') as file: + loaded_number, _ = pickle.load(file) + batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) + print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") + + elif zero_stage == 3: + # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format + def save_model_hook(models, weights, output_dir): + if accelerator.is_main_process: + with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: + pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) + + def load_model_hook(models, input_dir): + pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") + if os.path.exists(pkl_path): + with open(pkl_path, 'rb') as file: + loaded_number, _ = pickle.load(file) + batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) + print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") + else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): if accelerator.is_main_process: @@ -874,20 +928,6 @@ def main(): loaded_number, _ = pickle.load(file) batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - else: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - if accelerator.is_main_process: - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") accelerator.register_save_state_pre_hook(save_model_hook) accelerator.register_load_state_pre_hook(load_model_hook) @@ -1698,7 +1738,7 @@ def main(): ) # Predict the noise residual - with torch.cuda.amp.autocast(dtype=weight_dtype): + with torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): noise_pred = transformer3d( x=noisy_latents, context=prompt_embeds, @@ -1739,7 +1779,7 @@ def main(): # Backpropagate accelerator.backward(loss) if accelerator.sync_gradients: - if not args.use_deepspeed: + if not args.use_deepspeed and not args.use_fsdp: trainable_params_grads = [p.grad for p in trainable_params if p.grad is not None] trainable_params_total_norm = torch.norm(torch.stack([torch.norm(g.detach(), 2) for g in trainable_params_grads]), 2) max_grad_norm = linear_decay(args.max_grad_norm * args.initial_grad_norm_ratio, args.max_grad_norm, args.abnormal_norm_clip_start, global_step) @@ -1750,14 +1790,14 @@ def main(): else: actual_max_grad_norm = args.max_grad_norm - if not args.use_deepspeed and args.report_model_info and accelerator.is_main_process: + if not args.use_deepspeed and not args.use_fsdp and args.report_model_info and accelerator.is_main_process: if trainable_params_total_norm > 1 and global_step > args.abnormal_norm_clip_start: for name, param in transformer3d.named_parameters(): if param.requires_grad: writer.add_scalar(f'gradients/before_clip_norm/{name}', param.grad.norm(), global_step=global_step) norm_sum = accelerator.clip_grad_norm_(trainable_params, actual_max_grad_norm) - if not args.use_deepspeed and args.report_model_info and accelerator.is_main_process: + if not args.use_deepspeed and not args.use_fsdp and args.report_model_info and accelerator.is_main_process: writer.add_scalar(f'gradients/norm_sum', norm_sum, global_step=global_step) writer.add_scalar(f'gradients/actual_max_grad_norm', actual_max_grad_norm, global_step=global_step) optimizer.step() @@ -1775,7 +1815,7 @@ def main(): train_loss = 0.0 if global_step % args.checkpointing_steps == 0: - if args.use_deepspeed or accelerator.is_main_process: + if args.use_deepspeed or args.use_fsdp or accelerator.is_main_process: # _before_ saving state, check if this save would set us over the `checkpoints_total_limit` if args.checkpoints_total_limit is not None: checkpoints = os.listdir(args.output_dir) @@ -1857,7 +1897,7 @@ def main(): if args.use_ema: ema_transformer3d.copy_to(transformer3d.parameters()) - if args.use_deepspeed or accelerator.is_main_process: + if args.use_deepspeed or args.use_fsdp or accelerator.is_main_process: save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") accelerator.save_state(save_path) logger.info(f"Saved state to {save_path}") diff --git a/scripts/wan2.1_fun/train_control_lora.py b/scripts/wan2.1_fun/train_control_lora.py index f38e0fc..9e85399 100755 --- a/scripts/wan2.1_fun/train_control_lora.py +++ b/scripts/wan2.1_fun/train_control_lora.py @@ -561,6 +561,9 @@ def parse_args(): parser.add_argument( "--use_deepspeed", action="store_true", help="Whether or not to use deepspeed." ) + parser.add_argument( + "--use_fsdp", action="store_true", help="Whether or not to use fsdp." + ) parser.add_argument( "--low_vram", action="store_true", help="Whether enable low_vram mode." ) @@ -657,18 +660,28 @@ def main(): project_config=accelerator_project_config, ) deepspeed_plugin = accelerator.state.deepspeed_plugin + fsdp_plugin = accelerator.state.fsdp_plugin if deepspeed_plugin is not None: zero_stage = int(deepspeed_plugin.zero_stage) + fsdp_stage = 0 print(f"Using DeepSpeed Zero stage: {zero_stage}") + + args.use_deepspeed = True + elif fsdp_plugin is not None: + from torch.distributed.fsdp import ShardingStrategy + zero_stage = 0 + if fsdp_plugin.sharding_strategy is ShardingStrategy.FULL_SHARD: + fsdp_stage = 3 + elif fsdp_plugin.sharding_strategy is ShardingStrategy.SHARD_GRAD_OP: + fsdp_stage = 2 + print(f"Using FSDP stage: {fsdp_stage}") + + args.use_fsdp = True else: zero_stage = 0 + fsdp_stage = 0 print("DeepSpeed is not enabled.") - if zero_stage == 3: - accelerator_transformer3d = Accelerator( - gradient_accumulation_steps=args.gradient_accumulation_steps, - mixed_precision=args.mixed_precision, - project_config=accelerator_project_config, - ) + if accelerator.is_main_process: writer = SummaryWriter(log_dir=logging_dir) @@ -817,7 +830,45 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if zero_stage != 3: + if fsdp_stage != 0: + def save_model_hook(models, weights, output_dir): + accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) + if accelerator.is_main_process: + from safetensors.torch import save_file + + safetensor_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model.safetensors") + network_state_dict = {} + for key in accelerate_state_dict: + if "network" in key: + network_state_dict[key.replace("network.", "")] = accelerate_state_dict[key] + + save_file(network_state_dict, safetensor_save_path, metadata={"format": "pt"}) + + with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: + pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) + + def load_model_hook(models, input_dir): + pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") + if os.path.exists(pkl_path): + with open(pkl_path, 'rb') as file: + loaded_number, _ = pickle.load(file) + batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) + print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") + + elif zero_stage == 3: + def save_model_hook(models, weights, output_dir): + if accelerator.is_main_process: + with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: + pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) + + def load_model_hook(models, input_dir): + pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") + if os.path.exists(pkl_path): + with open(pkl_path, 'rb') as file: + loaded_number, _ = pickle.load(file) + batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) + print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") + else: def save_model_hook(models, weights, output_dir): if accelerator.is_main_process: safetensor_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model.safetensors") @@ -836,20 +887,6 @@ def main(): loaded_number, _ = pickle.load(file) batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - else: - def save_model_hook(models, weights, output_dir): - if accelerator.is_main_process: - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - accelerator.register_save_state_pre_hook(save_model_hook) accelerator.register_load_state_pre_hook(load_model_hook) @@ -1236,13 +1273,22 @@ def main(): ) # Prepare everything with our `accelerator`. - network, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( - network, optimizer, train_dataloader, lr_scheduler - ) - if zero_stage == 3: - transformer3d = accelerator_transformer3d.prepare( - transformer3d + if fsdp_stage != 0: + transformer3d.network = network + transformer3d = transformer3d.to(weight_dtype) + transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( + transformer3d, optimizer, train_dataloader, lr_scheduler ) + else: + network, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( + network, optimizer, train_dataloader, lr_scheduler + ) + + if zero_stage == 3: + from functools import partial + from videox_fun.dist import set_multi_gpus_devices, shard_model + shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype) + transformer3d = shard_fn(transformer3d) # Move text_encode and vae to gpu and cast to weight_dtype vae.to(accelerator.device, dtype=weight_dtype) @@ -1696,7 +1742,7 @@ def main(): ) # Predict the noise residual - with torch.cuda.amp.autocast(dtype=weight_dtype): + with torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): noise_pred = transformer3d( x=noisy_latents, context=prompt_embeds, @@ -1750,7 +1796,7 @@ def main(): train_loss = 0.0 if global_step % args.checkpointing_steps == 0: - if args.use_deepspeed or accelerator.is_main_process: + if args.use_deepspeed or args.use_fsdp or accelerator.is_main_process: # _before_ saving state, check if this save would set us over the `checkpoints_total_limit` if args.checkpoints_total_limit is not None: checkpoints = os.listdir(args.output_dir) diff --git a/scripts/wan2.1_fun/train_lora.py b/scripts/wan2.1_fun/train_lora.py index b72a61c..090f14f 100755 --- a/scripts/wan2.1_fun/train_lora.py +++ b/scripts/wan2.1_fun/train_lora.py @@ -639,6 +639,9 @@ def parse_args(): parser.add_argument( "--use_deepspeed", action="store_true", help="Whether or not to use deepspeed." ) + parser.add_argument( + "--use_fsdp", action="store_true", help="Whether or not to use fsdp." + ) parser.add_argument( "--low_vram", action="store_true", help="Whether enable low_vram mode." ) @@ -719,18 +722,28 @@ def main(): project_config=accelerator_project_config, ) deepspeed_plugin = accelerator.state.deepspeed_plugin + fsdp_plugin = accelerator.state.fsdp_plugin if deepspeed_plugin is not None: zero_stage = int(deepspeed_plugin.zero_stage) + fsdp_stage = 0 print(f"Using DeepSpeed Zero stage: {zero_stage}") + + args.use_deepspeed = True + elif fsdp_plugin is not None: + from torch.distributed.fsdp import ShardingStrategy + zero_stage = 0 + if fsdp_plugin.sharding_strategy is ShardingStrategy.FULL_SHARD: + fsdp_stage = 3 + elif fsdp_plugin.sharding_strategy is ShardingStrategy.SHARD_GRAD_OP: + fsdp_stage = 2 + print(f"Using FSDP stage: {fsdp_stage}") + + args.use_fsdp = True else: zero_stage = 0 + fsdp_stage = 0 print("DeepSpeed is not enabled.") - if zero_stage == 3: - accelerator_transformer3d = Accelerator( - gradient_accumulation_steps=args.gradient_accumulation_steps, - mixed_precision=args.mixed_precision, - project_config=accelerator_project_config, - ) + if accelerator.is_main_process: writer = SummaryWriter(log_dir=logging_dir) @@ -881,15 +894,34 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if zero_stage != 3: + if fsdp_stage != 0: + def save_model_hook(models, weights, output_dir): + accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) + if accelerator.is_main_process: + from safetensors.torch import save_file + + safetensor_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model.safetensors") + network_state_dict = {} + for key in accelerate_state_dict: + if "network" in key: + network_state_dict[key.replace("network.", "")] = accelerate_state_dict[key] + + save_file(network_state_dict, safetensor_save_path, metadata={"format": "pt"}) + + with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: + pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) + + def load_model_hook(models, input_dir): + pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") + if os.path.exists(pkl_path): + with open(pkl_path, 'rb') as file: + loaded_number, _ = pickle.load(file) + batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) + print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") + + elif zero_stage == 3: def save_model_hook(models, weights, output_dir): if accelerator.is_main_process: - safetensor_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model.safetensors") - save_model(safetensor_save_path, accelerator.unwrap_model(models[-1])) - if not args.use_deepspeed: - for _ in range(len(weights)): - weights.pop() - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) @@ -903,6 +935,12 @@ def main(): else: def save_model_hook(models, weights, output_dir): if accelerator.is_main_process: + safetensor_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model.safetensors") + save_model(safetensor_save_path, accelerator.unwrap_model(models[-1])) + if not args.use_deepspeed: + for _ in range(len(weights)): + weights.pop() + with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) @@ -1226,13 +1264,23 @@ def main(): ) # Prepare everything with our `accelerator`. - network, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( - network, optimizer, train_dataloader, lr_scheduler - ) - if zero_stage == 3: - transformer3d = accelerator_transformer3d.prepare( - transformer3d + if fsdp_stage != 0: + transformer3d.network = network + transformer3d = transformer3d.to(weight_dtype) + transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( + transformer3d, optimizer, train_dataloader, lr_scheduler ) + else: + network, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( + network, optimizer, train_dataloader, lr_scheduler + ) + + if zero_stage == 3: + from functools import partial + from videox_fun.dist import set_multi_gpus_devices, shard_model + shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype) + transformer3d = shard_fn(transformer3d) + # Move text_encode and vae to gpu and cast to weight_dtype vae.to(accelerator.device, dtype=weight_dtype) transformer3d.to(accelerator.device, dtype=weight_dtype) @@ -1653,7 +1701,7 @@ def main(): target_shape[1] ) # Predict the noise residual - with torch.cuda.amp.autocast(dtype=weight_dtype): + with torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): noise_pred = transformer3d( x=noisy_latents, context=prompt_embeds, @@ -1705,7 +1753,7 @@ def main(): train_loss = 0.0 if global_step % args.checkpointing_steps == 0: - if args.use_deepspeed or accelerator.is_main_process: + if args.use_deepspeed or args.use_fsdp or accelerator.is_main_process: # _before_ saving state, check if this save would set us over the `checkpoints_total_limit` if args.checkpoints_total_limit is not None: checkpoints = os.listdir(args.output_dir) diff --git a/videox_fun/dist/__init__.py b/videox_fun/dist/__init__.py index a42b71e..ac1612e 100755 --- a/videox_fun/dist/__init__.py +++ b/videox_fun/dist/__init__.py @@ -28,6 +28,17 @@ if importlib.util.find_spec("pai_fuser") is not None: from pai_fuser.core.rope import ENABLE_KERNEL, usp_fast_rope_apply_qk if ENABLE_KERNEL: - wan_xfuser.rope_apply_qk = usp_fast_rope_apply_qk - rope_apply_qk = usp_fast_rope_apply_qk - print("Import PAI Fast rope") \ No newline at end of file + import torch + from .wan_xfuser import rope_apply + + def adaptive_fast_usp_rope_apply_qk(q, k, grid_sizes, freqs): + if torch.is_grad_enabled(): + q = rope_apply(q, grid_sizes, freqs) + k = rope_apply(k, grid_sizes, freqs) + return q, k + else: + return usp_fast_rope_apply_qk(q, k, grid_sizes, freqs) + + wan_xfuser.rope_apply_qk = adaptive_fast_usp_rope_apply_qk + rope_apply_qk = adaptive_fast_usp_rope_apply_qk + print("Import PAI Fast rope") diff --git a/videox_fun/models/__init__.py b/videox_fun/models/__init__.py index 241e212..2878198 100755 --- a/videox_fun/models/__init__.py +++ b/videox_fun/models/__init__.py @@ -37,6 +37,16 @@ if importlib.util.find_spec("pai_fuser") is not None: from pai_fuser.core.rope import ENABLE_KERNEL, fast_rope_apply_qk if ENABLE_KERNEL: - wan_transformer3d.rope_apply_qk = fast_rope_apply_qk - rope_apply_qk = fast_rope_apply_qk - print("Import PAI Fast rope") \ No newline at end of file + from .wan_transformer3d import rope_apply + + def adaptive_fast_rope_apply_qk(q, k, grid_sizes, freqs): + if torch.is_grad_enabled(): + q = rope_apply(q, grid_sizes, freqs) + k = rope_apply(k, grid_sizes, freqs) + return q, k + else: + return fast_rope_apply_qk(q, k, grid_sizes, freqs) + + wan_transformer3d.rope_apply_qk = adaptive_fast_rope_apply_qk + rope_apply_qk = adaptive_fast_rope_apply_qk + print("Import PAI Fast rope") diff --git a/videox_fun/models/wan_image_encoder.py b/videox_fun/models/wan_image_encoder.py old mode 100644 new mode 100755 index 68c442a..0265743 --- a/videox_fun/models/wan_image_encoder.py +++ b/videox_fun/models/wan_image_encoder.py @@ -7,7 +7,7 @@ import torch.nn as nn import torch.nn.functional as F import torchvision.transforms as T -from .wan_transformer3d import attention +from .wan_transformer3d import attention, flash_attention from .wan_xlm_roberta import XLMRoberta from diffusers.configuration_utils import ConfigMixin from diffusers.loaders.single_file_model import FromOriginalModelMixin @@ -84,7 +84,7 @@ class SelfAttention(nn.Module): # compute attention p = self.attn_dropout if self.training else 0.0 - x = attention(q, k, v, dropout_p=p, causal=self.causal) + x = attention(q, k, v, dropout_p=p, causal=self.causal, attention_type="none") x = x.reshape(b, s, c) # output diff --git a/videox_fun/models/wan_transformer3d.py b/videox_fun/models/wan_transformer3d.py index 41e6c07..16bb296 100755 --- a/videox_fun/models/wan_transformer3d.py +++ b/videox_fun/models/wan_transformer3d.py @@ -183,8 +183,9 @@ def attention( deterministic=False, dtype=torch.bfloat16, fa_version=None, + attention_type=None, ): - attention_type = os.environ.get("VIDEOX_ATTENTION_TYPE", "FLASH_ATTENTION") + attention_type = os.environ.get("VIDEOX_ATTENTION_TYPE", "FLASH_ATTENTION") if attention_type is None else attention_type if torch.is_grad_enabled() and attention_type == "SAGE_ATTENTION": attention_type = "FLASH_ATTENTION" @@ -921,6 +922,7 @@ class WanTransformer3DModel(ModelMixin, ConfigMixin, FromOriginalModelMixin): y=None, y_camera=None, full_ref=None, + subject_ref=None, cond_flag=True, ): r""" @@ -974,6 +976,13 @@ class WanTransformer3DModel(ModelMixin, ConfigMixin, FromOriginalModelMixin): seq_len += full_ref.size(1) x = [torch.concat([_full_ref.unsqueeze(0), u], dim=1) for _full_ref, u in zip(full_ref, x)] + if subject_ref is not None: + subject_ref_frames = subject_ref.size(2) + subject_ref = self.patch_embedding(subject_ref).flatten(2).transpose(1, 2) + grid_sizes = torch.stack([torch.tensor([u[0] + subject_ref_frames, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device) + seq_len += subject_ref.size(1) + x = [torch.concat([u, _subject_ref.unsqueeze(0)], dim=1) for _subject_ref, u in zip(subject_ref, x)] + seq_lens = torch.tensor([u.size(1) for u in x], dtype=torch.long) if self.sp_world_size > 1: seq_len = int(math.ceil(seq_len / self.sp_world_size)) * self.sp_world_size @@ -1125,6 +1134,11 @@ class WanTransformer3DModel(ModelMixin, ConfigMixin, FromOriginalModelMixin): x = x[:, full_ref_length:] grid_sizes = torch.stack([torch.tensor([u[0] - 1, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device) + if subject_ref is not None: + subject_ref_length = subject_ref[0].size(1) + x = x[:, :-subject_ref_length] + grid_sizes = torch.stack([torch.tensor([u[0] - subject_ref_frames, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device) + # head x = self.head(x, e) diff --git a/videox_fun/pipeline/__init__.py b/videox_fun/pipeline/__init__.py index fc8afdc..2b01503 100755 --- a/videox_fun/pipeline/__init__.py +++ b/videox_fun/pipeline/__init__.py @@ -4,6 +4,7 @@ from .pipeline_cogvideox_fun_inpaint import CogVideoXFunInpaintPipeline from .pipeline_wan_fun import WanFunPipeline from .pipeline_wan_fun_inpaint import WanFunInpaintPipeline from .pipeline_wan_fun_control import WanFunControlPipeline +from .pipeline_wan_phantom import WanFunPhantomPipeline WanPipeline = WanFunPipeline WanI2VPipeline = WanFunInpaintPipeline @@ -17,4 +18,5 @@ if importlib.util.find_spec("pai_fuser") is not None: WanFunPipeline.__call__ = sparse_reset(WanFunPipeline.__call__) WanFunControlPipeline.__call__ = sparse_reset(WanFunControlPipeline.__call__) WanI2VPipeline.__call__ = sparse_reset(WanI2VPipeline.__call__) - WanPipeline.__call__ = sparse_reset(WanPipeline.__call__) \ No newline at end of file + WanPipeline.__call__ = sparse_reset(WanPipeline.__call__) + WanFunPhantomPipeline.__call__ = sparse_reset(WanFunPhantomPipeline.__call__) \ No newline at end of file diff --git a/videox_fun/pipeline/pipeline_wan_phantom.py b/videox_fun/pipeline/pipeline_wan_phantom.py new file mode 100644 index 0000000..2935089 --- /dev/null +++ b/videox_fun/pipeline/pipeline_wan_phantom.py @@ -0,0 +1,695 @@ +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, CLIPModel, + WanT5EncoderModel, WanTransformer3DModel) +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 WanFunPhantomPipeline(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->clip_image_encoder->transformer->vae" + + _callback_tensor_inputs = [ + "latents", + "prompt_embeds", + "negative_prompt_embeds", + ] + + def __init__( + self, + tokenizer: AutoTokenizer, + text_encoder: WanT5EncoderModel, + vae: AutoencoderKLWan, + transformer: WanTransformer3DModel, + 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.spacial_compression_ratio) + self.image_processor = VaeImageProcessor(vae_scale_factor=self.vae.spacial_compression_ratio) + self.mask_processor = VaeImageProcessor( + vae_scale_factor=self.vae.spacial_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 + ): + 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, + height // self.vae.spacial_compression_ratio, + width // self.vae.spacial_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 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, + 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) + + # 5. Prepare latents. + latent_channels = self.vae.config.latent_channels + latents = self.prepare_latents( + batch_size * num_videos_per_prompt, + latent_channels, + num_frames, + height, + width, + weight_dtype, + device, + generator, + latents, + ) + if comfyui_progressbar: + pbar.update(1) + + 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_latentes = torch.cat( + [ + self.prepare_control_latents( + None, + subject_ref_images[:, :, i:i+1], + batch_size, + height, + width, + weight_dtype, + device, + generator, + do_classifier_free_guidance + )[1] for i in range(video_length) + ], dim = 2 + ) + + if comfyui_progressbar: + pbar.update(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, (num_frames - 1) // self.vae.temporal_compression_ratio + 1, width // self.vae.spacial_compression_ratio, height // self.vae.spacial_compression_ratio) + 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) + + if subject_ref_images is not None: + subject_ref = ( + torch.cat( + [torch.zeros_like(subject_ref_images_latentes), subject_ref_images_latentes] + ) if do_classifier_free_guidance else subject_ref_images_latentes + ).to(device, weight_dtype) + else: + subject_ref = None + + # 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, + seq_len=seq_len, + subject_ref=subject_ref, + ) + + # 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 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) diff --git a/videox_fun/utils/utils.py b/videox_fun/utils/utils.py index a769b2a..3ca52e7 100755 --- a/videox_fun/utils/utils.py +++ b/videox_fun/utils/utils.py @@ -235,10 +235,40 @@ def get_video_to_video_latent(input_video_path, video_length, sample_size, fps=N ref_image = ref_image.unsqueeze(0).permute([3, 0, 1, 2]).unsqueeze(0) / 255 return input_video, input_video_mask, ref_image, clip_image -def get_image_latent(ref_image=None, sample_size=None): +def padding_image(images, new_width, new_height): + new_image = Image.new('RGB', (new_width, new_height), (255, 255, 255)) + + aspect_ratio = images.width / images.height + if new_width / new_height > 1: + if aspect_ratio > new_width / new_height: + new_img_width = new_width + new_img_height = int(new_img_width / aspect_ratio) + else: + new_img_height = new_height + new_img_width = int(new_img_height * aspect_ratio) + else: + if aspect_ratio > new_width / new_height: + new_img_width = new_width + new_img_height = int(new_img_width / aspect_ratio) + else: + new_img_height = new_height + new_img_width = int(new_img_height * aspect_ratio) + + resized_img = images.resize((new_img_width, new_img_height)) + + paste_x = (new_width - new_img_width) // 2 + paste_y = (new_height - new_img_height) // 2 + + new_image.paste(resized_img, (paste_x, paste_y)) + + return new_image + +def get_image_latent(ref_image=None, sample_size=None, padding=False): if ref_image is not None: if isinstance(ref_image, str): ref_image = Image.open(ref_image).convert("RGB") + if padding: + ref_image = padding_image(ref_image, sample_size[1], sample_size[0]) ref_image = ref_image.resize((sample_size[1], sample_size[0])) ref_image = torch.from_numpy(np.array(ref_image)) ref_image = ref_image.unsqueeze(0).permute([3, 0, 1, 2]).unsqueeze(0) / 255 From 72c2792139c95a80c1f714cfd2fa0f6d7c86d388 Mon Sep 17 00:00:00 2001 From: bubbliiiing <3323290568@qq.com> Date: Mon, 9 Jun 2025 09:42:32 +0000 Subject: [PATCH 2/4] Update Training --- scripts/wan2.1/train.py | 5 +++-- scripts/wan2.1/train_lora.py | 5 +++-- scripts/wan2.1_fun/train.py | 5 +++-- scripts/wan2.1_fun/train_control.py | 5 +++-- scripts/wan2.1_fun/train_control_lora.py | 5 +++-- scripts/wan2.1_fun/train_lora.py | 5 +++-- 6 files changed, 18 insertions(+), 12 deletions(-) diff --git a/scripts/wan2.1/train.py b/scripts/wan2.1/train.py index 2b9227f..4d3a827 100755 --- a/scripts/wan2.1/train.py +++ b/scripts/wan2.1/train.py @@ -759,8 +759,9 @@ def main(): log_with=args.report_to, project_config=accelerator_project_config, ) - deepspeed_plugin = accelerator.state.deepspeed_plugin - fsdp_plugin = accelerator.state.fsdp_plugin + + deepspeed_plugin = accelerator.state.deepspeed_plugin if hasattr(accelerator.state, "deepspeed_plugin") else None + fsdp_plugin = accelerator.state.fsdp_plugin if hasattr(accelerator.state, "fsdp_plugin") else None if deepspeed_plugin is not None: zero_stage = int(deepspeed_plugin.zero_stage) fsdp_stage = 0 diff --git a/scripts/wan2.1/train_lora.py b/scripts/wan2.1/train_lora.py index 42ca41c..50a1883 100755 --- a/scripts/wan2.1/train_lora.py +++ b/scripts/wan2.1/train_lora.py @@ -758,8 +758,9 @@ def main(): log_with=args.report_to, project_config=accelerator_project_config, ) - deepspeed_plugin = accelerator.state.deepspeed_plugin - fsdp_plugin = accelerator.state.fsdp_plugin + + deepspeed_plugin = accelerator.state.deepspeed_plugin if hasattr(accelerator.state, "deepspeed_plugin") else None + fsdp_plugin = accelerator.state.fsdp_plugin if hasattr(accelerator.state, "fsdp_plugin") else None if deepspeed_plugin is not None: zero_stage = int(deepspeed_plugin.zero_stage) fsdp_stage = 0 diff --git a/scripts/wan2.1_fun/train.py b/scripts/wan2.1_fun/train.py index 7743942..1863178 100755 --- a/scripts/wan2.1_fun/train.py +++ b/scripts/wan2.1_fun/train.py @@ -729,8 +729,9 @@ def main(): log_with=args.report_to, project_config=accelerator_project_config, ) - deepspeed_plugin = accelerator.state.deepspeed_plugin - fsdp_plugin = accelerator.state.fsdp_plugin + + deepspeed_plugin = accelerator.state.deepspeed_plugin if hasattr(accelerator.state, "deepspeed_plugin") else None + fsdp_plugin = accelerator.state.fsdp_plugin if hasattr(accelerator.state, "fsdp_plugin") else None if deepspeed_plugin is not None: zero_stage = int(deepspeed_plugin.zero_stage) fsdp_stage = 0 diff --git a/scripts/wan2.1_fun/train_control.py b/scripts/wan2.1_fun/train_control.py index 9886c14..23d01c6 100755 --- a/scripts/wan2.1_fun/train_control.py +++ b/scripts/wan2.1_fun/train_control.py @@ -661,8 +661,9 @@ def main(): log_with=args.report_to, project_config=accelerator_project_config, ) - deepspeed_plugin = accelerator.state.deepspeed_plugin - fsdp_plugin = accelerator.state.fsdp_plugin + + deepspeed_plugin = accelerator.state.deepspeed_plugin if hasattr(accelerator.state, "deepspeed_plugin") else None + fsdp_plugin = accelerator.state.fsdp_plugin if hasattr(accelerator.state, "fsdp_plugin") else None if deepspeed_plugin is not None: zero_stage = int(deepspeed_plugin.zero_stage) fsdp_stage = 0 diff --git a/scripts/wan2.1_fun/train_control_lora.py b/scripts/wan2.1_fun/train_control_lora.py index 9e85399..9a79d25 100755 --- a/scripts/wan2.1_fun/train_control_lora.py +++ b/scripts/wan2.1_fun/train_control_lora.py @@ -659,8 +659,9 @@ def main(): log_with=args.report_to, project_config=accelerator_project_config, ) - deepspeed_plugin = accelerator.state.deepspeed_plugin - fsdp_plugin = accelerator.state.fsdp_plugin + + deepspeed_plugin = accelerator.state.deepspeed_plugin if hasattr(accelerator.state, "deepspeed_plugin") else None + fsdp_plugin = accelerator.state.fsdp_plugin if hasattr(accelerator.state, "fsdp_plugin") else None if deepspeed_plugin is not None: zero_stage = int(deepspeed_plugin.zero_stage) fsdp_stage = 0 diff --git a/scripts/wan2.1_fun/train_lora.py b/scripts/wan2.1_fun/train_lora.py index 090f14f..e64e224 100755 --- a/scripts/wan2.1_fun/train_lora.py +++ b/scripts/wan2.1_fun/train_lora.py @@ -721,8 +721,9 @@ def main(): log_with=args.report_to, project_config=accelerator_project_config, ) - deepspeed_plugin = accelerator.state.deepspeed_plugin - fsdp_plugin = accelerator.state.fsdp_plugin + + deepspeed_plugin = accelerator.state.deepspeed_plugin if hasattr(accelerator.state, "deepspeed_plugin") else None + fsdp_plugin = accelerator.state.fsdp_plugin if hasattr(accelerator.state, "fsdp_plugin") else None if deepspeed_plugin is not None: zero_stage = int(deepspeed_plugin.zero_stage) fsdp_stage = 0 From a4476549e4d49aa5f2f5ba22a7cf46976cddb803 Mon Sep 17 00:00:00 2001 From: bubbliiiing <3323290568@qq.com> Date: Tue, 10 Jun 2025 03:54:11 +0000 Subject: [PATCH 3/4] Update Training Code --- scripts/wan2.1/train.py | 4 ++++ scripts/wan2.1/train_lora.py | 4 ++++ scripts/wan2.1_fun/train.py | 4 ++++ scripts/wan2.1_fun/train_control.py | 4 ++++ scripts/wan2.1_fun/train_control_lora.py | 4 ++++ scripts/wan2.1_fun/train_lora.py | 4 ++++ videox_fun/models/wan_transformer3d.py | 21 +++++++++++++++++---- 7 files changed, 41 insertions(+), 4 deletions(-) diff --git a/scripts/wan2.1/train.py b/scripts/wan2.1/train.py index 4d3a827..c5aa7e6 100755 --- a/scripts/wan2.1/train.py +++ b/scripts/wan2.1/train.py @@ -773,8 +773,12 @@ def main(): zero_stage = 0 if fsdp_plugin.sharding_strategy is ShardingStrategy.FULL_SHARD: fsdp_stage = 3 + elif fsdp_plugin.sharding_strategy is None: # The fsdp_plugin.sharding_strategy is None in FSDP 2. + fsdp_stage = 3 elif fsdp_plugin.sharding_strategy is ShardingStrategy.SHARD_GRAD_OP: fsdp_stage = 2 + else: + fsdp_stage = 0 print(f"Using FSDP stage: {fsdp_stage}") args.use_fsdp = True diff --git a/scripts/wan2.1/train_lora.py b/scripts/wan2.1/train_lora.py index 50a1883..7c021dc 100755 --- a/scripts/wan2.1/train_lora.py +++ b/scripts/wan2.1/train_lora.py @@ -772,8 +772,12 @@ def main(): zero_stage = 0 if fsdp_plugin.sharding_strategy is ShardingStrategy.FULL_SHARD: fsdp_stage = 3 + elif fsdp_plugin.sharding_strategy is None: # The fsdp_plugin.sharding_strategy is None in FSDP 2. + fsdp_stage = 3 elif fsdp_plugin.sharding_strategy is ShardingStrategy.SHARD_GRAD_OP: fsdp_stage = 2 + else: + fsdp_stage = 0 print(f"Using FSDP stage: {fsdp_stage}") args.use_fsdp = True diff --git a/scripts/wan2.1_fun/train.py b/scripts/wan2.1_fun/train.py index 1863178..cf89f6e 100755 --- a/scripts/wan2.1_fun/train.py +++ b/scripts/wan2.1_fun/train.py @@ -743,8 +743,12 @@ def main(): zero_stage = 0 if fsdp_plugin.sharding_strategy is ShardingStrategy.FULL_SHARD: fsdp_stage = 3 + elif fsdp_plugin.sharding_strategy is None: # The fsdp_plugin.sharding_strategy is None in FSDP 2. + fsdp_stage = 3 elif fsdp_plugin.sharding_strategy is ShardingStrategy.SHARD_GRAD_OP: fsdp_stage = 2 + else: + fsdp_stage = 0 print(f"Using FSDP stage: {fsdp_stage}") args.use_fsdp = True diff --git a/scripts/wan2.1_fun/train_control.py b/scripts/wan2.1_fun/train_control.py index 23d01c6..ef471de 100755 --- a/scripts/wan2.1_fun/train_control.py +++ b/scripts/wan2.1_fun/train_control.py @@ -675,8 +675,12 @@ def main(): zero_stage = 0 if fsdp_plugin.sharding_strategy is ShardingStrategy.FULL_SHARD: fsdp_stage = 3 + elif fsdp_plugin.sharding_strategy is None: # The fsdp_plugin.sharding_strategy is None in FSDP 2. + fsdp_stage = 3 elif fsdp_plugin.sharding_strategy is ShardingStrategy.SHARD_GRAD_OP: fsdp_stage = 2 + else: + fsdp_stage = 0 print(f"Using FSDP stage: {fsdp_stage}") args.use_fsdp = True diff --git a/scripts/wan2.1_fun/train_control_lora.py b/scripts/wan2.1_fun/train_control_lora.py index 9a79d25..b36c39b 100755 --- a/scripts/wan2.1_fun/train_control_lora.py +++ b/scripts/wan2.1_fun/train_control_lora.py @@ -673,8 +673,12 @@ def main(): zero_stage = 0 if fsdp_plugin.sharding_strategy is ShardingStrategy.FULL_SHARD: fsdp_stage = 3 + elif fsdp_plugin.sharding_strategy is None: # The fsdp_plugin.sharding_strategy is None in FSDP 2. + fsdp_stage = 3 elif fsdp_plugin.sharding_strategy is ShardingStrategy.SHARD_GRAD_OP: fsdp_stage = 2 + else: + fsdp_stage = 0 print(f"Using FSDP stage: {fsdp_stage}") args.use_fsdp = True diff --git a/scripts/wan2.1_fun/train_lora.py b/scripts/wan2.1_fun/train_lora.py index e64e224..5d392fe 100755 --- a/scripts/wan2.1_fun/train_lora.py +++ b/scripts/wan2.1_fun/train_lora.py @@ -735,8 +735,12 @@ def main(): zero_stage = 0 if fsdp_plugin.sharding_strategy is ShardingStrategy.FULL_SHARD: fsdp_stage = 3 + elif fsdp_plugin.sharding_strategy is None: # The fsdp_plugin.sharding_strategy is None in FSDP 2. + fsdp_stage = 3 elif fsdp_plugin.sharding_strategy is ShardingStrategy.SHARD_GRAD_OP: fsdp_stage = 2 + else: + fsdp_stage = 0 print(f"Using FSDP stage: {fsdp_stage}") args.use_fsdp = True diff --git a/videox_fun/models/wan_transformer3d.py b/videox_fun/models/wan_transformer3d.py index 16bb296..ebb55ba 100755 --- a/videox_fun/models/wan_transformer3d.py +++ b/videox_fun/models/wan_transformer3d.py @@ -2,6 +2,7 @@ # Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved. import glob +import importlib.metadata import json import math import os @@ -17,6 +18,7 @@ from diffusers.configuration_utils import ConfigMixin, register_to_config from diffusers.loaders.single_file_model import FromOriginalModelMixin from diffusers.models.modeling_utils import ModelMixin from diffusers.utils import is_torch_version, logging +from packaging import version from torch import nn from ..dist import (get_sequence_parallel_rank, @@ -60,6 +62,13 @@ except: sageattn = None SAGE_ATTENTION_AVAILABLE = False +try: + diffusers_version = importlib.metadata.version("diffusers") +except importlib.metadata.PackageNotFoundError: + diffusers_version = "0.0.0" + +USE_NEW_SIGNATURE = version.parse(diffusers_version) >= version.parse("0.33.1") + def flash_attention( q, k, @@ -843,7 +852,14 @@ class WanTransformer3DModel(ModelMixin, ConfigMixin, FromOriginalModelMixin): self.gradient_checkpointing = False self.sp_world_size = 1 self.sp_world_rank = 0 - + + if USE_NEW_SIGNATURE: + def _set_gradient_checkpointing(self, enable=False, gradient_checkpointing_func=None): + self.gradient_checkpointing = enable + else: + def _set_gradient_checkpointing(self, module, value=False): + self.gradient_checkpointing = value + def enable_teacache( self, coefficients, @@ -908,9 +924,6 @@ class WanTransformer3DModel(ModelMixin, ConfigMixin, FromOriginalModelMixin): block.self_attn.forward = types.MethodType( usp_attn_forward, block.self_attn) - def _set_gradient_checkpointing(self, module, value=False): - self.gradient_checkpointing = value - @cfg_skip() def forward( self, From 95e024c681d180f0887ef675240f9798f987db81 Mon Sep 17 00:00:00 2001 From: bubbliiiing <3323290568@qq.com> Date: Tue, 10 Jun 2025 05:02:41 +0000 Subject: [PATCH 4/4] Update Transformer3d --- videox_fun/models/wan_transformer3d.py | 20 ++++++-------------- 1 file changed, 6 insertions(+), 14 deletions(-) diff --git a/videox_fun/models/wan_transformer3d.py b/videox_fun/models/wan_transformer3d.py index ebb55ba..17796a4 100755 --- a/videox_fun/models/wan_transformer3d.py +++ b/videox_fun/models/wan_transformer3d.py @@ -2,7 +2,6 @@ # Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved. import glob -import importlib.metadata import json import math import os @@ -18,7 +17,6 @@ from diffusers.configuration_utils import ConfigMixin, register_to_config from diffusers.loaders.single_file_model import FromOriginalModelMixin from diffusers.models.modeling_utils import ModelMixin from diffusers.utils import is_torch_version, logging -from packaging import version from torch import nn from ..dist import (get_sequence_parallel_rank, @@ -62,13 +60,6 @@ except: sageattn = None SAGE_ATTENTION_AVAILABLE = False -try: - diffusers_version = importlib.metadata.version("diffusers") -except importlib.metadata.PackageNotFoundError: - diffusers_version = "0.0.0" - -USE_NEW_SIGNATURE = version.parse(diffusers_version) >= version.parse("0.33.1") - def flash_attention( q, k, @@ -853,12 +844,13 @@ class WanTransformer3DModel(ModelMixin, ConfigMixin, FromOriginalModelMixin): self.sp_world_size = 1 self.sp_world_rank = 0 - if USE_NEW_SIGNATURE: - def _set_gradient_checkpointing(self, enable=False, gradient_checkpointing_func=None): - self.gradient_checkpointing = enable + def _set_gradient_checkpointing(self, *args, **kwargs): + if "value" in kwargs: + self.gradient_checkpointing = kwargs["value"] + elif "enable" in kwargs: + self.gradient_checkpointing = kwargs["enable"] else: - def _set_gradient_checkpointing(self, module, value=False): - self.gradient_checkpointing = value + raise ValueError("Invalid set gradient checkpointing") def enable_teacache( self,