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']
|
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))
|
||||||
|
|||||||
Reference in New Issue
Block a user