本文发表于 655 天前,其中的信息可能已经事过境迁

模型训练完成之后,我们当然需要将模型保存下来,以便后续使用。

在上一章节代码中,我们已经提到了模型的保存与加载,这里我们再详细介绍一下。

模型保存

Pytorch提供了torch.save函数用于保存模型。

python
import torch
import torch.nn as nn

class MyModel(nn.Module):
    def __init__(self):
        super(MyModel, self).__init__()
        self.conv = nn.Conv2d(3, 16, 3, 1, 1)
        self.fc = nn.Linear(16 * 32 * 32, 10)

    def forward(self, x):
        x = self.conv(x)
        x = x.view(x.size(0), -1)
        x = self.fc(x)
        return x

model = MyModel()

# 保存模型
torch.save(model.state_dict(), 'model.pth')

model.state_dict()是一个字典类型,包含了模型的所有参数,torch.save函数将其保存到文件model.pth中。

当然,不止于存储为pytorch的格式,还可以存储为ONNX格式,以便于在其他框架中使用。

ONNX(Open Neural Network Exchange)是一种开放式的文件格式,用于存储和交换训练好的机器学习模型。它使得不同的人工智能框架(如PyTorch、TensorFlow)可以共享模型,促进了模型在不同平台之间的迁移和复用。

对于ONNX格式的模型保存,可以使用torch.onnx.export函数。

python
import torch
import torch.nn as nn

# 保存模型为ONNX格式
dummy_input = torch.randn(1, 3, 32, 32)
torch.onnx.export(model, dummy_input, 'model.onnx')

torch.onnx.export函数将模型保存为ONNX格式,dummy_input是一个输入样本,用于指定输入的形状。

模型加载

Pytorch提供了torch.load函数用于加载模型。

python
import torch
import torch.nn as nn

# 加载模型
model.load_state_dict(torch.load('model.pth'))

torch.load函数将保存的模型参数加载到模型中。

既然能导出为ONNX格式,那么我们也可以加载ONNX格式的模型。

python
import torch
import onnx

# 加载ONNX格式的模型
model = onnx.load('model.onnx')

onnx.load函数将ONNX格式的模型加载到内存中。

评论 隐私政策