pytorch 加载(.pth)格式的模型

有一些非常流行的网络如 resnet、squeezenet、densenet等在pytorch里面都有,包括网络结构和训练好的模型。

pytorch自带模型网址:

按官网加载预训练好的模型:

import torchvision.models as models

# pretrained=True就可以使用预训练的模型
resnet18 = models.resnet18(pretrained=True)
print(resnet18)

报错如下:

requests.exceptions.ConnectionError: (Connection aborted., TimeoutError(10060, 由于连接方在一段时间后没有正确答复或连接的主机没有反应,连接尝试失败。, None, 10060, None))

主要是因为代码会去远端下载模型的参数,而国内的网一般连接不上,这是我们需要手动去下载你要的预训练网络。

通过地址下载,地址有两种获取方式:

1.从报错里面获取,上述代码运行时会出现这样一行信息:

Downloading: "https://download.pytorch.org/models/resnet18-5c106cde.pth" to C:UsersLuo/.torchmodels
esnet18-5c106cde.pth

复制这个网址到浏览器,有可能打不开,去掉https://,直接输入download.pytorch.org/models/resnet18-5c106cde.pth就可以下载了。

2.从pytorch的github下找模型的地址:

找到对应模型名称点进去找地址

下载好后自行保存,我是直接存在pytorch models里面

接下来就是运行这个.pth文件。首先要判断是保存的整个网络结构加参数呢,还是只保存了参数,可以测试一下。这是我的模型是squeezenet1_1,你可以测试自己下载的模型

import torch
pthfile = rE:anacondaappenvsluoLibsite-packages	orchvisionmodelssqueezenet1_1.pth
net = torch.load(pthfile)
print(net)

结果为

很明显就是只保存了参数,这是我们要换个方法加载模型

import torch
import torchvision.models as models

# pretrained=True就可以使用预训练的模型
net = models.squeezenet1_1(pretrained=False)
pthfile = rE:anacondaappenvsluoLibsite-packages	orchvisionmodelssqueezenet1_1.pth
net.load_state_dict(torch.load(pthfile))
print(net)

结果;

这下就加载好预训练模型了

另外,还有一种情况,pretrained = False加载模型也要出错

比如

model = torchvision.models.segmentation.fcn_resnet50(pretrained=False)

运行会显示如下结果:意思是它会把模型下载到缓存里,但是由于网络问题有时他下载不下来,或者每次运行代码都要重新下载,浪费时间。

这时,我们就手动复制链接下载下来,存放在你的项目里面,然后再把它复制到缓存里面

这时,需要用终端命令,在notebook上可以这样操作:

大意就是,先创建缓存地址,同之前报错的地址一样,之后把下载的文件复制到这个路径下,这样他就不会重新去下载了。

经验分享 程序员 微信小程序 职场和发展