【深度学习】pytorch——实现CIFAR-10数据集的分类

笔记为自我总结整理的学习笔记,若有错误欢迎指出哟~

往期文章:
【深度学习】pytorch——快速入门

CIFAR-10分类

  • CIFAR-10简介
  • CIFAR-10数据集分类实现步骤
  • 一、数据加载及预处理
    • 实现数据加载及预处理
    • 归一化的理解
    • 访问数据集
      • Dataset对象
      • Dataloader对象
  • 二、定义网络
  • 三、定义损失函数和优化器(loss和optimizer)
  • 四、训练网络并更新网络参数
    • enumerate函数
  • 五、测试网络
    • 部分数据集(实际的label)
    • 部分数据集(预测的label)
    • 整个测试集

CIFAR-10简介

CIFAR-10是一个常用的图像分类数据集,每张图片都是 3×32×32,3通道彩色图片,分辨率为 32×32。

它包含了10个不同类别,每个类别有6000张图像,其中5000张用于训练,1000张用于测试。这10个类别分别为:飞机、汽车、鸟类、猫、鹿、狗、青蛙、马、船和卡车。

CIFAR-10分类任务是将这些图像正确地分类到它们所属的类别中。对于这个任务,可以使用深度学习模型,如卷积神经网络(CNN)来实现高效的分类。

CIFAR-10分类任务是一个比较典型的图像分类问题,在计算机视觉领域中被广泛使用,是检验深度学习模型表现的一个重要基准。

CIFAR-10数据集分类实现步骤

  1. 使用torchvision加载并预处理CIFAR-10数据集
  2. 定义网络
  3. 定义损失函数和优化器
  4. 训练网络并更新网络参数
  5. 测试网络

一、数据加载及预处理

实现数据加载及预处理

import torch as t
import torchvision as tv
import torchvision.transforms as transforms
from torchvision.transforms import ToPILImage
show = ToPILImage() # 可以把Tensor转成Image,方便可视化

# 第一次运行程序torchvision会自动下载CIFAR-10数据集,大约100M。
# 如果已经下载有CIFAR-10,可通过root参数指定

# 定义对数据的预处理
transform = transforms.Compose([
        transforms.ToTensor(), # 转为Tensor
        transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5)), # 归一化
                             ])

# 训练集
trainset = tv.datasets.CIFAR10(		# PyTorch提供的CIFAR-10数据集的类,用于加载CIFAR-10数据集。
                    root='D:/深度学习基础/pytorch/data/', 	# 设置数据集存储的根目录。
                    train=True, 	# 指定加载的是CIFAR-10的训练集。
                    download=True,	# 如果数据集尚未下载,设置为True会自动下载CIFAR-10数据集。
                    transform=transform)	# 设置数据集的预处理方式。

# 数据加载器
trainloader = t.utils.data.DataLoader(
                    trainset, 		# 指定了要加载的训练集数据,即CIFAR-10数据集。
                    batch_size=4,	# 每个小批量(batch)的大小是4,即每次会加载4张图片进行训练。
                    shuffle=True, 	# 在每个epoch训练开始前,会打乱训练集中数据的顺序,以增加训练效果。
                    num_workers=2)	# 使用2个进程来加载数据,以提高数据的加载速度。

# 测试集
testset = tv.datasets.CIFAR10(
                    'D:/深度学习基础/pytorch/data/',
                    train=False, 
                    download=True, 
                    transform=transform)

testloader = t.utils.data.DataLoader(
                    testset,
                    batch_size=4, 
                    shuffle=False,
                    num_workers=2)

classes = ('plane', 'car', 'bird', 'cat',
           'deer', 'dog', 'frog', 'horse', 'ship', 'truck')

