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
@@ -0,0 +1,28 @@
{
"bf16": {
"enabled": true
},
"train_micro_batch_size_per_gpu": 1,
"train_batch_size": "auto",
"gradient_accumulation_steps": "auto",
"gradient_clipping": "auto",
"steps_per_print": 2000,
"wall_clock_breakdown": false,
"zero_optimization": {
"stage": 3,
"overlap_comm": true,
"contiguous_gradients": true,
"reduce_bucket_size": 5e8,
"sub_group_size": 1e9,
"stage3_max_live_parameters": 1e9,
"stage3_max_reuse_distance": 1e9,
"stage3_gather_16bit_weights_on_model_save": "auto",
"offload_optimizer": {
"device": "cpu"
},
"offload_param": {
"device": "cpu"
}
}
}
+6 -1
View File
@@ -103,8 +103,13 @@ if __name__ == '__main__':
# Support TeaCache.
enable_teacache = True
# Recommended to be set between 0.05 and 0.20. A larger threshold can cache more steps, speeding up the inference process,
# Recommended to be set between 0.05 and 0.30. A larger threshold can cache more steps, speeding up the inference process,
# but it may cause slight differences between the generated content and the original content.
# # --------------------------------------------------------------------------------------------------- #
# | Model Name | threshold | Model Name | threshold | Model Name | threshold |
# | Wan2.1-T2V-1.3B | 0.05~0.10 | Wan2.1-T2V-14B | 0.10~0.15 | Wan2.1-I2V-14B-720P | 0.20~0.30 |
# | Wan2.1-I2V-14B-480P | 0.20~0.25 | Wan2.1-Fun-*-1.3B-* | 0.05~0.10 | Wan2.1-Fun-*-14B-* | 0.20~0.30 |
# # --------------------------------------------------------------------------------------------------- #
teacache_threshold = 0.10
# The number of steps to skip TeaCache at the beginning of the inference process, which can
# reduce the impact of TeaCache on generated video quality.
+6 -1
View File
@@ -119,8 +119,13 @@ if __name__ == '__main__':
# Support TeaCache.
enable_teacache = True
# Recommended to be set between 0.05 and 0.20. A larger threshold can cache more steps, speeding up the inference process,
# Recommended to be set between 0.05 and 0.30. A larger threshold can cache more steps, speeding up the inference process,
# but it may cause slight differences between the generated content and the original content.
# # --------------------------------------------------------------------------------------------------- #
# | Model Name | threshold | Model Name | threshold | Model Name | threshold |
# | Wan2.1-T2V-1.3B | 0.05~0.10 | Wan2.1-T2V-14B | 0.10~0.15 | Wan2.1-I2V-14B-720P | 0.20~0.30 |
# | Wan2.1-I2V-14B-480P | 0.20~0.25 | Wan2.1-Fun-*-1.3B-* | 0.05~0.10 | Wan2.1-Fun-*-14B-* | 0.20~0.30 |
# # --------------------------------------------------------------------------------------------------- #
teacache_threshold = 0.10
# The number of steps to skip TeaCache at the beginning of the inference process, which can
# reduce the impact of TeaCache on generated video quality.
+6 -1
View File
@@ -55,8 +55,13 @@ compile_dit = False
# Support TeaCache.
enable_teacache = True
# Recommended to be set between 0.05 and 0.20. A larger threshold can cache more steps, speeding up the inference process,
# Recommended to be set between 0.05 and 0.30. A larger threshold can cache more steps, speeding up the inference process,
# but it may cause slight differences between the generated content and the original content.
# # --------------------------------------------------------------------------------------------------- #
# | Model Name | threshold | Model Name | threshold | Model Name | threshold |
# | Wan2.1-T2V-1.3B | 0.05~0.10 | Wan2.1-T2V-14B | 0.10~0.15 | Wan2.1-I2V-14B-720P | 0.20~0.30 |
# | Wan2.1-I2V-14B-480P | 0.20~0.25 | Wan2.1-Fun-*-1.3B-* | 0.05~0.10 | Wan2.1-Fun-*-14B-* | 0.20~0.30 |
# # --------------------------------------------------------------------------------------------------- #
teacache_threshold = 0.10
# The number of steps to skip TeaCache at the beginning of the inference process, which can
# reduce the impact of TeaCache on generated video quality.
+6 -1
View File
@@ -54,8 +54,13 @@ compile_dit = False
# TeaCache config
enable_teacache = True
# Recommended to be set between 0.05 and 0.20. A larger threshold can cache more steps, speeding up the inference process,
# Recommended to be set between 0.05 and 0.30. A larger threshold can cache more steps, speeding up the inference process,
# but it may cause slight differences between the generated content and the original content.
# # --------------------------------------------------------------------------------------------------- #
# | Model Name | threshold | Model Name | threshold | Model Name | threshold |
# | Wan2.1-T2V-1.3B | 0.05~0.10 | Wan2.1-T2V-14B | 0.10~0.15 | Wan2.1-I2V-14B-720P | 0.20~0.30 |
# | Wan2.1-I2V-14B-480P | 0.20~0.25 | Wan2.1-Fun-*-1.3B-* | 0.05~0.10 | Wan2.1-Fun-*-14B-* | 0.20~0.30 |
# # --------------------------------------------------------------------------------------------------- #
teacache_threshold = 0.10
# The number of steps to skip TeaCache at the beginning of the inference process, which can
# reduce the impact of TeaCache on generated video quality.
+6 -1
View File
@@ -103,8 +103,13 @@ if __name__ == '__main__':
# Support TeaCache.
enable_teacache = True
# Recommended to be set between 0.05 and 0.20. A larger threshold can cache more steps, speeding up the inference process,
# Recommended to be set between 0.05 and 0.30. A larger threshold can cache more steps, speeding up the inference process,
# but it may cause slight differences between the generated content and the original content.
# # --------------------------------------------------------------------------------------------------- #
# | Model Name | threshold | Model Name | threshold | Model Name | threshold |
# | Wan2.1-T2V-1.3B | 0.05~0.10 | Wan2.1-T2V-14B | 0.10~0.15 | Wan2.1-I2V-14B-720P | 0.20~0.30 |
# | Wan2.1-I2V-14B-480P | 0.20~0.25 | Wan2.1-Fun-*-1.3B-* | 0.05~0.10 | Wan2.1-Fun-*-14B-* | 0.20~0.30 |
# # --------------------------------------------------------------------------------------------------- #
teacache_threshold = 0.10
# The number of steps to skip TeaCache at the beginning of the inference process, which can
# reduce the impact of TeaCache on generated video quality.
+6 -1
View File
@@ -120,8 +120,13 @@ if __name__ == '__main__':
# Support TeaCache.
enable_teacache = True
# Recommended to be set between 0.05 and 0.20. A larger threshold can cache more steps, speeding up the inference process,
# Recommended to be set between 0.05 and 0.30. A larger threshold can cache more steps, speeding up the inference process,
# but it may cause slight differences between the generated content and the original content.
# # --------------------------------------------------------------------------------------------------- #
# | Model Name | threshold | Model Name | threshold | Model Name | threshold |
# | Wan2.1-T2V-1.3B | 0.05~0.10 | Wan2.1-T2V-14B | 0.10~0.15 | Wan2.1-I2V-14B-720P | 0.20~0.30 |
# | Wan2.1-I2V-14B-480P | 0.20~0.25 | Wan2.1-Fun-*-1.3B-* | 0.05~0.10 | Wan2.1-Fun-*-14B-* | 0.20~0.30 |
# # --------------------------------------------------------------------------------------------------- #
teacache_threshold = 0.10
# The number of steps to skip TeaCache at the beginning of the inference process, which can
# reduce the impact of TeaCache on generated video quality.
@@ -134,8 +134,13 @@ if __name__ == '__main__':
# Support TeaCache.
enable_teacache = True
# Recommended to be set between 0.05 and 0.20. A larger threshold can cache more steps, speeding up the inference process,
# Recommended to be set between 0.05 and 0.30. A larger threshold can cache more steps, speeding up the inference process,
# but it may cause slight differences between the generated content and the original content.
# # --------------------------------------------------------------------------------------------------- #
# | Model Name | threshold | Model Name | threshold | Model Name | threshold |
# | Wan2.1-T2V-1.3B | 0.05~0.10 | Wan2.1-T2V-14B | 0.10~0.15 | Wan2.1-I2V-14B-720P | 0.20~0.30 |
# | Wan2.1-I2V-14B-480P | 0.20~0.25 | Wan2.1-Fun-*-1.3B-* | 0.05~0.10 | Wan2.1-Fun-*-14B-* | 0.20~0.30 |
# # --------------------------------------------------------------------------------------------------- #
teacache_threshold = 0.10
# The number of steps to skip TeaCache at the beginning of the inference process, which can
# reduce the impact of TeaCache on generated video quality.
+6 -1
View File
@@ -55,8 +55,13 @@ compile_dit = False
# Support TeaCache.
enable_teacache = True
# Recommended to be set between 0.05 and 0.20. A larger threshold can cache more steps, speeding up the inference process,
# Recommended to be set between 0.05 and 0.30. A larger threshold can cache more steps, speeding up the inference process,
# but it may cause slight differences between the generated content and the original content.
# # --------------------------------------------------------------------------------------------------- #
# | Model Name | threshold | Model Name | threshold | Model Name | threshold |
# | Wan2.1-T2V-1.3B | 0.05~0.10 | Wan2.1-T2V-14B | 0.10~0.15 | Wan2.1-I2V-14B-720P | 0.20~0.30 |
# | Wan2.1-I2V-14B-480P | 0.20~0.25 | Wan2.1-Fun-*-1.3B-* | 0.05~0.10 | Wan2.1-Fun-*-14B-* | 0.20~0.30 |
# # --------------------------------------------------------------------------------------------------- #
teacache_threshold = 0.10
# The number of steps to skip TeaCache at the beginning of the inference process, which can
# reduce the impact of TeaCache on generated video quality.
+6 -1
View File
@@ -55,8 +55,13 @@ compile_dit = False
# Support TeaCache.
enable_teacache = True
# Recommended to be set between 0.05 and 0.20. A larger threshold can cache more steps, speeding up the inference process,
# Recommended to be set between 0.05 and 0.30. A larger threshold can cache more steps, speeding up the inference process,
# but it may cause slight differences between the generated content and the original content.
# # --------------------------------------------------------------------------------------------------- #
# | Model Name | threshold | Model Name | threshold | Model Name | threshold |
# | Wan2.1-T2V-1.3B | 0.05~0.10 | Wan2.1-T2V-14B | 0.10~0.15 | Wan2.1-I2V-14B-720P | 0.20~0.30 |
# | Wan2.1-I2V-14B-480P | 0.20~0.25 | Wan2.1-Fun-*-1.3B-* | 0.05~0.10 | Wan2.1-Fun-*-14B-* | 0.20~0.30 |
# # --------------------------------------------------------------------------------------------------- #
teacache_threshold = 0.10
# The number of steps to skip TeaCache at the beginning of the inference process, which can
# reduce the impact of TeaCache on generated video quality.
+6 -1
View File
@@ -58,8 +58,13 @@ compile_dit = False
# Support TeaCache.
enable_teacache = True
# Recommended to be set between 0.05 and 0.20. A larger threshold can cache more steps, speeding up the inference process,
# Recommended to be set between 0.05 and 0.30. A larger threshold can cache more steps, speeding up the inference process,
# but it may cause slight differences between the generated content and the original content.
# # --------------------------------------------------------------------------------------------------- #
# | Model Name | threshold | Model Name | threshold | Model Name | threshold |
# | Wan2.1-T2V-1.3B | 0.05~0.10 | Wan2.1-T2V-14B | 0.10~0.15 | Wan2.1-I2V-14B-720P | 0.20~0.30 |
# | Wan2.1-I2V-14B-480P | 0.20~0.25 | Wan2.1-Fun-*-1.3B-* | 0.05~0.10 | Wan2.1-Fun-*-14B-* | 0.20~0.30 |
# # --------------------------------------------------------------------------------------------------- #
teacache_threshold = 0.10
# The number of steps to skip TeaCache at the beginning of the inference process, which can
# reduce the impact of TeaCache on generated video quality.
@@ -58,8 +58,13 @@ compile_dit = False
# Support TeaCache.
enable_teacache = True
# Recommended to be set between 0.05 and 0.20. A larger threshold can cache more steps, speeding up the inference process,
# Recommended to be set between 0.05 and 0.30. A larger threshold can cache more steps, speeding up the inference process,
# but it may cause slight differences between the generated content and the original content.
# # --------------------------------------------------------------------------------------------------- #
# | Model Name | threshold | Model Name | threshold | Model Name | threshold |
# | Wan2.1-T2V-1.3B | 0.05~0.10 | Wan2.1-T2V-14B | 0.10~0.15 | Wan2.1-I2V-14B-720P | 0.20~0.30 |
# | Wan2.1-I2V-14B-480P | 0.20~0.25 | Wan2.1-Fun-*-1.3B-* | 0.05~0.10 | Wan2.1-Fun-*-14B-* | 0.20~0.30 |
# # --------------------------------------------------------------------------------------------------- #
teacache_threshold = 0.10
# The number of steps to skip TeaCache at the beginning of the inference process, which can
# reduce the impact of TeaCache on generated video quality.
@@ -58,8 +58,13 @@ compile_dit = False
# Support TeaCache.
enable_teacache = True
# Recommended to be set between 0.05 and 0.20. A larger threshold can cache more steps, speeding up the inference process,
# Recommended to be set between 0.05 and 0.30. A larger threshold can cache more steps, speeding up the inference process,
# but it may cause slight differences between the generated content and the original content.
# # --------------------------------------------------------------------------------------------------- #
# | Model Name | threshold | Model Name | threshold | Model Name | threshold |
# | Wan2.1-T2V-1.3B | 0.05~0.10 | Wan2.1-T2V-14B | 0.10~0.15 | Wan2.1-I2V-14B-720P | 0.20~0.30 |
# | Wan2.1-I2V-14B-480P | 0.20~0.25 | Wan2.1-Fun-*-1.3B-* | 0.05~0.10 | Wan2.1-Fun-*-14B-* | 0.20~0.30 |
# # --------------------------------------------------------------------------------------------------- #
teacache_threshold = 0.10
# The number of steps to skip TeaCache at the beginning of the inference process, which can
# reduce the impact of TeaCache on generated video quality.
+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
+1 -1
View File
@@ -1023,7 +1023,7 @@ class WanTransformer3DModel(ModelMixin, ConfigMixin, FromOriginalModelMixin):
if self.teacache is not None:
if not self.should_calc:
previous_residual = self.teacache.previous_residual_cond if cond_flag else self.teacache.previous_residual_uncond
x = x + previous_residual.to(x.device)
x = x + previous_residual.to(x.device)[-x.size()[0]:,]
else:
ori_x = x.clone().cpu() if self.teacache.offload else x.clone()
+53 -1
View File
@@ -255,4 +255,56 @@ def timer(func):
end_time = time.time()
print(f"function {func.__name__} running for {end_time - start_time} seconds")
return result
return wrapper
return wrapper
def timer_record(model_name=""):
def decorator(func):
def wrapper(*args, **kwargs):
torch.cuda.synchronize()
start_time = time.time()
result = func(*args, **kwargs)
torch.cuda.synchronize()
end_time = time.time()
import torch.distributed as dist
if dist.is_initialized():
if dist.get_rank() == 0:
time_sum = end_time - start_time
print('# --------------------------------------------------------- #')
print(f'# {model_name} time: {time_sum}s')
print('# --------------------------------------------------------- #')
_write_to_excel(model_name, time_sum)
else:
time_sum = end_time - start_time
print('# --------------------------------------------------------- #')
print(f'# {model_name} time: {time_sum}s')
print('# --------------------------------------------------------- #')
_write_to_excel(model_name, time_sum)
return result
return wrapper
return decorator
def _write_to_excel(model_name, time_sum):
import pandas as pd
import os
row_env = os.environ.get(f"{model_name}_EXCEL_ROW", "1") # 默认第1行
col_env = os.environ.get(f"{model_name}_EXCEL_COL", "1") # 默认第A列
file_path = os.environ.get(f"EXCEL_FILE", "timing_records.xlsx") # 默认文件名
try:
df = pd.read_excel(file_path, sheet_name="Sheet1", header=None)
except FileNotFoundError:
df = pd.DataFrame()
row_idx = int(row_env)
col_idx = int(col_env)
if row_idx >= len(df):
df = pd.concat([df, pd.DataFrame([ [None] * (len(df.columns) if not df.empty else 0) ] * (row_idx - len(df) + 1))], ignore_index=True)
if col_idx >= len(df.columns):
df = pd.concat([df, pd.DataFrame(columns=range(len(df.columns), col_idx + 1))], axis=1)
df.iloc[row_idx, col_idx] = time_sum
df.to_excel(file_path, index=False, header=False, sheet_name="Sheet1")