From a9e5be78d950d017a25adbb9f31feccc82ccc2a1 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Stardust=C2=B7=E5=87=8F?= <2225664821@qq.com> Date: Mon, 21 Aug 2023 17:18:13 +0800 Subject: [PATCH] Update utils.py --- utils.py | 10 ++++++++-- 1 file changed, 8 insertions(+), 2 deletions(-) diff --git a/utils.py b/utils.py index 98e0b87..fe03ba7 100644 --- a/utils.py +++ b/utils.py @@ -22,6 +22,12 @@ def load_checkpoint(checkpoint_path, model, optimizer=None, skip_optimizer=False learning_rate = checkpoint_dict['learning_rate'] if optimizer is not None and not skip_optimizer and checkpoint_dict['optimizer'] is not None: optimizer.load_state_dict(checkpoint_dict['optimizer']) + else: + new_opt_dict = optimizer.state_dict() + new_opt_dict_params = new_opt_dict['param_groups'][0]['params'] + new_opt_dict['param_groups'] = checkpoint_dict['optimizer']['param_groups'] + new_opt_dict['param_groups'][0]['params'] = new_opt_dict_params + optimizer.load_state_dict(new_opt_dict) saved_state_dict = checkpoint_dict['model'] if hasattr(model, 'module'): state_dict = model.module.state_dict() @@ -38,9 +44,9 @@ def load_checkpoint(checkpoint_path, model, optimizer=None, skip_optimizer=False print("error, %s is not in the checkpoint" % k) new_state_dict[k] = v if hasattr(model, 'module'): - model.module.load_state_dict(new_state_dict) + model.module.load_state_dict(new_state_dict, strict=False) else: - model.load_state_dict(new_state_dict) + model.load_state_dict(new_state_dict, strict=False) print("load ") logger.info("Loaded checkpoint '{}' (iteration {})".format( checkpoint_path, iteration))