【动手学深度学习】读写文件

【动手学深度学习】读写文件

加载和保存张量

对于单个张量 我么可以直接调用load和save函数分别读写,这两个函数要求我们提供一个名称,save要求保存的变量作为输入

py 复制代码
import torch
from torch import nn
from torch.nn import functional as F

# 创建一个长度为4的张量
x = torch.arange(4)
torch.save(x, 'x-file')


x2 = torch.load('x-file')
print(x2)

存储一个张量列表,然后把他们写入内存

py 复制代码
y = torch.zeros(4)
torch.save([x,y],'x-files')
x2,y2 = torch.load('x-files')
(x2,y2)

我们甚至可以写入或者读取从字符串映射到张量的字典,方面读取权重

py 复制代码
# 创建张量字典  保存张量
mydict = {'x':x,'y':y}
torch.save(mydict,'mydict')

mydict2 = torch.load('mydict')
mydict2

加载和保存模型参数

深度学习框架提供内置函数来保存和加载整个网络,这里是保存模型的参数而不是保存整个模型

py 复制代码
class MLP(nn.Module):
    def __init__(self):
        super().__init__()
        self.hidden = nn.Linear(20,256)
        self.output = nn.Linear(256,10)

    def forward(self,X):
        return self.output(F.relu(self.hidden(X)))
    
net = MLP()
X = torch.randn(size= (2,20))
Y = net(X)

取出模型的参数保存在一个mlp.params文件中

py 复制代码
# 取出模型的参数保存在一个mlp.params文件中
torch.save(net.state_dict(),'mlp.params')

恢复模型,实例化原始多层感知机模型的一个备份,我们不需要随机初始化模型参数,而是直接读取文件中存储的参数

py 复制代码
clone = MLP()
clone.load_state_dict(torch.load('mlp.params'))
clone.eval()

比较两个对象的模型参数,那么输入相同的X 计算的输出应该相同

py 复制代码
Y_clone = clone(X)
Y_clone == Y
相关推荐
勾股导航1 小时前
大模型Skill
人工智能·python·机器学习
卷福同学3 小时前
【养虾日记】Openclaw操作浏览器自动化发文
人工智能·后端·算法
春日见4 小时前
如何入门端到端自动驾驶?
linux·人工智能·算法·机器学习·自动驾驶
光锥智能4 小时前
从自动驾驶到 AI 能力体系,元戎启行 GTC 发布基座模型新进展
人工智能
luoganttcc4 小时前
自动驾驶 世界模型 有哪些
人工智能·机器学习·自动驾驶
潘高4 小时前
10分钟教你手撸一个小龙虾(OpenClaw)
人工智能
禁默4 小时前
光学与机器视觉:解锁“机器之眼”的核心密码-《第五届光学与机器视觉国际学术会议(ICOMV 2026)》
人工智能·计算机视觉·光学
深小乐4 小时前
不是DeepSeek V4!这两个神秘的 Hunter 模型竟然来自小米
人工智能
laozhao4324 小时前
科大讯飞中标教育管理应用升级开发项目
大数据·人工智能
rainbow7242444 小时前
AI人才简历评估选型:技术面试、代码评审与项目复盘的综合运用方案
人工智能·面试·职场和发展