news 2026/8/2 5:12:14

从零开始运用PyTorch构建神经网络实现分辨蚂蚁和蜜蜂

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
从零开始运用PyTorch构建神经网络实现分辨蚂蚁和蜜蜂

在初学PyToch和神经网络时,面对众多的概念和函数,我们可能会很迷惑如雾里看花。这时最好的方法便是做一个简单的小项目理解这个项目的构成,虽然博主也是新手,但我们不妨一起通过PyTorch来编写一个可以区分蜜蜂和蚂蚁的net。这是蜜蜂和蚂蚁训练集测试集的下载地址,来源于B站视频评论区。

首先是我们需要的包:

import torch import torch.nn as nn import torch.nn.function as F import torchvision #读取图片 from torchvision import transforms

然后第一步是录入我们训练和测试用的照片。先设定好照片的统一格式:

transform=transforms.Compose([ transforms.Resize((32,32)), #像素统一转换为32*32, #如果想让照片变得清晰点记得多加几个卷积层 transforms.Tensor #统一转化为张量 ])

然后导入训练照片和测试照片:

trainset=torchvision.datasets.Imagefloder( root=r'训练照片储存地址',#Python中地址的'\'会和'\n'混淆所以一般在前面加个'r'或者用'\\' trainsform=transform ) trainloader=torch.utils.data.DataLoader( trainset, #定义如何给net图片 batch_size=4, #每次会给四张图片 shuffle=True, #每次完整的训练伦次会打乱照片顺序 num_wokers=0 #单进程运行 ) #同理定义测试集 testset=torchvision.datasets.Imagefloder( root=r'训练照片储存地址', trainsform=transform ) testloader=torch.utils.data.DataLoader( trainset, #定义如何给net图片 batch_size=4, #每次会给四张图片 shuffle=True, #每次完整的训练伦次会打乱照片顺序 num_wokers=0 #单进程运行 ) #定义种类 classes=('ants','bees')

接下来开始定义net:

class Net(nn.Module): def __init__(self): super(Net,self).__init__ #标准格式 self.conv1=nn.Conv2d(3,6,5) #定义第一个卷积层,3代表彩色照片,6代表6个卷积核, #5代表每个卷积核尺寸为5*5 self.pool=nn.MaxPool2d(2,2) #定义池化矩阵 self.conv2=nn.Conv2d(6,16,5) #定义第二个卷积层 self.fc1=nn.Linear(16*5*5,120)#16和最后一个卷积层输出对应,5*5由公式计算得到 self.fc2=nn.Linear(120,84) #这里先转换为120,再到84最后到2(ants和bees)是 self.fc3=nn.Linear(84,2) #公式写法,当然也可以直接由120到2 def forward(self,x): x=self.pool(F.relu(self.conv1(x))) #按顺序池化 x=self.pool(F.relu(self.conv2(x))) x=torch.flatten(x,1) #展平,fc只接收一维变量 x=F.relu(self.fc1(x)) x=F.relu(self.fc2(x)) x=self.fc3(x) return x #输出结果 net=Net()

在fc1中的5*5原自公式:

其中W指照片原尺寸,K代表卷积核大小(或池化大小),P代表填充(默认为零),S代表步长(经过卷积核一般为1,经过池化矩阵我们这里为2)。也就是说如果我们统一的图片尺寸是256*256,那么经过两组卷积池化后会尺寸变成61*61,即fc1的括号变成(16*61*61,120),这会相当麻烦,所以如果大家想取大像素还请适当增加卷积层个数。

接下来我们开始梯度下降:

criterion=nn.CrossEntropyLoss() optimizer=torch.optim.SGD(net.parameters(),lr=0.01,momentum=0.9)#优化器选择带动量的SGD print("开始训练") for epoch in range(训练轮数): for i , data in enumerate(trainloader,0): inputs,labels = data outputs=net(inputs) optimitzer.zero_grad() #每次梯度记得归零 loss=criter(outputs,labels) loss.backward optimizer.step() print("训练完成") #储存训练结束后的参数 PATH=r'\-.pth' #储存的地址,记得提前建好.pth文件 torch.save(net.state_dict(),PATH)

最后我们对训练完成后的net进行测试:

snet.load_state_dict(torch.load(PATH)) #读取上面储存的参数 net.eval() #结束学习(即下面运行测试集时不进行迭代) total=0 correct=0 with torch.no_grad(): #不计算梯度(节约算力) for data in testloader: images , labels = data outputs=net(images) _, predicted = torch.max(outputs,1)# '_,'指把返回的images值舍弃, #毕竟我们只在意labels对不对 #最后的1指的是一维 #不直接使用total=len(testset)的原因是电脑会舍弃总数对每组照片的余数 total +=labels.size(0) correct +=(predicted == labels).sum().itrm() print(f"正确率为:{100*correct//total}%。")

至此我们的分辨蜜蜂和蚂蚁的神经网络就此写完了,接下来我们可以随便测试一张照片,同时想要直接读取存好的地址不重新训练(浪费时间),可以这样写:

import os PATH=r'-\.pth' #训练后参数储存的地址 if os.path.exists(PATH): net.load_state_dict(torch.load(PATH)) else: #训练过程 #…………

然后:

from PIL import Image image_path=r'你想测试的照片地址' img=Image.open(image_path).convert('RGB') img_tensor=transfrom(img) #将我们选择的照片变为3*32*32的张量 img_tensor=img_tensor.unsqueeze(0) #在这个三维张量前再插入一个维度(批次1) #因为Conv2d只支持四维变量 with torch.no_grad: outputs=net(img) -, predicted = torch.max(outputs,1) result=predicted.item() print(classes[result]) #输出具体类别

最后是如何运用显卡(GPU)加速我们的训练:

device=torch.device("cuda:0") #首先把net转移进GPU net=Net() net.to(device) #最后记得把数据也转移进GPU,net和数据必须同时在CPU或者同时在GPU #……训练过程…… inputs = inputs.to(device) labels = labels.to(device) #……正常后续流程……

最后说说题外话,我感觉用PyCharm编写Python中的会方便很多,它会直接提醒你括号里的数字代表什么:

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/8/2 5:11:46

AIGC降重工具指南:本科生学术写作必备

1. 项目概述:本科生必备的AIGC降重工具指南在学术写作和课程作业中,如何合理使用AI生成内容(AIGC)同时避免抄袭风险,已成为当代大学生面临的新挑战。根据2023年高等教育学术诚信报告显示,超过67%的院校已开始使用专业工具检测AI生…

作者头像 李华
网站建设 2026/8/2 5:09:09

微模块数据中心:乐高式敏捷部署与高效运维实战解析

1. 项目概述:从“机房”到“乐高积木”的进化干了十几年IT基础设施,从早期的“小黑屋”式机房,到后来的标准化数据中心,再到如今遍地开花的微模块,我亲眼见证了数据中心形态的演变。今天聊的“微模块数据中心解决方案”…

作者头像 李华
网站建设 2026/8/2 5:07:09

老师傅请假了,报价谁出?一个 STEP,AI 3 分钟出价

传统机加工报价,到底卡在哪?看图费时。3D 结构复杂,特征多、理解困难,人工读图慢且易漏。工序漏算。孔、腔、螺纹、倒角逐个数,稍有遗漏,报价就不准。依赖老师傅。报价靠经验,老师傅会、新人不会…

作者头像 李华
网站建设 2026/8/2 5:00:31

AI驱动消费新范式:从意图识别到商业闭环的实战解析

1. 项目概述:一场由AI驱动的春节消费狂欢春节档向来是各大互联网平台营销的必争之地,但今年,战火从传统的电商、支付、本地生活,直接烧到了AI领域。当大家还在讨论大模型能写诗、能画画时,阿里旗下的通义千问直接甩出了…

作者头像 李华
网站建设 2026/8/2 4:59:40

W25QXX SPI Flash模块:嵌入式存储解决方案与实战应用

1. 项目概述:W25QXX DataFlash Board 是什么?如果你玩过ESP32、STM32或者树莓派Pico这类微控制器,肯定遇到过存储空间不够用的问题。程序代码、配置文件、图片、音频,这些数据动不动就几百K甚至几兆,芯片内置的那点Fla…

作者头像 李华
网站建设 2026/8/2 4:55:15

E103-W02 Wi-Fi模块AT指令全解析:从TCP/UDP到HTTP/HTTPS实战指南

1. 项目概述:从零上手E103-W02最近在做一个物联网小项目,需要把传感器数据通过Wi-Fi传到服务器,选型时看中了EBYTE的E103-W02模块。这玩意儿体积小、功耗低,最关键的是官方资料说它支持TCP、UDP、HTTP甚至云透传,功能挺…

作者头像 李华