Update Peft Lora && Update Readme (#376)

This commit is contained in:
Bubbliiiing
2025-11-20 17:44:22 +08:00
committed by GitHub
parent ccc3b1055e
commit 037a2e8360
79 changed files with 2499 additions and 849 deletions
+89 -29
View File
@@ -593,6 +593,9 @@ def parse_args():
default=64,
help=("The dimension of the LoRA update matrices."),
)
parser.add_argument(
"--use_peft_lora", action="store_true", help="Whether or not to use peft lora."
)
parser.add_argument(
"--train_text_encoder",
action="store_true",
@@ -768,6 +771,12 @@ def parse_args():
default=None,
help=("The module is not trained in loras. "),
)
parser.add_argument(
"--target_name",
type=str,
default=None,
help=("The module is trained in loras. "),
)
args = parser.parse_args()
env_local_rank = int(os.environ.get("LOCAL_RANK", -1))
@@ -954,16 +963,25 @@ def main():
transformer3d.requires_grad_(False)
# Lora will work with this...
network = create_network(
1.0,
args.rank,
args.network_alpha,
text_encoder,
transformer3d,
neuron_dropout=None,
skip_name=args.lora_skip_name,
)
network.apply_to(text_encoder, transformer3d, args.train_text_encoder and not args.training_with_video_token_length, True)
if args.use_peft_lora:
from peft import LoraConfig, inject_adapter_in_model, get_peft_model_state_dict
lora_config = LoraConfig(r=args.rank, lora_alpha=args.network_alpha, target_modules=args.target_name.split(","))
transformer3d = inject_adapter_in_model(lora_config, transformer3d)
network = None
else:
network = create_network(
1.0,
args.rank,
args.network_alpha,
text_encoder,
transformer3d,
neuron_dropout=None,
target_name=args.target_name,
skip_name=args.lora_skip_name,
)
network = network.to(weight_dtype)
network.apply_to(text_encoder, transformer3d, args.train_text_encoder and not args.training_with_video_token_length, True)
if args.transformer_path is not None:
print(f"From checkpoint: {args.transformer_path}")
@@ -999,13 +1017,14 @@ def main():
accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True)
if accelerator.is_main_process:
from safetensors.torch import save_file
safetensor_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model.safetensors")
network_state_dict = {}
for key in accelerate_state_dict:
if "network" in key:
network_state_dict[key.replace("network.", "")] = accelerate_state_dict[key].to(weight_dtype)
if args.use_peft_lora:
network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1]), accelerate_state_dict)
else:
network_state_dict = {}
for key in accelerate_state_dict:
if "network" in key:
network_state_dict[key.replace("network.", "")] = accelerate_state_dict[key].to(weight_dtype)
save_file(network_state_dict, safetensor_save_path, metadata={"format": "pt"})
with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file:
@@ -1020,8 +1039,18 @@ def main():
print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.")
elif zero_stage == 3:
# create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format
def save_model_hook(models, weights, output_dir):
accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True)
if accelerator.is_main_process:
from safetensors.torch import save_file
safetensor_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model.safetensors")
if args.use_peft_lora:
network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1]), accelerate_state_dict)
else:
network_state_dict = accelerate_state_dict
save_file(network_state_dict, safetensor_save_path, metadata={"format": "pt"})
with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file:
pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file)
@@ -1033,10 +1062,15 @@ def main():
batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0)
print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.")
else:
# create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format
def save_model_hook(models, weights, output_dir):
if accelerator.is_main_process:
safetensor_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model.safetensors")
save_model(safetensor_save_path, accelerator.unwrap_model(models[-1]))
if args.use_peft_lora:
save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(models[-1])))
else:
save_model(safetensor_save_path, accelerator.unwrap_model(models[-1]))
if not args.use_deepspeed:
for _ in range(len(weights)):
weights.pop()
@@ -1090,9 +1124,14 @@ def main():
else:
optimizer_cls = torch.optim.AdamW
logging.info("Add network parameters")
trainable_params = list(filter(lambda p: p.requires_grad, network.parameters()))
trainable_params_optim = network.prepare_optimizer_params(args.learning_rate / 2, args.learning_rate, args.learning_rate)
if args.use_peft_lora:
logging.info("Add peft parameters")
trainable_params = list(filter(lambda p: p.requires_grad, transformer3d.parameters()))
trainable_params_optim = list(filter(lambda p: p.requires_grad, transformer3d.parameters()))
else:
logging.info("Add network parameters")
trainable_params = list(filter(lambda p: p.requires_grad, network.parameters()))
trainable_params_optim = network.prepare_optimizer_params(args.learning_rate / 2, args.learning_rate, args.learning_rate)
if args.use_came:
optimizer = optimizer_cls(
@@ -1359,9 +1398,13 @@ def main():
)
# Prepare everything with our `accelerator`.
if fsdp_stage != 0:
if args.use_peft_lora:
transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare(
transformer3d, optimizer, train_dataloader, lr_scheduler
)
elif fsdp_stage != 0:
transformer3d.network = network
transformer3d = transformer3d.to(weight_dtype)
transformer3d = transformer3d.to(dtype=weight_dtype)
transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare(
transformer3d, optimizer, train_dataloader, lr_scheduler
)
@@ -1370,14 +1413,16 @@ def main():
network, optimizer, train_dataloader, lr_scheduler
)
if zero_stage == 3:
if zero_stage != 0 and not args.use_peft_lora:
from functools import partial
from videox_fun.dist import set_multi_gpus_devices, shard_model
shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype)
transformer3d = shard_fn(transformer3d)
if fsdp_stage != 0:
if fsdp_stage != 0 or zero_stage != 0:
from functools import partial
from videox_fun.dist import set_multi_gpus_devices, shard_model
shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype)
text_encoder = shard_fn(text_encoder)
@@ -1519,6 +1564,10 @@ def main():
def save_model(ckpt_file, unwrapped_nw):
os.makedirs(args.output_dir, exist_ok=True)
accelerator.print(f"\nsaving checkpoint: {ckpt_file}")
if isinstance(unwrapped_nw, dict):
from safetensors.torch import save_file
save_file(unwrapped_nw, ckpt_file, metadata={"format": "pt"})
return ckpt_file
unwrapped_nw.save_weights(ckpt_file, weight_dtype, None)
progress_bar = tqdm(
@@ -1897,9 +1946,14 @@ def main():
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
if not args.save_state:
safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors")
save_model(safetensor_save_path, accelerator.unwrap_model(network))
logger.info(f"Saved safetensor to {safetensor_save_path}")
if args.use_peft_lora:
safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors")
save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(transformer3d)))
logger.info(f"Saved safetensor to {safetensor_save_path}")
else:
safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors")
save_model(safetensor_save_path, accelerator.unwrap_model(network))
logger.info(f"Saved safetensor to {safetensor_save_path}")
else:
accelerator_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}")
accelerator.save_state(accelerator_save_path)
@@ -1948,8 +2002,14 @@ def main():
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
if not args.save_state:
safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors")
save_model(safetensor_save_path, accelerator.unwrap_model(network))
if args.use_peft_lora:
safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors")
save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(transformer3d)))
logger.info(f"Saved safetensor to {safetensor_save_path}")
else:
safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors")
save_model(safetensor_save_path, accelerator.unwrap_model(network))
logger.info(f"Saved safetensor to {safetensor_save_path}")
else:
accelerator_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}")
accelerator.save_state(accelerator_save_path)