Update utils.py
This commit is contained in:
10
utils.py
10
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))
|
||||
|
||||
Reference in New Issue
Block a user