动手学深度学习 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 | x = torch.ones(3) |

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

存储并读取一个从字符串映射到Tensor的字典。
1 | torch.save({'x': x, 'y': y}, 'xy_dict.pt') |

2. net.state_dict() 介绍
在 PyTorch 中,net.state_dict() 是一个 字典(dict)对象,它保存了模型中的所有 参数和缓冲区的名称和值。
1 | class MLP(nn.Module): |

以下代码能清晰查看 net.state_dict() 的键值结构。
1 | for key in net.state_dict(): |

你可以理解 net.state_dict() 仅包含了模型中所有的可学习参数即
nn.Parameter 而不包含类似
self.some_value = 123 等。
3. 保存和加载模型
PyTorch中保存和加载训练模型有两种常见的方法:
- 仅保存和加载模型参数(state_dict);
- 保存和加载整个模型。
3.1 保存和加载state_dict(推荐方式)
保存
1 | torch.save(net.state_dict(), PATH) # 推荐的文件后缀名是pt或pth |
加载
1 | net = TheModelClass() |
3.2 保存和加载整个模型
保存:
1 | torch.save(net, PATH) |
加载:
1 | net = torch.load(PATH) |