掌握PyTorch模型保存与加载:从训练到部署的完整指南

掌握PyTorch模型保存与加载:从训练到部署的完整指南

pytorch模型加载时,需要先定义模型结构,再加载保存的state_dict参数。这是因为pytorch通常只保存模型参数而非整个模型对象,以避免python对象序列化问题。本文将详细介绍如何分离模型的训练、保存与加载推理过程,并通过示例代码演示这一标准实践,帮助用户高效复用预训练模型。

在PyTorch中,将训练好的模型保存到磁盘并在后续加载进行推理是机器学习工作流中的常见需求。初学者常遇到的一个困惑是:加载模型时是否必须重新定义模型的完整结构?答案是肯定的,且这是PyTorch推荐的标准实践。本教程将深入探讨PyTorch的模型保存与加载机制,并提供清晰的示例代码,指导您如何正确地分离模型的训练、保存与推理过程。

理解PyTorch模型保存机制

PyTorch模型(nn.Module的实例)的保存通常有两种主要方式:

保存整个模型(不推荐):使用 torch.save(model, “model.pth”)。这种方法会保存整个模型对象,包括其结构和所有参数。然而,它依赖于Python的pickle模块进行序列化。当模型定义所在的类、包或文件结构发生变化时,或者在不同Python版本、PyTorch版本之间加载时,可能会遇到兼容性问题和序列化错误。因此,这种方法通常不被推荐用于生产环境或长期存储。保存模型的state_dict(推荐):使用 torch.save(model.state_dict(), “model.pth”)。state_dict是一个Python字典,它存储了模型中所有可学习参数(如权重和偏置)的映射。这种方式只保存参数,而模型的结构定义则需要独立存在。加载时,您需要先实例化一个具有相同结构的模型对象,然后将state_dict加载到这个新创建的对象中。这种方法更加健壮、灵活,且不易受环境变化的影响。

核心思想是: 模型结构(由nn.Module类定义)与模型参数(存储在state_dict中)是分离的。当您保存state_dict时,您只是保存了模型学到的“知识”,而模型的“骨架”——其架构定义——则需要在加载时重新提供。

模型训练与保存示例

为了演示这一过程,我们将使用一个简单的神经网络在FashionMNIST数据集上进行训练,并保存其state_dict。

Stable Diffusion 2.1 Demo Stable Diffusion 2.1 Demo

最新体验版 Stable Diffusion 2.1

Stable Diffusion 2.1 Demo 101 查看详情 Stable Diffusion 2.1 Demo

首先,我们需要设置环境、定义模型、数据加载器以及训练和测试函数。

# train_model.pyimport torchfrom torch import nnfrom torch.utils.data import DataLoaderfrom torchvision import datasetsfrom torchvision.transforms import ToTensor# 1. 准备数据training_data = datasets.FashionMNIST(    root="data",    train=True,    download=True,    transform=ToTensor(),)test_data = datasets.FashionMNIST(    root="data",    train=False,    download=True,    transform=ToTensor(),)batch_size = 64train_dataloader = DataLoader(training_data, batch_size=batch_size)test_dataloader = DataLoader(test_data, batch_size=batch_size)# 2. 获取设备device = (    "cuda"    if torch.cuda.is_available()    else "mps"    if torch.backends.mps.is_available()    else "cpu")print(f"Using {device} device")# 3. 定义模型class NeuralNetwork(nn.Module):    def __init__(self):        super().__init__()        self.flatten = nn.Flatten()        self.linear_relu_stack = nn.Sequential(            nn.Linear(28*28, 512),            nn.ReLU(),            nn.Linear(512, 512),            nn.ReLU(),            nn.Linear(512, 10)        )    def forward(self, x):        x = self.flatten(x)        logits = self.linear_relu_stack(x)        return logitsmodel = NeuralNetwork().to(device)print(model)# 4. 定义损失函数和优化器loss_fn = nn.CrossEntropyLoss()optimizer = torch.optim.SGD(model.parameters(), lr=1e-3)# 5. 训练函数def train(dataloader, model, loss_fn, optimizer):    size = len(dataloader.dataset)    model.train()    for batch, (X, y) in enumerate(dataloader):        X, y = X.to(device), y.to(device)        pred = model(X)        loss = loss_fn(pred, y)        optimizer.zero_grad()        loss.backward()        optimizer.step()        if batch % 100 == 0:            loss, current = loss.item(), (batch + 1) * len(X)            print(f"loss: {loss:>7f}  [{current:>5d}/{size:>5d}]")# 6. 测试函数def test(dataloader, model, loss_fn):    size = len(dataloader.dataset)    num_batches = len(dataloader)    model.eval()    test_loss, correct = 0, 0    with torch.no_grad():        for X, y in dataloader:            X, y = X.to(device), y.to(device)            pred = model(X)            test_loss += loss_fn(pred, y).item()            correct += (pred.argmax(1) == y).type(torch.float).sum().item()    test_loss /= num_batches    correct /= size    print(f"Test Error: n Accuracy: {(100*correct):>0.1f}%, Avg loss: {test_loss:>8f} n")# 7. 训练模型并保存epochs = 5for t in range(epochs):    print(f"Epoch {t+1}n-------------------------------")    train(train_dataloader, model, loss_fn, optimizer)    test(test_dataloader, model, loss_fn)print("Done training!")# 保存模型的state_dicttorch.save(model.state_dict(), "model.pth")print("Saved PyTorch Model State to model.pth")

