第五章 模型篇: 模型保存与加载

Pytorch学习笔记十七:模型保存加载 一、模型保存加载 当我们的模型训练好之后是需要保存下来,以备后续的使用,那么如何保存加载模型呢?下面就从三个方面来理解一下。 1、序列化反序列化 序列化是指内存中的某一对象保存到硬盘中,以二进制的形式存储下来,这就是一个序列化的过程;反序列化就是将硬盘中的存储的二进制数反序列化到内存中,得到一个相应的对象,这样就可以再次使用这个模型了。如下图所示: 序列化和反序列化的目的就是将模型保存并再次使用。 pytorch中序列化和反序列化的方法: torch.save(obj, f):obj表示对象,也就 阅读详情

参考教程
https://pytorch.org/tutorials/beginner/basics/saveloadrun_tutorial.html


训练好的模型,可以保存下来,用于后续的预测或者训练过程的重启。
为了便于理解模型保存和加载的过程,我们定义一个简单的小模型作为例子,进行后续的讲解。

这个模型里面包含一个名为self.p1的Parameter和一个名为conv1的卷积层。我们没有给模型定义forward()函数,是因为暂时不需要用到该方法。假如你想使用这个模型对数据进行前向传播,会返回 “NotImplementedError: Module [Model] is missing the required “forward” function”

import torch
import torch.nn as nn
class Model(nn.Module):
    def __init__(self):
        super().__init__()
        self.t1 = torch.randn((3,2))
        self.p1 = nn.Parameter(self.t1)
        self.conv1 = nn.Conv2d(1, 1, 5)
net = Model()

pytorch中的保存与加载

首先我们来看一下pytorch中的保存和加载的方法是怎么实现的。

torch.save()

参考文档:https://pytorch.org/docs/stable/generated/torch.save.html
首先来看一下torch.save()函数。

torch.save(obj, f, pickle_module=pickle, pickle_protocol=DEFAULT_PROTOCOL, _use_new_zipfile_serialization=True)

torch.save()函数传入的第一个参数,就是我们要保存的对象,它的类别要求是object,而没有限定在nn.Module()或者nn.Parameters()等等之间。说明它可以保存的类型是多种多样的,很灵活。
传入的第二个参数是f,f是一个file-like object或者文件路径,也就是我们想要保存的位置。
后面的几个参数可以不用管它,一般也不会用到。从参数名称可以看到,我们想要保存的object是以pickle的形式保存的。因为pickle支持多种数据类型。
在源码中给了两个使用torch.save的例子。

  >>> # xdoctest: +SKIP("makes cwd dirty")
        >>> # Save to file
        >>>
pytorch模型保存加载 PyTorch提供了三种种方式来保存加载模型,在这三种方式中,加载模型的代码和保存模型的代码必须相匹配,才能保证模型加载成功。通常情况下,使用第一种方式(保存加载模型状态字典)更加常见,因为它更轻量且不依赖于特定的模型类。 阅读详情

相关推荐

关于PyTorch继承nn.Module出现raise NotImplementedError的问题解决方案

问题描述:解决方法:NotImplementedError 错误:子类没有完成父类的接口,在此就是父类(nn.Module)中的 forward 方法在子类中没有定义,则会自动调用 nn.Module 中的forward方法,而 nn.Module 中的 forward 是 raise 将错误抛出。所以出现 NotImplementedError 错误。2.问题锁定在forward方法上:(1)没...

算法与编程之美 1551

NotImplementedError: Module is missing the required “forward“ function

2.def forward函数def __init__(self,config):一定要对齐。1.重写父类函数时,函数名称写错,我将。

qq_52190828的博客 1万+

第五章模型分词器

本章我们将介绍 Transformers 库中的两个重要组件:**模型**和**分词器**。 ## 5.1 模型 除了像之前使用 `AutoModel` 根据 checkpoint 自动加载模型以外,我们也可以直接使用模型对应的 `Model` 类,例如 BERT 对应的就是 `BertModel`: ``` from transformers import BertModel model = BertModel.from_pretrained("bert-base-cased") ``` 注意,

南七小僧的学海无涯 205

PyTorch继承nn.Module时,类函数forward出现raise NotImplementedError

NotImplementedError

弱就多努力的博客 3927

pytorch checkpoint_PyTorch专栏(七):模型保存加载那些事

作者 | News编辑 |安可出品 |磐创AI团队出品【磐创AI导读】:本文章讲解了PyTorch专栏的第三章中的保存加载模型。查看专栏历史文章,请点击下方蓝色字体进入相应链接阅读。查看关于本专栏的介绍:PyTorch专栏开。想要更多电子杂志的机器学习,深度学习资源,大家欢迎点击上方蓝字关注我们的公众号:磐创AI。专栏目录:第一章:PyTorch之简介下载PyTorch简介...

weixin_39611037的博客 1588

PyTorch专栏(七):模型保存加载那些事

作者 | News编辑 |安可出品 |磐创AI团队出品【磐创AI导读】:本文章讲解了PyTorch专栏的第三章中的保存加载模型。查看专栏历史文章,请点击下方蓝色...

TensorFlowNews 962

4.8 PyTorch模型保存加载

欢迎订阅本专栏:《PyTorch深度学习实践》 订阅地址:https://blog.csdn.net/sinat_33761963/category_9720080.html 第二章:认识Tensor的类型、创建、存储、api等,打好Tensor的基础,是进行PyTorch深度学习实践的重中之重的基础。 第三章:学习PyTorch如何读入各种外部数据 第四章:利用PyTorch从头到尾创建、训...

【人工智能】王小草的博客 1097

