Update Peft Lora && Update Readme (#376)

This commit is contained in:
Bubbliiiing
2025-11-20 17:44:22 +08:00
committed by GitHub
parent ccc3b1055e
commit 037a2e8360
79 changed files with 2499 additions and 849 deletions
+3
View File
@@ -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
+7 -3
View File
@@ -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フレーム/秒でトレーニングされています。|
</details>
# 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)の下でライセンスされています。
+3
View File
@@ -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
+6 -10
View File
@@ -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 \
@@ -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 \
+22 -9
View File
@@ -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"
```
+123 -44
View File
@@ -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()
+123 -47
View File
@@ -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()
+234 -85
View File
@@ -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)
+8
View File
@@ -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"
+15 -2
View File
@@ -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
+5
View File
@@ -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)
+12 -44
View File
@@ -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
+24 -41
View File
@@ -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
```
+5
View File
@@ -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)
+87 -29
View File
@@ -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)
+4
View File
@@ -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
+15 -2
View File
@@ -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
+15 -2
View File
@@ -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
+33 -2
View File
@@ -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
```
+5
View File
@@ -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)
+5
View File
@@ -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)
+87 -29
View File
@@ -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)
+4
View File
@@ -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"
+87 -29
View File
@@ -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)
+4
View File
@@ -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
+16 -3
View File
@@ -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
+16 -3
View File
@@ -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
+34 -3
View File
@@ -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
```
+37 -9
View File
@@ -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
```
+5
View File
@@ -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)
+5
View File
@@ -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)
+138 -61
View File
@@ -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")
+4
View File
@@ -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
+89 -29
View File
@@ -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)
+8
View File
@@ -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"
+16 -3
View File
@@ -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
+4 -2
View File
@@ -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
@@ -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
+34 -3
View File
@@ -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
```
+5
View File
@@ -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)
+5
View File
@@ -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)
+89 -29
View File
@@ -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)
+4
View File
@@ -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
+89 -29
View File
@@ -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)
+8
View File
@@ -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"
+3 -1
View File
@@ -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
+5
View File
@@ -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)
+15 -3
View File
@@ -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
+15 -2
View File
@@ -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
+16 -3
View File
@@ -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
+34 -3
View File
@@ -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
```
+38 -6
View File
@@ -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
```
+15 -2
View File
@@ -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
+5
View File
@@ -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)
+5
View File
@@ -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)
+87 -31
View File
@@ -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)
+4
View File
@@ -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
+5
View File
@@ -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)
+148 -61
View File
@@ -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()
+4
View File
@@ -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
+89 -29
View File
@@ -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)
+8
View File
@@ -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
+5
View File
@@ -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)
+93 -33
View File
@@ -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)
+4
View File
@@ -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
+16 -3
View File
@@ -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
+3 -1
View File
@@ -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
@@ -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
```
+34 -8
View File
@@ -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
```
+5
View File
@@ -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)
+5
View File
@@ -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)
+88 -29
View File
@@ -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()
+4 -1
View File
@@ -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
+89 -29
View File
@@ -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)
+4 -1
View File
@@ -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
+3 -1
View File
@@ -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
+5
View File
@@ -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)
+55 -39
View File
@@ -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