diff --git a/config/zero_stage3_config_cpu_offload.json b/config/zero_stage3_config_cpu_offload.json new file mode 100644 index 0000000..19f779d --- /dev/null +++ b/config/zero_stage3_config_cpu_offload.json @@ -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" + } + } +} + diff --git a/examples/wan2.1/post_infer_queue.py b/examples/wan2.1/post_infer_queue.py index 0de7eb0..872d4de 100755 --- a/examples/wan2.1/post_infer_queue.py +++ b/examples/wan2.1/post_infer_queue.py @@ -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. diff --git a/examples/wan2.1/post_infer_queue_i2v.py b/examples/wan2.1/post_infer_queue_i2v.py index 585fd7d..f453cc3 100755 --- a/examples/wan2.1/post_infer_queue_i2v.py +++ b/examples/wan2.1/post_infer_queue_i2v.py @@ -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. diff --git a/examples/wan2.1/predict_i2v.py b/examples/wan2.1/predict_i2v.py index c3ab137..3fec9fc 100755 --- a/examples/wan2.1/predict_i2v.py +++ b/examples/wan2.1/predict_i2v.py @@ -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. diff --git a/examples/wan2.1/predict_t2v.py b/examples/wan2.1/predict_t2v.py index 01a6dc4..a00a8dc 100755 --- a/examples/wan2.1/predict_t2v.py +++ b/examples/wan2.1/predict_t2v.py @@ -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. diff --git a/examples/wan2.1_fun/post_infer_queue.py b/examples/wan2.1_fun/post_infer_queue.py index e6e89b6..bc6e868 100755 --- a/examples/wan2.1_fun/post_infer_queue.py +++ b/examples/wan2.1_fun/post_infer_queue.py @@ -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. diff --git a/examples/wan2.1_fun/post_infer_queue_i2v.py b/examples/wan2.1_fun/post_infer_queue_i2v.py index cfb9a90..d24d20b 100755 --- a/examples/wan2.1_fun/post_infer_queue_i2v.py +++ b/examples/wan2.1_fun/post_infer_queue_i2v.py @@ -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. diff --git a/examples/wan2.1_fun/post_infer_queue_v2v_control.py b/examples/wan2.1_fun/post_infer_queue_v2v_control.py index 99e6c85..23a6c58 100755 --- a/examples/wan2.1_fun/post_infer_queue_v2v_control.py +++ b/examples/wan2.1_fun/post_infer_queue_v2v_control.py @@ -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. diff --git a/examples/wan2.1_fun/predict_i2v.py b/examples/wan2.1_fun/predict_i2v.py index f2b0b71..4706215 100755 --- a/examples/wan2.1_fun/predict_i2v.py +++ b/examples/wan2.1_fun/predict_i2v.py @@ -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. diff --git a/examples/wan2.1_fun/predict_t2v.py b/examples/wan2.1_fun/predict_t2v.py index 7a76ef4..8f7583f 100755 --- a/examples/wan2.1_fun/predict_t2v.py +++ b/examples/wan2.1_fun/predict_t2v.py @@ -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. diff --git a/examples/wan2.1_fun/predict_v2v_control.py b/examples/wan2.1_fun/predict_v2v_control.py index c1a58e6..515897a 100755 --- a/examples/wan2.1_fun/predict_v2v_control.py +++ b/examples/wan2.1_fun/predict_v2v_control.py @@ -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. diff --git a/examples/wan2.1_fun/predict_v2v_control_camera.py b/examples/wan2.1_fun/predict_v2v_control_camera.py index 24f79f4..34953bd 100755 --- a/examples/wan2.1_fun/predict_v2v_control_camera.py +++ b/examples/wan2.1_fun/predict_v2v_control_camera.py @@ -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. diff --git a/examples/wan2.1_fun/predict_v2v_control_ref.py b/examples/wan2.1_fun/predict_v2v_control_ref.py index 6442105..8c269d9 100755 --- a/examples/wan2.1_fun/predict_v2v_control_ref.py +++ b/examples/wan2.1_fun/predict_v2v_control_ref.py @@ -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. diff --git a/scripts/cogvideox_fun/train_lora.py b/scripts/cogvideox_fun/train_lora.py index 6f20d34..a0c6d0d 100755 --- a/scripts/cogvideox_fun/train_lora.py +++ b/scripts/cogvideox_fun/train_lora.py @@ -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 diff --git a/scripts/wan2.1/train_lora.py b/scripts/wan2.1/train_lora.py index c32038d..ba51018 100755 --- a/scripts/wan2.1/train_lora.py +++ b/scripts/wan2.1/train_lora.py @@ -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 diff --git a/scripts/wan2.1_fun/train_control_lora.py b/scripts/wan2.1_fun/train_control_lora.py index cf2484e..f38e0fc 100755 --- a/scripts/wan2.1_fun/train_control_lora.py +++ b/scripts/wan2.1_fun/train_control_lora.py @@ -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 diff --git a/scripts/wan2.1_fun/train_lora.py b/scripts/wan2.1_fun/train_lora.py index 16ace98..b72a61c 100755 --- a/scripts/wan2.1_fun/train_lora.py +++ b/scripts/wan2.1_fun/train_lora.py @@ -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 diff --git a/videox_fun/models/wan_transformer3d.py b/videox_fun/models/wan_transformer3d.py index 42d2b07..3b3fb12 100755 --- a/videox_fun/models/wan_transformer3d.py +++ b/videox_fun/models/wan_transformer3d.py @@ -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() diff --git a/videox_fun/utils/utils.py b/videox_fun/utils/utils.py index 07b16dc..a769b2a 100755 --- a/videox_fun/utils/utils.py +++ b/videox_fun/utils/utils.py @@ -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 \ No newline at end of file + 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")