Update Peft Lora && Update Readme (#376)
This commit is contained in:
@@ -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
@@ -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)の下でライセンスされています。
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 \
|
||||
|
||||
@@ -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
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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"
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
```
|
||||
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
```
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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"
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
```
|
||||
@@ -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
|
||||
```
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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"
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
```
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
@@ -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)
|
||||
|
||||
@@ -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"
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
```
|
||||
@@ -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
|
||||
```
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
```
|
||||
@@ -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
|
||||
```
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user