Pytorch保存我们训练好的模型,然后加载用于测试
第一种方法
(1)保存
torch.save(model.state_dict(), PATH)
# example
torch.save(resnet50.state_dict(),'ckp/model.pth')
(2)恢复
model = ModelClass(*args, **kwargs)
model.load_state_dict(torch.load(PATH))
#example
resnet=resnet50(pretrained=True)
resnet.load_state_dict(torch.load('ckp/model.pth'))
第二种方法
(1)保存
torch.save (model, PATH)
(2)恢复
model = torch.load(PATH)
本文详细介绍了如何使用PyTorch保存和加载模型的方法,包括两种常见方式:仅保存和加载模型参数(state_dict),以及保存和加载整个模型。通过具体示例展示了不同场景下模型保存与恢复的步骤。

2万+

被折叠的 条评论
为什么被折叠?