这段代码主要是使用PyTorch和torchvision库来加载并处理CIFAR-10数据集,其中包括训练集和测试集。

  1. import torch as timport torchvision as tv 导入了PyTorch和torchvision库。
  2. import torchvision.transforms as transforms 导入了torchvision.transforms模块,用于进行数据转换和增强操作。
  3. from torchvision.transforms import ToPILImage 导入了ToPILImage类,它可以将Tensor对象转换为PIL Image对象,以方便后续的可视化操作。
  4. show = ToPILImage() 创建一个ToPILImage对象,用于将张量(Tensor)对象转换为PIL Image对象,以便于后续的可视化操作。
  5. transform = transforms.Compose([...]) 定义对数据的预处理操作,将多个预处理操作组合在一起,形成一个数据预处理的管道。该管道首先使用transforms.ToTensor()函数将图像转换为张量(Tensor)对象,然后使用transforms.Normalize()函数对图像进行归一化操作,以便于后续的训练。
  6. trainset = tv.datasets.CIFAR10([...]) 使用tv.datasets.CIFAR10()函数加载CIFAR-10数据集,并指定数据集的存储位置、是否为训练集、是否需要下载等参数。还可以通过transform参数来指定对数据进行的预处理操作。
  7. trainloader = t.utils.data.DataLoader([...]) 使用PyTorch的DataLoader类来创建一个数据加载器,该加载器可以按照指定的批量大小将数据集分成小批量进行加载。可以指定加载器的参数,如批量大小、是否随机洗牌、使用的进程数等。
  8. testset = tv.datasets.CIFAR10([...])testloader = t.utils.data.DataLoader([...]) 与训练集的加载方式类似,只是将参数中的train改为False,表示这是测试集。
  9. classes = ('plane', 'car', 'bird', 'cat', 'deer', 'dog', 'frog', 'horse', 'ship', 'truck') 定义了CIFAR-10数据集中包含的10个类别。

注:tv.datasets.CIFAR10()函数会自动下载CIFAR-10数据集并存储到指定位置,如果已经下载过该数据集,可以通过root参数来指定数据集的存储位置,避免重复下载浪费时间和带宽。

归一化的理解

transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5)), # 归一化

transforms.Normalize()函数实现了对图像数据进行归一化操作。该函数的参数是均值和标准差,在CIFAR-10数据集中,每个像素有3个通道(R,G,B),因此传入的均值和标准差是一个长度为3的元组。这里(0.5, 0.5, 0.5)表示每个通道的均值为0.5,(0.5, 0.5, 0.5)表示每个通道的标准差也为0.5。具体地,对于每个像素的每个通道,该函数执行以下计算:

input[channel] = (input[channel] - mean[channel]) / std[channel]

其中,input[channel]表示一个像素的某个通道的像素值,mean[channel]std[channel]分别表示该通道的均值和标准差。通过这样的归一化操作,每个通道的像素值都将落在-1到1之间,从而便于模型的训练。

因此,这行代码的作用是对CIFAR-10数据集中的图像进行归一化,将每个通道的像素值映射到-1到1之间。

访问数据集

Dataset对象

Dataset对象是一个数据集,可以按下标访问,返回形如(data, label)的数据。

(data, label) = trainset[100]	# 从训练集中获取第100个样本的数据(图像)和标签。
print(classes[label])	

# (data + 1) / 2是为了还原被归一化的数据,将之前归一化的数据重新映射到0到1的范围内。
show((data + 1) / 2).resize((200, 200))

输出为:

ship
在这里插入图片描述

Dataloader对象

Dataloader是一个可迭代的对象,它将dataset返回的每一条数据拼接成一个batch,并提供多线程加速优化和数据打乱等操作。当程序对dataset的所有数据遍历完一遍之后,相应的对Dataloader也完成了一次迭代

