手写数字识别(mxnet官网例子)

手写数字识别
简介:通过MNIST数据集建立一个手写数字分类器。

(MNIST对于手写数据分类任务是一个广泛使用的数据集)。

1.前提:mxnet 0.10及以上、python、jupyter notebook(有时间可以jupyter notebook的用法,如:PPT的制作)
pip install requests jupyter ——python下jupyter notebook 的安装
2.加载数据集:
import mxnet as mx
mnist = mx.test_utils.get_mnist()
此时MXNET数据集已完全加载到内存中(注:此法对于大型数据集不适用)
考虑要素:快速高效地从源直接流数据+输入样本的顺序
图像通常用4维数组来表示:(batch_size,num_channels,width,height)
对于MNIST数据集,因为是28*28灰度图像,所以只有1个颜色通道,width=28,height=28,本例中batch=100(批处理100),即输入形状是(batch_size,1,28,28)
数据迭代器通过随机的调整输入来解决连续feed相同样本的问题。

测试数据的顺序无关紧要。

batch_size = 100
train_iter=,mnist['train_label'], batch_size, shuffle=True)
val_iter = , mnist['test_label'], batch_size)
——初始化MNIST数据集的数据迭代器(2个:训练数据+测试数据)
3.训练+预测:(2种方法)(CNN优于MLP)
1)传统深度神经网络结构——MLP(多层神经网络)
MLP——MXNET的符号接口
为输入的数据创建一个占位符变量
data =
data =
——将数据从4维变成2维(batch_size,num_channel*width*height) fc1 = , num_hidden=128)
act1 = , act_type="relu")
——第一个全连接层及相应的激活函数
fc2 = , num_hidden = 64)
act2 = , act_type="relu")
——第二个全连接层及相应的激活函数
(声明2个全连接层,每层有128个和64个神经元)
fc3 = , num_hidden=10)
——声明大小10的最终完全连接层
mlp = , name='softmax')
——softmax的交叉熵损失
MNIST的MLP网络结构
以上,已完成了数据迭代器和神经网络的申明,下面可以进行训练。

超参数:处理大小、学习速率
import logging
logging.getLogger().setLevel(logging.DEBUG) ——记录到标准输出
mlp_model = , context=mx.cpu())
——在CPU上创建一个可训练的模块
mlp_model.fit(train_iter ——训练数据
eval_data=val_iter, ——验证数据
optimizer='sgd', ——使用SGD训练
optimizer_params={'learning_rate':0.1}, ——使用
固定的学习速率
eval_metric='acc', ——训练过程中报告准确性
batch_end_callback=, 100), ——每批次100数据输
出的进展num_epoch=10) ——训练
至多通过10个数据
预测:
test_iter = , None, batch_size)
prob = mlp_model.predict(test_iter)
assert prob.shape == (10000, 10)
——计算每一个测试图像可能的预测得分(prob[i][j]第i个测试图像包含j输出类)
test_iter = , mnist['test_label'], batch_size) ——预测精度的方法
acc =
mlp_model.score(test_iter, acc)
print(acc)
assert acc.get()[1] > 0.96
如果一切顺利的话,我们将看到一个准确的值大约是0.96,这意味着我们能够准确地预测96%的测试图像中的数字。

2)卷积神经网络(CNN)
卷积层+池化层
data =
conv1 = , kernel=(5,5), num_filter=20)
tanh1 = , act_type="tanh")
pool1 = , pool_type="max", kernel=(2,2), stride=(2,2))
——第一个卷积层、池化层
conv2 = , kernel=(5,5), num_filter=50) ——第二个卷积层
tanh2 = , act_type="tanh")
pool2 = , pool_type="max", kernel=(2,2), stride=(2,2))
flatten = ——第一个全连接层
fc1 = , num_hidden=500)
tanh3 = , act_type="tanh")
fc2 = , num_hidden=10)——第二个全连接层
lenet = , name='softmax') ——Softmax损失
LeNet第一个卷积层+池化层
lenet_model = , context=mx.cpu())
——在GPU上创建一个可训练的模块
lenet_model.fit(train_iter,
eval_data=val_iter,
optimizer='sgd',
optimizer_params={'learning_rate':0.1},
eval_metric='acc',
batch_end_callback = , 100),
num_epoch=10)
——训练(同MLP)
test_iter = , None, batch_size)
prob = lenet_model.predict(test_iter)
test_iter = , mnist['test_label'], batch_size)
acc = ——预测LeNet的准确性
lenet_model.score(test_iter, acc)
print(acc)
assert acc.get()[1] > 0.98
使用CNN,我们能够正确地预测所有测试图像的98%左右。

