INNER CODE UNIT · Python

get_loss_optim

digantamisra98/Mish · PyTorch Benchmarks/train_imagenet.py:185

def get_loss_optim(model, device, lr, momentum, weight_decay):
    criterion = nn.CrossEntropyLoss().to(device)

    optimizer = torch.optim.SGD(
        model.parameters(), lr, momentum=momentum, weight_decay=weight_decay
    )
    return criterion, optimizer


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"]))

View source record →

📰 Research Paper
Loading…
⏳ Fetching content…