pytorch保存模型pth_PyTorch专栏(七):模型保存加载那些事

作者 | News编辑 |安可出品 |磐创AI团队出品【磐创AI导读】:本文章讲解了PyTorch专栏的第三章中的保存加载模型。查看专栏历史文章,请点击下方蓝色字体进入相应链接阅读。查看关于本专栏的介绍:PyTorch专栏开。想要更多电子杂志的机器学习,深度学习资源,大家欢迎点击上方蓝字关注我们的公众号:磐创AI。专栏目录:第一章:PyTorch之简介下载PyTorch简介...

weixin_29364297的博客 464

Tensorflow【实战Google深度学习框架】TensorFlow模型保存恢复加载

我们使用TensorFlow进行模型的训练,训练好的模型需要保存,预测阶段我们需要将模型进行加载还原使用,这就涉及TensorFlow模型保存恢复加载。 总结一下Tensorflow常用的模型保存方式。 文章目录保存checkpoint模型文件(.ckpt)模型保存模型加载还原完整代码 保存checkpoint模型文件(.ckpt) 首先,TensorFlow提供了一个非常方便的api,tf....

NJU phd 在读。 567

第五章:AI大模型的优化调参5.3 模型训练技巧5.3.2 早停法模型保存

背景介绍 在深度学习领域,大模型因其强大的表达能力和泛化能力在各个领域得到了广泛应用。然而,训练大模型需要大量的计算资源和时间,因此如何优化大模型的训练过程成为了一个重要的研究方向。本文将介绍大模型训练中的两个重要技巧:早停法和模型保存。 核心概念联系 早停法

AI天才研究院 786

第五章 PyTorch模型定义

场景典型写法优点局限或最短代码,一行顶多行;自动串行forward只能按固定顺序前向;不方便插入分支/多输入;“可迭代层列表”,支持动态循环、重复结构仍需手写forward,无法自动注册子层顺序名字可读,利于按键访问 / 条件分支ModuleList相同,需要自己控制执行逻辑# ---------- ① Sequential:完全按顺序 ----------seq_net = nn.Sequential( # 自动生成 forward;最省事。

weixin_72250436的博客 1805

PyTorch继承nn.Module,forward出现raise NotImplementedError

查看别人写的方法,基本就是forward拼写错误,要么就是tab错了,只有我这个大怨种,encoder的那个类,我忘记写forward了。。。所以检查一下是不是自己调用的每个类的forward是不是对了~

qq_53922287的博客 4604

保存载入模型的常用方法

一般使用torch.save()函数将模型保存起来,该函数一般经常模型的state_dict()方法联合使用。同理,加载模型使用的是torch.load()函数,该函数模型的load_state_dict()方法联合使用 1.保存模型 torch.save(model.state_dict(), "./model.pth") 执行完该命令后,会在本地目录中生成一个model.pth文件,该文件就是保存好的模型文件 2.载入模型 model.load_state_dict(torch.loa

weixin_48592695的博客 2582

pytorch模型保存加载总结

pytorch模型保存加载方式、打包保存tar、多卡训练遇到的问题、torch.jit、加载预训练模型保存模型加载精度损失

蜗牛博客 1万+

模型保存加载

模型保存加载深度学习中非常重要的一部分,通过保存加载模型可以实现模型的持久化,方便模型的重复使用和部署。在深度学习中,模型保存可以通过多种方式实现,常用的包括使用TensorFlow和PyTorch等框架提供的模型保存函数,以及使用第三方库(如Pickle和Joblib)来保存模型参数等。该函数可以保存整个模型,包括模型的结构和参数。上述代码中,我们使用tf.keras.models.load_model()函数加载了之前保存的my_model模型,并使用加载模型进行了预测。

互联网知识分享 925

从0开始的OpenGL学习(十七)-加载模型

阅读 7,352 本文主要解决一个问题: 如何在OpenGL中加载模型? 引言 学到现在,我们把盒子兄弟折磨得死去活来,虽说弄出了一些效果,但也总是感觉有点不给力,换个时髦的说法就是:用户体验不好。在实际的图形应用中,会有很多复杂并且有趣的模型,比我们的盒子强太多。但是,由于太复杂,我们不可能手动定义模型的顶点坐标、法线和纹理坐标等值。我们希望的是,直接把模型导入到应用中使用,把创建模型的工作交给专业的建模师去做。他们有很高端的工具,例如3DS Max、Maya等等。 这些3D建模工具十分强大

翻肚鱼儿的博客 1414

Pytorch加载模型

torchvision是PyTorch生态系统中的一个包,专门用于计算机视觉任务。它提供了一系列用于加载、处理和预处理图像和视频数据的工具,以及常用的计算机视觉模型。torchvision.models模块包含许多常用的预训练计算机视觉模型,例如ResNet、AlexNet、VGG等分类、分割等模型

m0_50460160的博客 1900
上一篇: 第四章 模型篇:模型训练与示例
下一篇: 第六章 番外篇:webdataset
江米江米
博客等级 码龄9年 80粉丝 49原创
评论
添加红包

请填写红包祝福语或标题

红包个数最小为10个

红包金额最低5元

当前余额3.43前往充值 >
需支付:10.00
成就一亿技术人!
领取后你会自动成为博主和红包主的粉丝 规则
hope_wisdom
发出的红包
实付
使用余额支付
点击重新获取
扫码支付
钱包余额 0

抵扣说明:

1.余额是钱包充值的虚拟货币,按照1:1的比例进行支付金额的抵扣。
2.余额无法直接购买下载,可以购买VIP、付费专栏及课程。

余额充值