尧图网络 高端网站定制 · 原创设计
免费咨询热线
400-888-6620
免费获取方案
【刘二老师】pytorch深度学习笔记【09多分类问题】
【刘二老师】pytorch深度学习笔记【09多分类问题】一、概念多分类问题会用到 softmax 分类器。十个输出之间要互相抑制要有竞争性输出的本质是个分布。要求输出是个分布的条件1每一个输出值大于零 2加一起的概率和为1因此在处理多分类问题时前面这些层还用sigmoid层来处理最后一层用softmax层因为softmax可以输出一个分布。二、Softmax层一Softmax 公式因为用了指数满足了输出每项都大于零。又除以了总和所以满足输出加一起和为1。二Softmax层实例架构讲解首先0.20.1-0.1分别是经过了线性变换的输出值。对这三个输出值取e的指数运算得到1.221.110.90。sum求和再divide除以他们的和得到0.380.340.28它们加一起为1。整个绿框就表示softmax函数。三softmax得到分布后损失函数怎么做Loss 【负的原始标签】× log预测输出Y^------- 与二分类的交叉熵形式相似。红色框范围为NLLLoss实现该Loss的计算过程的代码。pytorch中的交叉熵损失整个红框范围都是交叉熵损失。使用交叉熵损失时神经网络的最后一层不要做激活因为激活部分包含在交叉熵损失模块中了。y 需要是 LongTensor长整型张量。交叉熵损失示例代码第一组预测预测值类别分别为201看[ ]中哪个位置数字最大就是哪个类别与 Y一致我们就可以初步判断第一组的损失函数较小。第二组预测022 不太一致损失较大。三、代码及详细解析importtorchfromtorchvisionimporttransformsfromtorchvisionimportdatasetsfromtorch.utils.dataimportDataLoaderimporttorch.nn.functionalasFimporttorch.optimasoptim# prepare datasetbatch_size64transformtransforms.Compose([transforms.ToTensor(),transforms.Normalize((0.1307,),(0.3081,))])# 归一化,均值和方差train_datasetdatasets.MNIST(root../dataset/mnist/,trainTrue,downloadTrue,transformtransform)train_loaderDataLoader(train_dataset,shuffleTrue,batch_sizebatch_size)test_datasetdatasets.MNIST(root../dataset/mnist/,trainFalse,downloadTrue,transformtransform)test_loaderDataLoader(test_dataset,shuffleFalse,batch_sizebatch_size)# design model using classclassNet(torch.nn.Module):def__init__(self):super(Net,self).__init__()self.l1torch.nn.Linear(784,512)self.l2torch.nn.Linear(512,256)self.l3torch.nn.Linear(256,128)self.l4torch.nn.Linear(128,64)self.l5torch.nn.Linear(64,10)defforward(self,x):xx.view(-1,784)# -1其实就是自动获取mini_batchxF.relu(self.l1(x))xF.relu(self.l2(x))xF.relu(self.l3(x))xF.relu(self.l4(x))returnself.l5(x)# 最后一层不做激活不进行非线性变换modelNet()# construct loss and optimizercriteriontorch.nn.CrossEntropyLoss()optimizeroptim.SGD(model.parameters(),lr0.01,momentum0.5)#冲量值设为0.5优化训练过程。# training cycle forward, backward, updatedeftrain(epoch):running_loss0.0forbatch_idx,datainenumerate(train_loader,0):# 获得一个批次的数据和标签inputs,targetdata optimizer.zero_grad()# 获得模型预测结果(64, 10)outputsmodel(inputs)# 交叉熵代价函数outputs(64,10),target64losscriterion(outputs,target)loss.backward()optimizer.step()running_lossloss.item()ifbatch_idx%300299:print([%d, %5d] loss: %.3f%(epoch1,batch_idx1,running_loss/300))running_loss0.0deftest():correct0total0withtorch.no_grad():#表示下面的代码不会计算梯度fordataintest_loader:images,labelsdata outputsmodel(images)#输出每一个样本是个矩阵矩阵里每一行最大值的下标拿出来对应的就是它的分类。_,predictedtorch.max(outputs.data,dim1)# dim 1 列是第0个维度行是第1个维度totallabels.size(0)#labels是一个N×1的矩阵N1size0就代表取N的值。correct(predictedlabels).sum().item()# 推测predicted与label之间的比较真就是1假就是0。print(accuracy on test set: %d %% %(100*correct/total))if__name____main__:#封装进if函数forepochinrange(10):train(epoch)test()一importtorchfromtorchvisionimporttransformsfromtorchvisionimportdatasetsfromtorch.utils.dataimportDataLoaderimporttorch.nn.functionalasFimporttorch.optimasoptimtransforms针对图像做各种各样处理对数据进行原始处理的工具。torch.nn.functional全连接层的激活不再像之前用sigmoid而是用更流行的r elu( ) 作为激活函数。torch.optim优化器包二batch_size64transformtransforms.Compose([transforms.ToTensor(),transforms.Normalize((0.1307,),(0.3081,))])# 归一化,均值和方差train_datasetdatasets.MNIST(root../dataset/mnist/,trainTrue,downloadTrue,transformtransform)train_loaderDataLoader(train_dataset,shuffleTrue,batch_sizebatch_size)test_datasetdatasets.MNIST(root../dataset/mnist/,trainFalse,downloadTrue,transformtransform)test_loaderDataLoader(test_dataset,shuffleFalse,batch_sizebatch_size)先做一个batch_size因为后面有datasets 和 dataloader。除了多了个transform剩下模块与上节课的一样。使用transform模块的原因下面的过程要通过transform的ToTensor来实现。神经网络在处理时希望输入数值比较小最好在-11之间最好能遵从正态分布。所以要把原始的{0……255}转变为图像张量取值为 [01]也是把单通道转变为多通道。C×W×H1×28×28Normalize归一化Nomalize均值标准差满足N01正态分布给神经网络进行训练。三classNet(torch.nn.Module):def__init__(self):super(Net,self).__init__()self.l1torch.nn.Linear(784,512)self.l2torch.nn.Linear(512,256)self.l3torch.nn.Linear(256,128)self.l4torch.nn.Linear(128,64)self.l5torch.nn.Linear(64,10)defforward(self,x):xx.view(-1,784)# -1其实就是自动获取mini_batchxF.relu(self.l1(x))xF.relu(self.l2(x))xF.relu(self.l3(x))xF.relu(self.l4(x))returnself.l5(x)# 最后一层不做激活不进行非线性变换modelNet()全连接网络要求输入是矩阵所以要把12828三阶的张量变成一阶的向量。如何把张量变成向量把图像的每一行拼起来构成一串就做出了向量所以每一行需要28×28784个元素。把张量变为向量的代码x x.view(-1, 784)view函数改变张量的形状变成二阶张量也就是矩阵。该矩阵有784列-1是占位符代表N这个维度N的大小交给 PyTorch 自动计算。【一批多少个图片样本每张图片784个像素】Linear输入输出relu 对每一层算出的结果进行激活最后一层 self.15(x) 不激活直接接到后面的softmax。四训练把一轮循环封装到函数deftrain(epoch):running_loss0.0forbatch_idx,datainenumerate(train_loader,0):# 获得一个批次的数据和标签inputs,targetdata optimizer.zero_grad()# 获得模型预测结果(64, 10)outputsmodel(inputs)# 交叉熵代价函数outputs(64,10),target64losscriterion(outputs,target)loss.backward()optimizer.step()running_lossloss.item()#取item把值拿出来否则要构建计算图。ifbatch_idx%300299:print([%d, %5d] loss: %.3f%(epoch1,batch_idx1,running_loss/300))running_loss0.0X存到 inputY 存到 target每300轮输出一次 running_loss五deftest():correct0total0withtorch.no_grad():#表示下面的代码不会计算梯度fordataintest_loader:images,labelsdata outputsmodel(images)#输出每一个样本是个矩阵矩阵里每一行最大值的下标拿出来对应的就是它的分类。_,predictedtorch.max(outputs.data,dim1)# dim 1 列是第0个维度行是第1个维度totallabels.size(0)#labels是一个N×1的矩阵N1size0就代表取N的值。correct(predictedlabels).sum().item()# 推测predicted与label之间的比较真就是1假就是0。print(accuracy on test set: %d %% %(100*correct/total))max函数predicted函数
RELATED

