feat: 优化日志打印

This commit is contained in:
源文雨
2023-09-04 00:47:33 +08:00
parent 5a6f824537
commit 391dea85de
3 changed files with 23 additions and 10 deletions

View File

@@ -11,8 +11,7 @@ import torch
MATPLOTLIB_FLAG = False
logging.basicConfig(stream=sys.stdout, level=logging.DEBUG)
logger = logging
logger = logging.getLogger(__name__)
def load_checkpoint(checkpoint_path, model, optimizer=None, skip_optimizer=False):
@@ -42,13 +41,12 @@ def load_checkpoint(checkpoint_path, model, optimizer=None, skip_optimizer=False
new_state_dict[k] = saved_state_dict[k]
assert saved_state_dict[k].shape == v.shape, (saved_state_dict[k].shape, v.shape)
except:
print("error, %s is not in the checkpoint" % k)
logger.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, strict=False)
else:
model.load_state_dict(new_state_dict, strict=False)
print("load ")
logger.info("Loaded checkpoint '{}' (iteration {})".format(
checkpoint_path, iteration))
return model, optimizer, learning_rate, iteration