Update utils.py

This commit is contained in:
Stardust·减
2023-08-21 17:18:13 +08:00
committed by GitHub
parent 12c9fb9ee8
commit a9e5be78d9

View File

@@ -22,6 +22,12 @@ def load_checkpoint(checkpoint_path, model, optimizer=None, skip_optimizer=False
learning_rate = checkpoint_dict['learning_rate'] learning_rate = checkpoint_dict['learning_rate']
if optimizer is not None and not skip_optimizer and checkpoint_dict['optimizer'] is not None: if optimizer is not None and not skip_optimizer and checkpoint_dict['optimizer'] is not None:
optimizer.load_state_dict(checkpoint_dict['optimizer']) 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'] saved_state_dict = checkpoint_dict['model']
if hasattr(model, 'module'): if hasattr(model, 'module'):
state_dict = model.module.state_dict() 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) print("error, %s is not in the checkpoint" % k)
new_state_dict[k] = v new_state_dict[k] = v
if hasattr(model, 'module'): if hasattr(model, 'module'):
model.module.load_state_dict(new_state_dict) model.module.load_state_dict(new_state_dict, strict=False)
else: else:
model.load_state_dict(new_state_dict) model.load_state_dict(new_state_dict, strict=False)
print("load ") print("load ")
logger.info("Loaded checkpoint '{}' (iteration {})".format( logger.info("Loaded checkpoint '{}' (iteration {})".format(
checkpoint_path, iteration)) checkpoint_path, iteration))