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)

View source record →

📰 Research Paper
Loading…
⏳ Fetching content…