附:完整代码
1)MLP
2)CNN。

合集下载

lenet5应用实例

lenet5应用实例

lenet5应用实例
LeNet-5是一种经典的卷积神经网络架构,最初由Yann LeCun 等人在1998年提出,用于手写数字识别任务。

它被广泛应用于数字识别、图像分类等领域。

下面我将从多个角度给出LeNet-5的应用实例。

1. 手写数字识别,LeNet-5最初是为了解决手写数字识别问题而设计的。

它在MNIST数据集上取得了很好的效果,成为了早期数字识别任务的经典模型。

LeNet-5的应用实例包括自动识别支票上的手写数字、银行卡上的卡号识别等。

2. 物体识别,除了手写数字识别外,LeNet-5也被应用于物体识别任务。

通过对图像进行卷积和池化操作,LeNet-5可以有效地提取图像特征,从而用于识别不同类别的物体,例如交通标志、人脸识别等。

3. 文字识别,LeNet-5的卷积结构也使得它适用于文字识别任务。

例如,可以将LeNet-5应用于识别车牌上的文字、自动识别手写的地址信息等。

4. 医学图像分析,LeNet-5也被应用于医学图像分析领域,例如X光片的识别、病理图像的分析等。

通过LeNet-5对医学图像进行特征提取和分类,可以帮助医生进行疾病诊断和治疗。

5. 智能驾驶,LeNet-5的应用还延伸到智能驾驶领域,例如通过LeNet-5对道路标志、行人、车辆等进行识别,以实现自动驾驶和交通管理等功能。

总之,LeNet-5作为卷积神经网络的先驱之一,在数字识别、物体识别、文字识别、医学图像分析、智能驾驶等领域都有着广泛的应用实例。

其经典的网络结构和有效的特征提取能力使得它成为了深度学习领域的重要里程碑之一。

使用神经网络进行手写数字识别的方法

使用神经网络进行手写数字识别的方法

使用神经网络进行手写数字识别的方法随着人工智能的发展,神经网络在图像识别领域发挥了重要作用。

其中,手写数字识别是神经网络应用的一个重要方向。

本文将介绍使用神经网络进行手写数字识别的方法。

一、神经网络的基本原理神经网络是一种模仿人脑神经元网络结构和工作方式的计算模型。

它由输入层、隐藏层和输出层组成,每一层都由多个神经元节点组成。

神经网络通过对输入数据进行加权和激活函数处理,从而输出预测结果。

在手写数字识别中,我们可以将每个手写数字图像作为输入数据,每个像素点的灰度值作为输入特征。

神经网络通过学习大量已标记的手写数字图像,调整权重和偏置,从而实现对手写数字的准确识别。

二、数据预处理在使用神经网络进行手写数字识别之前,需要对数据进行预处理。

首先,我们需要将手写数字图像转换为灰度图像,以减少输入特征的维度。

其次,对图像进行归一化处理,将像素值缩放到0到1之间,以便神经网络更好地学习和处理数据。

除了对图像进行处理,还需要对标签进行处理。

手写数字识别通常使用独热编码(One-Hot Encoding)对标签进行表示。

例如,对于数字0,其独热编码为[1, 0, 0, 0, 0, 0, 0, 0, 0, 0],对于数字1,其独热编码为[0, 1, 0, 0, 0, 0, 0, 0, 0, 0],以此类推。

三、神经网络的构建在构建神经网络时,我们可以选择不同的网络结构和参数设置。

常见的神经网络结构包括多层感知机(Multilayer Perceptron,MLP)、卷积神经网络(Convolutional Neural Network,CNN)等。

以多层感知机为例,我们可以选择输入层节点数、隐藏层节点数、隐藏层数量和输出层节点数等。

通过调整网络结构和参数,可以提高神经网络的准确率和泛化能力。

四、神经网络的训练神经网络的训练是指通过大量的已标记数据,调整网络的权重和偏置,使其能够准确地预测未标记数据的标签。

