Update Peft Lora && Update Readme (#376)
This commit is contained in:
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user