运行上述代码后,您将得到一个名为 model.pth 的文件,其中包含了训练好的模型参数。

模型加载与推理示例

现在,假设我们希望在一个完全独立的脚本中加载 model.pth 文件并进行推理。这个脚本不需要知道模型是如何训练的,但它必须知道模型的结构定义。

# inference_model.pyimport torchfrom torch import nnfrom torchvision import datasetsfrom torchvision.transforms import ToTensor# 1. 获取设备 (与训练时保持一致)device = (    "cuda"    if torch.cuda.is_available()    else "mps"    if torch.backends.mps.is_available()    else "cpu"

以上就是掌握PyTorch模型保存与加载:从训练到部署的完整指南的详细内容,更多请关注创想鸟其它相关文章!

版权声明:本文内容由互联网用户自发贡献,该文观点仅代表作者本人。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。
如发现本站有涉嫌抄袭侵权/违法违规的内容, 请发送邮件至 chuangxiangniao@163.com 举报,一经查实,本站将立刻删除。
发布者:程序猿,转转请注明出处:https://www.chuangxiangniao.com/p/918418.html

(0)
打赏 微信扫一扫 微信扫一扫 支付宝扫一扫 支付宝扫一扫
能上qq但无法访问网页是怎么回事
上一篇 2025年11月29日 06:33:02
Win11版本区别对照表_Win11各个版本怎么区分
下一篇 2025年11月29日 06:33:07