dataiter = iter(trainloader)
images, labels = next(dataiter) # 返回4张图片及标签
print(','.join('%11s'%classes[labels[j]] for j in range(4)))
show(tv.utils.make_grid((images+1)/2)).resize((400,100))
  • 使用iter(trainloader)将训练数据加载器转换成一个迭代器对象dataiter

  • 使用next(dataiter)从迭代器中获取下一个批次的数据。这里假设每个批次的大小为4,所以imageslabels分别是一个包含4张图片和对应标签的张量。

  • 通过一个循环遍历了这4张图片的标签,并使用classes[labels[j]]将每个标签转换为对应的类别名称。classes是一个包含CIFAR-10数据集各个类别名称的列表。

  • 使用tv.utils.make_grid()函数将这4张图片拼接成一张网格图,并通过(images+1)/2将像素值从[-1, 1]的范围映射到[0, 1]的范围。使用show()函数显示图像,并调用resize()对图像进行调整大小,再使用print()输出调整大小后的图像。

输出为:
cat, truck, plane, deer
在这里插入图片描述

二、定义网络

LeNet网络,self.conv1第一个参数为3通道,因为CIFAR-10是3通道彩图

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.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 = F.max_pool2d(F.relu(self.conv1(x)), (2, 2)) 
        x = F.max_pool2d(F.relu(self.conv2(x)), 2) 
        x = x.view(x.size()[0], -1) 	# -1表示会自适应的调整剩余的维度
        x = F.relu(self.fc1(x))
        x = F.relu(self.fc2(x))
        x = self.fc3(x)        
        return x


net = Net()
print(net)

输出为:

Net(
  (conv1): Conv2d(3, 6, kernel_size=(5, 5), stride=(1, 1))
  (conv2): Conv2d(6, 16, kernel_size=(5, 5), stride=(1, 1))
  (fc1): Linear(in_features=400, out_features=120, bias=True)
  (fc2): Linear(in_features=120, out_features=84, bias=True)
  (fc3): Linear(in_features=84, out_features=10, bias=True)
)

模型包含以下层:

  1. self.conv1: 输入通道数为3,输出通道数为6,卷积核大小为5x5的卷积层。
  2. self.conv2: 输入通道数为6,输出通道数为16,卷积核大小为5x5的卷积层。
  3. self.fc1: 输入大小为16x5x5,输出大小为120的全连接层。
  4. self.fc2: 输入大小为120,输出大小为84的全连接层。
  5. self.fc3: 输入大小为84,输出大小为10的全连接层。

模型的前向传播函数(forward):

  1. 先经过第一个卷积层,然后应用ReLU激活函数和2x2的最大池化操作。
  2. 再经过第二个卷积层,同样应用ReLU激活函数和2x2的最大池化操作。
  3. 通过x.view(x.size()[0], -1)将特征张量x展平为一维向量,以便输入全连接层。
  4. 依次经过两个全连接层,并使用ReLU激活函数进行非线性变换。
  5. 最后一层是一个全连接层,输出大小为10,对应CIFAR-10数据集的10个类别。这里没有使用激活函数,因为该模型将其输出直接作为分类的得分。

总体而言,该模型由两个卷积层和三个全连接层组成,用于对CIFAR-10数据集进行图像分类。

三、定义损失函数和优化器(loss和optimizer)

from torch import optim
criterion = nn.CrossEntropyLoss() # 交叉熵损失函数
optimizer = optim.SGD(net.parameters(), lr=0.001, momentum=0.9)
  • nn.CrossEntropyLoss()创建了一个交叉熵损失函数的实例,用于计算分类任务中的损失。交叉熵损失函数通常用于多类别分类问题,它将模型的输出与真实标签进行比较,并计算出一个数值作为损失值,用来衡量模型预测与真实标签之间的差异。

  • optim.SGD(net.parameters(), lr=0.001, momentum=0.9)创建了一个随机梯度下降(SGD)优化器的实例。

    net.parameters()表示要优化的模型参数,即神经网络中的权重和偏置。

    lr=0.001是学习率(learning rate),控制每次参数更新的步长大小。

    momentum=0.9表示动量(momentum)参数,用于加速优化过程并避免陷入局部最优解。

四、训练网络并更新网络参数

