尧图网络 高端网站定制 · 原创设计
免费咨询热线
400-888-6620
免费获取方案
训练一个分类器
目录一、关于数据二、训练分类器1、加载并标准化CIFAR102、定义卷积网络3、定义损失函数和优化器4、训练网络5、在测试数据上测试神经网络三、在GPU上训练一、关于数据通常当你需要处理图像、文本、音频或视频数据时你可以使用标准Python包来将数据加载到NumPy数组中。然后将数组转化为torch.*Tensor。对于图像可以使用包PillowOpenCV对于音频可以使用包scipylibrosa对于文本可以使用NLTKSpaCy特别是对于视觉我们创建了一个名为torchvision的包它包含用于常见数据集的数据加载器如ImagenetCIFAR10MNIST等以及用于图像的数据转换器即torchvision.datasets和torch.utils.data.DataLoader。对于本教程我们将使用CIFAR10数据集。它包含10个类 ‘airplane’, ‘automobile’, ‘bird’, ‘cat’, ‘deer’, ‘dog’, ‘frog’, ‘horse’, ‘ship’, ‘truck’。图片的大小为3x32x32即尺寸为32×32像素的3通道彩色图像。二、训练分类器我们将会依次执行以下步骤使用torchvision加载和标准化CIFAR10训练和测试数据集定义卷积神经网络定义损失函数在训练数据上训练网络在测试数据上测试网络1、加载并标准化CIFAR10import torch import torchvision import torchvision.transforms as transformstorchvision数据集的输出是范围[0,1]的PILImage图像。 我们将它们转换为归一化范围的tensor[-1,1]。transform transforms.Compose( [transforms.ToTensor(), transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))]) trainset torchvision.datasets.CIFAR10(root./data, trainTrue, downloadTrue, transformtransform) trainloader torch.utils.data.DataLoader(trainset, batch_size4, shuffleTrue, num_workers2) testset torchvision.datasets.CIFAR10(root./data, trainFalse, downloadTrue, transformtransform) testloader torch.utils.data.DataLoader(testset, batch_size4, shuffleFalse, num_workers2) classes (plane, car, bird, cat, deer, dog, frog, horse, ship, truck)输出Downloading https://www.cs.toronto.edu/~kriz/cifar-10-python.tar.gz to ./data/cifar-10-python.tar.gz Extracting ./data/cifar-10-python.tar.gz to ./data Files already downloaded and verified让我们展示一些训练图片import matplotlib.pyplot as plt import numpy as np # functions to show an image def imshow(img): img img / 2 0.5 # unnormalize npimg img.numpy() plt.imshow(np.transpose(npimg, (1, 2, 0))) plt.show() # get some random training images dataiter iter(trainloader) images, labels dataiter.next() # show images imshow(torchvision.utils.make_grid(images)) # print labels print( .join(%5s % classes[labels[j]] for j in range(4)))输出frog ship cat plane2、定义卷积网络import torch.nn as nn import torch.nn.functional as F class Net(nn.Module): def __init__(self): super(Net, self).__init__() self.conv1 nn.Conv2d(3, 6, 5) self.pool nn.MaxPool2d(2, 2) self.conv2 nn.Conv2d(6, 16, 5) self.fc1 nn.Linear(16 * 5 * 5, 120) self.fc2 nn.Linear(120, 84) self.fc3 nn.Linear(84, 10) def forward(self, x): x self.pool(F.relu(self.conv1(x))) x self.pool(F.relu(self.conv2(x))) x x.view(-1, 16 * 5 * 5) x F.relu(self.fc1(x)) x F.relu(self.fc2(x)) x self.fc3(x) return x net Net()3、定义损失函数和优化器让我们使用分类交叉熵损失函数和带动量的SGD。import torch.optim as optim criterion nn.CrossEntropyLoss() optimizer optim.SGD(net.parameters(), lr0.001, momentum0.9)4、训练网络我们只需循环遍历数据迭代器并将输入提供给神经网络并进行优化。for epoch in range(2): # loop over the dataset multiple times running_loss 0.0 for i, data in enumerate(trainloader, 0): # get the inputs; data is a list of [inputs, labels] inputs, labels data # zero the parameter gradients optimizer.zero_grad() # forward backward optimize outputs net(inputs) loss criterion(outputs, labels) loss.backward() optimizer.step() # print statistics running_loss loss.item() if i % 2000 1999: # print every 2000 mini-batches print([%d, %5d] loss: %.3f % (epoch 1, i 1, running_loss / 2000)) running_loss 0.0 print(Finished Training)输出[1, 2000] loss: 2.169 [1, 4000] loss: 1.808 [1, 6000] loss: 1.659 [1, 8000] loss: 1.553 [1, 10000] loss: 1.488 [1, 12000] loss: 1.455 [2, 2000] loss: 1.379 [2, 4000] loss: 1.346 [2, 6000] loss: 1.320 [2, 8000] loss: 1.305 [2, 10000] loss: 1.275 [2, 12000] loss: 1.262 Finished Training5、在测试数据上测试神经网络我们已经在训练数据集上训练了两次。 但我们需要检查神经网络是否已经学到了什么。我们将通过预测神经网络输出的类标签来检查这一点并根据真实情况进行检查。 如果预测正确我们将样本添加到正确预测列表中。我们首先展示测试集中的一些图片dataiter iter(testloader) images, labels dataiter.next() # print images imshow(torchvision.utils.make_grid(images)) print(GroundTruth: , .join(%5s % classes[labels[j]] for j in range(4)))输出GroundTruth: cat ship ship plane现在我们来看一看神经网络认为这些图片是什么。outputs是10个类的能量。 一个类的能量越高网络认为图像是特定类的可能性越大。 那么让我们得到最高能量的索引outputs net(images) _, predicted torch.max(outputs, 1) print(Predicted: , .join(%5s % classes[predicted[j]] for j in range(4)))输出Predicted: cat plane plane ship查看神经网络在整个数据集上的表现correct 0 total 0 with torch.no_grad(): for data in testloader: images, labels data outputs net(images) _, predicted torch.max(outputs.data, 1) total labels.size(0) correct (predicted labels).sum().item() print(Accuracy of the network on the 10000 test images: %d %% % ( 100 * correct / total))输出Accuracy of the network on the 10000 test images: 54 %查看表现最好的类和最坏的类class_correct list(0. for i in range(10)) class_total list(0. for i in range(10)) with torch.no_grad(): for data in testloader: images, labels data outputs net(images) _, predicted torch.max(outputs, 1) c (predicted labels).squeeze() for i in range(4): label labels[i] class_correct[label] c[i].item() class_total[label] 1 for i in range(10): print(Accuracy of %5s : %2d %% % ( classes[i], 100 * class_correct[i] / class_total[i]))输出Accuracy of plane : 46 % Accuracy of car : 63 % Accuracy of bird : 50 % Accuracy of cat : 37 % Accuracy of deer : 40 % Accuracy of dog : 51 % Accuracy of frog : 70 % Accuracy of horse : 48 % Accuracy of ship : 76 % Accuracy of truck : 64 %三、在GPU上训练如果我们有可用的CUDA我们首先将我们的设备定义为第一个可见的cuda设备device torch.device(cuda:0 if torch.cuda.is_available() else cpu) # Assuming that we are on a CUDA machine, this should print a CUDA device: print(device)输出cuda:0本节的其余部分假定设备是CUDA设备。然后这些方法将递归遍历所有模块并将其参数和缓冲区转换为CUDA tensornet.to(device)还必须将每一步的输入和目标发送到GPUinputs, labels data[0].to(device), data[1].to(device)
RELATED

