PyTorch深度学习入门笔记(小土堆)P26-32

发布时间:2026/8/6 19:52:05

PyTorch深度学习入门笔记(小土堆)P26-32
PyTorch深度学习入门笔记P26-32ZZHow(ZZHow1024)参考课程【PyTorch深度学习快速入门教程【小土堆】】[https://www.bilibili.com/video/BV1hE411t7RN]P26. 完整的模型训练套路一训练部分model.pyimporttorchfromtorchimportnnfromtorch.nnimportSequential# 搭建神经网络classMyModel(torch.nn.Module):def__init__(self):super(MyModel,self).__init__()self.modelSequential(nn.Conv2d(in_channels3,out_channels32,kernel_size5,stride1,padding2),nn.MaxPool2d(kernel_size2),nn.Conv2d(in_channels32,out_channels32,kernel_size5,stride1,padding2),nn.MaxPool2d(kernel_size2),nn.Conv2d(in_channels32,out_channels64,kernel_size5,stride1,padding2),nn.MaxPool2d(kernel_size2),nn.Flatten(),nn.Linear(in_features64*4*4,out_features64),nn.Linear(in_features64,out_features10),)defforward(self,x):xself.model(x)returnx# 测试神经网络模型结构的正确性if__name____main__:modelMyModel()inputtorch.ones([64,3,32,32])outputmodel(input)print(output.shape)train.pyimporttorchimporttorchvision.datasetsfrommodelimportMyModel# 准备数据集train_datatorchvision.datasets.CIFAR10(dataset,trainTrue,transformtorchvision.transforms.ToTensor(),downloadTrue)test_datatorchvision.datasets.CIFAR10(dataset,trainFalse,transformtorchvision.transforms.ToTensor(),downloadTrue)# 获取数据集的长度train_data_sizelen(train_data)test_data_sizelen(test_data)print(f训练数据集的长度为{train_data_size})print(f测试数据集的长度为{test_data_size})# 使用 Dataloader 加载数据集train_dataloadertorch.utils.data.DataLoader(train_data,batch_size64)test_dataloadertorch.utils.data.DataLoader(test_data,batch_size64)# 创建网络模型modelMyModel()# 损失函数loss_fntorch.nn.CrossEntropyLoss()# 优化器learning_rate1e-2optimizertorch.optim.SGD(model.parameters(),lrlearning_rate)# 设置训练网络的参数total_train_step0# 训练次数total_test_step0# 测试次数epoch10# 训练轮次foriinrange(epoch):print(f---第{i1}轮训练开始---)# 训练步骤开始fordataintrain_dataloader:images,targetsdata outputsmodel(images)lossloss_fn(outputs,targets)# 优化器优化模型optimizer.zero_grad()loss.backward()optimizer.step()total_train_step1print(f训练次数{total_train_step}Loss{loss.item()})P27. 完整的模型训练套路二测试验证部分# 测试步骤开始total_test_loss0# 总测试 Losstotal_accuracy0# 总正确率withtorch.no_grad():fordataintest_dataloader:images,targetsdata outputsmodel(images)lossloss_fn(outputs,targets)total_test_lossloss.item()accuracy(outputs.argmax(1)targets).sum()total_accuracyaccuracy writer.add_scalar(test_loss,total_test_loss,total_test_step)writer.add_scalar(test_accuracy,total_accuracy/test_data_size,total_test_step)print(f测试集上的总 Loss{total_test_loss})print(f测试集上的总 正确率{total_accuracy/test_data_size})torch.save(model.state_dict(),os.path.join(model,fmodel_{i}.pth))print(f模型已保存文件名model_{i}.pth)total_test_step1P28. 完整的模型训练套路三训练步骤开始时model.train()测试步骤开始时model.eval()案例演示model.py和train.pyP29. 利用GPU训练一方式一在网络模型、数据输入标注和损失函数后加上.cuda()# 创建网络模型modelMyModel()iftorch.cuda.is_available():modelmodel.cuda()# 损失函数loss_fntorch.nn.CrossEntropyLoss()iftorch.cuda.is_available():loss_fnloss_fn.cuda()# 数据输入标注images,targetsdataiftorch.cuda.is_available():imagesimages.cuda()targetstargets.cuda()案例演示train_gpu_1.pyP30. 利用GPU训练二方式二在网络模型、数据输入标注和损失函数后通过.to(device)转移到对应设备# 训练设备devicecpuiftorch.cuda.is_available():devicecudaeliftorch.mps.is_available():devicempsprint(f训练设备{device})# 创建网络模型modelMyModel()model.to(device)# 损失函数loss_fntorch.nn.CrossEntropyLoss()loss_fnloss_fn.to(device)# 数据输入标注images,targetsdata imagesimages.to(device)targetstargets.to(device)案例演示train_gpu_2.pyP31. 完整的模型验证套路test.pyimportosimporttorchimporttorchvisionfromPILimportImagefrommodelimportMyModel# 测试图片名称image_namedog.png# 测试模型名称model_namemodel_29.pth# 测试设备devicecpuiftorch.cuda.is_available():devicecudaeliftorch.mps.is_available():devicempsprint(f测试设备{device})# 测试图片路径image_pathos.path.join(images,image_name)imageImage.open(image_path)imageimage.convert(RGB)print(image)# 图片预处理transformtorchvision.transforms.Compose([torchvision.transforms.Resize((32,32)),torchvision.transforms.ToTensor()])imagetransform(image)imagetorch.reshape(image,(1,3,32,32))print(image.shape)# 加载模型modelMyModel()model.load_state_dict(torch.load(os.path.join(model,model_name),map_locationtorch.device(device)))# 开始测试model.eval()withtorch.no_grad():outputmodel(image)print(output)print(output.argmax(1))注意若训练模型的设备与当前加载加载模型的设备不一致时需要在torch.load()时指定map_locationtorch.device(device)。案例演示test.py

相关新闻

软考网络工程师|第 5 章 TCP/UDP 完整备考笔记

软考网络工程师|第 5 章 TCP/UDP 完整备考笔记

2026/8/6 19:52:05

一、TCP 与 UDP 基础对比★★★★1 核心特性总览维度TCP(传输控制协议)UDP(用户数据报协议)连接属性面向连接,传输前建立连接无连接,直接发送报文可靠性可靠传输,重传、确认、有序不可靠尽力交付…

2026年青岛做城市生命线安全工程建设的厂家有哪些?

2026年青岛做城市生命线安全工程建设的厂家有哪些?

2026/8/6 19:52:05

青岛是沿海城市,地下管线规模大,燃气、供水管网建设标准高,汛期防潮防涝压力较大。海风带来的高湿盐雾环境对监测设备提出了特殊要求,既要防得住潮气,也要测得准管网数据,城市生命线安全工程的需求稳定而务…

2026年合肥做城市生命线安全工程建设的厂家有哪些?

2026年合肥做城市生命线安全工程建设的厂家有哪些?

2026/8/6 19:52:05

合肥作为长三角副中心城市,城市规模快速扩张,燃气管网、排水管网延伸迅速。作为「清华方案合肥模式」的诞生地,合肥在城市生命线安全工程建设上起步早、标准高,燃气、桥梁、供水、排水等专项监测的需求持续释放,示范效…

深入理解Xiaomi-Robotics-0-LIBERO的动作生成机制:从输入编码到控制命令的全流程

深入理解Xiaomi-Robotics-0-LIBERO的动作生成机制:从输入编码到控制命令的全流程

2026/8/6 21:02:08

深入理解Xiaomi-Robotics-0-LIBERO的动作生成机制:从输入编码到控制命令的全流程 【免费下载链接】Xiaomi-Robotics-0-LIBERO 项目地址: https://ai.gitcode.com/hf_mirrors/XiaomiRobotics/Xiaomi-Robotics-0-LIBERO Xiaomi-Robotics-0-LIBERO是小米机器人…

njs学习资源大全:从官方示例到社区最佳实践

njs学习资源大全:从官方示例到社区最佳实践

2026/8/6 21:02:08

njs学习资源大全:从官方示例到社区最佳实践 【免费下载链接】njs-examples NGINX JavaScript examples 项目地址: https://gitcode.com/gh_mirrors/nj/njs-examples njs(NGINX JavaScript)是一个强大的工具,它允许开发者使…

【AI在线咨询落地实战指南】:20年IT专家亲授5大避坑法则与3周上线速成路径

【AI在线咨询落地实战指南】:20年IT专家亲授5大避坑法则与3周上线速成路径

2026/8/6 21:02:08

更多请点击: https://codechina.net 第一章:AI在线咨询落地的核心价值与战略定位 AI在线咨询已从技术概念演进为关键业务基础设施,其核心价值不仅体现在响应效率提升,更在于重构客户信任路径与服务成本结构。当企业将AI咨询能力嵌…

发电企业的数据,为什么“存得多”却“用得少”?

发电企业的数据,为什么“存得多”却“用得少”?

2026/8/6 21:02:08

在数字化转型的浪潮中,发电企业无疑是走在最前列的行业之一。SIS系统实时采集机组运行数据,MIS系统管理设备台账和检修记录,燃料系统跟踪煤质化验和库存变化,财务系统核算成本和经营指标,大大小小七八套系统&#xff0…

CodeFlow高级技巧:3种方式自定义架构图展示效果提升代码理解效率

CodeFlow高级技巧:3种方式自定义架构图展示效果提升代码理解效率

2026/8/6 21:02:08

CodeFlow高级技巧:3种方式自定义架构图展示效果提升代码理解效率 【免费下载链接】codeflow Paste any GitHub URL → interactive architecture map. See how files connect, find what breaks if you change something. No install, no accounts — runs entirely…

Stats.js前端性能监控实战与Three.js优化指南

Stats.js前端性能监控实战与Three.js优化指南

2026/8/6 20:52:07

1. Stats.js 插件核心价值解析Stats.js 是前端性能监控领域的一个轻量级工具库,专门用于实时显示网页运行时的关键性能指标。作为 Three.js 等 WebGL 框架的黄金搭档,它通过浮动面板直观展示 FPS(帧率)、MS(渲染耗时&a…

ncmdumpGUI:一键解锁网易云音乐ncm文件的终极解决方案

ncmdumpGUI:一键解锁网易云音乐ncm文件的终极解决方案

2026/8/6 19:19:00

ncmdumpGUI:一键解锁网易云音乐ncm文件的终极解决方案 【免费下载链接】ncmdumpGUI C#版本网易云音乐ncm文件格式转换,Windows图形界面版本 项目地址: https://gitcode.com/gh_mirrors/nc/ncmdumpGUI 你是否曾经从网易云音乐下载了心爱的歌曲&am…

分布式配置中心选型实战:Nacos与Consul在创业场景下的对比

分布式配置中心选型实战:Nacos与Consul在创业场景下的对比

2026/8/5 6:02:27

分布式配置中心选型实战:Nacos与Consul在创业场景下的对比工程导读:本文深入讨论 分布式配置中心选型实战:Nacos与Consul在创业场景下的对比 在生产工程实践中的核心落地方案。基于 分布式架构与微服务设计 视角,剖析实际痛点、架…

MoneyPrinterPlus实战指南:AI视频批量生成与自动化发布完整解决方案

MoneyPrinterPlus实战指南:AI视频批量生成与自动化发布完整解决方案

2026/8/5 8:19:55

MoneyPrinterPlus实战指南:AI视频批量生成与自动化发布完整解决方案 【免费下载链接】MoneyPrinterPlus AI一键批量生成各类短视频,自动批量混剪短视频,自动把视频发布到抖音,快手,小红书,视频号上,赚钱从来没有这么容易过! 支持本地语音模型chatTTS,fasterwhisper,…

Unity相机抖动插件Camera-Shake集成与应用实战指南

Unity相机抖动插件Camera-Shake集成与应用实战指南

2026/8/6 0:00:51

1. 项目概述与核心价值最近在做一个动作游戏,需要给主角的重击和爆炸场景加点料,让打击感更足。我第一时间就想到了给相机加个抖动效果,毕竟这是提升玩家沉浸感最简单直接的手段之一。自己手写一个也不是不行,但时间成本高&#x…

Cocos Creator 3.7微信小游戏开发:从架构设计到提审上线的全流程实战指南

Cocos Creator 3.7微信小游戏开发:从架构设计到提审上线的全流程实战指南

2026/8/6 0:00:51

1. 项目概述:为什么需要一份3.7版本的专属适配指南?如果你是一位使用Cocos Creator开发微信小游戏的开发者,并且项目正运行在3.7版本上,那么你很可能已经感受到了那份“甜蜜的烦恼”。一方面,Cocos Creator 3.7是一个功…

AI编程实战:从Prompt工程到工具链集成,打造高效开发工作流

AI编程实战:从Prompt工程到工具链集成,打造高效开发工作流

2026/8/6 0:00:51

1. 项目概述:一次开源AI编程课程的深度重构 最近,我把自己的开源AI编程课程《Claude Code》做了一次从里到外的大更新。如果你对利用Claude、Codex这类大模型来辅助编程感兴趣,或者正在寻找一个能跟上最新AI编码工具迭代节奏的学习路径&#…

摆脱论文困扰!盘点2026年全网爆红的的AI论文写作工具

摆脱论文困扰!盘点2026年全网爆红的的AI论文写作工具

2026/8/6 5:43:30

一天写完毕业论文在2026年已不再是天方夜谭。2026年最炸裂、实测能大幅提速的AI论文写作工具,覆盖选题构思、文献整理、内容生成、格式排版等核心场景,真正帮你高效搞定论文难题。 一、全流程王者:一站式搞定论文全链路(一天定稿首…

导师推荐!2026最新AI论文工具测评与实用推荐

导师推荐!2026最新AI论文工具测评与实用推荐

2026/8/4 14:25:14

2026年真正好用的AI论文工具,核心看生成的论文质量、低AI味、格式正确、学术适配四大指标。综合实测,千笔AI、ThouPen、豆包、DeepSeek、Grammarly 是当前最值得推荐的梯队,覆盖从免费到付费、从中文到英文、从文科到理工的全场景需求。 一、…

告别游戏崩溃:XCOM 2模组管理器的智能革命

告别游戏崩溃:XCOM 2模组管理器的智能革命

2026/8/4 15:11:03

告别游戏崩溃:XCOM 2模组管理器的智能革命 【免费下载链接】xcom2-launcher The Alternative Mod Launcher (AML) is a replacement for the default game launchers from XCOM 2 and XCOM Chimera Squad. 项目地址: https://gitcode.com/gh_mirrors/xc/xcom2-lau…