t.set_num_threads(8)	# 设置线程数为 8,以加速训练过程。
for epoch in range(2):  	# 指定训练的轮数为 2 轮(epoch),即遍历整个数据集两次。
    
    running_loss = 0.0		# 记录当前训练阶段的损失值
    for i, data in enumerate(trainloader, 0):
        
        # 输入数据
        inputs, labels = data
        
        # 梯度清零
        optimizer.zero_grad()		# 每个 batch 开始时,将优化器的梯度缓存清零,以避免梯度累积
        
        # forward + backward 
        outputs = net(inputs)
        loss = criterion(outputs, labels)	# 进行前向传播,然后计算损失函数 loss
        loss.backward()   	# 自动计算损失函数相对于模型参数的梯度
        
        # 更新参数 
        optimizer.step()	# 使用优化器 optimizer 来更新模型的权重和偏置,以最小化损失函数
        
        # 打印log信息
        # loss 是一个scalar,需要使用loss.item()来获取数值,不能使用loss[0]
        running_loss += loss.item()
        if i % 2000 == 1999: # 每2000个batch打印一下训练状态
            print('[%d, %5d] loss: %.3f' \
                  % (epoch+1, i+1, running_loss / 2000))
            running_loss = 0.0
print('Finished Training')

输出结果:

[1,  2000] loss: 2.247
[1,  4000] loss: 1.974
[1,  6000] loss: 1.753
[1,  8000] loss: 1.605
[1, 10000] loss: 1.527
[1, 12000] loss: 1.472
[2,  2000] loss: 1.424
[2,  4000] loss: 1.386
[2,  6000] loss: 1.331
[2,  8000] loss: 1.303
[2, 10000] loss: 1.300
[2, 12000] loss: 1.275
Finished Training

enumerate函数

enumerate是Python内置函数之一,用于将一个可迭代的对象(如列表、元组、字符串等)组合为一个索引序列。它返回一个枚举对象,包含了原始对象中的元素以及对应的索引值。

enumerate函数的一般语法如下:

enumerate(iterable, start=0)

其中,iterable是要进行枚举的可迭代对象,start是可选参数,表示起始的索引值,默认为0。

下面是一个简单的例子来说明enumerate函数的用法:

fruits = ['apple', 'banana', 'cherry']
for index, fruit in enumerate(fruits):
    print(index, fruit)

输出结果:

0 apple
1 banana
2 cherry

在上述示例中,enumerate函数将列表fruits中的元素与对应的索引值配对,然后通过for循环依次取出每个元素和索引值进行打印。

在机器学习或深度学习中,enumerate函数常常与循环结合使用,用于遍历数据集或批次数据,并同时获取数据的索引值。这在模型训练过程中很有用,可以方便地记录当前处理的数据的位置信息。

五、测试网络

部分数据集(实际的label)

dataiter = iter(testloader)
images, labels = next(dataiter) # 一个batch返回4张图片
print('实际的label: ', ' '.join(\
            '%08s'%classes[labels[j]] for j in range(4)))
show(tv.utils.make_grid(images+1)/2).resize((400,100))

输出结果:

实际的label:  cat     ship      ship      plane

在这里插入图片描述

部分数据集(预测的label)

# 计算图片在每个类别上的分数
outputs = net(images)
# 得分最高的那个类
_, predicted = t.max(outputs.data, 1)

print('预测结果: ', ' '.join('%5s'\
            % classes[predicted[j]] for j in range(4)))

输出结果:

预测结果:  cat      car       ship        plane

在这里插入图片描述

整个测试集

correct = 0 # 预测正确的图片数
total = 0 # 总共的图片数

# 使用 torch.no_grad() 上下文管理器,表示在测试过程中不需要计算梯度,以提高速度和节约内存
with t.no_grad():
    for data in testloader:
        images, labels = data
        outputs = net(images)
        _, predicted = t.max(outputs, 1)
        total += labels.size(0)
        correct += (predicted == labels).sum()

print('10000张测试集中的准确率为: %d %%' % (100 * correct / total))

输出结果:

10000张测试集中的准确率为: 54 %

训练的准确率远比随机猜测(准确率10%)好。

本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若转载,请注明出处:/a/114864.html