训练神经网络通常采用反向传播算法(Backpropagation),该算法通过计算预测结果与实际标签之间的误差,然后根据误差调整网络的权重和偏置。

使用卷积神经网络进行手写数字识别的技巧

使用卷积神经网络进行手写数字识别的技巧

使用卷积神经网络进行手写数字识别的技巧手写数字识别是计算机视觉领域的一个重要任务。

近年来,卷积神经网络(Convolutional Neural Network,CNN)在图像分类问题上取得了显著的突破,也成为手写数字识别的主要方法之一。

本文将介绍使用卷积神经网络进行手写数字识别的一些关键技巧。

首先,准备数据集是进行手写数字识别的基础。

MNIST(Modified National Institute of Standards and Technology)是一个常用的手写数字数据集,包含了大量的手写数字图像。

可以使用它作为训练和测试的数据集。

准备好数据集后,我们可以开始构建卷积神经网络模型。

其次,设计合适的卷积神经网络结构是关键。

在手写数字识别任务中,常用的卷积神经网络结构包括LeNet-5、AlexNet、VGGNet和ResNet等。

这些网络结构具有不同的层数和参数量,可以根据任务需求选择适合的网络结构。

一般情况下,浅层网络结构如LeNet-5适用于简单的手写数字识别,而深层网络如ResNet适用于更复杂的手写数字识别任务。

然后,正确设置卷积神经网络的超参数也是非常重要的。

超参数包括学习率、批量大小、卷积核大小、卷积核数量等。

学习率决定了模型在每次迭代中更新的程度,过大或过小都可能导致模型无法收敛或过拟合。

批量大小决定了模型每次训练时使用的样本数量,过大可能导致内存不足,过小可能导致梯度估计不准确。

卷积核大小和数量决定了模型对输入的特征提取能力,需要根据数据集的大小和复杂程度进行合理的选择。

可以通过尝试不同的超参数组合并评估模型性能来选择最优的超参数。

接下来,数据增强是提升模型性能的一个有效方法。

数据增强指的是通过对训练数据进行随机的图像变换来增加数据样本的数量和多样性。

常用的数据增强方法包括旋转、平移、缩放、翻转、加噪声等。

这样可以提高模型的泛化能力,并减轻过拟合的风险。

此外,正则化技术也可以帮助抑制模型的过拟合。

基于 lenet 手写数字体识别实验总结

基于 lenet 手写数字体识别实验总结

基于LeNet的手写数字识别实验是计算机视觉领域中一个经典的实例,通过对MNIST数据集进行处理和分析,使用LeNet-5神经网络模型实现对手写数字(0-9)的识别。

以下是对该实验的总结:1. 数据集介绍MNIST数据集是计算机视觉领域的经典入门数据集,包含了60,000个训练样本和10,000个测试样本。

这些数字已经过尺寸标准化并位于图像中心,图像是固定大小(28x28像素)。

数据集分为训练集、验证集和测试集,方便进行模型训练和性能评估。

2. LeNet-5模型LeNet-5是一种卷积神经网络模型,由Yann LeCun于1998年提出。

尽管其提出时间较早,但在手写数字识别任务上取得了显著的成功。

实验中,我们采用LeNet-5模型对MNIST数据集进行处理。

3. 模型结构LeNet-5模型包括两个卷积层和三个全连接层。

卷积层分别包含6个和16个卷积核,卷积核大小为5x5。

每个卷积层之后跟着一个最大池化层,池化核大小为2x2。

全连接层分别具有64、120和84个神经元。

最后,模型输出10个神经元,对应10个数字类别。

4. 实验流程实验中,首先对数据集进行预处理,将图像缩放到28x28像素。

然后,将数据集划分为训练集、验证集和测试集。

接着,构建LeNet-5模型并使用训练集进行训练。

在训练过程中,采用交叉熵损失函数和随机梯度下降(SGD)优化器。

最后,使用验证集评估模型性能,并选取最优模型在测试集上进行测试。

5. 实验结果经过训练,LeNet-5模型在MNIST数据集上取得了较好的识别效果。

在测试集上,模型对数字的识别准确率达到了98.89%。

实验结果表明,尽管LeNet-5模型相对简单,但在手写数字识别任务上具有较高的准确率。

