PyTorch网络模型的保存与加载总共分为4种情况,大家要牢记。首先,我们要定义一个模型,以便下文所使用。
import torch.nn as nn
import torch.nn.functional as F
class MyModel(nn.Module):
def __init__(self):
super().__init__()
self.nn1 = nn.Linear(2, 3)
self.nn2 = nn.Linear(3, 6)
def forward(self, x):
x = F.relu(self.nn1(x))
return F.relu(self.nn2(x))
model = MyModel()
torch.save(model, "model.pth")
model = torch.load("model.pth", weights_only=False)
PyTorch模型将学习到的参数存储在一个名为state_dict的内部状态字典中,这些参数可以通过torch.save方法进行持久化。
在加载模型权重时,我们首先需要实例化模型类,因为该类定义了网络的结构。
# 保存模型
torch.save(model.state_dict(), "model_weights.pth")
# 加载模型
model2 = MyModel()
model2.load_state_dict(torch.load("model_weights.pth", weights_only=True))
model2.eval()
# 定义优化器
optimizer = optim.SGD(model.parameters(), lr=0.001, momentum=0.9)
# 保存优化器
EPOCH = 5
PATH = "model.pt"
LOSS = 0.4
torch.save({
"epoch": EPOCH,
"model_state_dict": net.state_dict(),
"optimizer_state_dict": optimizer.state_dict(),
"loss": LOSS,
}, PATH)
# 加载优化器
model = MyModel()
optimizer = optim.SGD(model.parameters(), lr=0.001, momentum=0.9)
checkpoint = torch.load(PATH, weights_only=True)
model.load_state_dict(checkpoint["model_state_dict"])
optimizer.load_state_dict(checkpoint["optimizer_state_dict"])
epoch = checkpoint["epoch"]
loss = checkpoint["loss"]
model.eval()
# - or -
model.train()
PATH = "model.pt"
torch.save(model.state_dict(), PATH)
device = torch.device("cpu")
model = MyModel()
model.load_state_dict(torch.load(PATH, map_location=device, weights_only=True))
# Save
torch.save(net.state_dict(), PATH)
# Load
device = torch.device("cuda")
model = MyModel()
# Choose whatever GPU device number you want
model.load_state_dict(torch.load(PATH, map_location="cuda:0"))
# Make sure to call input = input.to(device) on any input tensors that you feed to the model
model.to(device)
torch.save(net.state_dict(), PATH)
device = torch.device("cuda")
model = MyModel()
model.load_state_dict(torch.load(PATH))
model.to(device)