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))