6. 实验总结基于LeNet的手写数字识别实验展示了卷积神经网络在计算机视觉领域的应用。

通过搭建LeNet-5模型并对MNIST数据集进行处理,实验证明了卷积神经网络在识别手写数字方面的有效性。

手写数字识别的卷积神经网络实现方法

手写数字识别的卷积神经网络实现方法

手写数字识别的卷积神经网络实现方法手写数字识别一直是计算机视觉领域热门的问题之一,而卷积神经网络(CNN)是实现手写数字识别的有效算法。

在这篇文章中,我们将讨论如何使用卷积神经网络实现手写数字识别。

卷积神经网络是什么?卷积神经网络是一种深度学习模型,它专门用于处理图像。

它具有优秀的特征提取能力和分类能力。

在卷积神经网络中,包含多个卷积层、池化层和全连接层等结构。

其中,卷积层通常用来提取图像中的特征,池化层用于减小特征图的大小,而全连接层则用于分类。

卷积神经网络实现手写数字识别的步骤1. 数据集准备在实现手写数字识别之前,我们需要一个包含手写数字的数据集。

一个常用的数据集是MNIST数据集,它包含大量的手写数字样本。

这些样本被标记为0-9之间的数字,位置规范。

2. 数据集预处理在使用卷积神经网络之前,我们需要对数据集进行预处理。

这通常包括将数据集划分为训练集、验证集和测试集;将图像转化为正确的数据格式,例如灰度图像可以转化为$28*28$的矩阵,RGB彩色图像可以转化为$28*28*3$的张量等。

3. 搭建卷积神经网络模型在准备好数据集后,我们需要搭建卷积神经网络模型。

一个简单的实现可以包括以下层次结构:- 输入层 -接受图像数据;- 卷积层 - 通过逐层学习提取图像中更高级别的特征;- 池化层 - 减小卷积层输出的特征图大小,以减少模型的计算时间和消耗;- Dropout层 - 减少过拟合风险;- Flatten层 - 将卷积层的输出矩阵转换为一维向量;- 全连接层 - 用于分类输出。

4. 模型训练和优化搭建好卷积神经网络模型后,我们需要对它进行训练和优化。

我们可以使用反向传播算法来更新神经网络的权重和偏差。

在训练过程中,我们可以使用准确度、损失函数等指标来监控网络的性能。

为了进一步优化网络性能,我们可以使用正则化技术、批量归一化和调整学习率等方法。

5. 模型评估和预测最后,我们使用验证集来评估模型的性能,并使用测试集进行预测。

分类-MNIST(手写数字识别)

分类-MNIST(手写数字识别)

分类-MNIST(⼿写数字识别)这是学习《Hands-On Machine Learning with Scikit-Learn and TensorFlow》的笔记,如果此笔记对该书有侵权内容,请联系我,将其删除。

这⾥⾯的内容⽬前条理还不是特别清析,后⾯有时间会更新整理⼀下。

下⾯的代码运⾏环境为jupyter + python3.6获取数据# from sklearn.datasets import fetch_mldata# from sklearn import datasets# mnist = fetch_mldata('MNIST original')# mnist好像下载不到它的数据,直接从⽹上找到它的数据,放到当⾯⽬录下的\datasets\mldata⽬录下。

MNIST data的百度⽹盘链接: 提取码: 9dq2,如果链接失效,可在下⾯评论区告知我,或者⾃⼰去⽹上找⼀样的,相信各位⼩伙伴的能⼒呀。

输⼊如下代码:from sklearn.datasets import fetch_mldatafrom sklearn import datasetsimport numpy as npmnist = fetch_mldata('mnist-original', data_home = './datasets/')mnist上⾯的代码中的data_home表⽰你的数据集的⽂件路径,写的是⼀个相对路径,如果你没有将你的数据集放在你当前代码的⽬录下,你可能需要使⽤绝对路径。

输出:{'DESCR': ' dataset: mnist-original','COL_NAMES': ['label', 'data'],'target': array([0., 0., 0., ..., 9., 9., 9.]),'data': array([[0, 0, 0, ..., 0, 0, 0],[0, 0, 0, ..., 0, 0, 0],[0, 0, 0, ..., 0, 0, 0],...,[0, 0, 0, ..., 0, 0, 0],[0, 0, 0, ..., 0, 0, 0],[0, 0, 0, ..., 0, 0, 0]], dtype=uint8)}可以看出,我们成功读到了它的数据,⽹上有很多的说法是错误的,没有办法读成功,只有这个才是正解 。