相关推荐

SpringBoot+Vue养老院管理系统设计与实现

SpringBoot+Vue养老院管理系统设计与实现

1. 项目背景与核心价值 养老院管理系统作为智慧养老领域的关键数字化工具,其优化设计直接关系到养老机构的运营效率和服务质量。这个毕业设计项目从实际业务场景出发,通过完整的源码实现、部署文档和技术讲解,构建了一套可落地的解决方案。 …

📅 2026/9/18 11:19:05
BetterNCM-Installer:3分钟解决网易云插件安装难题

BetterNCM-Installer:3分钟解决网易云插件安装难题

BetterNCM-Installer:3分钟解决网易云插件安装难题 【免费下载链接】BetterNCM-Installer 一键安装 Better 系软件 项目地址: https://gitcode.com/gh_mirrors/be/BetterNCM-Installer 还在为网易云音乐插件安装而头疼吗?你是否曾经因为复杂的安装…

📅 2026/9/15 13:48:50
PHP项目敏感信息加密管理:git-crypt与Composer协同实战

PHP项目敏感信息加密管理:git-crypt与Composer协同实战

1. 项目概述:当PHP依赖项需要“上锁”时在开发一个涉及敏感配置(比如数据库密码、第三方API密钥、支付网关令牌)的PHP项目时,我们常常会面临一个两难境地。一方面,我们希望将这些配置纳入版本控制系统(如Gi…

📅 2026/8/24 14:52:15
MORE NEWS

更多资讯

📰

BCT脑网络分析实战指南:从MATLAB安装到可发表指标计算

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

📰

SQLite迁移PostgreSQL全攻略:脚本写法与踩坑避坑指南

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

📰

Flutter鸿蒙应用稳定性排查:黑屏、白屏与OOM闪退实战指南

做 Flutter 鸿蒙应用的稳定性问题排查,说句实话,比纯 Android 上要绕不少。同一个 App 在 Android 上跑得好好的,一适配到鸿蒙,线上就开始反馈“打不开”“白屏”“用一会儿就闪退”。黑屏、白屏、OOM 闪退、内存持续增长&#xf…

📰

CSS Grid 核心概念与实战:从网格线到响应式布局一文讲透

做前端这几年,如果让我评一个“文档看过、真到用时就怂”的 CSS 模块,CSS Grid 网格布局绝对排第一。不少人都在教程里见它炫技,什么九宫格、双飞翼、瀑布流,看起来无所不能,可真到自己写页面,手指头还是习…

📰

CANN ops-cv 非连续 Tensor 完全指南:基于 (shape, strides, offset) 的内存视图表示

CANN ops-cv 非连续 Tensor 完全指南:基于 (shape, strides, offset) 的内存视图表示 【免费下载链接】ops-cv 本项目是CANN提供的图像处理、目标检测相关的算子库,实现网络在NPU上加速计算。 项目地址: https://gitcode.com/cann/ops-cv 非连续 …

📰

什么是折腾一个优化:渐进式工程优化方法论

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

TODAY

今日更新

THIS WEEK

本周精选

THIS MONTH

本月热门

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

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

📞 💬