Fix bug in lora training && Update zero3 cpu offload && Teacache thresholds (#209)

This commit is contained in:
Bubbliiiing
2025-05-13 20:54:20 +08:00
committed by GitHub
parent f26f0a809b
commit ed4f10beef
19 changed files with 377 additions and 36 deletions
+55 -5
View File
@@ -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
+58 -7
View File
@@ -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
+55 -5
View File
@@ -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
+55 -5
View File
@@ -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