实验一--手写数字识别

大数据应用实例1、下面我们来做个大数据人工智能实例:手写数字识别2、我们使用python语言进行代码的编写,使用pycharm开发工具对其实例进行编写,3、下面我们首先来看我们的样本数据:4、我们使用python语言来对其数据进行机器学习和识别代码如下:import numpy as npimport pandas as pdimport matplotlib.pyplot as pltfrom sklearn.neighbors import KNeighborsClassifier# 1、数据读取data = plt.imread('./data/0/0_1.bmp')# plt.imshow(data)# plt.show()# x_tain=[]# for i in range(1,501):# x_tain.append(plt.imread('./data/0/0_%d.bmp'%(i)))x_tain =[]x_test =[]y_tain=[]y_test=[]for i in range(0,10):for j in range(1,501):if j < 451: #将数据保存到训练数据中x_tain.append(plt.imread('./data/%d/%d_%d.bmp'%(i,i,j)).reshape(-1) ) #reshape 可以降维也就是矩阵变化y_tain.append(i) #append 是读进来的数据进行存储的意思else: #保存到预测数据中x_test.append(plt.imread('./data/%d/%d_%d.bmp'%(i,i,j)).reshape(-1)) y_test.append(i)# 2、数据转换成x_tain,y_tain= np.array(x_tain),np.array(y_tain)# print(x_tain.shape,len(y_tain),len(x_test))# 3、机器学习knn = KNeighborsClassifier() #构造分类器knn.fit(x_tain,y_tain)y_ = knn.predict(x_test) #进行预测的结果# print(len(y_[::10]),'\n',y_test[::10])gl=knn.score(x_test,y_test)print('准确率为:',gl)# 3、图片绘制plt.figure(figsize=(13,15))img = x_test[::10]img1 = y_test[::10]yimg = y_[::10]for i in range(50):plt.subplot(5,10,i+1)plt.imshow(img[i].reshape(28,28))plt.title('预测数据:%d'%(yimg[i])+'\n真实数据:%d'%(img1[i]))plt.rcParams['font.sans-serif'] = ['SimHei'] # 设置字体为SimHei显示中文plt.rcParams['axes.unicode_minus'] = False # 设置正常显示符号plt.show()'''import matplotlib.ticker as tickerfig=plt.figure()ax = fig.add_subplot(111)ax.yaxis.set_major_locator(ticker.NullLocator())'''5、结果展示。

手写数字识别


Better AI, Better Life!
KNN算法(K最近邻居算法)
1
1
11
?
1 1
1
3 3 3
2 2
3 3
2

2
2
2
3 3
2
• 假设已知一些确定了类别的数据,对于一个 未知数据,它的类别由K个最相似的邻居投 票决定。
• 如果K个最相似邻居的大多数属于一个类别, 那么这个未知数据就属于这个类别。
假设K=3: 识别为3 假设K=5: 识别为5
[0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 81 131 152 194 194 225 98 5 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 14 209 253 242 242 242 251 254 183 9 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 25 228 95 0 0 0 113 250 254 219 24 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 01141 0 0 0 0 0 80 210 254 16 7 22 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 075 254 254 113 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 28 223 54 161 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 124 254 223 20 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 082 254 254 30 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 082 254 254 30 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 82 254 254 30 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 115 254 229 22 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 023 220 254 161 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 12 82 82 82 82 82 93 254 254 122 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 054 173 254 254 254 254 254 234 254 224 24 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 010 192 254 192 76 37 151 254 254 254 199 60 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 135 254 22 111 023 137 254 254 254 254 243 136 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 149 254 105 0 73 139 254 254 134 24 128 148 18 122 0 0 0 0 0 0 0 0 0 0 0 0 0 0 149 254 164 113 210 254 254 135 4 0 0 0 7 2 0 0 0 0 0 0 0 0 0 0 0 0 0 0 92 254 254 254 254 22845 3 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 13 206 255 255 144 4 0 0 0 0 0 0 0 0 0 0 0 0 0 000000000000000000000000000000000000000000000000 000000000000000000000000000000000000000000000000 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0]

