From 037a2e8360586bbca4adf935f807aebdb8caaeda Mon Sep 17 00:00:00 2001 From: Bubbliiiing <47347516+bubbliiiing@users.noreply.github.com> Date: Thu, 20 Nov 2025 17:44:22 +0800 Subject: [PATCH] Update Peft Lora && Update Readme (#376) --- README.md | 3 + README_ja-JP.md | 10 +- README_zh-CN.md | 3 + scripts/cogvideox_fun/README_TRAIN.md | 16 +- scripts/cogvideox_fun/README_TRAIN_CONTROL.md | 13 +- scripts/cogvideox_fun/README_TRAIN_LORA.md | 31 +- scripts/cogvideox_fun/train.py | 167 ++++++--- scripts/cogvideox_fun/train_control.py | 170 +++++++--- scripts/cogvideox_fun/train_lora.py | 319 +++++++++++++----- scripts/cogvideox_fun/train_lora.sh | 8 + scripts/fantasytalking/README_TRAIN.md | 17 +- scripts/fantasytalking/train.py | 5 + scripts/flux/README_TRAIN.md | 56 +-- scripts/flux/README_TRAIN_LORA.md | 65 ++-- scripts/flux/train.py | 5 + scripts/flux/train_lora.py | 116 +++++-- scripts/flux/train_lora.sh | 4 + scripts/qwenimage/README_TRAIN.md | 17 +- scripts/qwenimage/README_TRAIN_EDIT.md | 17 +- scripts/qwenimage/README_TRAIN_LORA.md | 35 +- scripts/qwenimage/train.py | 5 + scripts/qwenimage/train_edit.py | 5 + scripts/qwenimage/train_edit_lora.py | 116 +++++-- scripts/qwenimage/train_edit_lora.sh | 4 + scripts/qwenimage/train_lora.py | 116 +++++-- scripts/qwenimage/train_lora.sh | 4 + scripts/wan2.1/README_TRAIN.md | 19 +- scripts/wan2.1/README_TRAIN_DISTILL.md | 19 +- scripts/wan2.1/README_TRAIN_DISTILL_LORA.md | 37 +- scripts/wan2.1/README_TRAIN_LORA.md | 46 ++- scripts/wan2.1/train.py | 5 + scripts/wan2.1/train_distill.py | 5 + scripts/wan2.1/train_distill_lora.py | 199 +++++++---- scripts/wan2.1/train_distill_lora.sh | 4 + scripts/wan2.1/train_lora.py | 118 +++++-- scripts/wan2.1/train_lora.sh | 8 + scripts/wan2.1_fun/README_TRAIN.md | 19 +- scripts/wan2.1_fun/README_TRAIN_CONTROL.md | 6 +- .../wan2.1_fun/README_TRAIN_CONTROL_LORA.md | 24 +- scripts/wan2.1_fun/README_TRAIN_LORA.md | 37 +- scripts/wan2.1_fun/train.py | 5 + scripts/wan2.1_fun/train_control.py | 5 + scripts/wan2.1_fun/train_control_lora.py | 118 +++++-- scripts/wan2.1_fun/train_control_lora.sh | 4 + scripts/wan2.1_fun/train_lora.py | 118 +++++-- scripts/wan2.1_fun/train_lora.sh | 8 + scripts/wan2.1_vace/README_TRAIN.md | 4 +- scripts/wan2.1_vace/train.py | 5 + scripts/wan2.2/README_TRAIN.md | 18 +- scripts/wan2.2/README_TRAIN_ANIMATE.md | 17 +- scripts/wan2.2/README_TRAIN_DISTILL.md | 19 +- scripts/wan2.2/README_TRAIN_DISTILL_LORA.md | 37 +- scripts/wan2.2/README_TRAIN_LORA.md | 44 ++- scripts/wan2.2/README_TRAIN_S2V.md | 17 +- scripts/wan2.2/train.py | 5 + scripts/wan2.2/train_animate.py | 5 + scripts/wan2.2/train_animate_lora.py | 118 +++++-- scripts/wan2.2/train_animate_lora.sh | 4 + scripts/wan2.2/train_distill.py | 5 + scripts/wan2.2/train_distill_lora.py | 209 ++++++++---- scripts/wan2.2/train_distill_lora.sh | 4 + scripts/wan2.2/train_lora.py | 118 +++++-- scripts/wan2.2/train_lora.sh | 8 + scripts/wan2.2/train_s2v.py | 5 + scripts/wan2.2/train_s2v_lora.py | 126 +++++-- scripts/wan2.2/train_s2v_lora.sh | 4 + scripts/wan2.2_fun/README_TRAIN.md | 19 +- scripts/wan2.2_fun/README_TRAIN_CONTROL.md | 4 +- .../wan2.2_fun/README_TRAIN_CONTROL_LORA.md | 22 +- scripts/wan2.2_fun/README_TRAIN_LORA.md | 42 ++- scripts/wan2.2_fun/train.py | 5 + scripts/wan2.2_fun/train_control.py | 5 + scripts/wan2.2_fun/train_control_lora.py | 117 +++++-- scripts/wan2.2_fun/train_control_lora.sh | 5 +- scripts/wan2.2_fun/train_lora.py | 118 +++++-- scripts/wan2.2_fun/train_lora.sh | 5 +- scripts/wan2.2_vace_fun/README_TRAIN.md | 4 +- scripts/wan2.2_vace_fun/train.py | 5 + videox_fun/utils/lora_utils.py | 94 +++--- 79 files changed, 2499 insertions(+), 849 deletions(-) diff --git a/README.md b/README.md index 782e06e..4aeb6e1 100755 --- a/README.md +++ b/README.md @@ -652,6 +652,9 @@ V1.1: - EasyAnimate: https://github.com/aigc-apps/EasyAnimate - Wan2.1: https://github.com/Wan-Video/Wan2.1/ - Wan2.2: https://github.com/Wan-Video/Wan2.2/ +- Diffusers: https://github.com/huggingface/diffusers +- Qwen-Image: https://github.com/QwenLM/Qwen-Image +- Self-Forcing: https://github.com/guandeh17/Self-Forcing - ComfyUI-KJNodes: https://github.com/kijai/ComfyUI-KJNodes - ComfyUI-EasyAnimateWrapper: https://github.com/kijai/ComfyUI-EasyAnimateWrapper - ComfyUI-CameraCtrl-Wrapper: https://github.com/chaojie/ComfyUI-CameraCtrl-Wrapper diff --git a/README_ja-JP.md b/README_ja-JP.md index cc6f82c..23c8fc7 100755 --- a/README_ja-JP.md +++ b/README_ja-JP.md @@ -647,14 +647,18 @@ V1.1: | CogVideoX-Fun-5b-InP | 20.0 GB | [🤗リンク](https://huggingface.co/alibaba-pai/CogVideoX-Fun-5b-InP)| [😄リンク](https://modelscope.cn/models/PAI/CogVideoX-Fun-5b-InP)| 公式のグラフ生成ビデオモデルは、複数の解像度(512、768、1024、1280)でビデオを予測できます。49フレーム、8フレーム/秒でトレーニングされています。| -# TODOリスト -- 日本語をサポート。 - # 参考文献 - CogVideo: https://github.com/THUDM/CogVideo/ - EasyAnimate: https://github.com/aigc-apps/EasyAnimate - Wan2.1: https://github.com/Wan-Video/Wan2.1/ - Wan2.2: https://github.com/Wan-Video/Wan2.2/ +- Diffusers: https://github.com/huggingface/diffusers +- Qwen-Image: https://github.com/QwenLM/Qwen-Image +- Self-Forcing: https://github.com/guandeh17/Self-Forcing +- ComfyUI-KJNodes: https://github.com/kijai/ComfyUI-KJNodes +- ComfyUI-EasyAnimateWrapper: https://github.com/kijai/ComfyUI-EasyAnimateWrapper +- ComfyUI-CameraCtrl-Wrapper: https://github.com/chaojie/ComfyUI-CameraCtrl-Wrapper +- CameraCtrl: https://github.com/hehao13/CameraCtrl # ライセンス このプロジェクトは[Apache License (Version 2.0)](https://github.com/modelscope/modelscope/blob/master/LICENSE)の下でライセンスされています。 diff --git a/README_zh-CN.md b/README_zh-CN.md index e86abac..e4a0f2a 100755 --- a/README_zh-CN.md +++ b/README_zh-CN.md @@ -642,6 +642,9 @@ V1.1: - EasyAnimate: https://github.com/aigc-apps/EasyAnimate - Wan2.1: https://github.com/Wan-Video/Wan2.1/ - Wan2.2: https://github.com/Wan-Video/Wan2.2/ +- Diffusers: https://github.com/huggingface/diffusers +- Qwen-Image: https://github.com/QwenLM/Qwen-Image +- Self-Forcing: https://github.com/guandeh17/Self-Forcing - ComfyUI-KJNodes: https://github.com/kijai/ComfyUI-KJNodes - ComfyUI-EasyAnimateWrapper: https://github.com/kijai/ComfyUI-EasyAnimateWrapper - ComfyUI-CameraCtrl-Wrapper: https://github.com/chaojie/ComfyUI-CameraCtrl-Wrapper diff --git a/scripts/cogvideox_fun/README_TRAIN.md b/scripts/cogvideox_fun/README_TRAIN.md index 3f3e324..fec7675 100755 --- a/scripts/cogvideox_fun/README_TRAIN.md +++ b/scripts/cogvideox_fun/README_TRAIN.md @@ -2,7 +2,7 @@ The default training commands for the different versions are as follows: -We can choose whether to use deepspeed in CogVideoX-Fun, which can save a lot of video memory. +We can choose whether to use deepspeed and fsdp in CogVideoX-Fun, which can save a lot of video memory. Some parameters in the sh file can be confusing, and they are explained in this document: @@ -61,12 +61,11 @@ accelerate launch --mixed_precision="bf16" scripts/cogvideox_fun/train.py \ --random_hw_adapt \ --training_with_video_token_length \ --enable_bucket \ - --use_ema \ --train_mode="inpaint" \ --trainable_modules "." ``` -CogVideoX-Fun with deepspeed: +CogVideoX-Fun with Deepspeed Zero-2: ```sh export MODEL_NAME="models/Diffusion_Transformer/CogVideoX-Fun-2b-InP" export DATASET_NAME="datasets/internal_datasets/" @@ -110,7 +109,8 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con --trainable_modules "." ``` -CogVideoX-Fun with multi machines: +With FSDP: + ```sh export MODEL_NAME="models/Diffusion_Transformer/CogVideoX-Fun-2b-InP" export DATASET_NAME="datasets/internal_datasets/" @@ -120,11 +120,7 @@ export DATASET_META_NAME="datasets/internal_datasets/metadata.json" # export NCCL_P2P_DISABLE=1 NCCL_DEBUG=INFO -NUM_PROCESS=$((WORLD_SIZE * 8)) - -echo "MASTER_ADDR: ${MASTER_ADDR} MASTER_PORT: ${MASTER_PORT} NUM_PROCESS: ${NUM_PROCESS}" - -accelerate launch --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK scripts/cogvideox_fun/train.py \ +accelerate launch --mixed_precision="bf16" --use_fsdp --fsdp_auto_wrap_policy TRANSFORMER_BASED_WRAP --fsdp_transformer_layer_cls_to_wrap CogVideoXBlock --fsdp_sharding_strategy "FULL_SHARD" --fsdp_state_dict_type=SHARDED_STATE_DICT --fsdp_backward_prefetch "BACKWARD_PRE" --fsdp_cpu_ram_efficient_loading False scripts/cogvideox_fun/train.py \ --pretrained_model_name_or_path=$MODEL_NAME \ --train_data_dir=$DATASET_NAME \ --train_data_meta=$DATASET_META_NAME \ @@ -133,7 +129,7 @@ accelerate launch --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_POR --token_sample_size=512 \ --video_sample_stride=3 \ --video_sample_n_frames=49 \ - --train_batch_size=4 \ + --train_batch_size=1 \ --video_repeat=1 \ --gradient_accumulation_steps=1 \ --dataloader_num_workers=8 \ diff --git a/scripts/cogvideox_fun/README_TRAIN_CONTROL.md b/scripts/cogvideox_fun/README_TRAIN_CONTROL.md index 0dd6260..9957cee 100755 --- a/scripts/cogvideox_fun/README_TRAIN_CONTROL.md +++ b/scripts/cogvideox_fun/README_TRAIN_CONTROL.md @@ -2,7 +2,7 @@ The default training commands for the different versions are as follows: -We can choose whether to use deepspeed in CogVideoX-Fun, which can save a lot of video memory. +We can choose whether to use deepspeed and fsdp in CogVideoX-Fun, which can save a lot of video memory. The metadata_control.json is a little different from normal json in CogVideoX-Fun, you need to add a control_file_path, and [DWPose](https://github.com/IDEA-Research/DWPose) is suggested as tool to generate control file. @@ -83,7 +83,7 @@ accelerate launch --mixed_precision="bf16" scripts/cogvideox_fun/train_control.p --trainable_modules "." ``` -CogVideoX-Fun with deepspeed: +CogVideoX-Fun with Deepspeed Zero-2: ```sh export MODEL_NAME="models/Diffusion_Transformer/CogVideoX-Fun-V1.1-2b-Pose" export DATASET_NAME="datasets/internal_datasets/" @@ -126,7 +126,8 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con --trainable_modules "." ``` -CogVideoX-Fun with multi machines: +With FSDP: + ```sh export MODEL_NAME="models/Diffusion_Transformer/CogVideoX-Fun-V1.1-2b-Pose" export DATASET_NAME="datasets/internal_datasets/" @@ -136,11 +137,7 @@ export DATASET_META_NAME="datasets/internal_datasets/metadata_control.json" # export NCCL_P2P_DISABLE=1 NCCL_DEBUG=INFO -NUM_PROCESS=$((WORLD_SIZE * 8)) - -echo "MASTER_ADDR: ${MASTER_ADDR} MASTER_PORT: ${MASTER_PORT} NUM_PROCESS: ${NUM_PROCESS}" - -accelerate launch --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK scripts/cogvideox_fun/train.py \ +accelerate launch --mixed_precision="bf16" --use_fsdp --fsdp_auto_wrap_policy TRANSFORMER_BASED_WRAP --fsdp_transformer_layer_cls_to_wrap CogVideoXBlock --fsdp_sharding_strategy "FULL_SHARD" --fsdp_state_dict_type=SHARDED_STATE_DICT --fsdp_backward_prefetch "BACKWARD_PRE" --fsdp_cpu_ram_efficient_loading False scripts/cogvideox_fun/train_control.py \ --pretrained_model_name_or_path=$MODEL_NAME \ --train_data_dir=$DATASET_NAME \ --train_data_meta=$DATASET_META_NAME \ diff --git a/scripts/cogvideox_fun/README_TRAIN_LORA.md b/scripts/cogvideox_fun/README_TRAIN_LORA.md index 17d00d6..1d05f9e 100755 --- a/scripts/cogvideox_fun/README_TRAIN_LORA.md +++ b/scripts/cogvideox_fun/README_TRAIN_LORA.md @@ -1,6 +1,6 @@ ## Lora Training Code -We can choose whether to use deepspeed in CogVideoX-Fun, which can save a lot of video memory. +We can choose whether to use deepspeed and fsdp in CogVideoX-Fun, which can save a lot of video memory. Some parameters in the sh file can be confusing, and they are explained in this document: @@ -19,6 +19,10 @@ Some parameters in the sh file can be confusing, and they are explained in this - These resolutions combined with their corresponding lengths allow the model to generate videos of different sizes. - `train_mode` is used to specify the training mode, which can be either normal or i2v. Since CogVideoX-Fun uses the inpaint model to achieve image-to-video generation, the default is set to inpaint mode. If you only wish to achieve text-to-video generation, you can remove this line, and it will default to the text-to-video mode. - `resume_from_checkpoint` is used to set the training should be resumed from a previous checkpoint. Use a path or `"latest"` to automatically select the last available checkpoint. +- `target_name` represents the components/modules to which LoRA will be applied, separated by commas. +- `use_peft_lora` indicates whether to use the PEFT module for adding LoRA. Using this module will be more memory-efficient. +- `rank` means the dimension of the LoRA update matrices. +- `network_alpha` means the scale of the LoRA update matrices. CogVideoX-Fun without deepspeed: @@ -58,11 +62,15 @@ accelerate launch --mixed_precision="bf16" scripts/cogvideox_fun/train_lora.py \ --random_hw_adapt \ --training_with_video_token_length \ --enable_bucket \ + --rank=64 \ + --network_alpha=32 \ + --target_name="to_q,to_k,to_v,ff.0,ff.2" \ + --use_peft_lora \ --low_vram \ --train_mode="inpaint" ``` -CogVideoX-Fun with deepspeed: +CogVideoX-Fun with Deepspeed Zero-2: ```sh export MODEL_NAME="models/Diffusion_Transformer/CogVideoX-Fun-2b-InP" export DATASET_NAME="datasets/internal_datasets/" @@ -99,12 +107,16 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con --random_hw_adapt \ --training_with_video_token_length \ --enable_bucket \ - --use_deepspeed \ + --rank=64 \ + --network_alpha=32 \ + --target_name="to_q,to_k,to_v,ff.0,ff.2" \ + --use_peft_lora \ --low_vram \ --train_mode="inpaint" ``` -CogVideoX-Fun with multi machines: +With FSDP: + ```sh export MODEL_NAME="models/Diffusion_Transformer/CogVideoX-Fun-2b-InP" export DATASET_NAME="datasets/internal_datasets/" @@ -114,11 +126,7 @@ export DATASET_META_NAME="datasets/internal_datasets/metadata.json" # export NCCL_P2P_DISABLE=1 NCCL_DEBUG=INFO -NUM_PROCESS=$((WORLD_SIZE * 8)) - -echo "MASTER_ADDR: ${MASTER_ADDR} MASTER_PORT: ${MASTER_PORT} NUM_PROCESS: ${NUM_PROCESS}" - -accelerate launch --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK scripts/cogvideox_fun/train.py \ +accelerate launch --mixed_precision="bf16" --use_fsdp --fsdp_auto_wrap_policy TRANSFORMER_BASED_WRAP --fsdp_transformer_layer_cls_to_wrap CogVideoXBlock --fsdp_sharding_strategy "FULL_SHARD" --fsdp_state_dict_type=SHARDED_STATE_DICT --fsdp_backward_prefetch "BACKWARD_PRE" --fsdp_cpu_ram_efficient_loading False scripts/cogvideox_fun/train_lora.py \ --pretrained_model_name_or_path=$MODEL_NAME \ --train_data_dir=$DATASET_NAME \ --train_data_meta=$DATASET_META_NAME \ @@ -145,5 +153,10 @@ accelerate launch --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_POR --random_hw_adapt \ --training_with_video_token_length \ --enable_bucket \ + --rank=64 \ + --network_alpha=32 \ + --target_name="to_q,to_k,to_v,ff.0,ff.2" \ + --use_peft_lora \ + --low_vram \ --train_mode="inpaint" ``` \ No newline at end of file diff --git a/scripts/cogvideox_fun/train.py b/scripts/cogvideox_fun/train.py index 1ef9969..b199440 100755 --- a/scripts/cogvideox_fun/train.py +++ b/scripts/cogvideox_fun/train.py @@ -659,6 +659,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." ) @@ -728,6 +731,40 @@ def main(): log_with=args.report_to, project_config=accelerator_project_config, ) + + 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 + print(f"Using DeepSpeed Zero stage: {zero_stage}") + + args.use_deepspeed = True + if zero_stage == 3: + print(f"Auto set save_state to True because zero_stage == 3") + args.save_state = 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 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 + if fsdp_stage == 3: + print(f"Auto set save_state to True because fsdp_stage == 3") + args.save_state = True + else: + zero_stage = 0 + fsdp_stage = 0 + print("DeepSpeed is not enabled.") + if accelerator.is_main_process: writer = SummaryWriter(log_dir=logging_dir) @@ -868,51 +905,93 @@ 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 - def save_model_hook(models, weights, output_dir): - if accelerator.is_main_process: + 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") + accelerate_state_dict = {k: v.to(dtype=weight_dtype) for k, v in accelerate_state_dict.items()} + 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): + 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}.") + 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: + if args.use_ema: + ema_transformer3d.save_pretrained(os.path.join(output_dir, "transformer_ema")) + + models[0].save_pretrained(os.path.join(output_dir, "transformer")) + if not args.use_deepspeed: + 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) + + def load_model_hook(models, input_dir): if args.use_ema: - ema_transformer3d.save_pretrained(os.path.join(output_dir, "transformer_ema")) + ema_path = os.path.join(input_dir, "transformer_ema") + _, ema_kwargs = CogVideoXTransformer3DModel.load_config(ema_path, return_unused_kwargs=True) + load_model = CogVideoXTransformer3DModel.from_pretrained( + input_dir, subfolder="transformer_ema", + ) + load_model = EMAModel(load_model.parameters(), model_cls=CogVideoXTransformer3DModel, model_config=load_model.config) + load_model.load_state_dict(ema_kwargs) - models[0].save_pretrained(os.path.join(output_dir, "transformer")) - if not args.use_deepspeed: - weights.pop() + ema_transformer3d.load_state_dict(load_model.state_dict()) + ema_transformer3d.to(accelerator.device) + del load_model - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) + for i in range(len(models)): + # pop models so that they are not loaded again + model = models.pop() - def load_model_hook(models, input_dir): - if args.use_ema: - ema_path = os.path.join(input_dir, "transformer_ema") - _, ema_kwargs = CogVideoXTransformer3DModel.load_config(ema_path, return_unused_kwargs=True) - load_model = CogVideoXTransformer3DModel.from_pretrained( - input_dir, subfolder="transformer_ema" - ) - load_model = EMAModel(load_model.parameters(), model_cls=CogVideoXTransformer3DModel, model_config=load_model.config) - load_model.load_state_dict(ema_kwargs) + # load diffusers style into model + load_model = CogVideoXTransformer3DModel.from_pretrained( + input_dir, subfolder="transformer" + ) + model.register_to_config(**load_model.config) - ema_transformer3d.load_state_dict(load_model.state_dict()) - ema_transformer3d.to(accelerator.device) - del load_model + model.load_state_dict(load_model.state_dict()) + del load_model - for i in range(len(models)): - # pop models so that they are not loaded again - model = models.pop() - - # load diffusers style into model - load_model = CogVideoXTransformer3DModel.from_pretrained( - input_dir, subfolder="transformer" - ) - model.register_to_config(**load_model.config) - - model.load_state_dict(load_model.state_dict()) - del load_model - - 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}.") + 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) @@ -1628,7 +1707,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) @@ -1639,14 +1718,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() @@ -1664,7 +1743,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) @@ -1745,7 +1824,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: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() diff --git a/scripts/cogvideox_fun/train_control.py b/scripts/cogvideox_fun/train_control.py index 3ebdb18..0bd223f 100755 --- a/scripts/cogvideox_fun/train_control.py +++ b/scripts/cogvideox_fun/train_control.py @@ -669,6 +669,40 @@ def main(): log_with=args.report_to, project_config=accelerator_project_config, ) + + 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 + print(f"Using DeepSpeed Zero stage: {zero_stage}") + + args.use_deepspeed = True + if zero_stage == 3: + print(f"Auto set save_state to True because zero_stage == 3") + args.save_state = 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 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 + if fsdp_stage == 3: + print(f"Auto set save_state to True because fsdp_stage == 3") + args.save_state = True + else: + zero_stage = 0 + fsdp_stage = 0 + print("DeepSpeed is not enabled.") + if accelerator.is_main_process: writer = SummaryWriter(log_dir=logging_dir) @@ -742,13 +776,13 @@ def main(): # across multiple gpus and only UNet2DConditionModel will get ZeRO sharded. with ContextManagers(deepspeed_zero_init_disabled_context_manager()): text_encoder = T5EncoderModel.from_pretrained( - args.pretrained_model_name_or_path, subfolder="text_encoder", revision=args.revision, variant=args.variant, - torch_dtype=weight_dtype + args.pretrained_model_name_or_path, subfolder="text_encoder", torch_dtype=weight_dtype ) - + text_encoder = text_encoder.eval() vae = AutoencoderKLCogVideoX.from_pretrained( args.pretrained_model_name_or_path, subfolder="vae", revision=args.revision, variant=args.variant ) + vae = vae.eval() transformer3d = CogVideoXTransformer3DModel.from_pretrained( args.pretrained_model_name_or_path, subfolder="transformer" @@ -809,51 +843,93 @@ 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 - def save_model_hook(models, weights, output_dir): - if accelerator.is_main_process: + 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") + accelerate_state_dict = {k: v.to(dtype=weight_dtype) for k, v in accelerate_state_dict.items()} + 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): + 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}.") + 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: + if args.use_ema: + ema_transformer3d.save_pretrained(os.path.join(output_dir, "transformer_ema")) + + models[0].save_pretrained(os.path.join(output_dir, "transformer")) + if not args.use_deepspeed: + 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) + + def load_model_hook(models, input_dir): if args.use_ema: - ema_transformer3d.save_pretrained(os.path.join(output_dir, "transformer_ema")) + ema_path = os.path.join(input_dir, "transformer_ema") + _, ema_kwargs = CogVideoXTransformer3DModel.load_config(ema_path, return_unused_kwargs=True) + load_model = CogVideoXTransformer3DModel.from_pretrained( + input_dir, subfolder="transformer_ema", + ) + load_model = EMAModel(load_model.parameters(), model_cls=CogVideoXTransformer3DModel, model_config=load_model.config) + load_model.load_state_dict(ema_kwargs) - models[0].save_pretrained(os.path.join(output_dir, "transformer")) - if not args.use_deepspeed: - weights.pop() + ema_transformer3d.load_state_dict(load_model.state_dict()) + ema_transformer3d.to(accelerator.device) + del load_model - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) + for i in range(len(models)): + # pop models so that they are not loaded again + model = models.pop() - def load_model_hook(models, input_dir): - if args.use_ema: - ema_path = os.path.join(input_dir, "transformer_ema") - _, ema_kwargs = CogVideoXTransformer3DModel.load_config(ema_path, return_unused_kwargs=True) - load_model = CogVideoXTransformer3DModel.from_pretrained( - input_dir, subfolder="transformer_ema" - ) - load_model = EMAModel(load_model.parameters(), model_cls=CogVideoXTransformer3DModel, model_config=load_model.config) - load_model.load_state_dict(ema_kwargs) + # load diffusers style into model + load_model = CogVideoXTransformer3DModel.from_pretrained( + input_dir, subfolder="transformer" + ) + model.register_to_config(**load_model.config) - ema_transformer3d.load_state_dict(load_model.state_dict()) - ema_transformer3d.to(accelerator.device) - del load_model + model.load_state_dict(load_model.state_dict()) + del load_model - for i in range(len(models)): - # pop models so that they are not loaded again - model = models.pop() - - # load diffusers style into model - load_model = CogVideoXTransformer3DModel.from_pretrained( - input_dir, subfolder="transformer" - ) - model.register_to_config(**load_model.config) - - model.load_state_dict(load_model.state_dict()) - del load_model - - 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}.") + 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) @@ -1518,7 +1594,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) @@ -1529,14 +1605,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() @@ -1554,7 +1630,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) @@ -1635,7 +1711,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: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() diff --git a/scripts/cogvideox_fun/train_lora.py b/scripts/cogvideox_fun/train_lora.py index e00cffc..8332585 100755 --- a/scripts/cogvideox_fun/train_lora.py +++ b/scripts/cogvideox_fun/train_lora.py @@ -553,6 +553,9 @@ def parse_args(): default=64, help=("The dimension of the LoRA update matrices."), ) + parser.add_argument( + "--use_peft_lora", action="store_true", help="Whether or not to use peft lora." + ) parser.add_argument( "--train_text_encoder", action="store_true", @@ -673,6 +676,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." ) @@ -685,6 +691,12 @@ def parse_args(): ' (default), `"inpaint"`.' ), ) + parser.add_argument( + "--target_name", + type=str, + default=None, + help=("The module is trained in loras. "), + ) args = parser.parse_args() env_local_rank = int(os.environ.get("LOCAL_RANK", -1)) @@ -726,6 +738,40 @@ def main(): log_with=args.report_to, project_config=accelerator_project_config, ) + + 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 + print(f"Using DeepSpeed Zero stage: {zero_stage}") + + args.use_deepspeed = True + if zero_stage == 3: + print(f"Auto set save_state to True because zero_stage == 3") + args.save_state = 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 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 + if fsdp_stage == 3: + print(f"Auto set save_state to True because fsdp_stage == 3") + args.save_state = True + else: + zero_stage = 0 + fsdp_stage = 0 + print("DeepSpeed is not enabled.") + if accelerator.is_main_process: writer = SummaryWriter(log_dir=logging_dir) @@ -817,16 +863,25 @@ def main(): transformer3d.requires_grad_(False) # Lora will work with this... - network = create_network( - 1.0, - args.rank, - args.network_alpha, - text_encoder, - transformer3d, - neuron_dropout=None, - add_lora_in_attn_temporal=True, - ) - network.apply_to(text_encoder, transformer3d, args.train_text_encoder and not args.training_with_video_token_length, True) + if args.use_peft_lora: + from peft import LoraConfig, inject_adapter_in_model, get_peft_model_state_dict + lora_config = LoraConfig(r=args.rank, lora_alpha=args.network_alpha, target_modules=args.target_name.split(",")) + transformer3d = inject_adapter_in_model(lora_config, transformer3d) + + network = None + else: + network = create_network( + 1.0, + args.rank, + args.network_alpha, + text_encoder, + transformer3d, + neuron_dropout=None, + target_name=args.target_name, + skip_name=args.lora_skip_name, + ) + network = network.to(weight_dtype) + network.apply_to(text_encoder, transformer3d, args.train_text_encoder and not args.training_with_video_token_length, True) if args.transformer_path is not None: print(f"From checkpoint: {args.transformer_path}") @@ -857,24 +912,79 @@ 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 - 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() + 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") + if args.use_peft_lora: + network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1]), accelerate_state_dict) + else: + network_state_dict = {} + for key in accelerate_state_dict: + if "network" in key: + network_state_dict[key.replace("network.", "")] = accelerate_state_dict[key].to(weight_dtype) + 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) + 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}.") + 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): + 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") + if args.use_peft_lora: + network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1]), accelerate_state_dict) + else: + network_state_dict = accelerate_state_dict + 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}.") + 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: + safetensor_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model.safetensors") + if args.use_peft_lora: + save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(models[-1]))) + else: + 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) + + 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) @@ -914,9 +1024,14 @@ def main(): else: optimizer_cls = torch.optim.AdamW - logging.info("Add network parameters") - trainable_params = list(filter(lambda p: p.requires_grad, network.parameters())) - trainable_params_optim = network.prepare_optimizer_params(args.learning_rate / 2, args.learning_rate, args.learning_rate) + if args.use_peft_lora: + logging.info("Add peft parameters") + trainable_params = list(filter(lambda p: p.requires_grad, transformer3d.parameters())) + trainable_params_optim = list(filter(lambda p: p.requires_grad, transformer3d.parameters())) + else: + logging.info("Add network parameters") + trainable_params = list(filter(lambda p: p.requires_grad, network.parameters())) + trainable_params_optim = network.prepare_optimizer_params(args.learning_rate / 2, args.learning_rate, args.learning_rate) if args.use_came: optimizer = optimizer_cls( @@ -1159,9 +1274,27 @@ def main(): ) # Prepare everything with our `accelerator`. - network, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( - network, optimizer, train_dataloader, lr_scheduler - ) + if args.use_peft_lora: + transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( + transformer3d, optimizer, train_dataloader, lr_scheduler + ) + elif fsdp_stage != 0: + transformer3d.network = network + transformer3d = transformer3d.to(dtype=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 != 0 and not args.use_peft_lora: + 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, module_to_wrapper=transformer3d.transformer_blocks) + transformer3d = shard_fn(transformer3d) # Move text_encode and vae to gpu and cast to weight_dtype vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) @@ -1236,57 +1369,58 @@ def main(): first_epoch = global_step // num_update_steps_per_epoch print(f"Load pkl from {pkl_path}. Get first_epoch = {first_epoch}.") - from safetensors.torch import load_file - state_dict = load_file(os.path.join(checkpoint_folder_path, "lora_diffusion_pytorch_model.safetensors"), device=str(accelerator.device)) - m, u = accelerator.unwrap_model(network).load_state_dict(state_dict, strict=False) - print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") + if zero_stage != 3 and not args.use_fsdp: + from safetensors.torch import load_file + state_dict = load_file(os.path.join(checkpoint_folder_path, "lora_diffusion_pytorch_model.safetensors"), device=str(accelerator.device)) + m, u = accelerator.unwrap_model(network).load_state_dict(state_dict, strict=False) + print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") - optimizer_file_pt = os.path.join(checkpoint_folder_path, "optimizer.pt") - optimizer_file_bin = os.path.join(checkpoint_folder_path, "optimizer.bin") - optimizer_file_to_load = None + optimizer_file_pt = os.path.join(checkpoint_folder_path, "optimizer.pt") + optimizer_file_bin = os.path.join(checkpoint_folder_path, "optimizer.bin") + optimizer_file_to_load = None - if os.path.exists(optimizer_file_pt): - optimizer_file_to_load = optimizer_file_pt - elif os.path.exists(optimizer_file_bin): - optimizer_file_to_load = optimizer_file_bin + if os.path.exists(optimizer_file_pt): + optimizer_file_to_load = optimizer_file_pt + elif os.path.exists(optimizer_file_bin): + optimizer_file_to_load = optimizer_file_bin - if optimizer_file_to_load: - try: - accelerator.print(f"Loading optimizer state from {optimizer_file_to_load}") - optimizer_state = torch.load(optimizer_file_to_load, map_location=accelerator.device) - optimizer.load_state_dict(optimizer_state) - accelerator.print("Optimizer state loaded successfully.") - except Exception as e: - accelerator.print(f"Failed to load optimizer state from {optimizer_file_to_load}: {e}") - - scheduler_file_pt = os.path.join(checkpoint_folder_path, "scheduler.pt") - scheduler_file_bin = os.path.join(checkpoint_folder_path, "scheduler.bin") - scheduler_file_to_load = None - - if os.path.exists(scheduler_file_pt): - scheduler_file_to_load = scheduler_file_pt - elif os.path.exists(scheduler_file_bin): - scheduler_file_to_load = scheduler_file_bin - - if scheduler_file_to_load: - try: - accelerator.print(f"Loading scheduler state from {scheduler_file_to_load}") - scheduler_state = torch.load(scheduler_file_to_load, map_location=accelerator.device) - lr_scheduler.load_state_dict(scheduler_state) - accelerator.print("Scheduler state loaded successfully.") - except Exception as e: - accelerator.print(f"Failed to load scheduler state from {scheduler_file_to_load}: {e}") - - if hasattr(accelerator, 'scaler') and accelerator.scaler is not None: - scaler_file = os.path.join(checkpoint_folder_path, "scaler.pt") - if os.path.exists(scaler_file): + if optimizer_file_to_load: try: - accelerator.print(f"Loading GradScaler state from {scaler_file}") - scaler_state = torch.load(scaler_file, map_location=accelerator.device) - accelerator.scaler.load_state_dict(scaler_state) - accelerator.print("GradScaler state loaded successfully.") + accelerator.print(f"Loading optimizer state from {optimizer_file_to_load}") + optimizer_state = torch.load(optimizer_file_to_load, map_location=accelerator.device) + optimizer.load_state_dict(optimizer_state) + accelerator.print("Optimizer state loaded successfully.") except Exception as e: - accelerator.print(f"Failed to load GradScaler state: {e}") + accelerator.print(f"Failed to load optimizer state from {optimizer_file_to_load}: {e}") + + scheduler_file_pt = os.path.join(checkpoint_folder_path, "scheduler.pt") + scheduler_file_bin = os.path.join(checkpoint_folder_path, "scheduler.bin") + scheduler_file_to_load = None + + if os.path.exists(scheduler_file_pt): + scheduler_file_to_load = scheduler_file_pt + elif os.path.exists(scheduler_file_bin): + scheduler_file_to_load = scheduler_file_bin + + if scheduler_file_to_load: + try: + accelerator.print(f"Loading scheduler state from {scheduler_file_to_load}") + scheduler_state = torch.load(scheduler_file_to_load, map_location=accelerator.device) + lr_scheduler.load_state_dict(scheduler_state) + accelerator.print("Scheduler state loaded successfully.") + except Exception as e: + accelerator.print(f"Failed to load scheduler state from {scheduler_file_to_load}: {e}") + + if hasattr(accelerator, 'scaler') and accelerator.scaler is not None: + scaler_file = os.path.join(checkpoint_folder_path, "scaler.pt") + if os.path.exists(scaler_file): + try: + accelerator.print(f"Loading GradScaler state from {scaler_file}") + scaler_state = torch.load(scaler_file, map_location=accelerator.device) + accelerator.scaler.load_state_dict(scaler_state) + accelerator.print("GradScaler state loaded successfully.") + except Exception as e: + accelerator.print(f"Failed to load GradScaler state: {e}") else: accelerator.load_state(checkpoint_folder_path) @@ -1299,6 +1433,10 @@ def main(): def save_model(ckpt_file, unwrapped_nw): os.makedirs(args.output_dir, exist_ok=True) accelerator.print(f"\nsaving checkpoint: {ckpt_file}") + if isinstance(unwrapped_nw, dict): + from safetensors.torch import save_file + save_file(unwrapped_nw, ckpt_file, metadata={"format": "pt"}) + return ckpt_file unwrapped_nw.save_weights(ckpt_file, weight_dtype, None) progress_bar = tqdm( @@ -1638,7 +1776,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) @@ -1662,9 +1800,14 @@ def main(): torch.cuda.empty_cache() torch.cuda.ipc_collect() if not args.save_state: - safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") - save_model(safetensor_save_path, accelerator.unwrap_model(network)) - logger.info(f"Saved safetensor to {safetensor_save_path}") + if args.use_peft_lora: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(transformer3d))) + logger.info(f"Saved safetensor to {safetensor_save_path}") + else: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, accelerator.unwrap_model(network)) + logger.info(f"Saved safetensor to {safetensor_save_path}") else: accelerator_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") accelerator.save_state(accelerator_save_path) @@ -1706,13 +1849,19 @@ def main(): # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() - if args.use_deepspeed or accelerator.is_main_process: + if args.use_deepspeed or args.use_fsdp or accelerator.is_main_process: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() if not args.save_state: - safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") - save_model(safetensor_save_path, accelerator.unwrap_model(network)) + if args.use_peft_lora: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(transformer3d))) + logger.info(f"Saved safetensor to {safetensor_save_path}") + else: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, accelerator.unwrap_model(network)) + logger.info(f"Saved safetensor to {safetensor_save_path}") else: accelerator_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") accelerator.save_state(accelerator_save_path) diff --git a/scripts/cogvideox_fun/train_lora.sh b/scripts/cogvideox_fun/train_lora.sh index a128ce7..3243e53 100755 --- a/scripts/cogvideox_fun/train_lora.sh +++ b/scripts/cogvideox_fun/train_lora.sh @@ -33,6 +33,10 @@ accelerate launch --mixed_precision="bf16" scripts/cogvideox_fun/train_lora.py \ --random_hw_adapt \ --training_with_video_token_length \ --enable_bucket \ + --rank=64 \ + --network_alpha=32 \ + --target_name="to_q,to_k,to_v,ff.0,ff.2" \ + --use_peft_lora \ --low_vram \ --train_mode="inpaint" @@ -71,5 +75,9 @@ accelerate launch --mixed_precision="bf16" scripts/cogvideox_fun/train_lora.py \ # --random_hw_adapt \ # --training_with_video_token_length \ # --enable_bucket \ +# --rank=64 \ +# --network_alpha=32 \ +# --target_name="to_q,to_k,to_v,ff.0,ff.2" \ +# --use_peft_lora \ # --low_vram \ # --train_mode="inpaint" \ No newline at end of file diff --git a/scripts/fantasytalking/README_TRAIN.md b/scripts/fantasytalking/README_TRAIN.md index fbe85e3..138c569 100755 --- a/scripts/fantasytalking/README_TRAIN.md +++ b/scripts/fantasytalking/README_TRAIN.md @@ -33,6 +33,17 @@ Some parameters in the sh file can be confusing, and they are explained in this - These resolutions combined with their corresponding lengths allow the model to generate videos of different sizes. - `resume_from_checkpoint` is used to set the training should be resumed from a previous checkpoint. Use a path or `"latest"` to automatically select the last available checkpoint. +When train model with multi machines, please set the params as follows: +```sh +export MASTER_ADDR="your master address" +export MASTER_PORT=10086 +export WORLD_SIZE=1 # The number of machines +export NUM_PROCESS=8 # The number of processes, such as WORLD_SIZE * 8 +export RANK=0 # The rank of this machine + +accelerate launch --mixed_precision="bf16" --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK scripts/xxx/xxx.py +``` + FantasyTalking without deepspeed: ```sh @@ -78,7 +89,7 @@ accelerate launch --mixed_precision="bf16" scripts/fantasytalking/train.py \ --trainable_modules "processor." "proj_model." ``` -FantasyTalking with deepspeed zero-2: +FantasyTalking with Deepspeed Zero-2: ```sh export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-I2V-14B-720P" @@ -123,7 +134,9 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con --trainable_modules "processor." "proj_model." ``` -FantasyTalking with deepspeed zero-3: +DeepSpeed Zero-3 is not highly recommended at the moment. In this repository, using FSDP has fewer errors and is more stable. + +FantasyTalking with DeepSpeed Zero-3: ```sh python scripts/zero_to_bf16.py output_dir/checkpoint-{our-num-steps} output_dir/checkpoint-{your-num-steps}-outputs --max_shard_size 80GB --safe_serialization diff --git a/scripts/fantasytalking/train.py b/scripts/fantasytalking/train.py index b73effe..6787b99 100644 --- a/scripts/fantasytalking/train.py +++ b/scripts/fantasytalking/train.py @@ -919,7 +919,12 @@ def main(): 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): + 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) diff --git a/scripts/flux/README_TRAIN.md b/scripts/flux/README_TRAIN.md index b091fff..3308bff 100755 --- a/scripts/flux/README_TRAIN.md +++ b/scripts/flux/README_TRAIN.md @@ -9,6 +9,17 @@ Some parameters in the sh file can be confusing, and they are explained in this - For example, when `random_hw_adapt` is enabled, `image_sample_size=1024`, the resolution of image inputs for training is `512x512` to `1024x1024` - `resume_from_checkpoint` is used to set the training should be resumed from a previous checkpoint. Use a path or `"latest"` to automatically select the last available checkpoint. +When train model with multi machines, please set the params as follows: +```sh +export MASTER_ADDR="your master address" +export MASTER_PORT=10086 +export WORLD_SIZE=1 # The number of machines +export NUM_PROCESS=8 # The number of processes, such as WORLD_SIZE * 8 +export RANK=0 # The rank of this machine + +accelerate launch --mixed_precision="bf16" --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK scripts/xxx/xxx.py +``` + Without deepspeed: Training flux without DeepSpeed may result in insufficient GPU memory. @@ -47,7 +58,7 @@ accelerate launch --mixed_precision="bf16" scripts/flux/train.py \ --trainable_modules "." ``` -With deepspeed zero-2: +With Deepspeed Zero-2: ```sh export MODEL_NAME="models/Diffusion_Transformer/FLUX.1-dev" @@ -84,49 +95,6 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con --trainable_modules "." ``` -Deepspeed zero-3: - -After training, you can use the following command to get the final model: -```sh -python scripts/zero_to_bf16.py output_dir/checkpoint-{our-num-steps} output_dir/checkpoint-{your-num-steps}-outputs --max_shard_size 80GB --safe_serialization -``` - -Training shell command is as follows: -```sh -export MODEL_NAME="models/Diffusion_Transformer/FLUX.1-dev" -export DATASET_NAME="datasets/internal_datasets/" -export DATASET_META_NAME="datasets/internal_datasets/metadata.json" -# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA. -# 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_launcher standard scripts/flux/train.py \ - --pretrained_model_name_or_path=$MODEL_NAME \ - --train_data_dir=$DATASET_NAME \ - --train_data_meta=$DATASET_META_NAME \ - --train_batch_size=1 \ - --image_sample_size=1024 \ - --gradient_accumulation_steps=1 \ - --dataloader_num_workers=8 \ - --num_train_epochs=100 \ - --checkpointing_steps=50 \ - --learning_rate=2e-05 \ - --lr_scheduler="constant_with_warmup" \ - --lr_warmup_steps=100 \ - --seed=42 \ - --output_dir="output_dir" \ - --gradient_checkpointing \ - --mixed_precision="bf16" \ - --adam_weight_decay=3e-2 \ - --adam_epsilon=1e-10 \ - --vae_mini_batch=1 \ - --max_grad_norm=0.05 \ - --enable_bucket \ - --uniform_sampling \ - --trainable_modules "." -``` - With FSDP: ```sh diff --git a/scripts/flux/README_TRAIN_LORA.md b/scripts/flux/README_TRAIN_LORA.md index 021c491..476795f 100755 --- a/scripts/flux/README_TRAIN_LORA.md +++ b/scripts/flux/README_TRAIN_LORA.md @@ -8,6 +8,21 @@ Some parameters in the sh file can be confusing, and they are explained in this - `random_hw_adapt` is used to enable automatic height and width scaling for images. When `random_hw_adapt` is enabled, the training images will have their height and width set to `image_sample_size` as the maximum and `512` as the minimum. - For example, when `random_hw_adapt` is enabled, `image_sample_size=1024`, the resolution of image inputs for training is `512x512` to `1024x1024` - `resume_from_checkpoint` is used to set the training should be resumed from a previous checkpoint. Use a path or `"latest"` to automatically select the last available checkpoint. +- `target_name` represents the components/modules to which LoRA will be applied, separated by commas. +- `use_peft_lora` indicates whether to use the PEFT module for adding LoRA. Using this module will be more memory-efficient. +- `rank` means the dimension of the LoRA update matrices. +- `network_alpha` means the scale of the LoRA update matrices. + +When train model with multi machines, please set the params as follows: +```sh +export MASTER_ADDR="your master address" +export MASTER_PORT=10086 +export WORLD_SIZE=1 # The number of machines +export NUM_PROCESS=8 # The number of processes, such as WORLD_SIZE * 8 +export RANK=0 # The rank of this machine + +accelerate launch --mixed_precision="bf16" --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK scripts/xxx/xxx.py +``` Without deepspeed: @@ -41,10 +56,14 @@ accelerate launch --mixed_precision="bf16" scripts/flux/train_lora.py \ --vae_mini_batch=1 \ --max_grad_norm=0.05 \ --enable_bucket \ + --rank=64 \ + --network_alpha=32 \ + --target_name="to_q,to_k,to_v,ff.0,ff.2,ff_context.0,ff_context.2" \ + --use_peft_lora \ --uniform_sampling ``` -With deepspeed zero-2: +With Deepspeed Zero-2: ```sh export MODEL_NAME="models/Diffusion_Transformer/FLUX.1-dev" @@ -75,46 +94,10 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con --vae_mini_batch=1 \ --max_grad_norm=0.05 \ --enable_bucket \ - --uniform_sampling -``` - -Deepspeed zero-3: - -After training, you can use the following command to get the final model: -```sh -python scripts/zero_to_bf16.py output_dir/checkpoint-{our-num-steps} output_dir/checkpoint-{your-num-steps}-outputs --max_shard_size 80GB --safe_serialization -``` - -Training shell command is as follows: -```sh -export MODEL_NAME="models/Diffusion_Transformer/FLUX.1-dev" -export DATASET_NAME="datasets/internal_datasets/" -export DATASET_META_NAME="datasets/internal_datasets/metadata.json" -# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA. -# 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_launcher standard scripts/flux/train_lora.py \ - --pretrained_model_name_or_path=$MODEL_NAME \ - --train_data_dir=$DATASET_NAME \ - --train_data_meta=$DATASET_META_NAME \ - --train_batch_size=1 \ - --image_sample_size=1024 \ - --gradient_accumulation_steps=1 \ - --dataloader_num_workers=8 \ - --num_train_epochs=100 \ - --checkpointing_steps=50 \ - --learning_rate=1e-04 \ - --seed=42 \ - --output_dir="output_dir" \ - --gradient_checkpointing \ - --mixed_precision="bf16" \ - --adam_weight_decay=3e-2 \ - --adam_epsilon=1e-10 \ - --vae_mini_batch=1 \ - --max_grad_norm=0.05 \ - --enable_bucket \ + --rank=64 \ + --network_alpha=32 \ + --target_name="to_q,to_k,to_v,ff.0,ff.2,ff_context.0,ff_context.2" \ + --use_peft_lora \ --uniform_sampling ``` diff --git a/scripts/flux/train.py b/scripts/flux/train.py index 02bac96..1610c7f 100644 --- a/scripts/flux/train.py +++ b/scripts/flux/train.py @@ -989,7 +989,12 @@ def main(): 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): + 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) diff --git a/scripts/flux/train_lora.py b/scripts/flux/train_lora.py index d447c37..1bb399d 100644 --- a/scripts/flux/train_lora.py +++ b/scripts/flux/train_lora.py @@ -577,6 +577,9 @@ def parse_args(): default=64, help=("The dimension of the LoRA update matrices."), ) + parser.add_argument( + "--use_peft_lora", action="store_true", help="Whether or not to use peft lora." + ) parser.add_argument( "--train_text_encoder", action="store_true", @@ -710,6 +713,12 @@ def parse_args(): default=3.5, help="the FLUX.1 dev variant is a guidance distilled model", ) + parser.add_argument( + "--target_name", + type=str, + default=None, + help=("The module is trained in loras. "), + ) args = parser.parse_args() env_local_rank = int(os.environ.get("LOCAL_RANK", -1)) @@ -894,16 +903,25 @@ def main(): transformer3d.requires_grad_(False) # Lora will work with this... - network = create_network( - 1.0, - args.rank, - args.network_alpha, - text_encoder, - transformer3d, - neuron_dropout=None, - skip_name=args.lora_skip_name, - ) - network.apply_to(text_encoder, transformer3d, args.train_text_encoder, True) + if args.use_peft_lora: + from peft import LoraConfig, inject_adapter_in_model, get_peft_model_state_dict + lora_config = LoraConfig(r=args.rank, lora_alpha=args.network_alpha, target_modules=args.target_name.split(",")) + transformer3d = inject_adapter_in_model(lora_config, transformer3d) + + network = None + else: + network = create_network( + 1.0, + args.rank, + args.network_alpha, + text_encoder, + transformer3d, + neuron_dropout=None, + target_name=args.target_name, + skip_name=args.lora_skip_name, + ) + network = network.to(weight_dtype) + network.apply_to(text_encoder, transformer3d, args.train_text_encoder and not args.training_with_video_token_length, True) if args.transformer_path is not None: print(f"From checkpoint: {args.transformer_path}") @@ -939,13 +957,14 @@ def main(): 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].to(weight_dtype) - + if args.use_peft_lora: + network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1]), accelerate_state_dict) + else: + network_state_dict = {} + for key in accelerate_state_dict: + if "network" in key: + network_state_dict[key.replace("network.", "")] = accelerate_state_dict[key].to(weight_dtype) 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: @@ -960,8 +979,18 @@ def main(): 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): + 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") + if args.use_peft_lora: + network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1]), accelerate_state_dict) + else: + network_state_dict = accelerate_state_dict + 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) @@ -973,10 +1002,15 @@ def main(): 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: 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 args.use_peft_lora: + save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(models[-1]))) + else: + save_model(safetensor_save_path, accelerator.unwrap_model(models[-1])) + if not args.use_deepspeed: for _ in range(len(weights)): weights.pop() @@ -1030,9 +1064,14 @@ def main(): else: optimizer_cls = torch.optim.AdamW - logging.info("Add network parameters") - trainable_params = list(filter(lambda p: p.requires_grad, network.parameters())) - trainable_params_optim = network.prepare_optimizer_params(args.learning_rate / 2, args.learning_rate, args.learning_rate) + if args.use_peft_lora: + logging.info("Add peft parameters") + trainable_params = list(filter(lambda p: p.requires_grad, transformer3d.parameters())) + trainable_params_optim = list(filter(lambda p: p.requires_grad, transformer3d.parameters())) + else: + logging.info("Add network parameters") + trainable_params = list(filter(lambda p: p.requires_grad, network.parameters())) + trainable_params_optim = network.prepare_optimizer_params(args.learning_rate / 2, args.learning_rate, args.learning_rate) if args.use_came: optimizer = optimizer_cls( @@ -1252,9 +1291,13 @@ def main(): ) # Prepare everything with our `accelerator`. - if fsdp_stage != 0: + if args.use_peft_lora: + transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( + transformer3d, optimizer, train_dataloader, lr_scheduler + ) + elif fsdp_stage != 0: transformer3d.network = network - transformer3d = transformer3d.to(weight_dtype) + transformer3d = transformer3d.to(dtype=weight_dtype) transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( transformer3d, optimizer, train_dataloader, lr_scheduler ) @@ -1263,14 +1306,14 @@ def main(): network, optimizer, train_dataloader, lr_scheduler ) - if zero_stage == 3: + if zero_stage != 0 and not args.use_peft_lora: 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, module_to_wrapper=transformer3d.transformer_blocks) transformer3d = shard_fn(transformer3d) - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage != 0: from functools import partial from videox_fun.dist import set_multi_gpus_devices, shard_model @@ -1417,6 +1460,10 @@ def main(): def save_model(ckpt_file, unwrapped_nw): os.makedirs(args.output_dir, exist_ok=True) accelerator.print(f"\nsaving checkpoint: {ckpt_file}") + if isinstance(unwrapped_nw, dict): + from safetensors.torch import save_file + save_file(unwrapped_nw, ckpt_file, metadata={"format": "pt"}) + return ckpt_file unwrapped_nw.save_weights(ckpt_file, weight_dtype, None) progress_bar = tqdm( @@ -1644,9 +1691,14 @@ def main(): torch.cuda.empty_cache() torch.cuda.ipc_collect() if not args.save_state: - safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") - save_model(safetensor_save_path, accelerator.unwrap_model(network)) - logger.info(f"Saved safetensor to {safetensor_save_path}") + if args.use_peft_lora: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(transformer3d))) + logger.info(f"Saved safetensor to {safetensor_save_path}") + else: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, accelerator.unwrap_model(network)) + logger.info(f"Saved safetensor to {safetensor_save_path}") else: accelerator_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") accelerator.save_state(accelerator_save_path) @@ -1697,8 +1749,14 @@ def main(): torch.cuda.empty_cache() torch.cuda.ipc_collect() if not args.save_state: - safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") - save_model(safetensor_save_path, accelerator.unwrap_model(network)) + if args.use_peft_lora: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(transformer3d))) + logger.info(f"Saved safetensor to {safetensor_save_path}") + else: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, accelerator.unwrap_model(network)) + logger.info(f"Saved safetensor to {safetensor_save_path}") else: accelerator_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") accelerator.save_state(accelerator_save_path) diff --git a/scripts/flux/train_lora.sh b/scripts/flux/train_lora.sh index deb8f92..5080571 100644 --- a/scripts/flux/train_lora.sh +++ b/scripts/flux/train_lora.sh @@ -26,4 +26,8 @@ accelerate launch --mixed_precision="bf16" scripts/flux/train_lora.py \ --vae_mini_batch=1 \ --max_grad_norm=0.05 \ --enable_bucket \ + --rank=64 \ + --network_alpha=32 \ + --target_name="to_q,to_k,to_v,ff.0,ff.2,ff_context.0,ff_context.2" \ + --use_peft_lora \ --uniform_sampling \ No newline at end of file diff --git a/scripts/qwenimage/README_TRAIN.md b/scripts/qwenimage/README_TRAIN.md index c77b5b6..94ff338 100755 --- a/scripts/qwenimage/README_TRAIN.md +++ b/scripts/qwenimage/README_TRAIN.md @@ -9,6 +9,17 @@ Some parameters in the sh file can be confusing, and they are explained in this - For example, when `random_hw_adapt` is enabled, `image_sample_size=1024`, the resolution of image inputs for training is `512x512` to `1024x1024` - `resume_from_checkpoint` is used to set the training should be resumed from a previous checkpoint. Use a path or `"latest"` to automatically select the last available checkpoint. +When train model with multi machines, please set the params as follows: +```sh +export MASTER_ADDR="your master address" +export MASTER_PORT=10086 +export WORLD_SIZE=1 # The number of machines +export NUM_PROCESS=8 # The number of processes, such as WORLD_SIZE * 8 +export RANK=0 # The rank of this machine + +accelerate launch --mixed_precision="bf16" --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK scripts/xxx/xxx.py +``` + Without deepspeed: Training qwen-image without DeepSpeed may result in insufficient GPU memory. @@ -47,7 +58,7 @@ accelerate launch --mixed_precision="bf16" scripts/qwenimage/train.py \ --trainable_modules "." ``` -With deepspeed zero-2: +With Deepspeed Zero-2: ```sh export MODEL_NAME="models/Diffusion_Transformer/Qwen-Image" @@ -84,7 +95,9 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con --trainable_modules "." ``` -Deepspeed zero-3: +DeepSpeed Zero-3 is not highly recommended at the moment. In this repository, using FSDP has fewer errors and is more stable. + +DeepSpeed Zero-3: After training, you can use the following command to get the final model: ```sh diff --git a/scripts/qwenimage/README_TRAIN_EDIT.md b/scripts/qwenimage/README_TRAIN_EDIT.md index cd913e9..98631cb 100755 --- a/scripts/qwenimage/README_TRAIN_EDIT.md +++ b/scripts/qwenimage/README_TRAIN_EDIT.md @@ -35,6 +35,17 @@ Some parameters in the sh file can be confusing, and they are explained in this - `qwen_image_edit` is for Qwen-Image-Edit. - `qwen_image_edit_plus` is for Qwen-Image-Edit-2509 +When train model with multi machines, please set the params as follows: +```sh +export MASTER_ADDR="your master address" +export MASTER_PORT=10086 +export WORLD_SIZE=1 # The number of machines +export NUM_PROCESS=8 # The number of processes, such as WORLD_SIZE * 8 +export RANK=0 # The rank of this machine + +accelerate launch --mixed_precision="bf16" --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK scripts/xxx/xxx.py +``` + Without deepspeed: Training qwen-image-edit without DeepSpeed may result in insufficient GPU memory. @@ -74,7 +85,7 @@ accelerate launch --mixed_precision="bf16" scripts/qwenimage/train_edit.py \ --train_mode "qwen_image_edit" ``` -With deepspeed zero-2: +With Deepspeed Zero-2: ```sh export MODEL_NAME="models/Diffusion_Transformer/Qwen-Image" @@ -112,7 +123,9 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con --train_mode "qwen_image_edit" ``` -Deepspeed zero-3: +DeepSpeed Zero-3 is not highly recommended at the moment. In this repository, using FSDP has fewer errors and is more stable. + +DeepSpeed Zero-3: After training, you can use the following command to get the final model: ```sh diff --git a/scripts/qwenimage/README_TRAIN_LORA.md b/scripts/qwenimage/README_TRAIN_LORA.md index e09f2ca..2804919 100755 --- a/scripts/qwenimage/README_TRAIN_LORA.md +++ b/scripts/qwenimage/README_TRAIN_LORA.md @@ -8,6 +8,21 @@ Some parameters in the sh file can be confusing, and they are explained in this - `random_hw_adapt` is used to enable automatic height and width scaling for images. When `random_hw_adapt` is enabled, the training images will have their height and width set to `image_sample_size` as the maximum and `512` as the minimum. - For example, when `random_hw_adapt` is enabled, `image_sample_size=1024`, the resolution of image inputs for training is `512x512` to `1024x1024` - `resume_from_checkpoint` is used to set the training should be resumed from a previous checkpoint. Use a path or `"latest"` to automatically select the last available checkpoint. +- `target_name` represents the components/modules to which LoRA will be applied, separated by commas. +- `use_peft_lora` indicates whether to use the PEFT module for adding LoRA. Using this module will be more memory-efficient. +- `rank` means the dimension of the LoRA update matrices. +- `network_alpha` means the scale of the LoRA update matrices. + +When train model with multi machines, please set the params as follows: +```sh +export MASTER_ADDR="your master address" +export MASTER_PORT=10086 +export WORLD_SIZE=1 # The number of machines +export NUM_PROCESS=8 # The number of processes, such as WORLD_SIZE * 8 +export RANK=0 # The rank of this machine + +accelerate launch --mixed_precision="bf16" --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK scripts/xxx/xxx.py +``` Without deepspeed: @@ -41,10 +56,14 @@ accelerate launch --mixed_precision="bf16" scripts/qwenimage/train_lora.py \ --vae_mini_batch=1 \ --max_grad_norm=0.05 \ --enable_bucket \ + --rank=64 \ + --network_alpha=32 \ + --target_name="to_q,to_k,to_v,img_mod.1,txt_mod.1,img_mlp.0,img_mlp.2,txt_mlp.0,txt_mlp.2" \ + --use_peft_lora \ --uniform_sampling ``` -With deepspeed zero-2: +With Deepspeed Zero-2: ```sh export MODEL_NAME="models/Diffusion_Transformer/Qwen-Image" @@ -75,10 +94,18 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con --vae_mini_batch=1 \ --max_grad_norm=0.05 \ --enable_bucket \ + --rank=64 \ + --network_alpha=32 \ + --target_name="to_q,to_k,to_v,img_mod.1,txt_mod.1,img_mlp.0,img_mlp.2,txt_mlp.0,txt_mlp.2" \ + --use_peft_lora \ --uniform_sampling ``` -Deepspeed zero-3: +DeepSpeed Zero-3 is not highly recommended at the moment. In this repository, using FSDP has fewer errors and is more stable. + +It is known that DeepSpeed Zero-3 is not compatible with PEFT. + +DeepSpeed Zero-3: After training, you can use the following command to get the final model: ```sh @@ -149,5 +176,9 @@ accelerate launch --mixed_precision="bf16" --use_fsdp --fsdp_auto_wrap_policy TR --vae_mini_batch=1 \ --max_grad_norm=0.05 \ --enable_bucket \ + --rank=64 \ + --network_alpha=32 \ + --target_name="to_q,to_k,to_v,img_mod.1,txt_mod.1,img_mlp.0,img_mlp.2,txt_mlp.0,txt_mlp.2" \ + --use_peft_lora \ --uniform_sampling ``` \ No newline at end of file diff --git a/scripts/qwenimage/train.py b/scripts/qwenimage/train.py index b98bbac..c943e5d 100644 --- a/scripts/qwenimage/train.py +++ b/scripts/qwenimage/train.py @@ -839,7 +839,12 @@ def main(): 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): + 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) diff --git a/scripts/qwenimage/train_edit.py b/scripts/qwenimage/train_edit.py index 4f72e58..787eaae 100644 --- a/scripts/qwenimage/train_edit.py +++ b/scripts/qwenimage/train_edit.py @@ -871,7 +871,12 @@ def main(): 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): + 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) diff --git a/scripts/qwenimage/train_edit_lora.py b/scripts/qwenimage/train_edit_lora.py index 065128b..f9c4181 100644 --- a/scripts/qwenimage/train_edit_lora.py +++ b/scripts/qwenimage/train_edit_lora.py @@ -490,6 +490,9 @@ def parse_args(): default=64, help=("The dimension of the LoRA update matrices."), ) + parser.add_argument( + "--use_peft_lora", action="store_true", help="Whether or not to use peft lora." + ) parser.add_argument( "--train_text_encoder", action="store_true", @@ -601,6 +604,12 @@ def parse_args(): default=None, help=("The module is not trained in loras. "), ) + parser.add_argument( + "--target_name", + type=str, + default=None, + help=("The module is trained in loras. "), + ) args = parser.parse_args() env_local_rank = int(os.environ.get("LOCAL_RANK", -1)) @@ -795,16 +804,25 @@ def main(): transformer3d.requires_grad_(False) # Lora will work with this... - network = create_network( - 1.0, - args.rank, - args.network_alpha, - text_encoder, - transformer3d, - neuron_dropout=None, - skip_name=args.lora_skip_name, - ) - network.apply_to(text_encoder, transformer3d, args.train_text_encoder, True) + if args.use_peft_lora: + from peft import LoraConfig, inject_adapter_in_model, get_peft_model_state_dict + lora_config = LoraConfig(r=args.rank, lora_alpha=args.network_alpha, target_modules=args.target_name.split(",")) + transformer3d = inject_adapter_in_model(lora_config, transformer3d) + + network = None + else: + network = create_network( + 1.0, + args.rank, + args.network_alpha, + text_encoder, + transformer3d, + neuron_dropout=None, + target_name=args.target_name, + skip_name=args.lora_skip_name, + ) + network = network.to(weight_dtype) + network.apply_to(text_encoder, transformer3d, args.train_text_encoder and not args.training_with_video_token_length, True) if args.transformer_path is not None: print(f"From checkpoint: {args.transformer_path}") @@ -840,13 +858,14 @@ def main(): 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].to(weight_dtype) - + if args.use_peft_lora: + network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1]), accelerate_state_dict) + else: + network_state_dict = {} + for key in accelerate_state_dict: + if "network" in key: + network_state_dict[key.replace("network.", "")] = accelerate_state_dict[key].to(weight_dtype) 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: @@ -861,8 +880,18 @@ def main(): 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): + 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") + if args.use_peft_lora: + network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1]), accelerate_state_dict) + else: + network_state_dict = accelerate_state_dict + 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) @@ -874,10 +903,15 @@ def main(): 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: 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 args.use_peft_lora: + save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(models[-1]))) + else: + save_model(safetensor_save_path, accelerator.unwrap_model(models[-1])) + if not args.use_deepspeed: for _ in range(len(weights)): weights.pop() @@ -931,9 +965,14 @@ def main(): else: optimizer_cls = torch.optim.AdamW - logging.info("Add network parameters") - trainable_params = list(filter(lambda p: p.requires_grad, network.parameters())) - trainable_params_optim = network.prepare_optimizer_params(args.learning_rate / 2, args.learning_rate, args.learning_rate) + if args.use_peft_lora: + logging.info("Add peft parameters") + trainable_params = list(filter(lambda p: p.requires_grad, transformer3d.parameters())) + trainable_params_optim = list(filter(lambda p: p.requires_grad, transformer3d.parameters())) + else: + logging.info("Add network parameters") + trainable_params = list(filter(lambda p: p.requires_grad, network.parameters())) + trainable_params_optim = network.prepare_optimizer_params(args.learning_rate / 2, args.learning_rate, args.learning_rate) if args.use_came: optimizer = optimizer_cls( @@ -1183,9 +1222,13 @@ def main(): ) # Prepare everything with our `accelerator`. - if fsdp_stage != 0: + if args.use_peft_lora: + transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( + transformer3d, optimizer, train_dataloader, lr_scheduler + ) + elif fsdp_stage != 0: transformer3d.network = network - transformer3d = transformer3d.to(weight_dtype) + transformer3d = transformer3d.to(dtype=weight_dtype) transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( transformer3d, optimizer, train_dataloader, lr_scheduler ) @@ -1194,14 +1237,14 @@ def main(): network, optimizer, train_dataloader, lr_scheduler ) - if zero_stage == 3: + if zero_stage != 0 and not args.use_peft_lora: 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, module_to_wrapper=transformer3d.transformer_blocks) transformer3d = shard_fn(transformer3d) - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage != 0: from functools import partial from videox_fun.dist import set_multi_gpus_devices, shard_model @@ -1345,6 +1388,10 @@ def main(): def save_model(ckpt_file, unwrapped_nw): os.makedirs(args.output_dir, exist_ok=True) accelerator.print(f"\nsaving checkpoint: {ckpt_file}") + if isinstance(unwrapped_nw, dict): + from safetensors.torch import save_file + save_file(unwrapped_nw, ckpt_file, metadata={"format": "pt"}) + return ckpt_file unwrapped_nw.save_weights(ckpt_file, weight_dtype, None) progress_bar = tqdm( @@ -1711,9 +1758,14 @@ def main(): torch.cuda.empty_cache() torch.cuda.ipc_collect() if not args.save_state: - safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") - save_model(safetensor_save_path, accelerator.unwrap_model(network)) - logger.info(f"Saved safetensor to {safetensor_save_path}") + if args.use_peft_lora: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(transformer3d))) + logger.info(f"Saved safetensor to {safetensor_save_path}") + else: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, accelerator.unwrap_model(network)) + logger.info(f"Saved safetensor to {safetensor_save_path}") else: accelerator_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") accelerator.save_state(accelerator_save_path) @@ -1760,8 +1812,14 @@ def main(): torch.cuda.empty_cache() torch.cuda.ipc_collect() if not args.save_state: - safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") - save_model(safetensor_save_path, accelerator.unwrap_model(network)) + if args.use_peft_lora: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(transformer3d))) + logger.info(f"Saved safetensor to {safetensor_save_path}") + else: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, accelerator.unwrap_model(network)) + logger.info(f"Saved safetensor to {safetensor_save_path}") else: accelerator_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") accelerator.save_state(accelerator_save_path) diff --git a/scripts/qwenimage/train_edit_lora.sh b/scripts/qwenimage/train_edit_lora.sh index 9ff174f..6450449 100644 --- a/scripts/qwenimage/train_edit_lora.sh +++ b/scripts/qwenimage/train_edit_lora.sh @@ -27,4 +27,8 @@ accelerate launch --mixed_precision="bf16" scripts/qwenimage/train_edit_lora.py --max_grad_norm=0.05 \ --enable_bucket \ --uniform_sampling \ + --rank=64 \ + --network_alpha=32 \ + --target_name="to_q,to_k,to_v,img_mod.1,txt_mod.1,img_mlp.0,img_mlp.2,txt_mlp.0,txt_mlp.2" \ + --use_peft_lora \ --train_mode "qwen_image_edit" \ No newline at end of file diff --git a/scripts/qwenimage/train_lora.py b/scripts/qwenimage/train_lora.py index d9e58e2..0055b52 100644 --- a/scripts/qwenimage/train_lora.py +++ b/scripts/qwenimage/train_lora.py @@ -460,6 +460,9 @@ def parse_args(): default=64, help=("The dimension of the LoRA update matrices."), ) + parser.add_argument( + "--use_peft_lora", action="store_true", help="Whether or not to use peft lora." + ) parser.add_argument( "--train_text_encoder", action="store_true", @@ -578,6 +581,12 @@ def parse_args(): default=None, help=("The module is not trained in loras. "), ) + parser.add_argument( + "--target_name", + type=str, + default=None, + help=("The module is trained in loras. "), + ) args = parser.parse_args() env_local_rank = int(os.environ.get("LOCAL_RANK", -1)) @@ -756,16 +765,25 @@ def main(): transformer3d.requires_grad_(False) # Lora will work with this... - network = create_network( - 1.0, - args.rank, - args.network_alpha, - text_encoder, - transformer3d, - neuron_dropout=None, - skip_name=args.lora_skip_name, - ) - network.apply_to(text_encoder, transformer3d, args.train_text_encoder, True) + if args.use_peft_lora: + from peft import LoraConfig, inject_adapter_in_model, get_peft_model_state_dict + lora_config = LoraConfig(r=args.rank, lora_alpha=args.network_alpha, target_modules=args.target_name.split(",")) + transformer3d = inject_adapter_in_model(lora_config, transformer3d) + + network = None + else: + network = create_network( + 1.0, + args.rank, + args.network_alpha, + text_encoder, + transformer3d, + neuron_dropout=None, + target_name=args.target_name, + skip_name=args.lora_skip_name, + ) + network = network.to(weight_dtype) + network.apply_to(text_encoder, transformer3d, args.train_text_encoder and not args.training_with_video_token_length, True) if args.transformer_path is not None: print(f"From checkpoint: {args.transformer_path}") @@ -801,13 +819,14 @@ def main(): 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].to(weight_dtype) - + if args.use_peft_lora: + network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1]), accelerate_state_dict) + else: + network_state_dict = {} + for key in accelerate_state_dict: + if "network" in key: + network_state_dict[key.replace("network.", "")] = accelerate_state_dict[key].to(weight_dtype) 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: @@ -822,8 +841,18 @@ def main(): 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): + 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") + if args.use_peft_lora: + network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1]), accelerate_state_dict) + else: + network_state_dict = accelerate_state_dict + 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) @@ -835,10 +864,15 @@ def main(): 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: 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 args.use_peft_lora: + save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(models[-1]))) + else: + save_model(safetensor_save_path, accelerator.unwrap_model(models[-1])) + if not args.use_deepspeed: for _ in range(len(weights)): weights.pop() @@ -892,9 +926,14 @@ def main(): else: optimizer_cls = torch.optim.AdamW - logging.info("Add network parameters") - trainable_params = list(filter(lambda p: p.requires_grad, network.parameters())) - trainable_params_optim = network.prepare_optimizer_params(args.learning_rate / 2, args.learning_rate, args.learning_rate) + if args.use_peft_lora: + logging.info("Add peft parameters") + trainable_params = list(filter(lambda p: p.requires_grad, transformer3d.parameters())) + trainable_params_optim = list(filter(lambda p: p.requires_grad, transformer3d.parameters())) + else: + logging.info("Add network parameters") + trainable_params = list(filter(lambda p: p.requires_grad, network.parameters())) + trainable_params_optim = network.prepare_optimizer_params(args.learning_rate / 2, args.learning_rate, args.learning_rate) if args.use_came: optimizer = optimizer_cls( @@ -1131,9 +1170,13 @@ def main(): ) # Prepare everything with our `accelerator`. - if fsdp_stage != 0: + if args.use_peft_lora: + transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( + transformer3d, optimizer, train_dataloader, lr_scheduler + ) + elif fsdp_stage != 0: transformer3d.network = network - transformer3d = transformer3d.to(weight_dtype) + transformer3d = transformer3d.to(dtype=weight_dtype) transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( transformer3d, optimizer, train_dataloader, lr_scheduler ) @@ -1142,14 +1185,14 @@ def main(): network, optimizer, train_dataloader, lr_scheduler ) - if zero_stage == 3: + if zero_stage != 0 and not args.use_peft_lora: 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, module_to_wrapper=transformer3d.transformer_blocks) transformer3d = shard_fn(transformer3d) - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage != 0: from functools import partial from videox_fun.dist import set_multi_gpus_devices, shard_model @@ -1293,6 +1336,10 @@ def main(): def save_model(ckpt_file, unwrapped_nw): os.makedirs(args.output_dir, exist_ok=True) accelerator.print(f"\nsaving checkpoint: {ckpt_file}") + if isinstance(unwrapped_nw, dict): + from safetensors.torch import save_file + save_file(unwrapped_nw, ckpt_file, metadata={"format": "pt"}) + return ckpt_file unwrapped_nw.save_weights(ckpt_file, weight_dtype, None) progress_bar = tqdm( @@ -1530,9 +1577,14 @@ def main(): torch.cuda.empty_cache() torch.cuda.ipc_collect() if not args.save_state: - safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") - save_model(safetensor_save_path, accelerator.unwrap_model(network)) - logger.info(f"Saved safetensor to {safetensor_save_path}") + if args.use_peft_lora: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(transformer3d))) + logger.info(f"Saved safetensor to {safetensor_save_path}") + else: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, accelerator.unwrap_model(network)) + logger.info(f"Saved safetensor to {safetensor_save_path}") else: accelerator_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") accelerator.save_state(accelerator_save_path) @@ -1579,8 +1631,14 @@ def main(): torch.cuda.empty_cache() torch.cuda.ipc_collect() if not args.save_state: - safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") - save_model(safetensor_save_path, accelerator.unwrap_model(network)) + if args.use_peft_lora: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(transformer3d))) + logger.info(f"Saved safetensor to {safetensor_save_path}") + else: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, accelerator.unwrap_model(network)) + logger.info(f"Saved safetensor to {safetensor_save_path}") else: accelerator_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") accelerator.save_state(accelerator_save_path) diff --git a/scripts/qwenimage/train_lora.sh b/scripts/qwenimage/train_lora.sh index cd4dd82..c273ad3 100644 --- a/scripts/qwenimage/train_lora.sh +++ b/scripts/qwenimage/train_lora.sh @@ -26,4 +26,8 @@ accelerate launch --mixed_precision="bf16" scripts/qwenimage/train_lora.py \ --vae_mini_batch=1 \ --max_grad_norm=0.05 \ --enable_bucket \ + --rank=64 \ + --network_alpha=32 \ + --target_name="to_q,to_k,to_v,img_mod.1,txt_mod.1,img_mlp.0,img_mlp.2,txt_mlp.0,txt_mlp.2" \ + --use_peft_lora \ --uniform_sampling \ No newline at end of file diff --git a/scripts/wan2.1/README_TRAIN.md b/scripts/wan2.1/README_TRAIN.md index 049c88e..d2c769d 100755 --- a/scripts/wan2.1/README_TRAIN.md +++ b/scripts/wan2.1/README_TRAIN.md @@ -22,6 +22,17 @@ Some parameters in the sh file can be confusing, and they are explained in this - `train_mode` is used to specify the training mode, which can be either normal or i2v. Since Wan uses the inpaint model to achieve image-to-video generation, the default is set to inpaint mode. If you only wish to achieve text-to-video generation, you can remove this line, and it will default to the text-to-video mode. - `resume_from_checkpoint` is used to set the training should be resumed from a previous checkpoint. Use a path or `"latest"` to automatically select the last available checkpoint. +When train model with multi machines, please set the params as follows: +```sh +export MASTER_ADDR="your master address" +export MASTER_PORT=10086 +export WORLD_SIZE=1 # The number of machines +export NUM_PROCESS=8 # The number of processes, such as WORLD_SIZE * 8 +export RANK=0 # The rank of this machine + +accelerate launch --mixed_precision="bf16" --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK scripts/xxx/xxx.py +``` + Wan T2V without deepspeed: Wan without DeepSpeed is more suitable for 1.3B Wan, as using it with 14B Wan may result in insufficient GPU memory. @@ -70,9 +81,9 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.1/train.py \ --trainable_modules "." ``` -Wan T2V with deepspeed zero-2: +Wan T2V with Deepspeed Zero-2: -Wan with DeepSpeed Zero-2 is suitable for training 1.3B Wan and 14B Wan at low resolutions, but training 14B Wan at high resolutions may still result in insufficient GPU memory. +Wan with Deepspeed Zero-2 is suitable for training 1.3B Wan and 14B Wan at low resolutions, but training 14B Wan at high resolutions may still result in insufficient GPU memory. ```sh export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-T2V-14B" @@ -120,7 +131,9 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con --trainable_modules "." ``` -Wan T2V with deepspeed zero-3: +DeepSpeed Zero-3 is not highly recommended at the moment. In this repository, using FSDP has fewer errors and is more stable. + +Wan T2V with DeepSpeed Zero-3: Wan with DeepSpeed Zero-3 is suitable for 14B Wan at high resolutions. After training, you can use the following command to get the final model: ```sh diff --git a/scripts/wan2.1/README_TRAIN_DISTILL.md b/scripts/wan2.1/README_TRAIN_DISTILL.md index f8b38c4..785a57e 100755 --- a/scripts/wan2.1/README_TRAIN_DISTILL.md +++ b/scripts/wan2.1/README_TRAIN_DISTILL.md @@ -22,6 +22,17 @@ Some parameters in the sh file can be confusing, and they are explained in this - `train_mode` is used to specify the training mode, which can be either normal or i2v. Since Wan uses the inpaint model to achieve image-to-video generation, the default is set to inpaint mode. If you only wish to achieve text-to-video generation, you can remove this line, and it will default to the text-to-video mode. - `resume_from_checkpoint` is used to set the training should be resumed from a previous checkpoint. Use a path or `"latest"` to automatically select the last available checkpoint. +When train model with multi machines, please set the params as follows: +```sh +export MASTER_ADDR="your master address" +export MASTER_PORT=10086 +export WORLD_SIZE=1 # The number of machines +export NUM_PROCESS=8 # The number of processes, such as WORLD_SIZE * 8 +export RANK=0 # The rank of this machine + +accelerate launch --mixed_precision="bf16" --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK scripts/xxx/xxx.py +``` + Wan distill without deepspeed: Wan distill without DeepSpeed and FSDP is more suitable for 1.3B Wan, as using it with 14B Wan may result in insufficient GPU memory. @@ -71,9 +82,9 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.1/train_distill.py \ --low_vram ``` -Wan distill with deepspeed zero-2: +Wan distill with Deepspeed Zero-2: -Wan with DeepSpeed Zero-2 is suitable for training 1.3B Wan and 14B Wan at low resolutions, but training 14B Wan at high resolutions may still result in insufficient GPU memory. +Wan with Deepspeed Zero-2 is suitable for training 1.3B Wan and 14B Wan at low resolutions, but training 14B Wan at high resolutions may still result in insufficient GPU memory. ```sh export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-T2V-1.3B/" @@ -121,7 +132,9 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con --low_vram ``` -Wan distill with deepspeed zero-3: +DeepSpeed Zero-3 is not highly recommended at the moment. In this repository, using FSDP has fewer errors and is more stable. + +Wan distill with DeepSpeed Zero-3: Wan with DeepSpeed Zero-3 is suitable for 14B Wan at high resolutions. After training, you can use the following command to get the final model: ```sh diff --git a/scripts/wan2.1/README_TRAIN_DISTILL_LORA.md b/scripts/wan2.1/README_TRAIN_DISTILL_LORA.md index fb0ae51..5da4d08 100755 --- a/scripts/wan2.1/README_TRAIN_DISTILL_LORA.md +++ b/scripts/wan2.1/README_TRAIN_DISTILL_LORA.md @@ -21,6 +21,21 @@ Some parameters in the sh file can be confusing, and they are explained in this - These resolutions combined with their corresponding lengths allow the model to generate videos of different sizes. - `train_mode` is used to specify the training mode, which can be either normal or i2v. Since Wan uses the inpaint model to achieve image-to-video generation, the default is set to inpaint mode. If you only wish to achieve text-to-video generation, you can remove this line, and it will default to the text-to-video mode. - `resume_from_checkpoint` is used to set the training should be resumed from a previous checkpoint. Use a path or `"latest"` to automatically select the last available checkpoint. +- `target_name` represents the components/modules to which LoRA will be applied, separated by commas. +- `use_peft_lora` indicates whether to use the PEFT module for adding LoRA. Using this module will be more memory-efficient. +- `rank` means the dimension of the LoRA update matrices. +- `network_alpha` means the scale of the LoRA update matrices. + +When train model with multi machines, please set the params as follows: +```sh +export MASTER_ADDR="your master address" +export MASTER_PORT=10086 +export WORLD_SIZE=1 # The number of machines +export NUM_PROCESS=8 # The number of processes, such as WORLD_SIZE * 8 +export RANK=0 # The rank of this machine + +accelerate launch --mixed_precision="bf16" --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK scripts/xxx/xxx.py +``` Wan distill without deepspeed: @@ -64,13 +79,17 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.1/train_distill_lora.py --training_with_video_token_length \ --enable_bucket \ --uniform_sampling \ + --rank=64 \ + --network_alpha=32 \ + --target_name="q,k,v,ffn.0,ffn.2" \ + --use_peft_lora \ --train_mode="normal" \ --low_vram ``` -Wan distill with deepspeed zero-2: +Wan distill with Deepspeed Zero-2: -Wan with DeepSpeed Zero-2 is suitable for training 1.3B Wan and 14B Wan at low resolutions, but training 14B Wan at high resolutions may still result in insufficient GPU memory. +Wan with Deepspeed Zero-2 is suitable for training 1.3B Wan and 14B Wan at low resolutions, but training 14B Wan at high resolutions may still result in insufficient GPU memory. ```sh export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-T2V-1.3B/" @@ -111,11 +130,19 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con --training_with_video_token_length \ --enable_bucket \ --uniform_sampling \ + --rank=64 \ + --network_alpha=32 \ + --target_name="q,k,v,ffn.0,ffn.2" \ + --use_peft_lora \ --train_mode="normal" \ --low_vram ``` -Wan distill with deepspeed zero-3: +DeepSpeed Zero-3 is not highly recommended at the moment. In this repository, using FSDP has fewer errors and is more stable. + +It is known that DeepSpeed Zero-3 is not compatible with PEFT. + +Wan distill with DeepSpeed Zero-3: Wan with DeepSpeed Zero-3 is suitable for 14B Wan at high resolutions. After training, you can use the following command to get the final model: ```sh @@ -208,6 +235,10 @@ accelerate launch --mixed_precision="bf16" --use_fsdp --fsdp_auto_wrap_policy TR --training_with_video_token_length \ --enable_bucket \ --uniform_sampling \ + --rank=64 \ + --network_alpha=32 \ + --target_name="q,k,v,ffn.0,ffn.2" \ + --use_peft_lora \ --train_mode="normal" \ --low_vram ``` \ No newline at end of file diff --git a/scripts/wan2.1/README_TRAIN_LORA.md b/scripts/wan2.1/README_TRAIN_LORA.md index 69f59f3..f04d036 100755 --- a/scripts/wan2.1/README_TRAIN_LORA.md +++ b/scripts/wan2.1/README_TRAIN_LORA.md @@ -19,6 +19,21 @@ Some parameters in the sh file can be confusing, and they are explained in this - These resolutions combined with their corresponding lengths allow the model to generate videos of different sizes. - `train_mode` is used to specify the training mode, which can be either normal or i2v. Since Wan uses the inpaint model to achieve image-to-video generation, the default is set to inpaint mode. If you only wish to achieve text-to-video generation, you can remove this line, and it will default to the text-to-video mode. - `resume_from_checkpoint` is used to set the training should be resumed from a previous checkpoint. Use a path or `"latest"` to automatically select the last available checkpoint. +- `target_name` represents the components/modules to which LoRA will be applied, separated by commas. +- `use_peft_lora` indicates whether to use the PEFT module for adding LoRA. Using this module will be more memory-efficient. +- `rank` means the dimension of the LoRA update matrices. +- `network_alpha` means the scale of the LoRA update matrices. + +When train model with multi machines, please set the params as follows: +```sh +export MASTER_ADDR="your master address" +export MASTER_PORT=10086 +export WORLD_SIZE=1 # The number of machines +export NUM_PROCESS=8 # The number of processes, such as WORLD_SIZE * 8 +export RANK=0 # The rank of this machine + +accelerate launch --mixed_precision="bf16" --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK scripts/xxx/xxx.py +``` Wan T2V without deepspeed: @@ -61,12 +76,16 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.1/train_lora.py \ --training_with_video_token_length \ --enable_bucket \ --uniform_sampling \ + --rank=64 \ + --network_alpha=32 \ + --target_name="q,k,v,ffn.0,ffn.2" \ + --use_peft_lora \ --low_vram ``` -Wan T2V with deepspeed zero-2: +Wan T2V with Deepspeed Zero-2: -Wan with DeepSpeed Zero-2 is suitable for training 1.3B Wan and 14B Wan at low resolutions, but training 14B Wan at high resolutions may still result in insufficient GPU memory. +Wan with Deepspeed Zero-2 is suitable for training 1.3B Wan and 14B Wan at low resolutions, but training 14B Wan at high resolutions may still result in insufficient GPU memory. ```sh export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-T2V-14B" @@ -106,11 +125,18 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con --training_with_video_token_length \ --enable_bucket \ --uniform_sampling \ - --use_deepspeed \ - --low_vram + --rank=64 \ + --network_alpha=32 \ + --target_name="q,k,v,ffn.0,ffn.2" \ + --use_peft_lora \ + --low_vram ``` -Wan T2V with deepspeed zero-3: +DeepSpeed Zero-3 is not highly recommended at the moment. In this repository, using FSDP has fewer errors and is more stable. + +It is known that DeepSpeed Zero-3 is not compatible with PEFT. + +Wan T2V with DeepSpeed Zero-3: Wan with DeepSpeed Zero-3 is suitable for 14B Wan at high resolutions. You must set save_state to True to save the model. After training, you can use the following command to get the final model: ```sh @@ -156,8 +182,7 @@ accelerate launch --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag --training_with_video_token_length \ --enable_bucket \ --uniform_sampling \ - --use_deepspeed \ - --low_vram + --low_vram ``` Wan T2V with FSDP: @@ -201,6 +226,9 @@ accelerate launch --mixed_precision="bf16" --use_fsdp --fsdp_auto_wrap_policy TR --training_with_video_token_length \ --enable_bucket \ --uniform_sampling \ - --use_deepspeed \ - --low_vram + --rank=64 \ + --network_alpha=32 \ + --target_name="q,k,v,ffn.0,ffn.2" \ + --use_peft_lora \ + --low_vram ``` \ No newline at end of file diff --git a/scripts/wan2.1/train.py b/scripts/wan2.1/train.py index 7316fe4..9e0a954 100755 --- a/scripts/wan2.1/train.py +++ b/scripts/wan2.1/train.py @@ -995,7 +995,12 @@ def main(): 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): + 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) diff --git a/scripts/wan2.1/train_distill.py b/scripts/wan2.1/train_distill.py index 65027e6..2123167 100644 --- a/scripts/wan2.1/train_distill.py +++ b/scripts/wan2.1/train_distill.py @@ -1008,7 +1008,12 @@ def main(): 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): + 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) diff --git a/scripts/wan2.1/train_distill_lora.py b/scripts/wan2.1/train_distill_lora.py index 6142a10..62ea9c3 100644 --- a/scripts/wan2.1/train_distill_lora.py +++ b/scripts/wan2.1/train_distill_lora.py @@ -580,6 +580,9 @@ def parse_args(): default=64, help=("The dimension of the LoRA update matrices."), ) + parser.add_argument( + "--use_peft_lora", action="store_true", help="Whether or not to use peft lora." + ) parser.add_argument( "--train_text_encoder", action="store_true", @@ -746,6 +749,12 @@ def parse_args(): default=None, help=("The module is not trained in loras. "), ) + parser.add_argument( + "--target_name", + type=str, + default=None, + help=("The module is trained in loras. "), + ) args = parser.parse_args() env_local_rank = int(os.environ.get("LOCAL_RANK", -1)) @@ -951,27 +960,40 @@ def main(): clip_image_encoder.requires_grad_(False) # Lora will work with this... - network = create_network( - 1.0, - args.rank, - args.network_alpha, - text_encoder, - generator_transformer3d, - neuron_dropout=None, - skip_name=args.lora_skip_name, - ) - network.apply_to(text_encoder, generator_transformer3d, args.train_text_encoder and not args.training_with_video_token_length, True) + if args.use_peft_lora: + from peft import LoraConfig, inject_adapter_in_model, get_peft_model_state_dict + lora_config = LoraConfig(r=args.rank, lora_alpha=args.network_alpha, target_modules=args.target_name.split(",")) + generator_transformer3d = inject_adapter_in_model(lora_config, generator_transformer3d) - fake_score_network = create_network( - 1.0, - args.rank, - args.network_alpha, - None, - fake_score_transformer3d, - neuron_dropout=None, - skip_name=args.lora_skip_name, - ) - fake_score_network.apply_to(None, fake_score_transformer3d, False, True) + fake_score_lora_config = LoraConfig(r=args.rank, lora_alpha=args.network_alpha, target_modules=args.target_name.split(",")) + fake_score_transformer3d = inject_adapter_in_model(fake_score_lora_config, fake_score_transformer3d) + + network = None + fake_score_network = None + else: + network = create_network( + 1.0, + args.rank, + args.network_alpha, + text_encoder, + generator_transformer3d, + neuron_dropout=None, + target_name=args.target_name, + skip_name=args.lora_skip_name, + ) + network.apply_to(text_encoder, generator_transformer3d, args.train_text_encoder and not args.training_with_video_token_length, True) + + fake_score_network = create_network( + 1.0, + args.rank, + args.network_alpha, + None, + fake_score_transformer3d, + neuron_dropout=None, + target_name=args.target_name, + skip_name=args.lora_skip_name, + ) + fake_score_network.apply_to(None, fake_score_transformer3d, False, True) if args.transformer_path is not None: print(f"From checkpoint: {args.transformer_path}") @@ -1007,13 +1029,14 @@ def main(): 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].to(weight_dtype) - + if args.use_peft_lora: + network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1]), accelerate_state_dict) + else: + network_state_dict = {} + for key in accelerate_state_dict: + if "network" in key: + network_state_dict[key.replace("network.", "")] = accelerate_state_dict[key].to(weight_dtype) 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: @@ -1030,7 +1053,16 @@ def main(): 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): + 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") + if args.use_peft_lora: + network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1]), accelerate_state_dict) + else: + network_state_dict = accelerate_state_dict + 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) @@ -1046,7 +1078,11 @@ def main(): 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 args.use_peft_lora: + save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(models[-1]))) + else: + save_model(safetensor_save_path, accelerator.unwrap_model(models[-1])) + if not args.use_deepspeed: for _ in range(len(weights)): weights.pop() @@ -1105,14 +1141,22 @@ def main(): else: optimizer_cls = torch.optim.AdamW + if args.use_peft_lora: + logging.info("Add peft parameters") + trainable_params = list(filter(lambda p: p.requires_grad, generator_transformer3d.parameters())) + trainable_params_optim = list(filter(lambda p: p.requires_grad, generator_transformer3d.parameters())) - logging.info("Add network parameters") - trainable_params = list(filter(lambda p: p.requires_grad, network.parameters())) - trainable_params_optim = network.prepare_optimizer_params(args.learning_rate / 2, args.learning_rate, args.learning_rate) + logging.info("Add fake score peft parameters") + fake_trainable_params = list(filter(lambda p: p.requires_grad, fake_score_transformer3d.parameters())) + fake_trainable_params_optim = list(filter(lambda p: p.requires_grad, fake_score_transformer3d.parameters())) + else: + logging.info("Add network parameters") + trainable_params = list(filter(lambda p: p.requires_grad, network.parameters())) + trainable_params_optim = network.prepare_optimizer_params(args.learning_rate / 2, args.learning_rate, args.learning_rate) - logging.info("Add fake_score_network parameters") - fake_trainable_params = list(filter(lambda p: p.requires_grad, fake_score_network.parameters())) - fake_trainable_params_optim = fake_score_network.prepare_optimizer_params(args.learning_rate / 2, args.learning_rate, args.learning_rate) + logging.info("Add fake_score_network parameters") + fake_trainable_params = list(filter(lambda p: p.requires_grad, fake_score_network.parameters())) + fake_trainable_params_optim = fake_score_network.prepare_optimizer_params(args.learning_rate / 2, args.learning_rate, args.learning_rate) if args.use_came: optimizer = optimizer_cls( @@ -1488,14 +1532,21 @@ def main(): ) # Prepare everything with our `accelerator`. - if fsdp_stage != 0: + if args.use_peft_lora: + generator_transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( + generator_transformer3d, optimizer, train_dataloader, lr_scheduler + ) + fake_score_transformer3d, critic_optimizer, fake_score_lr_scheduler = accelerator_fake_score_transformer3d.prepare( + fake_score_transformer3d, critic_optimizer, fake_score_lr_scheduler + ) + elif fsdp_stage != 0: generator_transformer3d.network = network - generator_transformer3d = generator_transformer3d.to(weight_dtype) + generator_transformer3d = generator_transformer3d.to(dtype=weight_dtype) generator_transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( generator_transformer3d, optimizer, train_dataloader, lr_scheduler ) fake_score_transformer3d.network = fake_score_network - fake_score_transformer3d = fake_score_transformer3d.to(weight_dtype) + fake_score_transformer3d = fake_score_transformer3d.to(dtype=weight_dtype) fake_score_transformer3d, critic_optimizer, fake_score_lr_scheduler = accelerator_fake_score_transformer3d.prepare( fake_score_transformer3d, critic_optimizer, fake_score_lr_scheduler ) @@ -1507,33 +1558,26 @@ def main(): fake_score_network, critic_optimizer, fake_score_lr_scheduler ) - if zero_stage == 3: + if zero_stage != 0 and not args.use_peft_lora: 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) generator_transformer3d = shard_fn(generator_transformer3d) - - 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) fake_score_transformer3d = shard_fn(fake_score_transformer3d) + if fsdp_stage != 0 or zero_stage != 0: 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) real_score_transformer3d = shard_fn(real_score_transformer3d) - if fsdp_stage != 0 or zero_stage != 0: - 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) text_encoder = shard_fn(text_encoder) # Move text_encode and vae to gpu and cast to weight_dtype vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + generator_transformer3d.to(accelerator.device, dtype=weight_dtype) + fake_score_transformer3d.to(accelerator.device, dtype=weight_dtype) real_score_transformer3d.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) if not args.enable_text_encoder_in_dataloader: text_encoder.to(accelerator.device if not args.low_vram else "cpu") @@ -1613,6 +1657,16 @@ def main(): else: initial_global_step = 0 + # function for saving/removing + def save_model(ckpt_file, unwrapped_nw): + os.makedirs(args.output_dir, exist_ok=True) + accelerator.print(f"\nsaving checkpoint: {ckpt_file}") + if isinstance(unwrapped_nw, dict): + from safetensors.torch import save_file + save_file(unwrapped_nw, ckpt_file, metadata={"format": "pt"}) + return ckpt_file + unwrapped_nw.save_weights(ckpt_file, weight_dtype, None) + progress_bar = tqdm( range(0, args.max_train_steps), initial=initial_global_step, @@ -2197,9 +2251,13 @@ def main(): return final_loss denoising_loss = custom_mse_loss(fake_score_denoised_output, critic_noise - fake_score_denoised_pred) - avg_denoising_loss = accelerator.gather(denoising_loss.repeat(args.train_batch_size)).mean() + avg_denoising_loss = accelerator_fake_score_transformer3d.gather(denoising_loss.repeat(args.train_batch_size)).mean() train_denoising_loss += avg_denoising_loss.item() / args.gradient_accumulation_steps - + + if args.low_vram: + generator_transformer3d = generator_transformer3d.to("cpu") + torch.cuda.empty_cache() + accelerator_fake_score_transformer3d.backward(denoising_loss) if accelerator_fake_score_transformer3d.sync_gradients: accelerator_fake_score_transformer3d.clip_grad_norm_(fake_trainable_params, args.max_grad_norm) @@ -2246,13 +2304,22 @@ def main(): torch.cuda.empty_cache() torch.cuda.ipc_collect() if not args.save_state: - safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") - save_model(safetensor_save_path, accelerator.unwrap_model(network)) - logger.info(f"Saved safetensor to {safetensor_save_path}") + if args.use_peft_lora: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(generator_transformer3d))) + logger.info(f"Saved safetensor to {safetensor_save_path}") - safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}-fake_score.safetensors") - save_model(safetensor_save_path, accelerator.unwrap_model(fake_score_network)) - logger.info(f"Saved safetensor to {safetensor_save_path}") + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}-fake_score.safetensors") + save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(fake_score_transformer3d))) + logger.info(f"Saved safetensor to {safetensor_save_path}") + else: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, accelerator.unwrap_model(network)) + logger.info(f"Saved safetensor to {safetensor_save_path}") + + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}-fake_score.safetensors") + save_model(safetensor_save_path, accelerator.unwrap_model(fake_score_network)) + logger.info(f"Saved safetensor to {safetensor_save_path}") else: save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") fake_score_save_path = os.path.join(save_path, "fake_score") @@ -2302,19 +2369,29 @@ def main(): accelerator.wait_for_everyone() if accelerator.is_main_process: generator_transformer3d = unwrap_model(generator_transformer3d) + fake_score_transformer3d = unwrap_model(fake_score_transformer3d) if args.use_deepspeed or args.use_fsdp or accelerator.is_main_process: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() if not args.save_state: - safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") - save_model(safetensor_save_path, accelerator.unwrap_model(network)) - logger.info(f"Saved safetensor to {safetensor_save_path}") + if args.use_peft_lora: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(generator_transformer3d))) + logger.info(f"Saved safetensor to {safetensor_save_path}") - safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}-fake_score.safetensors") - save_model(safetensor_save_path, accelerator.unwrap_model(fake_score_network)) - logger.info(f"Saved safetensor to {safetensor_save_path}") + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}-fake_score.safetensors") + save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(fake_score_transformer3d))) + logger.info(f"Saved safetensor to {safetensor_save_path}") + else: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, accelerator.unwrap_model(network)) + logger.info(f"Saved safetensor to {safetensor_save_path}") + + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}-fake_score.safetensors") + save_model(safetensor_save_path, accelerator.unwrap_model(fake_score_network)) + logger.info(f"Saved safetensor to {safetensor_save_path}") else: save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") fake_score_save_path = os.path.join(save_path, "fake_score") diff --git a/scripts/wan2.1/train_distill_lora.sh b/scripts/wan2.1/train_distill_lora.sh index e1b8197..2d55cb1 100644 --- a/scripts/wan2.1/train_distill_lora.sh +++ b/scripts/wan2.1/train_distill_lora.sh @@ -36,6 +36,10 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.1/train_distill_lora.py --training_with_video_token_length \ --enable_bucket \ --uniform_sampling \ + --rank=64 \ + --network_alpha=32 \ + --target_name="q,k,v,ffn.0,ffn.2" \ + --use_peft_lora \ --train_mode="normal" \ --low_vram diff --git a/scripts/wan2.1/train_lora.py b/scripts/wan2.1/train_lora.py index b97771d..2ff6a44 100755 --- a/scripts/wan2.1/train_lora.py +++ b/scripts/wan2.1/train_lora.py @@ -557,6 +557,9 @@ def parse_args(): default=64, help=("The dimension of the LoRA update matrices."), ) + parser.add_argument( + "--use_peft_lora", action="store_true", help="Whether or not to use peft lora." + ) parser.add_argument( "--train_text_encoder", action="store_true", @@ -724,6 +727,12 @@ def parse_args(): default=None, help=("The module is not trained in loras. "), ) + parser.add_argument( + "--target_name", + type=str, + default=None, + help=("The module is trained in loras. "), + ) args = parser.parse_args() env_local_rank = int(os.environ.get("LOCAL_RANK", -1)) @@ -910,16 +919,25 @@ def main(): clip_image_encoder.requires_grad_(False) # Lora will work with this... - network = create_network( - 1.0, - args.rank, - args.network_alpha, - text_encoder, - transformer3d, - neuron_dropout=None, - skip_name=args.lora_skip_name, - ) - network.apply_to(text_encoder, transformer3d, args.train_text_encoder and not args.training_with_video_token_length, True) + if args.use_peft_lora: + from peft import LoraConfig, inject_adapter_in_model, get_peft_model_state_dict + lora_config = LoraConfig(r=args.rank, lora_alpha=args.network_alpha, target_modules=args.target_name.split(",")) + transformer3d = inject_adapter_in_model(lora_config, transformer3d) + + network = None + else: + network = create_network( + 1.0, + args.rank, + args.network_alpha, + text_encoder, + transformer3d, + neuron_dropout=None, + target_name=args.target_name, + skip_name=args.lora_skip_name, + ) + network = network.to(weight_dtype) + network.apply_to(text_encoder, transformer3d, args.train_text_encoder and not args.training_with_video_token_length, True) if args.transformer_path is not None: print(f"From checkpoint: {args.transformer_path}") @@ -955,13 +973,14 @@ def main(): 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].to(weight_dtype) - + if args.use_peft_lora: + network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1]), accelerate_state_dict) + else: + network_state_dict = {} + for key in accelerate_state_dict: + if "network" in key: + network_state_dict[key.replace("network.", "")] = accelerate_state_dict[key].to(weight_dtype) 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: @@ -976,8 +995,18 @@ def main(): 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): + 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") + if args.use_peft_lora: + network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1]), accelerate_state_dict) + else: + network_state_dict = accelerate_state_dict + 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) @@ -989,10 +1018,15 @@ def main(): 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: 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 args.use_peft_lora: + save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(models[-1]))) + else: + save_model(safetensor_save_path, accelerator.unwrap_model(models[-1])) + if not args.use_deepspeed: for _ in range(len(weights)): weights.pop() @@ -1046,9 +1080,14 @@ def main(): else: optimizer_cls = torch.optim.AdamW - logging.info("Add network parameters") - trainable_params = list(filter(lambda p: p.requires_grad, network.parameters())) - trainable_params_optim = network.prepare_optimizer_params(args.learning_rate / 2, args.learning_rate, args.learning_rate) + if args.use_peft_lora: + logging.info("Add peft parameters") + trainable_params = list(filter(lambda p: p.requires_grad, transformer3d.parameters())) + trainable_params_optim = list(filter(lambda p: p.requires_grad, transformer3d.parameters())) + else: + logging.info("Add network parameters") + trainable_params = list(filter(lambda p: p.requires_grad, network.parameters())) + trainable_params_optim = network.prepare_optimizer_params(args.learning_rate / 2, args.learning_rate, args.learning_rate) if args.use_came: optimizer = optimizer_cls( @@ -1313,9 +1352,13 @@ def main(): ) # Prepare everything with our `accelerator`. - if fsdp_stage != 0: + if args.use_peft_lora: + transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( + transformer3d, optimizer, train_dataloader, lr_scheduler + ) + elif fsdp_stage != 0: transformer3d.network = network - transformer3d = transformer3d.to(weight_dtype) + transformer3d = transformer3d.to(dtype=weight_dtype) transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( transformer3d, optimizer, train_dataloader, lr_scheduler ) @@ -1324,14 +1367,16 @@ def main(): network, optimizer, train_dataloader, lr_scheduler ) - if zero_stage == 3: + if zero_stage != 0 and not args.use_peft_lora: 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) - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage != 0: 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) text_encoder = shard_fn(text_encoder) @@ -1475,6 +1520,10 @@ def main(): def save_model(ckpt_file, unwrapped_nw): os.makedirs(args.output_dir, exist_ok=True) accelerator.print(f"\nsaving checkpoint: {ckpt_file}") + if isinstance(unwrapped_nw, dict): + from safetensors.torch import save_file + save_file(unwrapped_nw, ckpt_file, metadata={"format": "pt"}) + return ckpt_file unwrapped_nw.save_weights(ckpt_file, weight_dtype, None) progress_bar = tqdm( @@ -1829,9 +1878,14 @@ def main(): torch.cuda.empty_cache() torch.cuda.ipc_collect() if not args.save_state: - safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") - save_model(safetensor_save_path, accelerator.unwrap_model(network)) - logger.info(f"Saved safetensor to {safetensor_save_path}") + if args.use_peft_lora: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(transformer3d))) + logger.info(f"Saved safetensor to {safetensor_save_path}") + else: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, accelerator.unwrap_model(network)) + logger.info(f"Saved safetensor to {safetensor_save_path}") else: accelerator_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") accelerator.save_state(accelerator_save_path) @@ -1882,8 +1936,14 @@ def main(): torch.cuda.empty_cache() torch.cuda.ipc_collect() if not args.save_state: - safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") - save_model(safetensor_save_path, accelerator.unwrap_model(network)) + if args.use_peft_lora: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(transformer3d))) + logger.info(f"Saved safetensor to {safetensor_save_path}") + else: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, accelerator.unwrap_model(network)) + logger.info(f"Saved safetensor to {safetensor_save_path}") else: accelerator_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") accelerator.save_state(accelerator_save_path) diff --git a/scripts/wan2.1/train_lora.sh b/scripts/wan2.1/train_lora.sh index 3dda4fb..3738616 100755 --- a/scripts/wan2.1/train_lora.sh +++ b/scripts/wan2.1/train_lora.sh @@ -35,6 +35,10 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.1/train_lora.py \ --training_with_video_token_length \ --enable_bucket \ --uniform_sampling \ + --rank=64 \ + --network_alpha=32 \ + --target_name="q,k,v,ffn.0,ffn.2" \ + --use_peft_lora \ --low_vram # # Training command for I2V @@ -74,5 +78,9 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.1/train_lora.py \ # --training_with_video_token_length \ # --enable_bucket \ # --uniform_sampling \ +# --rank=64 \ +# --network_alpha=32 \ +# --target_name="q,k,v,ffn.0,ffn.2" \ +# --use_peft_lora \ # --low_vram \ # --train_mode="i2v" \ No newline at end of file diff --git a/scripts/wan2.1_fun/README_TRAIN.md b/scripts/wan2.1_fun/README_TRAIN.md index 2465953..61eebd1 100755 --- a/scripts/wan2.1_fun/README_TRAIN.md +++ b/scripts/wan2.1_fun/README_TRAIN.md @@ -20,6 +20,17 @@ Some parameters in the sh file can be confusing, and they are explained in this - `train_mode` is used to specify the training mode, which can be either normal or inpaint. Since Wan uses the inpaint model to achieve image-to-video generation, the default is set to inpaint mode. If you only wish to achieve text-to-video generation, you can remove this line, and it will default to the text-to-video mode. - `resume_from_checkpoint` is used to set the training should be resumed from a previous checkpoint. Use a path or `"latest"` to automatically select the last available checkpoint. +When train model with multi machines, please set the params as follows: +```sh +export MASTER_ADDR="your master address" +export MASTER_PORT=10086 +export WORLD_SIZE=1 # The number of machines +export NUM_PROCESS=8 # The number of processes, such as WORLD_SIZE * 8 +export RANK=0 # The rank of this machine + +accelerate launch --mixed_precision="bf16" --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK scripts/xxx/xxx.py +``` + Wan T2V without deepspeed: Wan without DeepSpeed is more suitable for 1.3B Wan, as using it with 14B Wan may result in insufficient GPU memory. @@ -68,9 +79,9 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.1_fun/train.py \ --trainable_modules "." ``` -Wan T2V with deepspeed zero-2: +Wan T2V with Deepspeed Zero-2: -Wan with DeepSpeed Zero-2 is suitable for training 1.3B Wan and 14B Wan at low resolutions, but training 14B Wan at high resolutions may still result in insufficient GPU memory. +Wan with Deepspeed Zero-2 is suitable for training 1.3B Wan and 14B Wan at low resolutions, but training 14B Wan at high resolutions may still result in insufficient GPU memory. ```sh export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-V1.1-14B-InP" @@ -118,7 +129,9 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con --trainable_modules "." ``` -Wan T2V with deepspeed zero-3: +DeepSpeed Zero-3 is not highly recommended at the moment. In this repository, using FSDP has fewer errors and is more stable. + +Wan T2V with DeepSpeed Zero-3: Wan with DeepSpeed Zero-3 is suitable for 14B Wan at high resolutions. After training, you can use the following command to get the final model: ```sh diff --git a/scripts/wan2.1_fun/README_TRAIN_CONTROL.md b/scripts/wan2.1_fun/README_TRAIN_CONTROL.md index b1023f7..45077dd 100755 --- a/scripts/wan2.1_fun/README_TRAIN_CONTROL.md +++ b/scripts/wan2.1_fun/README_TRAIN_CONTROL.md @@ -154,7 +154,9 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con --trainable_modules "." ``` -Wan-Fun-Control-V1.1 with deepspeed zero-3: +DeepSpeed Zero-3 is not highly recommended at the moment. In this repository, using FSDP has fewer errors and is more stable. + +Wan-Fun-Control-V1.1 with DeepSpeed Zero-3: Wan with DeepSpeed Zero-3 is suitable for 14B Wan at high resolutions. After training, you can use the following command to get the final model: ```sh @@ -359,7 +361,7 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con --trainable_modules "." ``` -Wan-Fun-Control-V1.0 with deepspeed zero-3: +Wan-Fun-Control-V1.0 with DeepSpeed Zero-3: Wan with DeepSpeed Zero-3 is suitable for 14B Wan at high resolutions. After training, you can use the following command to get the final model: ```sh diff --git a/scripts/wan2.1_fun/README_TRAIN_CONTROL_LORA.md b/scripts/wan2.1_fun/README_TRAIN_CONTROL_LORA.md index 0fecd87..cf82ff9 100755 --- a/scripts/wan2.1_fun/README_TRAIN_CONTROL_LORA.md +++ b/scripts/wan2.1_fun/README_TRAIN_CONTROL_LORA.md @@ -45,6 +45,10 @@ Some parameters in the sh file can be confusing, and they are explained in this - `first_frame` is used in V1.0 because V1.0 supports using a specified start frame as the control image. The Control-Camera models use the first frame as the control image. - `random` is used in V1.1 because V1.1 supports both using a specified start frame and a reference image as the control image. - `add_full_ref_image_in_self_attention` determines whether to include the reference image in self-attention. This option is used in V1.1, as it supports using a reference image as the control image. It should not be used in V1.0 and Control-Camera models. +- `target_name` represents the components/modules to which LoRA will be applied, separated by commas. +- `use_peft_lora` indicates whether to use the PEFT module for adding LoRA. Using this module will be more memory-efficient. +- `rank` means the dimension of the LoRA update matrices. +- `network_alpha` means the scale of the LoRA update matrices. When train model with multi machines, please set the params as follows: ```sh @@ -99,6 +103,10 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.1_fun/train_control_lora --train_mode="control_ref" \ --control_ref_image="random" \ --add_full_ref_image_in_self_attention \ + --rank=64 \ + --network_alpha=32 \ + --target_name="q,k,v,ffn.0,ffn.2" \ + --use_peft_lora \ --low_vram ``` @@ -145,10 +153,18 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con --train_mode="control_ref" \ --control_ref_image="random" \ --add_full_ref_image_in_self_attention \ + --rank=64 \ + --network_alpha=32 \ + --target_name="q,k,v,ffn.0,ffn.2" \ + --use_peft_lora \ --low_vram ``` -Wan-Fun-Control-V1.1 with deepspeed zero-3: +DeepSpeed Zero-3 is not highly recommended at the moment. In this repository, using FSDP has fewer errors and is more stable. + +It is known that DeepSpeed Zero-3 is not compatible with PEFT. + +Wan-Fun-Control-V1.1 with DeepSpeed Zero-3: Wan with DeepSpeed Zero-3 is suitable for 14B Wan at high resolutions. You must set save_state to True to save the model. After training, you can use the following command to get the final model: ```sh @@ -248,6 +264,10 @@ accelerate launch --mixed_precision="bf16" --use_fsdp --fsdp_auto_wrap_policy TR --train_mode="control_ref" \ --control_ref_image="random" \ --add_full_ref_image_in_self_attention \ + --rank=64 \ + --network_alpha=32 \ + --target_name="q,k,v,ffn.0,ffn.2" \ + --use_peft_lora \ --low_vram ``` @@ -343,7 +363,7 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con --low_vram ``` -Wan-Fun-Control-V1.0 with deepspeed zero-3: +Wan-Fun-Control-V1.0 with DeepSpeed Zero-3: Wan with DeepSpeed Zero-3 is suitable for 14B Wan at high resolutions. You must set save_state to True to save the model. After training, you can use the following command to get the final model: ```sh diff --git a/scripts/wan2.1_fun/README_TRAIN_LORA.md b/scripts/wan2.1_fun/README_TRAIN_LORA.md index e76f655..9673306 100755 --- a/scripts/wan2.1_fun/README_TRAIN_LORA.md +++ b/scripts/wan2.1_fun/README_TRAIN_LORA.md @@ -19,6 +19,21 @@ Some parameters in the sh file can be confusing, and they are explained in this - These resolutions combined with their corresponding lengths allow the model to generate videos of different sizes. - `train_mode` is used to specify the training mode, which can be either normal or inpaint. Since Wan uses the inpaint model to achieve image-to-video generation, the default is set to inpaint mode. If you only wish to achieve text-to-video generation, you can remove this line, and it will default to the text-to-video mode. - `resume_from_checkpoint` is used to set the training should be resumed from a previous checkpoint. Use a path or `"latest"` to automatically select the last available checkpoint. +- `target_name` represents the components/modules to which LoRA will be applied, separated by commas. +- `use_peft_lora` indicates whether to use the PEFT module for adding LoRA. Using this module will be more memory-efficient. +- `rank` means the dimension of the LoRA update matrices. +- `network_alpha` means the scale of the LoRA update matrices. + +When train model with multi machines, please set the params as follows: +```sh +export MASTER_ADDR="your master address" +export MASTER_PORT=10086 +export WORLD_SIZE=1 # The number of machines +export NUM_PROCESS=8 # The number of processes, such as WORLD_SIZE * 8 +export RANK=0 # The rank of this machine + +accelerate launch --mixed_precision="bf16" --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK scripts/xxx/xxx.py +``` Wan T2V without deepspeed: @@ -62,12 +77,16 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.1_fun/train_lora.py \ --enable_bucket \ --uniform_sampling \ --train_mode="inpaint" \ + --rank=64 \ + --network_alpha=32 \ + --target_name="q,k,v,ffn.0,ffn.2" \ + --use_peft_lora \ --low_vram ``` -Wan T2V with deepspeed zero-2: +Wan T2V with Deepspeed Zero-2: -Wan with DeepSpeed Zero-2 is suitable for training 1.3B Wan and 14B Wan at low resolutions, but training 14B Wan at high resolutions may still result in insufficient GPU memory. +Wan with Deepspeed Zero-2 is suitable for training 1.3B Wan and 14B Wan at low resolutions, but training 14B Wan at high resolutions may still result in insufficient GPU memory. ```sh export MODEL_NAME="models/Diffusion_Transformer/Wan2.1-Fun-V1.1-14B-InP" @@ -109,10 +128,18 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con --uniform_sampling \ --use_deepspeed \ --train_mode="inpaint" \ + --rank=64 \ + --network_alpha=32 \ + --target_name="q,k,v,ffn.0,ffn.2" \ + --use_peft_lora \ --low_vram ``` -Wan T2V with deepspeed zero-3: +DeepSpeed Zero-3 is not highly recommended at the moment. In this repository, using FSDP has fewer errors and is more stable. + +It is known that DeepSpeed Zero-3 is not compatible with PEFT. + +Wan T2V with DeepSpeed Zero-3: Wan with DeepSpeed Zero-3 is suitable for 14B Wan at high resolutions. You must set save_state to True to save the model. After training, you can use the following command to get the final model: ```sh @@ -208,5 +235,9 @@ accelerate launch --mixed_precision="bf16" --use_fsdp --fsdp_auto_wrap_policy TR --save_state \ --use_deepspeed \ --train_mode="inpaint" \ + --rank=64 \ + --network_alpha=32 \ + --target_name="q,k,v,ffn.0,ffn.2" \ + --use_peft_lora \ --low_vram ``` \ No newline at end of file diff --git a/scripts/wan2.1_fun/train.py b/scripts/wan2.1_fun/train.py index b6936b3..7c2a4fa 100755 --- a/scripts/wan2.1_fun/train.py +++ b/scripts/wan2.1_fun/train.py @@ -959,7 +959,12 @@ def main(): 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): + 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) diff --git a/scripts/wan2.1_fun/train_control.py b/scripts/wan2.1_fun/train_control.py index 688a48d..85a05a2 100755 --- a/scripts/wan2.1_fun/train_control.py +++ b/scripts/wan2.1_fun/train_control.py @@ -894,7 +894,12 @@ def main(): 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): + 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) diff --git a/scripts/wan2.1_fun/train_control_lora.py b/scripts/wan2.1_fun/train_control_lora.py index ae5387b..804c543 100755 --- a/scripts/wan2.1_fun/train_control_lora.py +++ b/scripts/wan2.1_fun/train_control_lora.py @@ -455,6 +455,9 @@ def parse_args(): default=64, help=("The dimension of the LoRA update matrices."), ) + parser.add_argument( + "--use_peft_lora", action="store_true", help="Whether or not to use peft lora." + ) parser.add_argument( "--train_text_encoder", action="store_true", @@ -632,6 +635,12 @@ def parse_args(): default=None, help=("The module is not trained in loras. "), ) + parser.add_argument( + "--target_name", + type=str, + default=None, + help=("The module is trained in loras. "), + ) args = parser.parse_args() env_local_rank = int(os.environ.get("LOCAL_RANK", -1)) @@ -816,16 +825,25 @@ def main(): clip_image_encoder.requires_grad_(False) # Lora will work with this... - network = create_network( - 1.0, - args.rank, - args.network_alpha, - text_encoder, - transformer3d, - neuron_dropout=None, - skip_name=args.lora_skip_name, - ) - network.apply_to(text_encoder, transformer3d, args.train_text_encoder and not args.training_with_video_token_length, True) + if args.use_peft_lora: + from peft import LoraConfig, inject_adapter_in_model, get_peft_model_state_dict + lora_config = LoraConfig(r=args.rank, lora_alpha=args.network_alpha, target_modules=args.target_name.split(",")) + transformer3d = inject_adapter_in_model(lora_config, transformer3d) + + network = None + else: + network = create_network( + 1.0, + args.rank, + args.network_alpha, + text_encoder, + transformer3d, + neuron_dropout=None, + target_name=args.target_name, + skip_name=args.lora_skip_name, + ) + network = network.to(weight_dtype) + network.apply_to(text_encoder, transformer3d, args.train_text_encoder and not args.training_with_video_token_length, True) if args.transformer_path is not None: print(f"From checkpoint: {args.transformer_path}") @@ -859,13 +877,14 @@ def main(): 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].to(weight_dtype) - + if args.use_peft_lora: + network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1]), accelerate_state_dict) + else: + network_state_dict = {} + for key in accelerate_state_dict: + if "network" in key: + network_state_dict[key.replace("network.", "")] = accelerate_state_dict[key].to(weight_dtype) 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: @@ -880,8 +899,18 @@ def main(): 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): + 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") + if args.use_peft_lora: + network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1]), accelerate_state_dict) + else: + network_state_dict = accelerate_state_dict + 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) @@ -893,10 +922,15 @@ def main(): 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: 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 args.use_peft_lora: + save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(models[-1]))) + else: + save_model(safetensor_save_path, accelerator.unwrap_model(models[-1])) + if not args.use_deepspeed: for _ in range(len(weights)): weights.pop() @@ -950,9 +984,14 @@ def main(): else: optimizer_cls = torch.optim.AdamW - logging.info("Add network parameters") - trainable_params = list(filter(lambda p: p.requires_grad, network.parameters())) - trainable_params_optim = network.prepare_optimizer_params(args.learning_rate / 2, args.learning_rate, args.learning_rate) + if args.use_peft_lora: + logging.info("Add peft parameters") + trainable_params = list(filter(lambda p: p.requires_grad, transformer3d.parameters())) + trainable_params_optim = list(filter(lambda p: p.requires_grad, transformer3d.parameters())) + else: + logging.info("Add network parameters") + trainable_params = list(filter(lambda p: p.requires_grad, network.parameters())) + trainable_params_optim = network.prepare_optimizer_params(args.learning_rate / 2, args.learning_rate, args.learning_rate) if args.use_came: optimizer = optimizer_cls( @@ -1326,9 +1365,13 @@ def main(): ) # Prepare everything with our `accelerator`. - if fsdp_stage != 0: + if args.use_peft_lora: + transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( + transformer3d, optimizer, train_dataloader, lr_scheduler + ) + elif fsdp_stage != 0: transformer3d.network = network - transformer3d = transformer3d.to(weight_dtype) + transformer3d = transformer3d.to(dtype=weight_dtype) transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( transformer3d, optimizer, train_dataloader, lr_scheduler ) @@ -1337,14 +1380,16 @@ def main(): network, optimizer, train_dataloader, lr_scheduler ) - if zero_stage == 3: + if zero_stage != 0 and not args.use_peft_lora: 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) - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage != 0: 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) text_encoder = shard_fn(text_encoder) @@ -1487,6 +1532,10 @@ def main(): def save_model(ckpt_file, unwrapped_nw): os.makedirs(args.output_dir, exist_ok=True) accelerator.print(f"\nsaving checkpoint: {ckpt_file}") + if isinstance(unwrapped_nw, dict): + from safetensors.torch import save_file + save_file(unwrapped_nw, ckpt_file, metadata={"format": "pt"}) + return ckpt_file unwrapped_nw.save_weights(ckpt_file, weight_dtype, None) progress_bar = tqdm( @@ -1879,9 +1928,14 @@ def main(): torch.cuda.empty_cache() torch.cuda.ipc_collect() if not args.save_state: - safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") - save_model(safetensor_save_path, accelerator.unwrap_model(network)) - logger.info(f"Saved safetensor to {safetensor_save_path}") + if args.use_peft_lora: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(transformer3d))) + logger.info(f"Saved safetensor to {safetensor_save_path}") + else: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, accelerator.unwrap_model(network)) + logger.info(f"Saved safetensor to {safetensor_save_path}") else: accelerator_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") accelerator.save_state(accelerator_save_path) @@ -1932,8 +1986,14 @@ def main(): torch.cuda.empty_cache() torch.cuda.ipc_collect() if not args.save_state: - safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") - save_model(safetensor_save_path, accelerator.unwrap_model(network)) + if args.use_peft_lora: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(transformer3d))) + logger.info(f"Saved safetensor to {safetensor_save_path}") + else: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, accelerator.unwrap_model(network)) + logger.info(f"Saved safetensor to {safetensor_save_path}") else: accelerator_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") accelerator.save_state(accelerator_save_path) diff --git a/scripts/wan2.1_fun/train_control_lora.sh b/scripts/wan2.1_fun/train_control_lora.sh index 2dd0f94..24c9de5 100755 --- a/scripts/wan2.1_fun/train_control_lora.sh +++ b/scripts/wan2.1_fun/train_control_lora.sh @@ -38,4 +38,8 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.1_fun/train_control_lora --train_mode="control_ref" \ --control_ref_image="random" \ --add_full_ref_image_in_self_attention \ + --rank=64 \ + --network_alpha=32 \ + --target_name="q,k,v,ffn.0,ffn.2" \ + --use_peft_lora \ --low_vram \ No newline at end of file diff --git a/scripts/wan2.1_fun/train_lora.py b/scripts/wan2.1_fun/train_lora.py index 8f400de..6fc114f 100755 --- a/scripts/wan2.1_fun/train_lora.py +++ b/scripts/wan2.1_fun/train_lora.py @@ -526,6 +526,9 @@ def parse_args(): default=64, help=("The dimension of the LoRA update matrices."), ) + parser.add_argument( + "--use_peft_lora", action="store_true", help="Whether or not to use peft lora." + ) parser.add_argument( "--train_text_encoder", action="store_true", @@ -687,6 +690,12 @@ def parse_args(): default=None, help=("The module is not trained in loras. "), ) + parser.add_argument( + "--target_name", + type=str, + default=None, + help=("The module is trained in loras. "), + ) args = parser.parse_args() env_local_rank = int(os.environ.get("LOCAL_RANK", -1)) @@ -873,16 +882,25 @@ def main(): clip_image_encoder.requires_grad_(False) # Lora will work with this... - network = create_network( - 1.0, - args.rank, - args.network_alpha, - text_encoder, - transformer3d, - neuron_dropout=None, - skip_name=args.lora_skip_name, - ) - network.apply_to(text_encoder, transformer3d, args.train_text_encoder and not args.training_with_video_token_length, True) + if args.use_peft_lora: + from peft import LoraConfig, inject_adapter_in_model, get_peft_model_state_dict + lora_config = LoraConfig(r=args.rank, lora_alpha=args.network_alpha, target_modules=args.target_name.split(",")) + transformer3d = inject_adapter_in_model(lora_config, transformer3d) + + network = None + else: + network = create_network( + 1.0, + args.rank, + args.network_alpha, + text_encoder, + transformer3d, + neuron_dropout=None, + target_name=args.target_name, + skip_name=args.lora_skip_name, + ) + network = network.to(weight_dtype) + network.apply_to(text_encoder, transformer3d, args.train_text_encoder and not args.training_with_video_token_length, True) if args.transformer_path is not None: print(f"From checkpoint: {args.transformer_path}") @@ -918,13 +936,14 @@ def main(): 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].to(weight_dtype) - + if args.use_peft_lora: + network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1]), accelerate_state_dict) + else: + network_state_dict = {} + for key in accelerate_state_dict: + if "network" in key: + network_state_dict[key.replace("network.", "")] = accelerate_state_dict[key].to(weight_dtype) 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: @@ -939,8 +958,18 @@ def main(): 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): + 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") + if args.use_peft_lora: + network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1]), accelerate_state_dict) + else: + network_state_dict = accelerate_state_dict + 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) @@ -952,10 +981,15 @@ def main(): 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: 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 args.use_peft_lora: + save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(models[-1]))) + else: + save_model(safetensor_save_path, accelerator.unwrap_model(models[-1])) + if not args.use_deepspeed: for _ in range(len(weights)): weights.pop() @@ -1009,9 +1043,14 @@ def main(): else: optimizer_cls = torch.optim.AdamW - logging.info("Add network parameters") - trainable_params = list(filter(lambda p: p.requires_grad, network.parameters())) - trainable_params_optim = network.prepare_optimizer_params(args.learning_rate / 2, args.learning_rate, args.learning_rate) + if args.use_peft_lora: + logging.info("Add peft parameters") + trainable_params = list(filter(lambda p: p.requires_grad, transformer3d.parameters())) + trainable_params_optim = list(filter(lambda p: p.requires_grad, transformer3d.parameters())) + else: + logging.info("Add network parameters") + trainable_params = list(filter(lambda p: p.requires_grad, network.parameters())) + trainable_params_optim = network.prepare_optimizer_params(args.learning_rate / 2, args.learning_rate, args.learning_rate) if args.use_came: optimizer = optimizer_cls( @@ -1314,9 +1353,13 @@ def main(): ) # Prepare everything with our `accelerator`. - if fsdp_stage != 0: + if args.use_peft_lora: + transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( + transformer3d, optimizer, train_dataloader, lr_scheduler + ) + elif fsdp_stage != 0: transformer3d.network = network - transformer3d = transformer3d.to(weight_dtype) + transformer3d = transformer3d.to(dtype=weight_dtype) transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( transformer3d, optimizer, train_dataloader, lr_scheduler ) @@ -1325,14 +1368,16 @@ def main(): network, optimizer, train_dataloader, lr_scheduler ) - if zero_stage == 3: + if zero_stage != 0 and not args.use_peft_lora: 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) - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage != 0: 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) text_encoder = shard_fn(text_encoder) @@ -1476,6 +1521,10 @@ def main(): def save_model(ckpt_file, unwrapped_nw): os.makedirs(args.output_dir, exist_ok=True) accelerator.print(f"\nsaving checkpoint: {ckpt_file}") + if isinstance(unwrapped_nw, dict): + from safetensors.torch import save_file + save_file(unwrapped_nw, ckpt_file, metadata={"format": "pt"}) + return ckpt_file unwrapped_nw.save_weights(ckpt_file, weight_dtype, None) progress_bar = tqdm( @@ -1836,9 +1885,14 @@ def main(): torch.cuda.empty_cache() torch.cuda.ipc_collect() if not args.save_state: - safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") - save_model(safetensor_save_path, accelerator.unwrap_model(network)) - logger.info(f"Saved safetensor to {safetensor_save_path}") + if args.use_peft_lora: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(transformer3d))) + logger.info(f"Saved safetensor to {safetensor_save_path}") + else: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, accelerator.unwrap_model(network)) + logger.info(f"Saved safetensor to {safetensor_save_path}") else: accelerator_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") accelerator.save_state(accelerator_save_path) @@ -1889,8 +1943,14 @@ def main(): torch.cuda.empty_cache() torch.cuda.ipc_collect() if not args.save_state: - safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") - save_model(safetensor_save_path, accelerator.unwrap_model(network)) + if args.use_peft_lora: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(transformer3d))) + logger.info(f"Saved safetensor to {safetensor_save_path}") + else: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, accelerator.unwrap_model(network)) + logger.info(f"Saved safetensor to {safetensor_save_path}") else: accelerator_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") accelerator.save_state(accelerator_save_path) diff --git a/scripts/wan2.1_fun/train_lora.sh b/scripts/wan2.1_fun/train_lora.sh index 4749d3a..1445d3f 100755 --- a/scripts/wan2.1_fun/train_lora.sh +++ b/scripts/wan2.1_fun/train_lora.sh @@ -36,6 +36,10 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.1_fun/train_lora.py \ --enable_bucket \ --uniform_sampling \ --train_mode="inpaint" \ + --rank=64 \ + --network_alpha=32 \ + --target_name="q,k,v,ffn.0,ffn.2" \ + --use_peft_lora \ --low_vram # # Training command for T2V @@ -75,5 +79,9 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.1_fun/train_lora.py \ # --training_with_video_token_length \ # --enable_bucket \ # --uniform_sampling \ +# --rank=64 \ +# --network_alpha=32 \ +# --target_name="q,k,v,ffn.0,ffn.2" \ +# --use_peft_lora \ # --low_vram \ # --train_mode="normal" \ No newline at end of file diff --git a/scripts/wan2.1_vace/README_TRAIN.md b/scripts/wan2.1_vace/README_TRAIN.md index 8873fde..52aeaed 100755 --- a/scripts/wan2.1_vace/README_TRAIN.md +++ b/scripts/wan2.1_vace/README_TRAIN.md @@ -150,7 +150,9 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con --trainable_modules "vace" ``` -Wan-Fun-Control with deepspeed zero-3: +DeepSpeed Zero-3 is not highly recommended at the moment. In this repository, using FSDP has fewer errors and is more stable. + +Wan-Fun-Control with DeepSpeed Zero-3: Wan with DeepSpeed Zero-3 is suitable for 14B Wan at high resolutions. After training, you can use the following command to get the final model: ```sh diff --git a/scripts/wan2.1_vace/train.py b/scripts/wan2.1_vace/train.py index 6c45fe8..b47913e 100644 --- a/scripts/wan2.1_vace/train.py +++ b/scripts/wan2.1_vace/train.py @@ -880,7 +880,12 @@ def main(): 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): + 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) diff --git a/scripts/wan2.2/README_TRAIN.md b/scripts/wan2.2/README_TRAIN.md index 2d7552a..2a49238 100755 --- a/scripts/wan2.2/README_TRAIN.md +++ b/scripts/wan2.2/README_TRAIN.md @@ -23,6 +23,16 @@ Some parameters in the sh file can be confusing, and they are explained in this - `resume_from_checkpoint` is used to set the training should be resumed from a previous checkpoint. Use a path or `"latest"` to automatically select the last available checkpoint. - `boundary_type`: The Wan2.2 series includes two distinct models that handle different noise levels, specified via the `boundary_type` parameter. `low`: Corresponds to the **low noise model** (low_noise_model). `high`: Corresponds to the **high noise model**. (high_noise_model). `full`: Corresponds to the ti2v 5B model (single mode). +When train model with multi machines, please set the params as follows: +```sh +export MASTER_ADDR="your master address" +export MASTER_PORT=10086 +export WORLD_SIZE=1 # The number of machines +export NUM_PROCESS=8 # The number of processes, such as WORLD_SIZE * 8 +export RANK=0 # The rank of this machine + +accelerate launch --mixed_precision="bf16" --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK scripts/xxx/xxx.py +``` Wan2.2 T2V without deepspeed: @@ -73,9 +83,9 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.2/train.py \ --trainable_modules "." ``` -Wan T2V with deepspeed zero-2: +Wan T2V with Deepspeed Zero-2: -Wan with DeepSpeed Zero-2 is suitable for training Wan at low resolutions, but training 14B Wan at high resolutions may still result in insufficient GPU memory. +Wan with Deepspeed Zero-2 is suitable for training Wan at low resolutions, but training 14B Wan at high resolutions may still result in insufficient GPU memory. ```sh export MODEL_NAME="models/Diffusion_Transformer/Wan2.2-T2V-A14B" @@ -124,7 +134,9 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con --trainable_modules "." ``` -Wan T2V with deepspeed zero-3: +DeepSpeed Zero-3 is not highly recommended at the moment. In this repository, using FSDP has fewer errors and is more stable. + +Wan T2V with DeepSpeed Zero-3: Wan with DeepSpeed Zero-3 is suitable for 14B Wan at high resolutions. After training, you can use the following command to get the final model: ```sh diff --git a/scripts/wan2.2/README_TRAIN_ANIMATE.md b/scripts/wan2.2/README_TRAIN_ANIMATE.md index a00d1ec..e46b47e 100755 --- a/scripts/wan2.2/README_TRAIN_ANIMATE.md +++ b/scripts/wan2.2/README_TRAIN_ANIMATE.md @@ -52,6 +52,17 @@ Some parameters in the sh file can be confusing, and they are explained in this - These resolutions combined with their corresponding lengths allow the model to generate videos of different sizes. - `resume_from_checkpoint` is used to set the training should be resumed from a previous checkpoint. Use a path or `"latest"` to automatically select the last available checkpoint. +When train model with multi machines, please set the params as follows: +```sh +export MASTER_ADDR="your master address" +export MASTER_PORT=10086 +export WORLD_SIZE=1 # The number of machines +export NUM_PROCESS=8 # The number of processes, such as WORLD_SIZE * 8 +export RANK=0 # The rank of this machine + +accelerate launch --mixed_precision="bf16" --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK scripts/xxx/xxx.py +``` + Wan-Animate without deepspeed: ```sh @@ -99,7 +110,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.2/train_animate.py \ --trainable_modules "." ``` -Wan-Animate with deepspeed zero-2: +Wan-Animate with Deepspeed Zero-2: ```sh export MODEL_NAME="models/Diffusion_Transformer/Wan2.2-Animate-14B/" @@ -146,7 +157,9 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con --trainable_modules "." ``` -Wan-Animate with deepspeed zero-3: +DeepSpeed Zero-3 is not highly recommended at the moment. In this repository, using FSDP has fewer errors and is more stable. + +Wan-Animate with DeepSpeed Zero-3: ```sh python scripts/zero_to_bf16.py output_dir/checkpoint-{our-num-steps} output_dir/checkpoint-{your-num-steps}-outputs --max_shard_size 80GB --safe_serialization diff --git a/scripts/wan2.2/README_TRAIN_DISTILL.md b/scripts/wan2.2/README_TRAIN_DISTILL.md index 7fffaba..22f7fab 100755 --- a/scripts/wan2.2/README_TRAIN_DISTILL.md +++ b/scripts/wan2.2/README_TRAIN_DISTILL.md @@ -23,6 +23,17 @@ Some parameters in the sh file can be confusing, and they are explained in this - `resume_from_checkpoint` is used to set the training should be resumed from a previous checkpoint. Use a path or `"latest"` to automatically select the last available checkpoint. - `boundary_type`: The Wan2.2 series includes two distinct models that handle different noise levels, specified via the `boundary_type` parameter. `low`: Corresponds to the **low noise model** (low_noise_model). `high`: Corresponds to the **high noise model**. (high_noise_model). `full`: Corresponds to the ti2v 5B model (single mode). +When train model with multi machines, please set the params as follows: +```sh +export MASTER_ADDR="your master address" +export MASTER_PORT=10086 +export WORLD_SIZE=1 # The number of machines +export NUM_PROCESS=8 # The number of processes, such as WORLD_SIZE * 8 +export RANK=0 # The rank of this machine + +accelerate launch --mixed_precision="bf16" --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK scripts/xxx/xxx.py +``` + Wan distill without deepspeed: Wan distill without DeepSpeed and FSDP is more suitable for 1.3B Wan, as using it with 14B Wan may result in insufficient GPU memory. @@ -73,9 +84,9 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.2/train_distill.py \ --low_vram ``` -Wan distill with deepspeed zero-2: +Wan distill with Deepspeed Zero-2: -Wan with DeepSpeed Zero-2 is suitable for training 1.3B Wan and 14B Wan at low resolutions, but training 14B Wan at high resolutions may still result in insufficient GPU memory. +Wan with Deepspeed Zero-2 is suitable for training 1.3B Wan and 14B Wan at low resolutions, but training 14B Wan at high resolutions may still result in insufficient GPU memory. ```sh export MODEL_NAME="models/Diffusion_Transformer/Wan2.2-I2V-A14B" @@ -124,7 +135,9 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con --low_vram ``` -Wan distill with deepspeed zero-3: +DeepSpeed Zero-3 is not highly recommended at the moment. In this repository, using FSDP has fewer errors and is more stable. + +Wan distill with DeepSpeed Zero-3: Wan with DeepSpeed Zero-3 is suitable for 14B Wan at high resolutions. After training, you can use the following command to get the final model: ```sh diff --git a/scripts/wan2.2/README_TRAIN_DISTILL_LORA.md b/scripts/wan2.2/README_TRAIN_DISTILL_LORA.md index 6a0adb9..053a402 100755 --- a/scripts/wan2.2/README_TRAIN_DISTILL_LORA.md +++ b/scripts/wan2.2/README_TRAIN_DISTILL_LORA.md @@ -22,6 +22,21 @@ Some parameters in the sh file can be confusing, and they are explained in this - `train_mode` is used to specify the training mode, which can be either normal or i2v. Since Wan uses the inpaint model to achieve image-to-video generation, the default is set to inpaint mode. If you only wish to achieve text-to-video generation, you can remove this line, and it will default to the text-to-video mode. - `resume_from_checkpoint` is used to set the training should be resumed from a previous checkpoint. Use a path or `"latest"` to automatically select the last available checkpoint. - `boundary_type`: The Wan2.2 series includes two distinct models that handle different noise levels, specified via the `boundary_type` parameter. `low`: Corresponds to the **low noise model** (low_noise_model). `high`: Corresponds to the **high noise model**. (high_noise_model). `full`: Corresponds to the ti2v 5B model (single mode). +- `target_name` represents the components/modules to which LoRA will be applied, separated by commas. +- `use_peft_lora` indicates whether to use the PEFT module for adding LoRA. Using this module will be more memory-efficient. +- `rank` means the dimension of the LoRA update matrices. +- `network_alpha` means the scale of the LoRA update matrices. + +When train model with multi machines, please set the params as follows: +```sh +export MASTER_ADDR="your master address" +export MASTER_PORT=10086 +export WORLD_SIZE=1 # The number of machines +export NUM_PROCESS=8 # The number of processes, such as WORLD_SIZE * 8 +export RANK=0 # The rank of this machine + +accelerate launch --mixed_precision="bf16" --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK scripts/xxx/xxx.py +``` Wan distill without deepspeed: @@ -67,12 +82,16 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.2/train_distill_lora.py --uniform_sampling \ --boundary_type="low" \ --train_mode="i2v" \ + --rank=64 \ + --network_alpha=32 \ + --target_name="q,k,v,ffn.0,ffn.2" \ + --use_peft_lora \ --low_vram ``` -Wan distill with deepspeed zero-2: +Wan distill with Deepspeed Zero-2: -Wan with DeepSpeed Zero-2 is suitable for training 1.3B Wan and 14B Wan at low resolutions, but training 14B Wan at high resolutions may still result in insufficient GPU memory. +Wan with Deepspeed Zero-2 is suitable for training 1.3B Wan and 14B Wan at low resolutions, but training 14B Wan at high resolutions may still result in insufficient GPU memory. ```sh export MODEL_NAME="models/Diffusion_Transformer/Wan2.2-I2V-A14B" @@ -115,10 +134,18 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con --uniform_sampling \ --boundary_type="low" \ --train_mode="i2v" \ + --rank=64 \ + --network_alpha=32 \ + --target_name="q,k,v,ffn.0,ffn.2" \ + --use_peft_lora \ --low_vram ``` -Wan distill with deepspeed zero-3: +DeepSpeed Zero-3 is not highly recommended at the moment. In this repository, using FSDP has fewer errors and is more stable. + +It is known that DeepSpeed Zero-3 is not compatible with PEFT. + +Wan distill with DeepSpeed Zero-3: Wan with DeepSpeed Zero-3 is suitable for 14B Wan at high resolutions. After training, you can use the following command to get the final model: ```sh @@ -214,5 +241,9 @@ accelerate launch --mixed_precision="bf16" --use_fsdp --fsdp_auto_wrap_policy TR --uniform_sampling \ --boundary_type="low" \ --train_mode="i2v" \ + --rank=64 \ + --network_alpha=32 \ + --target_name="q,k,v,ffn.0,ffn.2" \ + --use_peft_lora \ --low_vram ``` \ No newline at end of file diff --git a/scripts/wan2.2/README_TRAIN_LORA.md b/scripts/wan2.2/README_TRAIN_LORA.md index ecd77c5..922b6ae 100755 --- a/scripts/wan2.2/README_TRAIN_LORA.md +++ b/scripts/wan2.2/README_TRAIN_LORA.md @@ -20,7 +20,21 @@ Some parameters in the sh file can be confusing, and they are explained in this - `train_mode` is used to specify the training mode, which can be either normal, i2v or ti2v. The t2v is used for 14B T2V model. The i2v is used for 14B I2V model. The ti2v is used in 5B TI2V model. - `resume_from_checkpoint` is used to set the training should be resumed from a previous checkpoint. Use a path or `"latest"` to automatically select the last available checkpoint. - `boundary_type`: The Wan2.2 series includes two distinct models that handle different noise levels, specified via the `boundary_type` parameter. `low`: Corresponds to the **low noise model** (low_noise_model). `high`: Corresponds to the **high noise model**. (high_noise_model). `full`: Corresponds to the ti2v 5B model (single mode). +- `target_name` represents the components/modules to which LoRA will be applied, separated by commas. +- `use_peft_lora` indicates whether to use the PEFT module for adding LoRA. Using this module will be more memory-efficient. +- `rank` means the dimension of the LoRA update matrices. +- `network_alpha` means the scale of the LoRA update matrices. +When train model with multi machines, please set the params as follows: +```sh +export MASTER_ADDR="your master address" +export MASTER_PORT=10086 +export WORLD_SIZE=1 # The number of machines +export NUM_PROCESS=8 # The number of processes, such as WORLD_SIZE * 8 +export RANK=0 # The rank of this machine + +accelerate launch --mixed_precision="bf16" --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK scripts/xxx/xxx.py +``` Wan2.2 T2V without deepspeed: @@ -64,13 +78,17 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.2/train_lora.py \ --enable_bucket \ --uniform_sampling \ --boundary_type="low" \ + --rank=64 \ + --network_alpha=32 \ + --target_name="q,k,v,ffn.0,ffn.2" \ + --use_peft_lora \ --train_mode="normal" \ --low_vram ``` -Wan T2V with deepspeed zero-2: +Wan T2V with Deepspeed Zero-2: -Wan with DeepSpeed Zero-2 is suitable for training Wan at low resolutions, but training 14B Wan at high resolutions may still result in insufficient GPU memory. +Wan with Deepspeed Zero-2 is suitable for training Wan at low resolutions, but training 14B Wan at high resolutions may still result in insufficient GPU memory. ```sh export MODEL_NAME="models/Diffusion_Transformer/Wan2.2-T2V-A14B" @@ -111,12 +129,19 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con --enable_bucket \ --uniform_sampling \ --boundary_type="low" \ + --rank=64 \ + --network_alpha=32 \ + --target_name="q,k,v,ffn.0,ffn.2" \ + --use_peft_lora \ --train_mode="normal" \ - --use_deepspeed \ - --low_vram + --low_vram ``` -Wan T2V with deepspeed zero-3: +DeepSpeed Zero-3 is not highly recommended at the moment. In this repository, using FSDP has fewer errors and is more stable. + +It is known that DeepSpeed Zero-3 is not compatible with PEFT. + +Wan T2V with DeepSpeed Zero-3: Wan with DeepSpeed Zero-3 is suitable for 14B Wan at high resolutions. You must set save_state to True to save the model. After training, you can use the following command to get the final model: ```sh @@ -164,7 +189,6 @@ accelerate launch --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag --uniform_sampling \ --boundary_type="low" \ --train_mode="normal" \ - --use_deepspeed \ --low_vram ``` @@ -210,6 +234,10 @@ accelerate launch --mixed_precision="bf16" --use_fsdp --fsdp_auto_wrap_policy TR --enable_bucket \ --uniform_sampling \ --boundary_type="low" \ + --rank=64 \ + --network_alpha=32 \ + --target_name="q,k,v,ffn.0,ffn.2" \ + --use_peft_lora \ --train_mode="normal" \ --low_vram ``` @@ -255,6 +283,10 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.2/train_lora.py \ --enable_bucket \ --uniform_sampling \ --boundary_type="full" \ + --rank=64 \ + --network_alpha=32 \ + --target_name="q,k,v,ffn.0,ffn.2" \ + --use_peft_lora \ --train_mode="ti2v" \ --low_vram ``` \ No newline at end of file diff --git a/scripts/wan2.2/README_TRAIN_S2V.md b/scripts/wan2.2/README_TRAIN_S2V.md index 1260f4c..4c377b5 100755 --- a/scripts/wan2.2/README_TRAIN_S2V.md +++ b/scripts/wan2.2/README_TRAIN_S2V.md @@ -34,6 +34,17 @@ Some parameters in the sh file can be confusing, and they are explained in this - These resolutions combined with their corresponding lengths allow the model to generate videos of different sizes. - `resume_from_checkpoint` is used to set the training should be resumed from a previous checkpoint. Use a path or `"latest"` to automatically select the last available checkpoint. +When train model with multi machines, please set the params as follows: +```sh +export MASTER_ADDR="your master address" +export MASTER_PORT=10086 +export WORLD_SIZE=1 # The number of machines +export NUM_PROCESS=8 # The number of processes, such as WORLD_SIZE * 8 +export RANK=0 # The rank of this machine + +accelerate launch --mixed_precision="bf16" --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK scripts/xxx/xxx.py +``` + Wan-S2V without deepspeed: ```sh @@ -78,7 +89,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.2/train_s2v.py \ --trainable_modules "." ``` -Wan-S2V with deepspeed zero-2: +Wan-S2V with Deepspeed Zero-2: ```sh export MODEL_NAME="models/Diffusion_Transformer/Wan2.2-S2V-14B" @@ -122,7 +133,9 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con --trainable_modules "." ``` -Wan-S2V with deepspeed zero-3: +DeepSpeed Zero-3 is not highly recommended at the moment. In this repository, using FSDP has fewer errors and is more stable. + +Wan-S2V with DeepSpeed Zero-3: ```sh python scripts/zero_to_bf16.py output_dir/checkpoint-{our-num-steps} output_dir/checkpoint-{your-num-steps}-outputs --max_shard_size 80GB --safe_serialization diff --git a/scripts/wan2.2/train.py b/scripts/wan2.2/train.py index 2e84d71..e42adf9 100644 --- a/scripts/wan2.2/train.py +++ b/scripts/wan2.2/train.py @@ -1033,7 +1033,12 @@ def main(): 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): + 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) diff --git a/scripts/wan2.2/train_animate.py b/scripts/wan2.2/train_animate.py index 6b988f7..aa5874d 100644 --- a/scripts/wan2.2/train_animate.py +++ b/scripts/wan2.2/train_animate.py @@ -966,7 +966,12 @@ def main(): 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): + 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) diff --git a/scripts/wan2.2/train_animate_lora.py b/scripts/wan2.2/train_animate_lora.py index cbd7976..98b9086 100644 --- a/scripts/wan2.2/train_animate_lora.py +++ b/scripts/wan2.2/train_animate_lora.py @@ -539,6 +539,9 @@ def parse_args(): default=64, help=("The dimension of the LoRA update matrices."), ) + parser.add_argument( + "--use_peft_lora", action="store_true", help="Whether or not to use peft lora." + ) parser.add_argument( "--train_text_encoder", action="store_true", @@ -699,6 +702,12 @@ def parse_args(): default=None, help=("The module is not trained in loras. "), ) + parser.add_argument( + "--target_name", + type=str, + default=None, + help=("The module is trained in loras. "), + ) args = parser.parse_args() env_local_rank = int(os.environ.get("LOCAL_RANK", -1)) @@ -891,17 +900,25 @@ def main(): clip_image_encoder.requires_grad_(False) # Lora will work with this... - network = create_network( - 1.0, - args.rank, - args.network_alpha, - text_encoder, - transformer3d, - neuron_dropout=None, - skip_name=args.lora_skip_name, - ) - network = network.to(weight_dtype) - network.apply_to(text_encoder, transformer3d, args.train_text_encoder and not args.training_with_video_token_length, True) + if args.use_peft_lora: + from peft import LoraConfig, inject_adapter_in_model, get_peft_model_state_dict + lora_config = LoraConfig(r=args.rank, lora_alpha=args.network_alpha, target_modules=args.target_name.split(",")) + transformer3d = inject_adapter_in_model(lora_config, transformer3d) + + network = None + else: + network = create_network( + 1.0, + args.rank, + args.network_alpha, + text_encoder, + transformer3d, + neuron_dropout=None, + target_name=args.target_name, + skip_name=args.lora_skip_name, + ) + network = network.to(weight_dtype) + network.apply_to(text_encoder, transformer3d, args.train_text_encoder and not args.training_with_video_token_length, True) if args.transformer_path is not None: print(f"From checkpoint: {args.transformer_path}") @@ -927,7 +944,6 @@ def main(): m, u = vae.load_state_dict(state_dict, strict=False) print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") - # `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 @@ -936,13 +952,14 @@ def main(): 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].to(weight_dtype) - + if args.use_peft_lora: + network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1]), accelerate_state_dict) + else: + network_state_dict = {} + for key in accelerate_state_dict: + if "network" in key: + network_state_dict[key.replace("network.", "")] = accelerate_state_dict[key].to(weight_dtype) 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: @@ -957,8 +974,18 @@ def main(): 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): + 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") + if args.use_peft_lora: + network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1]), accelerate_state_dict) + else: + network_state_dict = accelerate_state_dict + 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) @@ -970,10 +997,15 @@ def main(): 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: 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 args.use_peft_lora: + save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(models[-1]))) + else: + save_model(safetensor_save_path, accelerator.unwrap_model(models[-1])) + if not args.use_deepspeed: for _ in range(len(weights)): weights.pop() @@ -1027,9 +1059,14 @@ def main(): else: optimizer_cls = torch.optim.AdamW - logging.info("Add network parameters") - trainable_params = list(filter(lambda p: p.requires_grad, network.parameters())) - trainable_params_optim = network.prepare_optimizer_params(args.learning_rate / 2, args.learning_rate, args.learning_rate) + if args.use_peft_lora: + logging.info("Add peft parameters") + trainable_params = list(filter(lambda p: p.requires_grad, transformer3d.parameters())) + trainable_params_optim = list(filter(lambda p: p.requires_grad, transformer3d.parameters())) + else: + logging.info("Add network parameters") + trainable_params = list(filter(lambda p: p.requires_grad, network.parameters())) + trainable_params_optim = network.prepare_optimizer_params(args.learning_rate / 2, args.learning_rate, args.learning_rate) if args.use_came: optimizer = optimizer_cls( @@ -1355,9 +1392,13 @@ def main(): ) # Prepare everything with our `accelerator`. - if fsdp_stage != 0: + if args.use_peft_lora: + transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( + transformer3d, optimizer, train_dataloader, lr_scheduler + ) + elif fsdp_stage != 0: transformer3d.network = network - transformer3d = transformer3d.to(weight_dtype) + transformer3d = transformer3d.to(dtype=weight_dtype) transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( transformer3d, optimizer, train_dataloader, lr_scheduler ) @@ -1366,14 +1407,14 @@ def main(): network, optimizer, train_dataloader, lr_scheduler ) - if zero_stage == 3: + if zero_stage != 0 and not args.use_peft_lora: 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) - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage != 0: from functools import partial from videox_fun.dist import set_multi_gpus_devices, shard_model @@ -1462,6 +1503,10 @@ def main(): def save_model(ckpt_file, unwrapped_nw): os.makedirs(args.output_dir, exist_ok=True) accelerator.print(f"\nsaving checkpoint: {ckpt_file}") + if isinstance(unwrapped_nw, dict): + from safetensors.torch import save_file + save_file(unwrapped_nw, ckpt_file, metadata={"format": "pt"}) + return ckpt_file unwrapped_nw.save_weights(ckpt_file, weight_dtype, None) progress_bar = tqdm( @@ -1878,9 +1923,14 @@ def main(): torch.cuda.empty_cache() torch.cuda.ipc_collect() if not args.save_state: - safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") - save_model(safetensor_save_path, accelerator.unwrap_model(network)) - logger.info(f"Saved safetensor to {safetensor_save_path}") + if args.use_peft_lora: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(transformer3d))) + logger.info(f"Saved safetensor to {safetensor_save_path}") + else: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, accelerator.unwrap_model(network)) + logger.info(f"Saved safetensor to {safetensor_save_path}") else: accelerator_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") accelerator.save_state(accelerator_save_path) @@ -1928,8 +1978,14 @@ def main(): torch.cuda.empty_cache() torch.cuda.ipc_collect() if not args.save_state: - safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") - save_model(safetensor_save_path, accelerator.unwrap_model(network)) + if args.use_peft_lora: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(transformer3d))) + logger.info(f"Saved safetensor to {safetensor_save_path}") + else: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, accelerator.unwrap_model(network)) + logger.info(f"Saved safetensor to {safetensor_save_path}") else: accelerator_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") accelerator.save_state(accelerator_save_path) diff --git a/scripts/wan2.2/train_animate_lora.sh b/scripts/wan2.2/train_animate_lora.sh index e4ed2fd..e38a4b1 100644 --- a/scripts/wan2.2/train_animate_lora.sh +++ b/scripts/wan2.2/train_animate_lora.sh @@ -35,4 +35,8 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.2/train_animate_lora.py --enable_bucket \ --uniform_sampling \ --boundary_type="full" \ + --rank=64 \ + --network_alpha=32 \ + --target_name="q,k,v,ffn.0,ffn.2" \ + --use_peft_lora \ --low_vram \ No newline at end of file diff --git a/scripts/wan2.2/train_distill.py b/scripts/wan2.2/train_distill.py index ffb7724..9863483 100644 --- a/scripts/wan2.2/train_distill.py +++ b/scripts/wan2.2/train_distill.py @@ -1050,7 +1050,12 @@ def main(): 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): + 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) diff --git a/scripts/wan2.2/train_distill_lora.py b/scripts/wan2.2/train_distill_lora.py index 560d86a..f7d21d6 100644 --- a/scripts/wan2.2/train_distill_lora.py +++ b/scripts/wan2.2/train_distill_lora.py @@ -619,6 +619,9 @@ def parse_args(): default=64, help=("The dimension of the LoRA update matrices."), ) + parser.add_argument( + "--use_peft_lora", action="store_true", help="Whether or not to use peft lora." + ) parser.add_argument( "--train_text_encoder", action="store_true", @@ -793,6 +796,12 @@ def parse_args(): default=None, help=("The module is not trained in loras. "), ) + parser.add_argument( + "--target_name", + type=str, + default=None, + help=("The module is trained in loras. "), + ) args = parser.parse_args() env_local_rank = int(os.environ.get("LOCAL_RANK", -1)) @@ -997,27 +1006,40 @@ def main(): fake_score_transformer3d.requires_grad_(False) # Lora will work with this... - network = create_network( - 1.0, - args.rank, - args.network_alpha, - text_encoder, - generator_transformer3d, - neuron_dropout=None, - skip_name=args.lora_skip_name, - ) - network.apply_to(text_encoder, generator_transformer3d, args.train_text_encoder and not args.training_with_video_token_length, True) + if args.use_peft_lora: + from peft import LoraConfig, inject_adapter_in_model, get_peft_model_state_dict + lora_config = LoraConfig(r=args.rank, lora_alpha=args.network_alpha, target_modules=args.target_name.split(",")) + generator_transformer3d = inject_adapter_in_model(lora_config, generator_transformer3d) - fake_score_network = create_network( - 1.0, - args.rank, - args.network_alpha, - None, - fake_score_transformer3d, - neuron_dropout=None, - skip_name=args.lora_skip_name, - ) - fake_score_network.apply_to(None, fake_score_transformer3d, False, True) + fake_score_lora_config = LoraConfig(r=args.rank, lora_alpha=args.network_alpha, target_modules=args.target_name.split(",")) + fake_score_transformer3d = inject_adapter_in_model(fake_score_lora_config, fake_score_transformer3d) + + network = None + fake_score_network = None + else: + network = create_network( + 1.0, + args.rank, + args.network_alpha, + text_encoder, + generator_transformer3d, + neuron_dropout=None, + target_name=args.target_name, + skip_name=args.lora_skip_name, + ) + network.apply_to(text_encoder, generator_transformer3d, args.train_text_encoder and not args.training_with_video_token_length, True) + + fake_score_network = create_network( + 1.0, + args.rank, + args.network_alpha, + None, + fake_score_transformer3d, + neuron_dropout=None, + target_name=args.target_name, + skip_name=args.lora_skip_name, + ) + fake_score_network.apply_to(None, fake_score_transformer3d, False, True) if args.transformer_path is not None: print(f"From checkpoint: {args.transformer_path}") @@ -1053,13 +1075,14 @@ def main(): 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].to(weight_dtype) - + if args.use_peft_lora: + network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1]), accelerate_state_dict) + else: + network_state_dict = {} + for key in accelerate_state_dict: + if "network" in key: + network_state_dict[key.replace("network.", "")] = accelerate_state_dict[key].to(weight_dtype) 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: @@ -1076,7 +1099,16 @@ def main(): 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): + 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") + if args.use_peft_lora: + network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1]), accelerate_state_dict) + else: + network_state_dict = accelerate_state_dict + 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) @@ -1092,7 +1124,11 @@ def main(): 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 args.use_peft_lora: + save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(models[-1]))) + else: + save_model(safetensor_save_path, accelerator.unwrap_model(models[-1])) + if not args.use_deepspeed: for _ in range(len(weights)): weights.pop() @@ -1117,6 +1153,7 @@ def main(): if args.gradient_checkpointing: generator_transformer3d.enable_gradient_checkpointing() fake_score_transformer3d.enable_gradient_checkpointing() + real_score_transformer3d.enable_gradient_checkpointing() # Enable TF32 for faster training on Ampere GPUs, # cf https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices @@ -1150,14 +1187,22 @@ def main(): else: optimizer_cls = torch.optim.AdamW + if args.use_peft_lora: + logging.info("Add peft parameters") + trainable_params = list(filter(lambda p: p.requires_grad, generator_transformer3d.parameters())) + trainable_params_optim = list(filter(lambda p: p.requires_grad, generator_transformer3d.parameters())) - logging.info("Add network parameters") - trainable_params = list(filter(lambda p: p.requires_grad, network.parameters())) - trainable_params_optim = network.prepare_optimizer_params(args.learning_rate / 2, args.learning_rate, args.learning_rate) + logging.info("Add fake score peft parameters") + fake_trainable_params = list(filter(lambda p: p.requires_grad, fake_score_transformer3d.parameters())) + fake_trainable_params_optim = list(filter(lambda p: p.requires_grad, fake_score_transformer3d.parameters())) + else: + logging.info("Add network parameters") + trainable_params = list(filter(lambda p: p.requires_grad, network.parameters())) + trainable_params_optim = network.prepare_optimizer_params(args.learning_rate / 2, args.learning_rate, args.learning_rate) - logging.info("Add fake_score_network parameters") - fake_trainable_params = list(filter(lambda p: p.requires_grad, fake_score_network.parameters())) - fake_trainable_params_optim = fake_score_network.prepare_optimizer_params(args.learning_rate / 2, args.learning_rate, args.learning_rate) + logging.info("Add fake_score_network parameters") + fake_trainable_params = list(filter(lambda p: p.requires_grad, fake_score_network.parameters())) + fake_trainable_params_optim = fake_score_network.prepare_optimizer_params(args.learning_rate / 2, args.learning_rate, args.learning_rate) if args.use_came: optimizer = optimizer_cls( @@ -1533,14 +1578,21 @@ def main(): ) # Prepare everything with our `accelerator`. - if fsdp_stage != 0: + if args.use_peft_lora: + generator_transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( + generator_transformer3d, optimizer, train_dataloader, lr_scheduler + ) + fake_score_transformer3d, critic_optimizer, fake_score_lr_scheduler = accelerator_fake_score_transformer3d.prepare( + fake_score_transformer3d, critic_optimizer, fake_score_lr_scheduler + ) + elif fsdp_stage != 0: generator_transformer3d.network = network - generator_transformer3d = generator_transformer3d.to(weight_dtype) + generator_transformer3d = generator_transformer3d.to(dtype=weight_dtype) generator_transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( generator_transformer3d, optimizer, train_dataloader, lr_scheduler ) fake_score_transformer3d.network = fake_score_network - fake_score_transformer3d = fake_score_transformer3d.to(weight_dtype) + fake_score_transformer3d = fake_score_transformer3d.to(dtype=weight_dtype) fake_score_transformer3d, critic_optimizer, fake_score_lr_scheduler = accelerator_fake_score_transformer3d.prepare( fake_score_transformer3d, critic_optimizer, fake_score_lr_scheduler ) @@ -1552,33 +1604,26 @@ def main(): fake_score_network, critic_optimizer, fake_score_lr_scheduler ) - if zero_stage == 3: + if zero_stage != 0 and not args.use_peft_lora: 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) generator_transformer3d = shard_fn(generator_transformer3d) + fake_score_transformer3d = shard_fn(fake_score_transformer3d) - 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) - fake_score_network = shard_fn(fake_score_network) - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage != 0: 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) real_score_transformer3d = shard_fn(real_score_transformer3d) - if fsdp_stage != 0: - 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) text_encoder = shard_fn(text_encoder) # Move text_encode and vae to gpu and cast to weight_dtype vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + generator_transformer3d.to(accelerator.device, dtype=weight_dtype) + fake_score_transformer3d.to(accelerator.device, dtype=weight_dtype) real_score_transformer3d.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) if not args.enable_text_encoder_in_dataloader: text_encoder.to(accelerator.device if not args.low_vram else "cpu") @@ -1656,6 +1701,16 @@ def main(): else: initial_global_step = 0 + # function for saving/removing + def save_model(ckpt_file, unwrapped_nw): + os.makedirs(args.output_dir, exist_ok=True) + accelerator.print(f"\nsaving checkpoint: {ckpt_file}") + if isinstance(unwrapped_nw, dict): + from safetensors.torch import save_file + save_file(unwrapped_nw, ckpt_file, metadata={"format": "pt"}) + return ckpt_file + unwrapped_nw.save_weights(ckpt_file, weight_dtype, None) + progress_bar = tqdm( range(0, args.max_train_steps), initial=initial_global_step, @@ -2302,9 +2357,13 @@ def main(): return final_loss denoising_loss = custom_mse_loss(fake_score_denoised_output, critic_noise - fake_score_denoised_pred) - avg_denoising_loss = accelerator.gather(denoising_loss.repeat(args.train_batch_size)).mean() + avg_denoising_loss = accelerator_fake_score_transformer3d.gather(denoising_loss.repeat(args.train_batch_size)).mean() train_denoising_loss += avg_denoising_loss.item() / args.gradient_accumulation_steps - + + if args.low_vram: + generator_transformer3d = generator_transformer3d.to("cpu") + torch.cuda.empty_cache() + accelerator_fake_score_transformer3d.backward(denoising_loss) if accelerator_fake_score_transformer3d.sync_gradients: accelerator_fake_score_transformer3d.clip_grad_norm_(fake_trainable_params, args.max_grad_norm) @@ -2351,13 +2410,22 @@ def main(): torch.cuda.empty_cache() torch.cuda.ipc_collect() if not args.save_state: - safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") - save_model(safetensor_save_path, accelerator.unwrap_model(network)) - logger.info(f"Saved safetensor to {safetensor_save_path}") + if args.use_peft_lora: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(generator_transformer3d))) + logger.info(f"Saved safetensor to {safetensor_save_path}") - safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}-fake_score.safetensors") - save_model(safetensor_save_path, accelerator.unwrap_model(fake_score_network)) - logger.info(f"Saved safetensor to {safetensor_save_path}") + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}-fake_score.safetensors") + save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(fake_score_transformer3d))) + logger.info(f"Saved safetensor to {safetensor_save_path}") + else: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, accelerator.unwrap_model(network)) + logger.info(f"Saved safetensor to {safetensor_save_path}") + + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}-fake_score.safetensors") + save_model(safetensor_save_path, accelerator.unwrap_model(fake_score_network)) + logger.info(f"Saved safetensor to {safetensor_save_path}") else: save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") fake_score_save_path = os.path.join(save_path, "fake_score") @@ -2405,16 +2473,35 @@ def main(): accelerator.wait_for_everyone() if accelerator.is_main_process: generator_transformer3d = unwrap_model(generator_transformer3d) + fake_score_transformer3d = unwrap_model(fake_score_transformer3d) if args.use_deepspeed or args.use_fsdp or accelerator.is_main_process: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() - save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") - fake_score_save_path = os.path.join(save_path, "fake_score") - accelerator.save_state(save_path) - accelerator_fake_score_transformer3d.save_state(fake_score_save_path) - logger.info(f"Saved state to {save_path}") + if not args.save_state: + if args.use_peft_lora: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(generator_transformer3d))) + logger.info(f"Saved safetensor to {safetensor_save_path}") + + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}-fake_score.safetensors") + save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(fake_score_transformer3d))) + logger.info(f"Saved safetensor to {safetensor_save_path}") + else: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, accelerator.unwrap_model(network)) + logger.info(f"Saved safetensor to {safetensor_save_path}") + + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}-fake_score.safetensors") + save_model(safetensor_save_path, accelerator.unwrap_model(fake_score_network)) + logger.info(f"Saved safetensor to {safetensor_save_path}") + else: + save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") + fake_score_save_path = os.path.join(save_path, "fake_score") + accelerator.save_state(save_path) + accelerator_fake_score_transformer3d.save_state(fake_score_save_path) + logger.info(f"Saved state to {save_path}") accelerator.end_training() diff --git a/scripts/wan2.2/train_distill_lora.sh b/scripts/wan2.2/train_distill_lora.sh index 66b58e0..241daa3 100644 --- a/scripts/wan2.2/train_distill_lora.sh +++ b/scripts/wan2.2/train_distill_lora.sh @@ -37,5 +37,9 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.2/train_distill_lora.py --enable_bucket \ --uniform_sampling \ --boundary_type="low" \ + --rank=64 \ + --network_alpha=32 \ + --target_name="q,k,v,ffn.0,ffn.2" \ + --use_peft_lora \ --train_mode="i2v" \ --low_vram \ No newline at end of file diff --git a/scripts/wan2.2/train_lora.py b/scripts/wan2.2/train_lora.py index 5ca71c3..f608a4e 100755 --- a/scripts/wan2.2/train_lora.py +++ b/scripts/wan2.2/train_lora.py @@ -593,6 +593,9 @@ def parse_args(): default=64, help=("The dimension of the LoRA update matrices."), ) + parser.add_argument( + "--use_peft_lora", action="store_true", help="Whether or not to use peft lora." + ) parser.add_argument( "--train_text_encoder", action="store_true", @@ -768,6 +771,12 @@ def parse_args(): default=None, help=("The module is not trained in loras. "), ) + parser.add_argument( + "--target_name", + type=str, + default=None, + help=("The module is trained in loras. "), + ) args = parser.parse_args() env_local_rank = int(os.environ.get("LOCAL_RANK", -1)) @@ -954,16 +963,25 @@ def main(): transformer3d.requires_grad_(False) # Lora will work with this... - network = create_network( - 1.0, - args.rank, - args.network_alpha, - text_encoder, - transformer3d, - neuron_dropout=None, - skip_name=args.lora_skip_name, - ) - network.apply_to(text_encoder, transformer3d, args.train_text_encoder and not args.training_with_video_token_length, True) + if args.use_peft_lora: + from peft import LoraConfig, inject_adapter_in_model, get_peft_model_state_dict + lora_config = LoraConfig(r=args.rank, lora_alpha=args.network_alpha, target_modules=args.target_name.split(",")) + transformer3d = inject_adapter_in_model(lora_config, transformer3d) + + network = None + else: + network = create_network( + 1.0, + args.rank, + args.network_alpha, + text_encoder, + transformer3d, + neuron_dropout=None, + target_name=args.target_name, + skip_name=args.lora_skip_name, + ) + network = network.to(weight_dtype) + network.apply_to(text_encoder, transformer3d, args.train_text_encoder and not args.training_with_video_token_length, True) if args.transformer_path is not None: print(f"From checkpoint: {args.transformer_path}") @@ -999,13 +1017,14 @@ def main(): 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].to(weight_dtype) - + if args.use_peft_lora: + network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1]), accelerate_state_dict) + else: + network_state_dict = {} + for key in accelerate_state_dict: + if "network" in key: + network_state_dict[key.replace("network.", "")] = accelerate_state_dict[key].to(weight_dtype) 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: @@ -1020,8 +1039,18 @@ def main(): 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): + 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") + if args.use_peft_lora: + network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1]), accelerate_state_dict) + else: + network_state_dict = accelerate_state_dict + 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) @@ -1033,10 +1062,15 @@ def main(): 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: 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 args.use_peft_lora: + save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(models[-1]))) + else: + save_model(safetensor_save_path, accelerator.unwrap_model(models[-1])) + if not args.use_deepspeed: for _ in range(len(weights)): weights.pop() @@ -1090,9 +1124,14 @@ def main(): else: optimizer_cls = torch.optim.AdamW - logging.info("Add network parameters") - trainable_params = list(filter(lambda p: p.requires_grad, network.parameters())) - trainable_params_optim = network.prepare_optimizer_params(args.learning_rate / 2, args.learning_rate, args.learning_rate) + if args.use_peft_lora: + logging.info("Add peft parameters") + trainable_params = list(filter(lambda p: p.requires_grad, transformer3d.parameters())) + trainable_params_optim = list(filter(lambda p: p.requires_grad, transformer3d.parameters())) + else: + logging.info("Add network parameters") + trainable_params = list(filter(lambda p: p.requires_grad, network.parameters())) + trainable_params_optim = network.prepare_optimizer_params(args.learning_rate / 2, args.learning_rate, args.learning_rate) if args.use_came: optimizer = optimizer_cls( @@ -1359,9 +1398,13 @@ def main(): ) # Prepare everything with our `accelerator`. - if fsdp_stage != 0: + if args.use_peft_lora: + transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( + transformer3d, optimizer, train_dataloader, lr_scheduler + ) + elif fsdp_stage != 0: transformer3d.network = network - transformer3d = transformer3d.to(weight_dtype) + transformer3d = transformer3d.to(dtype=weight_dtype) transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( transformer3d, optimizer, train_dataloader, lr_scheduler ) @@ -1370,14 +1413,16 @@ def main(): network, optimizer, train_dataloader, lr_scheduler ) - if zero_stage == 3: + if zero_stage != 0 and not args.use_peft_lora: 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) - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage != 0: 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) text_encoder = shard_fn(text_encoder) @@ -1519,6 +1564,10 @@ def main(): def save_model(ckpt_file, unwrapped_nw): os.makedirs(args.output_dir, exist_ok=True) accelerator.print(f"\nsaving checkpoint: {ckpt_file}") + if isinstance(unwrapped_nw, dict): + from safetensors.torch import save_file + save_file(unwrapped_nw, ckpt_file, metadata={"format": "pt"}) + return ckpt_file unwrapped_nw.save_weights(ckpt_file, weight_dtype, None) progress_bar = tqdm( @@ -1897,9 +1946,14 @@ def main(): torch.cuda.empty_cache() torch.cuda.ipc_collect() if not args.save_state: - safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") - save_model(safetensor_save_path, accelerator.unwrap_model(network)) - logger.info(f"Saved safetensor to {safetensor_save_path}") + if args.use_peft_lora: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(transformer3d))) + logger.info(f"Saved safetensor to {safetensor_save_path}") + else: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, accelerator.unwrap_model(network)) + logger.info(f"Saved safetensor to {safetensor_save_path}") else: accelerator_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") accelerator.save_state(accelerator_save_path) @@ -1948,8 +2002,14 @@ def main(): torch.cuda.empty_cache() torch.cuda.ipc_collect() if not args.save_state: - safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") - save_model(safetensor_save_path, accelerator.unwrap_model(network)) + if args.use_peft_lora: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(transformer3d))) + logger.info(f"Saved safetensor to {safetensor_save_path}") + else: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, accelerator.unwrap_model(network)) + logger.info(f"Saved safetensor to {safetensor_save_path}") else: accelerator_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") accelerator.save_state(accelerator_save_path) diff --git a/scripts/wan2.2/train_lora.sh b/scripts/wan2.2/train_lora.sh index 9875f6c..d008bde 100755 --- a/scripts/wan2.2/train_lora.sh +++ b/scripts/wan2.2/train_lora.sh @@ -36,6 +36,10 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.2/train_lora.py \ --enable_bucket \ --uniform_sampling \ --boundary_type="low" \ + --rank=64 \ + --network_alpha=32 \ + --target_name="q,k,v,ffn.0,ffn.2" \ + --use_peft_lora \ --train_mode="normal" \ --low_vram @@ -80,5 +84,9 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.2/train_lora.py \ # --enable_bucket \ # --uniform_sampling \ # --boundary_type="low" \ +# --rank=64 \ +# --network_alpha=32 \ +# --target_name="q,k,v,ffn.0,ffn.2" \ +# --use_peft_lora \ # --train_mode="i2v" \ # --low_vram diff --git a/scripts/wan2.2/train_s2v.py b/scripts/wan2.2/train_s2v.py index ff5b410..84b3183 100644 --- a/scripts/wan2.2/train_s2v.py +++ b/scripts/wan2.2/train_s2v.py @@ -1000,7 +1000,12 @@ def main(): 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): + 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) diff --git a/scripts/wan2.2/train_s2v_lora.py b/scripts/wan2.2/train_s2v_lora.py index 04531d3..af48a47 100644 --- a/scripts/wan2.2/train_s2v_lora.py +++ b/scripts/wan2.2/train_s2v_lora.py @@ -567,6 +567,9 @@ def parse_args(): default=64, help=("The dimension of the LoRA update matrices."), ) + parser.add_argument( + "--use_peft_lora", action="store_true", help="Whether or not to use peft lora." + ) parser.add_argument( "--train_text_encoder", action="store_true", @@ -746,6 +749,12 @@ def parse_args(): default=None, help=("The module is not trained in loras. "), ) + parser.add_argument( + "--target_name", + type=str, + default=None, + help=("The module is trained in loras. "), + ) args = parser.parse_args() env_local_rank = int(os.environ.get("LOCAL_RANK", -1)) @@ -939,16 +948,25 @@ def main(): transformer3d.requires_grad_(False) # Lora will work with this... - network = create_network( - 1.0, - args.rank, - args.network_alpha, - text_encoder, - transformer3d, - neuron_dropout=None, - skip_name=args.lora_skip_name - ) - network.apply_to(text_encoder, transformer3d, args.train_text_encoder and not args.training_with_video_token_length, True) + if args.use_peft_lora: + from peft import LoraConfig, inject_adapter_in_model, get_peft_model_state_dict + lora_config = LoraConfig(r=args.rank, lora_alpha=args.network_alpha, target_modules=args.target_name.split(",")) + transformer3d = inject_adapter_in_model(lora_config, transformer3d) + + network = None + else: + network = create_network( + 1.0, + args.rank, + args.network_alpha, + text_encoder, + transformer3d, + neuron_dropout=None, + target_name=args.target_name, + skip_name=args.lora_skip_name, + ) + network = network.to(weight_dtype) + network.apply_to(text_encoder, transformer3d, args.train_text_encoder and not args.training_with_video_token_length, True) if args.transformer_path is not None: print(f"From checkpoint: {args.transformer_path}") @@ -975,7 +993,7 @@ def main(): m, u = vae.load_state_dict(state_dict, strict=False) print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") assert len(u) == 0 - + # `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 @@ -984,13 +1002,14 @@ def main(): 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].to(weight_dtype) - + if args.use_peft_lora: + network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1]), accelerate_state_dict) + else: + network_state_dict = {} + for key in accelerate_state_dict: + if "network" in key: + network_state_dict[key.replace("network.", "")] = accelerate_state_dict[key].to(weight_dtype) 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: @@ -1005,8 +1024,18 @@ def main(): 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): + 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") + if args.use_peft_lora: + network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1]), accelerate_state_dict) + else: + network_state_dict = accelerate_state_dict + 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) @@ -1018,10 +1047,15 @@ def main(): 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: 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 args.use_peft_lora: + save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(models[-1]))) + else: + save_model(safetensor_save_path, accelerator.unwrap_model(models[-1])) + if not args.use_deepspeed: for _ in range(len(weights)): weights.pop() @@ -1075,9 +1109,14 @@ def main(): else: optimizer_cls = torch.optim.AdamW - logging.info("Add network parameters") - trainable_params = list(filter(lambda p: p.requires_grad, network.parameters())) - trainable_params_optim = network.prepare_optimizer_params(args.learning_rate / 2, args.learning_rate, args.learning_rate) + if args.use_peft_lora: + logging.info("Add peft parameters") + trainable_params = list(filter(lambda p: p.requires_grad, transformer3d.parameters())) + trainable_params_optim = list(filter(lambda p: p.requires_grad, transformer3d.parameters())) + else: + logging.info("Add network parameters") + trainable_params = list(filter(lambda p: p.requires_grad, network.parameters())) + trainable_params_optim = network.prepare_optimizer_params(args.learning_rate / 2, args.learning_rate, args.learning_rate) if args.use_came: optimizer = optimizer_cls( @@ -1397,9 +1436,13 @@ def main(): ) # Prepare everything with our `accelerator`. - if fsdp_stage != 0: + if args.use_peft_lora: + transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( + transformer3d, optimizer, train_dataloader, lr_scheduler + ) + elif fsdp_stage != 0: transformer3d.network = network - transformer3d = transformer3d.to(weight_dtype) + transformer3d = transformer3d.to(dtype=weight_dtype) transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( transformer3d, optimizer, train_dataloader, lr_scheduler ) @@ -1408,23 +1451,20 @@ def main(): network, optimizer, train_dataloader, lr_scheduler ) - if zero_stage == 3: + if zero_stage != 0 and not args.use_peft_lora: 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) - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage != 0: 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) text_encoder = shard_fn(text_encoder) - if args.use_ema: - ema_transformer3d.to(accelerator.device) - # Move text_encode and vae to gpu and cast to weight_dtype vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) transformer3d.to(accelerator.device, dtype=weight_dtype) @@ -1503,6 +1543,15 @@ def main(): else: initial_global_step = 0 + def save_model(ckpt_file, unwrapped_nw): + os.makedirs(args.output_dir, exist_ok=True) + accelerator.print(f"\nsaving checkpoint: {ckpt_file}") + if isinstance(unwrapped_nw, dict): + from safetensors.torch import save_file + save_file(unwrapped_nw, ckpt_file, metadata={"format": "pt"}) + return ckpt_file + unwrapped_nw.save_weights(ckpt_file, weight_dtype, None) + progress_bar = tqdm( range(0, args.max_train_steps), initial=initial_global_step, @@ -1946,9 +1995,14 @@ def main(): torch.cuda.empty_cache() torch.cuda.ipc_collect() if not args.save_state: - safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") - save_model(safetensor_save_path, accelerator.unwrap_model(network)) - logger.info(f"Saved safetensor to {safetensor_save_path}") + if args.use_peft_lora: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(transformer3d))) + logger.info(f"Saved safetensor to {safetensor_save_path}") + else: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, accelerator.unwrap_model(network)) + logger.info(f"Saved safetensor to {safetensor_save_path}") else: accelerator_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") accelerator.save_state(accelerator_save_path) @@ -1995,8 +2049,14 @@ def main(): torch.cuda.empty_cache() torch.cuda.ipc_collect() if not args.save_state: - safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") - save_model(safetensor_save_path, accelerator.unwrap_model(network)) + if args.use_peft_lora: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(transformer3d))) + logger.info(f"Saved safetensor to {safetensor_save_path}") + else: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, accelerator.unwrap_model(network)) + logger.info(f"Saved safetensor to {safetensor_save_path}") else: accelerator_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") accelerator.save_state(accelerator_save_path) diff --git a/scripts/wan2.2/train_s2v_lora.sh b/scripts/wan2.2/train_s2v_lora.sh index e98454d..f52c617 100644 --- a/scripts/wan2.2/train_s2v_lora.sh +++ b/scripts/wan2.2/train_s2v_lora.sh @@ -36,4 +36,8 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.2/train_s2v_lora.py \ --uniform_sampling \ --boundary_type="full" \ --control_ref_image="random" \ + --rank=64 \ + --network_alpha=32 \ + --target_name="q,k,v,ffn.0,ffn.2" \ + --use_peft_lora \ --low_vram \ No newline at end of file diff --git a/scripts/wan2.2_fun/README_TRAIN.md b/scripts/wan2.2_fun/README_TRAIN.md index 2a2bab6..0c8cc9a 100755 --- a/scripts/wan2.2_fun/README_TRAIN.md +++ b/scripts/wan2.2_fun/README_TRAIN.md @@ -21,6 +21,17 @@ Some parameters in the sh file can be confusing, and they are explained in this - `resume_from_checkpoint` is used to set the training should be resumed from a previous checkpoint. Use a path or `"latest"` to automatically select the last available checkpoint. - `boundary_type`: The Wan2.2 series includes two distinct models that handle different noise levels, specified via the `boundary_type` parameter. `low`: Corresponds to the **low noise model** (low_noise_model). `high`: Corresponds to the **high noise model**. (high_noise_model). `full`: Corresponds to the ti2v 5B model (single mode). +When train model with multi machines, please set the params as follows: +```sh +export MASTER_ADDR="your master address" +export MASTER_PORT=10086 +export WORLD_SIZE=1 # The number of machines +export NUM_PROCESS=8 # The number of processes, such as WORLD_SIZE * 8 +export RANK=0 # The rank of this machine + +accelerate launch --mixed_precision="bf16" --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK scripts/wan2.2_fun/xxx.py +``` + Wan T2V without deepspeed: Training 14B Wan2.2 without DeepSpeed may result in insufficient GPU memory. @@ -70,9 +81,9 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.2_fun/train.py \ --trainable_modules "." ``` -Wan T2V with deepspeed zero-2: +Wan T2V with Deepspeed Zero-2: -Wan with DeepSpeed Zero-2 is suitable for training 14B Wan at low resolutions, but training 14B Wan at high resolutions may still result in insufficient GPU memory. +Wan with Deepspeed Zero-2 is suitable for training 14B Wan at low resolutions, but training 14B Wan at high resolutions may still result in insufficient GPU memory. ```sh export MODEL_NAME="models/Diffusion_Transformer/Wan2.2-Fun-A14B-InP" @@ -121,7 +132,9 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con --trainable_modules "." ``` -Wan T2V with deepspeed zero-3: +DeepSpeed Zero-3 is not highly recommended at the moment. In this repository, using FSDP has fewer errors and is more stable. + +Wan T2V with DeepSpeed Zero-3: Wan with DeepSpeed Zero-3 is suitable for 14B Wan at high resolutions. After training, you can use the following command to get the final model: ```sh diff --git a/scripts/wan2.2_fun/README_TRAIN_CONTROL.md b/scripts/wan2.2_fun/README_TRAIN_CONTROL.md index 6057d1c..0fdfa3c 100755 --- a/scripts/wan2.2_fun/README_TRAIN_CONTROL.md +++ b/scripts/wan2.2_fun/README_TRAIN_CONTROL.md @@ -160,7 +160,9 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con --trainable_modules "." ``` -Wan-Fun-Control with deepspeed zero-3: +DeepSpeed Zero-3 is not highly recommended at the moment. In this repository, using FSDP has fewer errors and is more stable. + +Wan-Fun-Control with DeepSpeed Zero-3: Wan with DeepSpeed Zero-3 is suitable for 14B Wan at high resolutions. After training, you can use the following command to get the final model: ```sh diff --git a/scripts/wan2.2_fun/README_TRAIN_CONTROL_LORA.md b/scripts/wan2.2_fun/README_TRAIN_CONTROL_LORA.md index 07e97b1..110e779 100755 --- a/scripts/wan2.2_fun/README_TRAIN_CONTROL_LORA.md +++ b/scripts/wan2.2_fun/README_TRAIN_CONTROL_LORA.md @@ -47,6 +47,10 @@ Some parameters in the sh file can be confusing, and they are explained in this - `add_full_ref_image_in_self_attention` determines whether to include the reference image in self-attention. This option is used in V1.1, as it supports using a reference image as the control image. It should not be used in V1.0 and Control-Camera models. - `add_inpaint_info` determines whether to incorporate inpaint information into the model training. When enabled, this allows the model to support specifying starting and ending images in the controls during generation. - `boundary_type`: The Wan2.2 series includes two distinct models that handle different noise levels, specified via the `boundary_type` parameter. `low`: Corresponds to the **low noise model** (low_noise_model). `high`: Corresponds to the **high noise model**. (high_noise_model). `full`: Corresponds to the ti2v 5B model (single mode). +- `target_name` represents the components/modules to which LoRA will be applied, separated by commas. +- `use_peft_lora` indicates whether to use the PEFT module for adding LoRA. Using this module will be more memory-efficient. +- `rank` means the dimension of the LoRA update matrices. +- `network_alpha` means the scale of the LoRA update matrices. When train model with multi machines, please set the params as follows: ```sh @@ -103,6 +107,10 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.2_fun/train_control_lora --control_ref_image="random" \ --add_inpaint_info \ --add_full_ref_image_in_self_attention \ + --rank=64 \ + --network_alpha=32 \ + --target_name="q,k,v,ffn.0,ffn.2" \ + --use_peft_lora \ --low_vram ``` @@ -151,10 +159,18 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con --control_ref_image="random" \ --add_inpaint_info \ --add_full_ref_image_in_self_attention \ + --rank=64 \ + --network_alpha=32 \ + --target_name="q,k,v,ffn.0,ffn.2" \ + --use_peft_lora \ --low_vram ``` -Wan-Fun-Control with deepspeed zero-3: +DeepSpeed Zero-3 is not highly recommended at the moment. In this repository, using FSDP has fewer errors and is more stable. + +It is known that DeepSpeed Zero-3 is not compatible with PEFT. + +Wan-Fun-Control with DeepSpeed Zero-3: Wan with DeepSpeed Zero-3 is suitable for 14B Wan at high resolutions. You must set save_state to True to save the model. After training, you can use the following command to get the final model: ```sh @@ -258,5 +274,9 @@ accelerate launch --mixed_precision="bf16" --use_fsdp --fsdp_auto_wrap_policy TR --control_ref_image="random" \ --add_inpaint_info \ --add_full_ref_image_in_self_attention \ + --rank=64 \ + --network_alpha=32 \ + --target_name="q,k,v,ffn.0,ffn.2" \ + --use_peft_lora \ --low_vram ``` \ No newline at end of file diff --git a/scripts/wan2.2_fun/README_TRAIN_LORA.md b/scripts/wan2.2_fun/README_TRAIN_LORA.md index b867025..f5de962 100755 --- a/scripts/wan2.2_fun/README_TRAIN_LORA.md +++ b/scripts/wan2.2_fun/README_TRAIN_LORA.md @@ -20,6 +20,21 @@ Some parameters in the sh file can be confusing, and they are explained in this - `train_mode` is used to specify the training mode, which can be either normal or inpaint. Since Wan uses the inpaint model to achieve image-to-video generation, the default is set to inpaint mode. If you only wish to achieve text-to-video generation, you can remove this line, and it will default to the text-to-video mode. - `resume_from_checkpoint` is used to set the training should be resumed from a previous checkpoint. Use a path or `"latest"` to automatically select the last available checkpoint. - `boundary_type`: The Wan2.2 series includes two distinct models that handle different noise levels, specified via the `boundary_type` parameter. `low`: Corresponds to the **low noise model** (low_noise_model). `high`: Corresponds to the **high noise model**. (high_noise_model). `full`: Corresponds to the ti2v 5B model (single mode). +- `target_name` represents the components/modules to which LoRA will be applied, separated by commas. +- `use_peft_lora` indicates whether to use the PEFT module for adding LoRA. Using this module will be more memory-efficient. +- `rank` means the dimension of the LoRA update matrices. +- `network_alpha` means the scale of the LoRA update matrices. + +When train model with multi machines, please set the params as follows: +```sh +export MASTER_ADDR="your master address" +export MASTER_PORT=10086 +export WORLD_SIZE=1 # The number of machines +export NUM_PROCESS=8 # The number of processes, such as WORLD_SIZE * 8 +export RANK=0 # The rank of this machine + +accelerate launch --mixed_precision="bf16" --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK scripts/wan2.2_fun/xxx.py +``` Wan T2V without deepspeed: @@ -63,13 +78,17 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.2_fun/train_lora.py \ --enable_bucket \ --uniform_sampling \ --boundary_type="low" \ + --rank=64 \ + --network_alpha=32 \ + --target_name="q,k,v,ffn.0,ffn.2" \ + --use_peft_lora \ --train_mode="inpaint" \ --low_vram ``` -Wan T2V with deepspeed zero-2: +Wan T2V with Deepspeed Zero-2: -Wan with DeepSpeed Zero-2 is suitable for training 14B Wan at low resolutions, but training 14B Wan at high resolutions may still result in insufficient GPU memory. +Wan with Deepspeed Zero-2 is suitable for training 14B Wan at low resolutions, but training 14B Wan at high resolutions may still result in insufficient GPU memory. ```sh export MODEL_NAME="models/Diffusion_Transformer/Wan2.2-Fun-A14B-InP" @@ -110,12 +129,19 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con --enable_bucket \ --uniform_sampling \ --boundary_type="low" \ - --use_deepspeed \ + --rank=64 \ + --network_alpha=32 \ + --target_name="q,k,v,ffn.0,ffn.2" \ + --use_peft_lora \ --train_mode="inpaint" \ --low_vram ``` -Wan T2V with deepspeed zero-3: +DeepSpeed Zero-3 is not highly recommended at the moment. In this repository, using FSDP has fewer errors and is more stable. + +It is known that DeepSpeed Zero-3 is not compatible with PEFT. + +Wan T2V with DeepSpeed Zero-3: Wan with DeepSpeed Zero-3 is suitable for 14B Wan at high resolutions. You must set save_state to True to save the model. After training, you can use the following command to get the final model: ```sh @@ -162,8 +188,6 @@ accelerate launch --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag --enable_bucket \ --uniform_sampling \ --boundary_type="low" \ - --save_state \ - --use_deepspeed \ --train_mode="inpaint" \ --low_vram ``` @@ -210,8 +234,10 @@ accelerate launch --mixed_precision="bf16" --use_fsdp --fsdp_auto_wrap_policy TR --enable_bucket \ --uniform_sampling \ --boundary_type="low" \ - --save_state \ - --use_deepspeed \ + --rank=64 \ + --network_alpha=32 \ + --target_name="q,k,v,ffn.0,ffn.2" \ + --use_peft_lora \ --train_mode="inpaint" \ --low_vram ``` \ No newline at end of file diff --git a/scripts/wan2.2_fun/train.py b/scripts/wan2.2_fun/train.py index 8ab36f2..42d311d 100644 --- a/scripts/wan2.2_fun/train.py +++ b/scripts/wan2.2_fun/train.py @@ -1012,7 +1012,12 @@ def main(): 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): + 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) diff --git a/scripts/wan2.2_fun/train_control.py b/scripts/wan2.2_fun/train_control.py index 8b19edb..6985e72 100644 --- a/scripts/wan2.2_fun/train_control.py +++ b/scripts/wan2.2_fun/train_control.py @@ -981,7 +981,12 @@ def main(): 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): + 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) diff --git a/scripts/wan2.2_fun/train_control_lora.py b/scripts/wan2.2_fun/train_control_lora.py index 56d77de..1f45630 100644 --- a/scripts/wan2.2_fun/train_control_lora.py +++ b/scripts/wan2.2_fun/train_control_lora.py @@ -526,6 +526,9 @@ def parse_args(): default=64, help=("The dimension of the LoRA update matrices."), ) + parser.add_argument( + "--use_peft_lora", action="store_true", help="Whether or not to use peft lora." + ) parser.add_argument( "--train_text_encoder", action="store_true", @@ -718,6 +721,12 @@ def parse_args(): default=None, help=("The module is not trained in loras. "), ) + parser.add_argument( + "--target_name", + type=str, + default=None, + help=("The module is trained in loras. "), + ) args = parser.parse_args() env_local_rank = int(os.environ.get("LOCAL_RANK", -1)) @@ -904,16 +913,25 @@ def main(): transformer3d.requires_grad_(False) # Lora will work with this... - network = create_network( - 1.0, - args.rank, - args.network_alpha, - text_encoder, - transformer3d, - neuron_dropout=None, - skip_name=args.lora_skip_name, - ) - network.apply_to(text_encoder, transformer3d, args.train_text_encoder and not args.training_with_video_token_length, True) + if args.use_peft_lora: + from peft import LoraConfig, inject_adapter_in_model, get_peft_model_state_dict + lora_config = LoraConfig(r=args.rank, lora_alpha=args.network_alpha, target_modules=args.target_name.split(",")) + transformer3d = inject_adapter_in_model(lora_config, transformer3d) + + network = None + else: + network = create_network( + 1.0, + args.rank, + args.network_alpha, + text_encoder, + transformer3d, + neuron_dropout=None, + target_name=args.target_name, + skip_name=args.lora_skip_name, + ) + network = network.to(weight_dtype) + network.apply_to(text_encoder, transformer3d, args.train_text_encoder and not args.training_with_video_token_length, True) if args.transformer_path is not None: print(f"From checkpoint: {args.transformer_path}") @@ -947,13 +965,14 @@ def main(): 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].to(weight_dtype) - + if args.use_peft_lora: + network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1]), accelerate_state_dict) + else: + network_state_dict = {} + for key in accelerate_state_dict: + if "network" in key: + network_state_dict[key.replace("network.", "")] = accelerate_state_dict[key].to(weight_dtype) 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: @@ -968,8 +987,18 @@ def main(): 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): + 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") + if args.use_peft_lora: + network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1]), accelerate_state_dict) + else: + network_state_dict = accelerate_state_dict + 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) @@ -981,10 +1010,15 @@ def main(): 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: 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 args.use_peft_lora: + save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(models[-1]))) + else: + save_model(safetensor_save_path, accelerator.unwrap_model(models[-1])) + if not args.use_deepspeed: for _ in range(len(weights)): weights.pop() @@ -1038,9 +1072,14 @@ def main(): else: optimizer_cls = torch.optim.AdamW - logging.info("Add network parameters") - trainable_params = list(filter(lambda p: p.requires_grad, network.parameters())) - trainable_params_optim = network.prepare_optimizer_params(args.learning_rate / 2, args.learning_rate, args.learning_rate) + if args.use_peft_lora: + logging.info("Add peft parameters") + trainable_params = list(filter(lambda p: p.requires_grad, transformer3d.parameters())) + trainable_params_optim = list(filter(lambda p: p.requires_grad, transformer3d.parameters())) + else: + logging.info("Add network parameters") + trainable_params = list(filter(lambda p: p.requires_grad, network.parameters())) + trainable_params_optim = network.prepare_optimizer_params(args.learning_rate / 2, args.learning_rate, args.learning_rate) if args.use_came: optimizer = optimizer_cls( @@ -1432,9 +1471,13 @@ def main(): ) # Prepare everything with our `accelerator`. - if fsdp_stage != 0: + if args.use_peft_lora: + transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( + transformer3d, optimizer, train_dataloader, lr_scheduler + ) + elif fsdp_stage != 0: transformer3d.network = network - transformer3d = transformer3d.to(weight_dtype) + transformer3d = transformer3d.to(dtype=weight_dtype) transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( transformer3d, optimizer, train_dataloader, lr_scheduler ) @@ -1443,14 +1486,14 @@ def main(): network, optimizer, train_dataloader, lr_scheduler ) - if zero_stage == 3: + if zero_stage != 0 and not args.use_peft_lora: 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) - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage != 0: from functools import partial from videox_fun.dist import set_multi_gpus_devices, shard_model @@ -1594,6 +1637,10 @@ def main(): def save_model(ckpt_file, unwrapped_nw): os.makedirs(args.output_dir, exist_ok=True) accelerator.print(f"\nsaving checkpoint: {ckpt_file}") + if isinstance(unwrapped_nw, dict): + from safetensors.torch import save_file + save_file(unwrapped_nw, ckpt_file, metadata={"format": "pt"}) + return ckpt_file unwrapped_nw.save_weights(ckpt_file, weight_dtype, None) progress_bar = tqdm( @@ -2049,9 +2096,14 @@ def main(): torch.cuda.empty_cache() torch.cuda.ipc_collect() if not args.save_state: - safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") - save_model(safetensor_save_path, accelerator.unwrap_model(network)) - logger.info(f"Saved safetensor to {safetensor_save_path}") + if args.use_peft_lora: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(transformer3d))) + logger.info(f"Saved safetensor to {safetensor_save_path}") + else: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, accelerator.unwrap_model(network)) + logger.info(f"Saved safetensor to {safetensor_save_path}") else: accelerator_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") accelerator.save_state(accelerator_save_path) @@ -2100,8 +2152,14 @@ def main(): torch.cuda.empty_cache() torch.cuda.ipc_collect() if not args.save_state: - safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") - save_model(safetensor_save_path, accelerator.unwrap_model(network)) + if args.use_peft_lora: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(transformer3d))) + logger.info(f"Saved safetensor to {safetensor_save_path}") + else: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, accelerator.unwrap_model(network)) + logger.info(f"Saved safetensor to {safetensor_save_path}") else: accelerator_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") accelerator.save_state(accelerator_save_path) @@ -2109,5 +2167,6 @@ def main(): accelerator.end_training() + if __name__ == "__main__": main() diff --git a/scripts/wan2.2_fun/train_control_lora.sh b/scripts/wan2.2_fun/train_control_lora.sh index 507216b..d7cb822 100644 --- a/scripts/wan2.2_fun/train_control_lora.sh +++ b/scripts/wan2.2_fun/train_control_lora.sh @@ -40,5 +40,8 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.2_fun/train_control_lora --add_inpaint_info \ --add_full_ref_image_in_self_attention \ --boundary_type="low" \ - --lora_skip_name="ffn" \ + --rank=64 \ + --network_alpha=32 \ + --target_name="q,k,v,ffn.0,ffn.2" \ + --use_peft_lora \ --low_vram \ No newline at end of file diff --git a/scripts/wan2.2_fun/train_lora.py b/scripts/wan2.2_fun/train_lora.py index 31b8b23..9eb6823 100644 --- a/scripts/wan2.2_fun/train_lora.py +++ b/scripts/wan2.2_fun/train_lora.py @@ -563,6 +563,9 @@ def parse_args(): default=64, help=("The dimension of the LoRA update matrices."), ) + parser.add_argument( + "--use_peft_lora", action="store_true", help="Whether or not to use peft lora." + ) parser.add_argument( "--train_text_encoder", action="store_true", @@ -732,6 +735,12 @@ def parse_args(): default=None, help=("The module is not trained in loras. "), ) + parser.add_argument( + "--target_name", + type=str, + default=None, + help=("The module is trained in loras. "), + ) args = parser.parse_args() env_local_rank = int(os.environ.get("LOCAL_RANK", -1)) @@ -918,16 +927,25 @@ def main(): transformer3d.requires_grad_(False) # Lora will work with this... - network = create_network( - 1.0, - args.rank, - args.network_alpha, - text_encoder, - transformer3d, - neuron_dropout=None, - skip_name=args.lora_skip_name, - ) - network.apply_to(text_encoder, transformer3d, args.train_text_encoder and not args.training_with_video_token_length, True) + if args.use_peft_lora: + from peft import LoraConfig, inject_adapter_in_model, get_peft_model_state_dict + lora_config = LoraConfig(r=args.rank, lora_alpha=args.network_alpha, target_modules=args.target_name.split(",")) + transformer3d = inject_adapter_in_model(lora_config, transformer3d) + + network = None + else: + network = create_network( + 1.0, + args.rank, + args.network_alpha, + text_encoder, + transformer3d, + neuron_dropout=None, + target_name=args.target_name, + skip_name=args.lora_skip_name, + ) + network = network.to(weight_dtype) + network.apply_to(text_encoder, transformer3d, args.train_text_encoder and not args.training_with_video_token_length, True) if args.transformer_path is not None: print(f"From checkpoint: {args.transformer_path}") @@ -963,13 +981,14 @@ def main(): 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].to(weight_dtype) - + if args.use_peft_lora: + network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1]), accelerate_state_dict) + else: + network_state_dict = {} + for key in accelerate_state_dict: + if "network" in key: + network_state_dict[key.replace("network.", "")] = accelerate_state_dict[key].to(weight_dtype) 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: @@ -984,8 +1003,18 @@ def main(): 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): + 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") + if args.use_peft_lora: + network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1]), accelerate_state_dict) + else: + network_state_dict = accelerate_state_dict + 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) @@ -997,10 +1026,15 @@ def main(): 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: 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 args.use_peft_lora: + save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(models[-1]))) + else: + save_model(safetensor_save_path, accelerator.unwrap_model(models[-1])) + if not args.use_deepspeed: for _ in range(len(weights)): weights.pop() @@ -1054,9 +1088,14 @@ def main(): else: optimizer_cls = torch.optim.AdamW - logging.info("Add network parameters") - trainable_params = list(filter(lambda p: p.requires_grad, network.parameters())) - trainable_params_optim = network.prepare_optimizer_params(args.learning_rate / 2, args.learning_rate, args.learning_rate) + if args.use_peft_lora: + logging.info("Add peft parameters") + trainable_params = list(filter(lambda p: p.requires_grad, transformer3d.parameters())) + trainable_params_optim = list(filter(lambda p: p.requires_grad, transformer3d.parameters())) + else: + logging.info("Add network parameters") + trainable_params = list(filter(lambda p: p.requires_grad, network.parameters())) + trainable_params_optim = network.prepare_optimizer_params(args.learning_rate / 2, args.learning_rate, args.learning_rate) if args.use_came: optimizer = optimizer_cls( @@ -1360,9 +1399,13 @@ def main(): ) # Prepare everything with our `accelerator`. - if fsdp_stage != 0: + if args.use_peft_lora: + transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( + transformer3d, optimizer, train_dataloader, lr_scheduler + ) + elif fsdp_stage != 0: transformer3d.network = network - transformer3d = transformer3d.to(weight_dtype) + transformer3d = transformer3d.to(dtype=weight_dtype) transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( transformer3d, optimizer, train_dataloader, lr_scheduler ) @@ -1371,14 +1414,16 @@ def main(): network, optimizer, train_dataloader, lr_scheduler ) - if zero_stage == 3: + if zero_stage != 0 and not args.use_peft_lora: 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) - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage != 0: 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) text_encoder = shard_fn(text_encoder) @@ -1520,6 +1565,10 @@ def main(): def save_model(ckpt_file, unwrapped_nw): os.makedirs(args.output_dir, exist_ok=True) accelerator.print(f"\nsaving checkpoint: {ckpt_file}") + if isinstance(unwrapped_nw, dict): + from safetensors.torch import save_file + save_file(unwrapped_nw, ckpt_file, metadata={"format": "pt"}) + return ckpt_file unwrapped_nw.save_weights(ckpt_file, weight_dtype, None) progress_bar = tqdm( @@ -1886,9 +1935,14 @@ def main(): torch.cuda.empty_cache() torch.cuda.ipc_collect() if not args.save_state: - safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") - save_model(safetensor_save_path, accelerator.unwrap_model(network)) - logger.info(f"Saved safetensor to {safetensor_save_path}") + if args.use_peft_lora: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(transformer3d))) + logger.info(f"Saved safetensor to {safetensor_save_path}") + else: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, accelerator.unwrap_model(network)) + logger.info(f"Saved safetensor to {safetensor_save_path}") else: accelerator_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") accelerator.save_state(accelerator_save_path) @@ -1937,8 +1991,14 @@ def main(): torch.cuda.empty_cache() torch.cuda.ipc_collect() if not args.save_state: - safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") - save_model(safetensor_save_path, accelerator.unwrap_model(network)) + if args.use_peft_lora: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(transformer3d))) + logger.info(f"Saved safetensor to {safetensor_save_path}") + else: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, accelerator.unwrap_model(network)) + logger.info(f"Saved safetensor to {safetensor_save_path}") else: accelerator_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") accelerator.save_state(accelerator_save_path) diff --git a/scripts/wan2.2_fun/train_lora.sh b/scripts/wan2.2_fun/train_lora.sh index fc3e06f..5f13f97 100644 --- a/scripts/wan2.2_fun/train_lora.sh +++ b/scripts/wan2.2_fun/train_lora.sh @@ -37,5 +37,8 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.2_fun/train_lora.py \ --uniform_sampling \ --train_mode="inpaint" \ --boundary_type="low" \ - --lora_skip_name="ffn" \ + --rank=64 \ + --network_alpha=32 \ + --target_name="q,k,v,ffn.0,ffn.2" \ + --use_peft_lora \ --low_vram diff --git a/scripts/wan2.2_vace_fun/README_TRAIN.md b/scripts/wan2.2_vace_fun/README_TRAIN.md index ca63aea..258dbd2 100755 --- a/scripts/wan2.2_vace_fun/README_TRAIN.md +++ b/scripts/wan2.2_vace_fun/README_TRAIN.md @@ -153,7 +153,9 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con --trainable_modules "vace" ``` -Wan-Fun-Control with deepspeed zero-3: +DeepSpeed Zero-3 is not highly recommended at the moment. In this repository, using FSDP has fewer errors and is more stable. + +Wan-Fun-Control with DeepSpeed Zero-3: Wan with DeepSpeed Zero-3 is suitable for 14B Wan at high resolutions. After training, you can use the following command to get the final model: ```sh diff --git a/scripts/wan2.2_vace_fun/train.py b/scripts/wan2.2_vace_fun/train.py index a14a08c..684bf82 100644 --- a/scripts/wan2.2_vace_fun/train.py +++ b/scripts/wan2.2_vace_fun/train.py @@ -932,7 +932,12 @@ def main(): 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): + 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) diff --git a/videox_fun/utils/lora_utils.py b/videox_fun/utils/lora_utils.py index 7ff11d4..94b0044 100755 --- a/videox_fun/utils/lora_utils.py +++ b/videox_fun/utils/lora_utils.py @@ -111,7 +111,6 @@ class LoRAModule(torch.nn.Module): return org_forwarded.to(weight_dtype) + lx.to(weight_dtype) * self.multiplier * scale - def addnet_hash_legacy(b): """Old model hash used by sd-webui-additional-networks for .safetensors format files""" m = hashlib.sha256() @@ -120,7 +119,6 @@ def addnet_hash_legacy(b): m.update(b.read(0x10000)) return m.hexdigest()[0:8] - def addnet_hash_safetensors(b): """New model hash used by sd-webui-additional-networks for .safetensors format files""" hash_sha256 = hashlib.sha256() @@ -137,7 +135,6 @@ def addnet_hash_safetensors(b): return hash_sha256.hexdigest() - def precalculate_safetensors_hashes(tensors, metadata): """Precalculate the model hashes needed by sd-webui-additional-networks to save time on indexing the model later.""" @@ -154,7 +151,6 @@ def precalculate_safetensors_hashes(tensors, metadata): legacy_hash = addnet_hash_legacy(b) return model_hash, legacy_hash - class LoRANetwork(torch.nn.Module): TRANSFORMER_TARGET_REPLACE_MODULE = [ "CogVideoXTransformer3DModel", "WanTransformer3DModel", \ @@ -207,18 +203,18 @@ class LoRANetwork(torch.nn.Module): is_conv2d = child_module.__class__.__name__ == "Conv2d" or child_module.__class__.__name__ == "LoRACompatibleConv" is_conv2d_1x1 = is_conv2d and child_module.kernel_size == (1, 1) - if skip_name is not None and skip_name in child_name: + skip_names = skip_name.split(',') if skip_name is not None else [] + target_names = target_name.split(',') if target_name is not None else [] + + skip_names = [name.strip() for name in skip_names if name.strip()] + target_names = [name.strip() for name in target_names if name.strip()] + + if skip_names and any(skip_n in child_name for skip_n in skip_names): continue - if target_name is not None: - target_name_in = False - if isinstance(target_name, str): - target_name_in = target_name in child_name - elif isinstance(target_name, list): - target_name_in = any([_target_name in child_name for _target_name in target_name]) - if not target_name_in: - continue - + if target_names and not any(target_n in child_name for target_n in target_names): + continue + if is_linear or is_conv2d: lora_name = prefix + "." + name + "." + child_name lora_name = lora_name.replace(".", "_") @@ -360,7 +356,7 @@ def create_network( transformer, neuron_dropout: Optional[float] = None, skip_name: str = None, - target_name = None, + target_name: str = None, **kwargs, ): if network_dim is None: @@ -382,6 +378,9 @@ def create_network( return network def merge_lora(pipeline, lora_path, multiplier, device='cpu', dtype=torch.float32, state_dict=None, transformer_only=False, sub_transformer_name="transformer"): + if lora_path is None: + return pipeline + LORA_PREFIX_TRANSFORMER = "lora_unet" LORA_PREFIX_TEXT_ENCODER = "lora_te" if state_dict is None: @@ -390,20 +389,27 @@ def merge_lora(pipeline, lora_path, multiplier, device='cpu', dtype=torch.float3 state_dict = state_dict updates = defaultdict(dict) for key, value in state_dict.items(): - if "diffusion_model" in key: - key = key.replace("diffusion_model.", "lora_unet__") - key = key.replace("blocks.", "blocks_") - key = key.replace(".self_attn.", "_self_attn_") - key = key.replace(".cross_attn.", "_cross_attn_") - key = key.replace(".ffn.", "_ffn_") if "lora_A" in key or "lora_B" in key: key = "lora_unet__" + key - key = key.replace("blocks.", "blocks_") - key = key.replace(".self_attn.", "_self_attn_") - key = key.replace(".cross_attn.", "_cross_attn_") - key = key.replace(".ffn.", "_ffn_") - key = key.replace(".lora_A.default.", ".lora_down.") - key = key.replace(".lora_B.default.", ".lora_up.") + key = key.replace(".", "_") + if key.endswith("_lora_up_weight"): + key = key[:-15] + ".lora_up.weight" + if key.endswith("_lora_down_weight"): + key = key[:-17] + ".lora_down.weight" + if key.endswith("_lora_A_default_weight"): + key = key[:-21] + ".lora_A.weight" + if key.endswith("_lora_B_default_weight"): + key = key[:-21] + ".lora_B.weight" + if key.endswith("_lora_A_weight"): + key = key[:-14] + ".lora_A.weight" + if key.endswith("_lora_B_weight"): + key = key[:-14] + ".lora_B.weight" + if key.endswith("_alpha"): + key = key[:-6] + ".alpha" + key = key.replace(".lora_A.default.", ".lora_down.") + key = key.replace(".lora_B.default.", ".lora_up.") + key = key.replace(".lora_A.", ".lora_down.") + key = key.replace(".lora_B.", ".lora_up.") layer, elem = key.split('.', 1) updates[layer][elem] = value @@ -504,6 +510,9 @@ def merge_lora(pipeline, lora_path, multiplier, device='cpu', dtype=torch.float3 # TODO: Refactor with merge_lora. def unmerge_lora(pipeline, lora_path, multiplier=1, device="cpu", dtype=torch.float32, sub_transformer_name="transformer"): + if lora_path is None: + return pipeline + """Unmerge state_dict in LoRANetwork from the pipeline in diffusers.""" LORA_PREFIX_UNET = "lora_unet" LORA_PREFIX_TEXT_ENCODER = "lora_te" @@ -511,20 +520,27 @@ def unmerge_lora(pipeline, lora_path, multiplier=1, device="cpu", dtype=torch.fl updates = defaultdict(dict) for key, value in state_dict.items(): - if "diffusion_model" in key: - key = key.replace("diffusion_model.", "lora_unet__") - key = key.replace("blocks.", "blocks_") - key = key.replace(".self_attn.", "_self_attn_") - key = key.replace(".cross_attn.", "_cross_attn_") - key = key.replace(".ffn.", "_ffn_") if "lora_A" in key or "lora_B" in key: key = "lora_unet__" + key - key = key.replace("blocks.", "blocks_") - key = key.replace(".self_attn.", "_self_attn_") - key = key.replace(".cross_attn.", "_cross_attn_") - key = key.replace(".ffn.", "_ffn_") - key = key.replace(".lora_A.default.", ".lora_down.") - key = key.replace(".lora_B.default.", ".lora_up.") + key = key.replace(".", "_") + if key.endswith("_lora_up_weight"): + key = key[:-15] + ".lora_up.weight" + if key.endswith("_lora_down_weight"): + key = key[:-17] + ".lora_down.weight" + if key.endswith("_lora_A_default_weight"): + key = key[:-21] + ".lora_A.weight" + if key.endswith("_lora_B_default_weight"): + key = key[:-21] + ".lora_B.weight" + if key.endswith("_lora_A_weight"): + key = key[:-14] + ".lora_A.weight" + if key.endswith("_lora_B_weight"): + key = key[:-14] + ".lora_B.weight" + if key.endswith("_alpha"): + key = key[:-6] + ".alpha" + key = key.replace(".lora_A.default.", ".lora_down.") + key = key.replace(".lora_B.default.", ".lora_up.") + key = key.replace(".lora_A.", ".lora_down.") + key = key.replace(".lora_B.", ".lora_up.") layer, elem = key.split('.', 1) updates[layer][elem] = value