如若内容造成侵权/违法违规/事实不符,请联系我们进行投诉反馈qq邮箱809451989@qq.com,一经查实,立即删除!

相关文章

【css3】涟漪动画

效果展示 dom代码 <div class"mapSelfTitle66"><div></div> </div> 样式代码 .mapSelfTitle66{width:120px;height:60px;position: relative;&>div{width:100%;height:100%;background: url("~/assets/images/video_show/err…

数据结构:邻接矩阵与邻接表

模型图 邻接矩阵 用于反应图中任意两点之间的关联&#xff0c;用二维数组表示比较方便 以行坐标为起点&#xff0c;列坐标为终点如果两个点之间有边&#xff0c;那么标记为绿色&#xff0c;如图&#xff1a; 适合表示稠密矩阵 邻接表 用一维数组 链表的形式表示&#xff…

离散数学实践(2)-编程实现关系性质的判断

*本文为博主本人校内的离散数学专业课的实践作业。由于实验步骤已经比较详细&#xff0c;故不再对该实验额外提供详解&#xff0c;本文仅提供填写的实验报告内容与代码部分&#xff0c;以供有需要的同学学习、参考。 -------------------------------------- 编程语言&#xff…

高效处理异常值的算法:One-class SVM模型的自动化方案

一、引言 数据清洗和异常值处理在数据分析和机器学习任务中扮演着关键的角色。清洗数据可以提高数据质量&#xff0c;消除噪声和错误&#xff0c;从而确保后续分析和建模的准确性和可靠性。而异常值则可能对数据分析结果产生严重影响&#xff0c;导致误导性的结论和决策。因此&…

人工智能与卫星:颠覆性技术融合开启太空新时代

人工智能与卫星&#xff1a;颠覆性技术融合开启太空新时代 摘要&#xff1a;本文将探讨人工智能与卫星技术的融合&#xff0c;并介绍其应用、发展和挑战。通过深入了解这一领域的前沿动态&#xff0c;我们将展望一个由智能卫星驱动的未来太空时代。 一、引言 近年来&#xf…

uniapp小程序砸金蛋抽奖

砸之前是金蛋png图片&#xff0c;点击砸完之后切换砸金蛋动效gif图片&#xff1b; 当前代码封装为砸金蛋的组件&#xff1b; vue代码&#xff1a; <template><view class"page" v-if"merchantInfo.cdn_static"><image class"bg&qu…

第6章_多表查询

文章目录 多表查询概述1 一个案例引发的多表连接1.1 案例说明1.2 笛卡尔积理解演示代码 2 多表查询分类讲解2.1 等值连接 & 非等值连接2.1.1 等值连接2.1.2 非等值连接 自连接 & 非自连接内连接与外连接演示代码 3 SQL99语法实现多表查询3.1 基本语法3.2 内连接&#x…

HTML脚本、字符实体、URL

HTML脚本&#xff1a; JavaScript 使 HTML 页面具有更强的动态和交互性。 <script> 标签用于定义客户端脚本&#xff0c;比如 JavaScript。<script> 元素既可包含脚本语句&#xff0c;也可通过 src 属性指向外部脚本文件。 JavaScript 最常用于图片操作、表单验…

【机器学习】几种常用的机器学习调参方法

在机器学习中&#xff0c;模型的性能往往受到模型的超参数、数据的质量、特征选择等因素影响。其中&#xff0c;模型的超参数调整是模型优化中最重要的环节之一。超参数&#xff08;Hyperparameters&#xff09;在机器学习算法中需要人为设定&#xff0c;它们不能直接从训练数据…

Locust:可能是一款最被低估的压测工具

01、Locust介绍 开源性能测试工具https://www.locust.io/&#xff0c;基于Python的性能压测工具&#xff0c;使用Python代码来定义用户行为&#xff0c;模拟百万计的并发用户访问。每个测试用户的行为由您定义&#xff0c;并且通过Web UI实时监控聚集过程。 压力发生器作为性…

Nacos报错Connection refused (Connection refused)(最后原因醉了,非常醉)

目录 一、问题产生二、排查思路1.nacos拒绝连接&#xff0c;排查思路&#xff1a;2.Nacos启动成功但是拒绝连接的几种原因&#xff1a; 三、实操过程&#xff08;着急解决问题直接看这个&#xff09;1.启动Nacos2.查看Nacos启动日志3.根据日志处理问题4.修改Nacos5.重启Nacos 一…

CSS基础知识点速览

1 基础认识 1.1 css的介绍 CSS:层叠样式表(Cascading style sheets) CSS作用&#xff1a; 给页面中的html标签设置样式 css写在style标签里&#xff0c;style标签一般在head标签里&#xff0c;位于head标签下。 <style>p{color: red;background-color: green;font-size…

Git客户端软件 Tower mac中文版特点说明

Tower mac是一款Mac OS X系统上的Git客户端软件&#xff0c;它提供了丰富的功能和工具&#xff0c;帮助用户更加方便地管理和使用Git版本控制系统。 Tower mac软件特点 1. 界面友好&#xff1a;Tower的界面友好&#xff0c;使用户能够轻松地掌握软件的使用方法。 2. 多种Git操…

探索主题建模:使用LDA分析文本主题

在数据分析和文本挖掘领域&#xff0c;主题建模是一种强大的工具&#xff0c;用于自动发现文本数据中的隐藏主题。Latent Dirichlet Allocation&#xff08;LDA&#xff09;是主题建模的一种常用技术。本文将介绍如何使用Python和Gensim库执行LDA主题建模&#xff0c;并探讨主题…

STM32F407的系统定时器

文章目录 系统定时器SysTick滴答定时器寄存器STK_CTRL 控制寄存器STK_LOAD 重载寄存器STK_VAL 当前值寄存器STK_CALRB 校准值寄存器 非系统初始化 Systick 定时器SysTick_InitSysTick_CLKSourceConfig delay_us寄存器delay_us库函数delay_xms短时delay_ms长时SysTick_Config 系…

Firefox修改缓存目录的方法

打开Firefox&#xff0c;在地址栏输入“about:config” 查找是否有 browser.cache.disk.parent_directory&#xff0c;如果没有就新建一个同名的字符串&#xff0c;然后修改值为你要存放Firefox浏览器缓存的目录地址&#xff08;E:\FirefoxCacheFiles&#xff09; 然后重新…

在python中加载tensorflow-probability模块和numpy模块

目录 操作步骤&#xff1a; 注意&#xff1a; 问题&#xff1a; 解决办法&#xff1a; 操作步骤&#xff1a; 在虚拟环境的文件夹中&#xff0c;找到Scripts文件夹&#xff0c;点击进去&#xff0c;找到地址栏&#xff0c;在地址栏中输入cmd&#xff0c;进入如下界面。 输…

详解Java经典数据结构——HashMap

Java 的 HashMap 是一个常用的基于哈希表的数据结构&#xff0c;它实现了 Map 接口&#xff0c;可以存储键值对。下面我们进行详细介绍&#xff1a; 基本结构&#xff1a;HashMap 底层是基于哈希表来实现的&#xff0c;每次插入一个键值对时&#xff0c;会先对该键进行 Hash 运…

思维训练3

题目描述1 Problem - A - Codeforces 题目分析 样例1解释&#xff1a; 对于此题&#xff0c;我们采用贪心的想法&#xff0c;从1到n块数越少越好&#xff0c;故刚好符合最少的块数即可&#xff0c;由于第1块与第n块是我们必须要走的路&#xff0c;所以我们可以根据这两块砖的…

C++--二叉搜索树初阶

前言&#xff1a;二叉搜索树是一种常用的数据结构&#xff0c;支持快速的查找、插入、删除操作&#xff0c;C中map和set的特性也是以二叉搜索树作为铺垫来实现的&#xff0c;而二叉搜索树也是一种树形结构&#xff0c;所以&#xff0c;在学习map和set之前&#xff0c;我们先来学…