超详细PyTorch实现手写数字识别器的示例代码

超详细PyTorch实现⼿写数字识别器的⽰例代码前⾔深度学习中有很多玩具数据,mnist就是其中⼀个,⼀个⼈能否⼊门深度学习往往就是以能否玩转mnist数据来判断的,在前⾯很多基础介绍后我们就可以来实现⼀个简单的⼿写数字识别的⽹络了数据的处理我们使⽤pytorch⾃带的包进⾏数据的预处理import torchimport torchvisionimport torchvision.transforms as transformsimport numpy as npimport matplotlib.pyplot as plttransform = pose([transforms.ToTensor(),transforms.Normalize((0.5), (0.5))])trainset = torchvision.datasets.MNIST(root='./data', train=True, download=True, transform=transform)trainloader = torch.utils.data.DataLoader(trainset, batch_size=32, shuffle=True,num_workers=2)注释:transforms.Normalize⽤于数据的标准化,具体实现mean:均值总和后除个数std:⽅差每个元素减去均值再平⽅再除个数norm_data = (tensor - mean) / std这⾥就直接将图⽚标准化到了-1到1的范围,标准化的原因就是因为如果某个数在数据中很⼤很⼤,就导致其权重较⼤,从⽽影响到其他数据,⽽本⾝我们的数据都是平等的,所以标准化后将数据分布到-1到1的范围,使得所有数据都不会有太⼤的权重导致⽹络出现巨⼤的波动trainloader现在是⼀个可迭代的对象,那么我们可以使⽤for循环进⾏遍历了,由于是使⽤yield返回的数据,为了节约内存观察⼀下数据def imshow(img):img = img / 2 + 0.5 # unnormalizenpimg = img.numpy()plt.imshow(np.transpose(npimg, (1, 2, 0)))plt.show()# torchvision.utils.make_grid 将图⽚进⾏拼接imshow(torchvision.utils.make_grid(iter(trainloader).next()[0]))构建⽹络from torch import nnimport torch.nn.functional as Fclass Net(nn.Module):def __init__(self):super(Net, self).__init__()self.conv1 = nn.Conv2d(in_channels=1, out_channels=28, kernel_size=5) # 14self.pool = nn.MaxPool2d(kernel_size=2, stride=2) # ⽆参数学习因此⽆需设置两个self.conv2 = nn.Conv2d(in_channels=28, out_channels=28*2, kernel_size=5) # 7self.fc1 = nn.Linear(in_features=28*2*4*4, out_features=1024)self.fc2 = nn.Linear(in_features=1024, out_features=10)def forward(self, inputs):x = self.pool(F.relu(self.conv1(inputs)))x = self.pool(F.relu(self.conv2(x)))x = x.view(inputs.size()[0],-1)x = F.relu(self.fc1(x))return self.fc2(x)下⾯是卷积的动态演⽰in_channels:为输⼊通道数彩⾊图⽚有3个通道⿊⽩有1个通道out_channels:输出通道数kernel_size:卷积核的⼤⼩stride:卷积的步长padding:外边距⼤⼩输出的size计算公式h = (h - kernel_size + 2*padding)/stride + 1w = (w - kernel_size + 2*padding)/stride + 1 MaxPool2d:是没有参数进⾏运算的实例化⽹络优化器,并且使⽤GPU进⾏训练net = Net()opt = torch.optim.Adam(params=net.parameters(), lr=0.001)device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")net.to(device)Net((conv1): Conv2d(1, 28, kernel_size=(5, 5), stride=(1, 1))(pool): MaxPool2d(kernel_size=2, stride=2, padding=0, dilation=1, ceil_mode=False)(conv2): Conv2d(28, 56, kernel_size=(5, 5), stride=(1, 1))(fc1): Linear(in_features=896, out_features=1024, bias=True)(fc2): Linear(in_features=1024, out_features=10, bias=True))训练主要代码for epoch in range(50):for images, labels in trainloader:images = images.to(device)labels = labels.to(device)pre_label = net(images)loss = F.cross_entropy(input=pre_label, target=labels).mean()pre_label = torch.argmax(pre_label, dim=1)acc = (pre_label==labels).sum()/torch.tensor(labels.size()[0], dtype=torch.float32)net.zero_grad()loss.backward()opt.step()print(acc.detach().cpu().numpy(), loss.detach().cpu().numpy())F.cross_entropy交叉熵函数源码中已经帮助我们实现了softmax因此不需要⾃⼰进⾏softmax操作了torch.argmax计算最⼤数所在索引值acc = (pre_label==labels).sum()/torch.tensor(labels.size()[0], dtype=torch.float32)# pre_label==labels 相同维度进⾏⽐较相同返回True不同的返回False,True为1 False为0, 即可获取到相等的个数,再除总个数,就得到了Accuracy准确度了预测testset = torchvision.datasets.MNIST(root='./data', train=False, download=True, transform=transform)testloader = torch.utils.data.DataLoader(testset, batch_size=128, shuffle=True,num_workers=2)images, labels = iter(testloader).next()images = images.to(device)labels = labels.to(device)with torch.no_grad():pre_label = net(images)pre_label = torch.argmax(pre_label, dim=1)acc = (pre_label==labels).sum()/torch.tensor(labels.size()[0], dtype=torch.float32)print(acc)总结本节我们了解了标准化数据·、卷积的原理、简答的构建了⼀个⽹络,并让它去识别⼿写体,也是对前⾯章节的总汇了到此这篇关于超详细PyTorch实现⼿写数字识别器的⽰例代码的⽂章就介绍到这了,更多相关PyTorch ⼿写数字识别器内容请搜索以前的⽂章或继续浏览下⾯的相关⽂章希望⼤家以后多多⽀持!。