相关推荐

  • 如何让你的电商前端快如闪电:SprykerTouch模块与Composer助力数据同步挑战

    Composer在线学习地址:学习地址 电商前端的“卡顿”之痛:数据同步的困境 想象一下,你正在运营一个繁忙的电商平台,商品价格、库存、描述等信息在后台(zed)频繁更新。用户在前端(yves)浏览商品时,他们期望看到的是最新、最准确的数据。然而,spryker 架构有一个核心设计原则:yves(前…

    用户投稿 2026年8月25日
    000
  • 中间件(Middleware)在Yii3中的应用

    在yii3中使用中间件是为了增强应用程序的灵活性和可维护性。中间件在请求处理前后执行特定操作,简化代码结构,提升扩展和维护的便捷性。 让我们先来回答一个关键问题:为什么在Yii3中使用中间件(Middleware)? 在Yii3中,中间件的使用主要是为了增强应用程序的灵活性和可维护性。中间件作为请求…

    2026年8月25日
    000
  • Swoole支持哪些网络协议(TCP/UDP/HTTP/WebSocket)?

    swoole支持tcp、udp、http和websocket协议。1.tcp:通过swooleserver类处理连接,适用于高性能服务器。2.udp:swooleserver类支持数据包收发,适用于快速响应应用。3.http:swoolehttpserver类适用于restful api和web应用…

    2026年8月25日
    000
  • 2025年全球智能手机出货量将同比增长1% IDC:苹果功不可没

    8月29日,idc发布《全球季度移动电话追踪报告》显示,2025年全球智能手机出货量预计同比增长1.0%,总量达12.4亿部,较此前预测的0.6%有所上调,主要驱动力来自ios设备出货量实现3.9%的增长。 尽管市场仍面临需求疲软等多重压力,但稳定的换机周期将支撑行业在2026年延续增长态势,预计2…

    2026年8月25日
    000
  • 如何利用AI工具来提高工作效率?

    以下是一些利用 AI 工具提高工作效率的方法:自动化日常任务:数据录入和处理:利用 AI 系统通过扫描、识别技术自动录入数据,例如银行用 AI 自动处理支票存款,识别支票信息并录入系统,节省人工录入时间和精力。文件整理与管理:借助 AI 的自然语言处理和机器学习技术,自动分类和整理文件。如法律行业中…

    用户投稿 2026年8月25日
    000
  • java中的consumer关键字用途 消费者Consumer的2个典型应用

    java中的consumer关键字用途 消费者Consumer的2个典型应用java中的consumer关键字用途 消费者Consumer的2个典型应用java中的consumer关键字用途 消费者Consumer的2个典型应用java中的consumer关键字用途 消费者Consumer的2个典型应用

    java中的consumer接口用于定义不返回结果的操作,其核心目的是简化代码并提升可读性与维护性。1. 它常用于集合的foreach方法,实现更简洁的遍历操作;2. 在stream api中通过peek和foreach方法支持中间处理与最终操作;3. 可自定义多参数consumer接口以满足特定需…

    2026年8月25日 用户投稿
    000
  • Word文档怎么设置文字环绕图片_Word图文混排环绕设置指南

    首先选择图片并设置环绕方式,如四周型或上下型;再通过“编辑环绕顶点”自定义路径;最后在布局选项中调整文字与图片间距,实现图文混排。 如果您在编辑Word文档时希望实现文字环绕图片的效果,以便让版面更加美观和专业,可以通过调整图片的环绕方式来实现图文混排。以下是具体操作步骤: 本文运行环境:联想小新A…

    2026年8月25日
    000
  • Swoole的C++底层源码解析

    学习swoole的底层源码是为了理解高性能网络服务器的工作原理和优化性能及架构设计。通过学习,1) 掌握c++++在高并发环境下的应用技巧,2) 理解事件驱动模型的精髓,3) 学习利用操作系统特性提升程序效率,4) 了解高效的异步i/o处理、协程调度和内存管理。 在深入探讨Swoole的C++底层源…

    2026年8月25日
    600
  • EasyControl— Tiamat AI 联合上海科大等开源的图像生成控制框架

    easycontrol:高效灵活的扩散模型控制框架 EasyControl是由Tiamat AI开源的基于扩散变换器(Diffusion Transformer,DiT)架构的图像生成控制框架。它通过轻量级LoRA模块独立处理条件信号,实现即插即用的功能,并兼容现有模型。 EasyControl支持…

    2026年8月25日
    000
  • Swoole 5.0新特性解读

    swoole 5.0的新特性包括:1)支持php 8的jit编译,提升性能;2)优化协程调度,减少上下文切换;3)引入新的异步i/o接口,简化大文件处理;4)改进内存管理,减少内存碎片。这些特性提升了开发效率和应用性能。 在Swoole 5.0发布后,很多开发者都迫不及待地想了解它的新特性和改进。这…

    2026年8月25日
    000
  • PHP多店铺电商平台痛点如何解决?Spryker/Store模块助你轻松管理多语言多货币配置

    可以通过一下地址学习composer:学习地址 在当今全球化的电商环境中,许多企业不再满足于单一市场。他们希望将业务扩展到不同的国家和地区,这意味着需要支持多种语言、多种货币,甚至不同的品牌和产品线。对于开发者而言,这无疑是一个巨大的挑战。 想象一下,你正在为一个雄心勃勃的电商平台构建后端系统。最初…

    用户投稿 2026年8月25日
    000
  • 搬别人图片AI上传视频号被提示优化如何申请?ai生成的视频会被限流吗?

    借助AI工具将他人图片转化为视频并发布到视频号,已成为一种常见的内容创作形式。然而,若因此收到“内容需优化”的系统提示,该如何应对?这类由AI生成的视频是否容易遭遇限流? 一、使用他人图片通过AI制作视频上传后被提示优化,怎样申请复核? 当你在视频号发布的内容出现“内容需优化”提示时,说明平台算法初…

    2026年8月25日
    000
  • Java中如何创建线程 详解三种创建线程的方式

    Java中如何创建线程 详解三种创建线程的方式Java中如何创建线程 详解三种创建线程的方式Java中如何创建线程 详解三种创建线程的方式Java中如何创建线程 详解三种创建线程的方式

    java中创建线程的核心方式有三种:实现runnable接口、继承thread类、使用executorservice。1.实现runnable接口是推荐方式,通过实现run()方法定义任务,再由thread执行,避免单继承限制并解耦任务与线程;2.继承thread类则直接重写run()方法,虽简单但…

    2026年8月25日 用户投稿
    500
  • 推荐一些具体的AI工具 豆包在线入口

    豆包是字节跳动推出的免费 AI 助手,提供多平台使用和登录方式。1. 豆包:字节跳动推出的免费 AI 助手,支持网页、iOS 和 Android 端,登录方式包括手机号和抖音账号,网页端入口为豆包官网。 ☞☞☞AI 智能聊天, 问答助手, AI 智能搜索, 免费无限量使用 DeepSeek R1 模…

    2026年8月25日
    000
  • 电脑提示“steam_api.dll没有被指定在windows上运行”的4个方法

    电脑提示“steam_api.dll没有被指定在windows上运行”的4个方法电脑提示“steam_api.dll没有被指定在windows上运行”的4个方法电脑提示“steam_api.dll没有被指定在windows上运行”的4个方法电脑提示“steam_api.dll没有被指定在windows上运行”的4个方法

    在运行游戏或部分软件时,有时会弹出提示:“steam_api.dll未被指定在windows上运行,或文件已损坏。”这通常是因为dll文件丢失、版本不兼容,或系统缺少必要的运行环境。下面介绍几种简单有效的解决方式。 方法一:使用“星空运行库修复大师”自动修复(首选) 对于不太熟悉手动操作的用户来说,…

    2026年8月25日 用户投稿
    300
  • Uniapp 中如何不拉伸不裁剪地展示图片?

    灵活展示图片:如何不拉伸不裁剪 在界面设计中,常常需要以原尺寸展示用户上传的图片。本文将介绍一种在 uniapp 框架中实现该功能的简单方法。 对于不同尺寸的图片,可以采用以下处理方式: 极端宽高比:撑满屏幕宽度或高度,再等比缩放居中。非极端宽高比:居中显示,若能撑满则撑满。 然而,如果需要不拉伸不…

    2025年12月24日
    600
  • 如何让小说网站控制台显示乱码,同时网页内容正常显示?

    如何在不影响用户界面的情况下实现控制台乱码? 当在小说网站上下载小说时,大家可能会遇到一个问题:网站上的文本在网页内正常显示,但是在控制台中却是乱码。如何实现此类操作,从而在不影响用户界面(UI)的情况下保持控制台乱码呢? 答案在于使用自定义字体。网站可以通过在服务器端配置自定义字体,并通过在客户端…

    2025年12月24日
    1000
  • 如何在地图上轻松创建气泡信息框?

    地图上气泡信息框的巧妙生成 地图上气泡信息框是一种常用的交互功能,它简便易用,能够为用户提供额外信息。本文将探讨如何借助地图库的功能轻松创建这一功能。 利用地图库的原生功能 大多数地图库,如高德地图,都提供了现成的信息窗体和右键菜单功能。这些功能可以通过以下途径实现: 高德地图 JS API 参考文…

    2025年12月24日
    400
  • 如何使用 scroll-behavior 属性实现元素scrollLeft变化时的平滑动画?

    如何实现元素scrollleft变化时的平滑动画效果? 在许多网页应用中,滚动容器的水平滚动条(scrollleft)需要频繁使用。为了让滚动动作更加自然,你希望给scrollleft的变化添加动画效果。 解决方案:scroll-behavior 属性 要实现scrollleft变化时的平滑动画效果…

    2025年12月24日
    000
  • 如何为滚动元素添加平滑过渡,使滚动条滑动时更自然流畅?

    给滚动元素平滑过渡 如何在滚动条属性(scrollleft)发生改变时为元素添加平滑的过渡效果? 解决方案:scroll-behavior 属性 为滚动容器设置 scroll-behavior 属性可以实现平滑滚动。 html 代码: click the button to slide right!…

    2025年12月24日
    700

发表回复

登录后才能评论
关注微信