Fix bug in lora training && Update zero3 cpu offload && Teacache thresholds (#209)
This commit is contained in:
@@ -1221,7 +1221,8 @@ def main():
|
||||
|
||||
initial_global_step = global_step
|
||||
|
||||
pkl_path = os.path.join(os.path.join(args.output_dir, path), "sampler_pos_start.pkl")
|
||||
checkpoint_folder_path = os.path.join(args.output_dir, path)
|
||||
pkl_path = os.path.join(checkpoint_folder_path, "sampler_pos_start.pkl")
|
||||
if os.path.exists(pkl_path):
|
||||
with open(pkl_path, 'rb') as file:
|
||||
_, first_epoch = pickle.load(file)
|
||||
@@ -1229,13 +1230,62 @@ 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, safe_open
|
||||
state_dict = load_file(os.path.join(os.path.join(args.output_dir, path), "lora_diffusion_pytorch_model.safetensors"))
|
||||
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)}")
|
||||
|
||||
accelerator.print(f"Resuming from checkpoint {path}")
|
||||
accelerator.load_state(os.path.join(args.output_dir, path))
|
||||
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 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):
|
||||
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)
|
||||
accelerator.print("accelerator.load_state() completed for zero_stage 3.")
|
||||
|
||||
else:
|
||||
initial_global_step = 0
|
||||
|
||||
|
||||
@@ -1289,7 +1289,8 @@ def main():
|
||||
|
||||
initial_global_step = global_step
|
||||
|
||||
pkl_path = os.path.join(os.path.join(args.output_dir, path), "sampler_pos_start.pkl")
|
||||
checkpoint_folder_path = os.path.join(args.output_dir, path)
|
||||
pkl_path = os.path.join(checkpoint_folder_path, "sampler_pos_start.pkl")
|
||||
if os.path.exists(pkl_path):
|
||||
with open(pkl_path, 'rb') as file:
|
||||
_, first_epoch = pickle.load(file)
|
||||
@@ -1297,13 +1298,63 @@ 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, safe_open
|
||||
state_dict = load_file(os.path.join(os.path.join(args.output_dir, path), "lora_diffusion_pytorch_model.safetensors"))
|
||||
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:
|
||||
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
|
||||
|
||||
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):
|
||||
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)
|
||||
accelerator.print("accelerator.load_state() completed for zero_stage 3.")
|
||||
|
||||
accelerator.print(f"Resuming from checkpoint {path}")
|
||||
accelerator.load_state(os.path.join(args.output_dir, path))
|
||||
else:
|
||||
initial_global_step = 0
|
||||
|
||||
|
||||
@@ -1307,7 +1307,8 @@ def main():
|
||||
|
||||
initial_global_step = global_step
|
||||
|
||||
pkl_path = os.path.join(os.path.join(args.output_dir, path), "sampler_pos_start.pkl")
|
||||
checkpoint_folder_path = os.path.join(args.output_dir, path)
|
||||
pkl_path = os.path.join(checkpoint_folder_path, "sampler_pos_start.pkl")
|
||||
if os.path.exists(pkl_path):
|
||||
with open(pkl_path, 'rb') as file:
|
||||
_, first_epoch = pickle.load(file)
|
||||
@@ -1316,13 +1317,62 @@ def main():
|
||||
print(f"Load pkl from {pkl_path}. Get first_epoch = {first_epoch}.")
|
||||
|
||||
if zero_stage != 3:
|
||||
from safetensors.torch import load_file, safe_open
|
||||
state_dict = load_file(os.path.join(os.path.join(args.output_dir, path), "lora_diffusion_pytorch_model.safetensors"))
|
||||
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)}")
|
||||
|
||||
accelerator.print(f"Resuming from checkpoint {path}")
|
||||
accelerator.load_state(os.path.join(args.output_dir, path))
|
||||
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 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):
|
||||
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)
|
||||
accelerator.print("accelerator.load_state() completed for zero_stage 3.")
|
||||
|
||||
else:
|
||||
initial_global_step = 0
|
||||
|
||||
|
||||
@@ -1296,7 +1296,8 @@ def main():
|
||||
|
||||
initial_global_step = global_step
|
||||
|
||||
pkl_path = os.path.join(os.path.join(args.output_dir, path), "sampler_pos_start.pkl")
|
||||
checkpoint_folder_path = os.path.join(args.output_dir, path)
|
||||
pkl_path = os.path.join(checkpoint_folder_path, "sampler_pos_start.pkl")
|
||||
if os.path.exists(pkl_path):
|
||||
with open(pkl_path, 'rb') as file:
|
||||
_, first_epoch = pickle.load(file)
|
||||
@@ -1305,13 +1306,62 @@ def main():
|
||||
print(f"Load pkl from {pkl_path}. Get first_epoch = {first_epoch}.")
|
||||
|
||||
if zero_stage != 3:
|
||||
from safetensors.torch import load_file, safe_open
|
||||
state_dict = load_file(os.path.join(os.path.join(args.output_dir, path), "lora_diffusion_pytorch_model.safetensors"))
|
||||
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)}")
|
||||
|
||||
accelerator.print(f"Resuming from checkpoint {path}")
|
||||
accelerator.load_state(os.path.join(args.output_dir, path))
|
||||
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 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):
|
||||
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)
|
||||
accelerator.print("accelerator.load_state() completed for zero_stage 3.")
|
||||
|
||||
else:
|
||||
initial_global_step = 0
|
||||
|
||||
|
||||
Reference in New Issue
Block a user