相关推荐

RISC-V IDE MounRiver Studio开发实战:TWEN32V RGB

RISC-V IDE MounRiver Studio开发实战:TWEN32V RGB

RISC-V IDE MounRiver Studio开发实战:TWEN32V RGB软件平台 Mounriver Studio,硬件平台TWENCH32V开发板。1、WS2812RGB RGB色彩模式是工业界的一种颜色标准,是通过对红、绿(G)、蓝(B)三个颜色通道的变化以及它们相互之间的叠加来得到各式各样…

📅 2026/9/15 13:09:12
将本地代码上传到gitee仓库

将本地代码上传到gitee仓库

在gitee中创建好项目,注意,不要选任何模板,不然有文件不再本地不太好处理1、选中项目文件夹;右键 -> Git Bash Here2、git init 初始化git3、git remote add origin 远程库地址4、git pull origin master 将码云上的仓库pull到…

📅 2026/8/24 9:58:43
大模型小白入门指南:收藏这份医疗垂直大模型学习资料,从原理到应用全解析!

大模型小白入门指南:收藏这份医疗垂直大模型学习资料,从原理到应用全解析!

本文介绍了大语言模型(LLM)的基本原理,包括其如何通过海量数据训练掌握语言规律和世界知识,以及如何通过预训练实现多任务处理。文章深入探讨了医疗垂直大模型的测试与优化方法,分析了模型在知识理解、应用及临床诊疗决…

