动手学深度学习 4.5 读取和存储

前言

从零开始学习ai文章系列计划是个人在《动手学深度学习》和《磨菇书》两本书的学习中的个人笔记,文章也会以课本中的章节分开,即每个章节一片笔记。我会尽量的把主要内容以及遇到的难点进行记录与解决,如果哪里有错误的欢迎指正。或者不清晰的可以直接查看原文部分。

《动手学深度学习》原文(课本):https://tangshusen.me/Dive-into-DL-PyTorch/#/

《动手学深度学习》代码:https://github.com/ShusenTang/Dive-into-DL-PyTorch

(由于有时候公式太多,可能会直接贴图片)

1. 读写Tensor

可以直接使用 save函数load函数 分别存储和读取Tensor。

下面的例子创建了Tensor变量x,并将其存在文件名同为x.pt的文件里。然后我们将数据从存储的文件读回内存。

1
2
3
4
5
x = torch.ones(3)
torch.save(x, 'x.pt')

x2 = torch.load('x.pt')
print(x2)

我们还可以存储一个Tensor列表并读回内存。(毕竟底层是pickle实现的)

1
2
3
4
5
6
7
8
import torch
x = torch.ones(3)
y = torch.zeros(4)

torch.save([x, y], 'xy.pt')

xy_list = torch.load('xy.pt')
print(xy_list)

存储并读取一个从字符串映射到Tensor的字典。

1
2
3
4
torch.save({'x': x, 'y': y}, 'xy_dict.pt')

xy = torch.load('xy_dict.pt')
print(xy)

2. net.state_dict() 介绍

在 PyTorch 中,net.state_dict() 是一个 字典(dict)对象,它保存了模型中的所有 参数和缓冲区的名称和值。

1
2
3
4
5
6
7
8
9
10
11
12
13
class MLP(nn.Module):
def __init__(self):
super(MLP, self).__init__()
self.hidden = nn.Linear(3, 2)
self.act = nn.ReLU()
self.output = nn.Linear(2, 1)

def forward(self, x):
a = self.act(self.hidden(x))
return self.output(a)

net = MLP()
print(net.state_dict())

以下代码能清晰查看 net.state_dict() 的键值结构。

1
2
for key in net.state_dict():
print(key,'\t',net.state_dict()[key])

你可以理解 net.state_dict() 仅包含了模型中所有的可学习参数即 nn.Parameter 而不包含类似 self.some_value = 123 等。

3. 保存和加载模型

PyTorch中保存和加载训练模型有两种常见的方法:

  1. 仅保存和加载模型参数(state_dict);
  2. 保存和加载整个模型。

3.1 保存和加载state_dict(推荐方式)

保存

1
torch.save(net.state_dict(), PATH) # 推荐的文件后缀名是pt或pth

加载

1
2
3
net = TheModelClass()

net.load_state_dict(torch.load(PATH))

3.2 保存和加载整个模型

保存:

1
torch.save(net, PATH)

加载:

1
net = torch.load(PATH)