Fix bug in lora training && Update zero3 cpu offload && Teacache thresholds (#209)
This commit is contained in:
@@ -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"
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user