INNER CODE UNIT · Python
get_model_checkpoint
digantamisra98/Mish · PyTorch Benchmarks/train_imagenet.py:194
def get_model_checkpoint(path, model, optimizer):
if os.path.isfile(path):
print("=> loading checkpoint '{}'".format(path))
checkpoint = torch.load(path)
start_epoch = checkpoint["epoch"]
model.load_state_dict(checkpoint["state_dict"])
if "optimizer" in checkpoint:
optimizer.load_state_dict(checkpoint["optimizer"])
print("=> loaded checkpoint '{}' (epoch {})".format(path, checkpoint["epoch"]))
return start_epoch
else:
print("=> no checkpoint found at '{}'".format(path))
exit()
def get_dataloader(root_dir, is_train, batch_size, workers):
dir_name = "train" if is_train else "val"
data_dir = os.path.join(root_dir, dir_name)