手写数字识别原理(一)

手写数字识别原理(一)手写数字识别原理解析1. 引言手写数字识别是一项经典的机器学习任务,其目标是通过计算机算法将手写的数字图像转换成对应的数字。

该技术在邮政编码识别、银行支票处理等领域有着广泛的应用。

本文将从浅入深,分析手写数字识别的相关原理。

2. 数据预处理在进行手写数字识别之前,我们首先需要对输入的图像进行预处理。

常见的预处理方法包括: - 图像灰度化:将彩色图像转化为灰度图像,减少处理的复杂性。

- 图像二值化:将灰度图像转化为黑白图像,便于提取特征。

- 图像平滑化:采用滤波器对图像进行平滑处理,去除噪声。

3. 特征提取特征提取是手写数字识别的关键步骤,通过提取有效的特征可以更好地描述图像。

常用的特征提取方法有: - 形状描述符:根据图像的形状进行特征提取,如轮廓面积、周长等。

- 纹理特征:通过分析图像的纹理信息来描述特征,如灰度共生矩阵、小波变换等。

- 直方图特征:将图像像素值的分布情况作为特征,如灰度直方图、颜色直方图等。

4. 分类模型为了将手写数字图像映射到对应的数字,我们需要训练一个分类模型。

常用的分类模型包括: - 支持向量机(SVM):通过构建超平面实现分类。

- 决策树:按照特征的不同取值划分样本,构建树形结构。

- 人工神经网络:通过多个神经元的连接实现分类。

5. 模型训练与评估模型训练是指通过已有的手写数字图像数据集对分类模型进行训练,使其能够泛化到未见过的图像。

模型评估是指使用独立于训练集的测试数据对训练好的模型进行性能评估。

常用的评估指标有: - 准确率:分类正确的样本数量占总样本数量的比例。

- 精确率:被分类器正确分类为正例的样本数量占被分类器分类为正例的样本总数的比例。

- 召回率:被分类器正确分类为正例的样本数量占真实正例的样本总数的比例。

6. 深度学习方法近年来,深度学习方法在手写数字识别领域取得了显著的成果。

深度学习模型如卷积神经网络(CNN)通过多层卷积与池化层对图像进行特征提取,再通过全连接层进行分类。

  1. 1、下载文档前请自行甄别文档内容的完整性,平台不提供额外的编辑、内容补充、找答案等附加服务。
  2. 2、"仅部分预览"的文档,不可在线预览部分如存在完整性等问题,可反馈申请退款(可完整预览的文档不适用该条件!)。
  3. 3、如文档侵犯您的权益,请联系客服反馈,我们会尽快为您处理(人工客服工作时间:9:00-18:30)。
相关文档
最新文档