📅 2026/8/24 9:58:43
MORE NEWS

更多资讯

📰

JSP+MySQL图书购物系统源码解析:MVC分层与部署避坑指南

简介:这是一套面向高校计算机相关专业学生的JavaWeb课程设计期末大作业资源,主题为基于JSP(MVC模式)与MySQL实现的网上图书购物系统,适合正在准备课程设计、期末大作业或需要JSP实战练手的新手参考。压缩包共75个文件&…

📰

智能体工程化实战:从Demo到业务级落地的完整路径

1. 从这期周报里我看到了什么:智能体不再只是演示这周的 GitHub Trending 榜单我翻来覆去看了好几遍,最大的感受就一句话:智能体项目正在集体从“能跑起来”往“能交付”的方向转。前两年大家关注的是哪个框架又出了新概念、哪个 Demo 能自动…

📰

自治系统与元进化闭环的工程化落地方法:从概念到可运行架构

这些年我参与过的“自治系统”项目,十有八九一开始都是这样:会议室白板上画着一个漂亮的闭环,感知、决策、执行、反馈,箭头绕一圈,旁边写着“元进化闭环”,听着特别高级。可三个月后再看,那圈箭…

📰

Jev模型接入Codex完整指南:从密钥申请到多文件重构实测

最近几天,不论你是刷技术社区还是看推荐流,应该都躲不开一个词:Jev。我最初以为又是某个营销号造出来的概念,直到身边几个做后端的朋友陆续开始聊“Jev密钥”“Jev在Codex里怎么配”,才意识到这东西的热度是实打实的。…

📰

Agent记忆海关:用AST扫描与双池隔离根治记忆污染

做 Agent 项目最怕什么?不是模型不够聪明,不是工具不够多,而是它记住了一堆不该记的东西,然后在某个关键时刻一本正经地拿错误信息去推理。我最近用 Python 3.14 重新搓了一套 Agent 记忆管理模块,代号"记忆海关&…

📰

JavaWeb问卷调查系统全解析:从部署到改造的实战指南

简介:这是一份基于JavaWeb的问卷调查系统完整源码与数据库压缩包,面向正在准备毕业设计、课程设计或期末大作业的Java学习者,也适合需要快速搭建在线问卷模块的开发者。包内包含完整的后端业务逻辑、前端交互页面以及数据库初始化脚本&#x…

TODAY

今日更新

THIS WEEK

本周精选

THIS MONTH

本月热门

读完文章,想聊聊您的网站?

告诉我们您的行业与需求,资深顾问一对一梳理方案与报价,全程免费。

📞 💬