ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

MATLAB实现卷积神经网络手写数字识别:例程拆解与调参实战

MATLAB实现卷积神经网络手写数字识别:例程拆解与调参实战 简介这是一份基于MATLAB的卷积神经网络手写数字识别例程完整演示了从图像数据读取、网络模型搭建、训练迭代到测试评估的整个流程尤其适合刚开始学习深度学习和图像分类的工程师与学生。压缩包内共19个文件整体体积仅20KB内容以m脚本为主包括网络结构初始化、前向传播、反向传播、梯度检查、样本扩充以及MNIST数据加载等功能模块同时附带MaxPooling的C源码与编译后的mexw64动态库用于实现池化加速另有一个说明文档辅助阅读。目前该资源已有173人学习使用是一种轻量但完整的参考资料。借助这组例程读者可以直接运行并观察手写数字的识别效果也可对照源码详细分析卷积层、池化层、全连接层、激活函数与权重更新等核心机制在此基础上调整网络结构或超参数还能将相同思路迁移至其他图像分类任务是理解并快速上手深度学习实验的实用入门材料。 上周帮一个刚入门深度学习的朋友调例程他翻遍了各种教程最后卡在了一个特别尴尬的地方别人给的代码都写着运行环境是PyTorch或者TensorFlow而他电脑里装的是MATLAB。他问我能不能用MATLAB跑卷积神经网络做手写数字识别我说当然能而且MATLAB在这件事上比你想的要省心得多——前提是你会看例程、会改例程而不是复制粘贴完就跑。这篇文章我不打算给你贴一段能直接抄的完整代码因为MATLAB官方工具箱里的trainNetwork全家桶已经把底层实现封装得足够好了。我更想做的是拆解一个可运行的CNN手写数字识别例程里每个模块到底在干什么、为什么这么设计、哪些参数值得你动、哪些地方是新手最容易踩的坑。1. 跑通例程前先搞清楚你手里有什么牌1.1 工具箱版本和数据集的准备手写数字识别是卷积神经网络的Hello World在MATLAB里跑这个任务最常见的路径有三条直接调官方Deep Learning Toolbox里的示例、自己写一个layers网络结构用trainNetwork训练、或者是老一点的版本里用nntool配合神经网络工具箱。我建议你直接走第二条因为trainNetwork这套接口从R2018b往后已经非常稳定社区里的例程绝大多数也是基于它写的。开始之前有两件事必须先确认。第一你的MATLAB版本是否包含Deep Learning Toolbox可以通过在命令行输入ver来查看如果列表里没有Deep Learning Toolbox后面所有代码都会在第一行就报错。第二数据集从哪来。MNIST数据集虽然经典但国内下载有时候不太顺畅替代方案是直接用MATLAB自带的digitDataset在命令行运行help digitDataset如果没报错说明本地已经有了一份现成的手写数字图片数据它是从英文手写字母库里裁剪出来的数字图像同样也是28x28灰度图足够用来做演示。实在不行还有一个兜底方案用imageDatastore指向你自己整理好的数字图片文件夹这个后面会说到。1.2 目录结构一个能跑的例程应该长什么样我拿我自己维护的一份标准例程目录给你做个参照mnist_cnn_demo/ ├── main_train.m % 主脚本加载数据、构建网络、训练、评估 ├── network_layers.m % 函数文件返回CNN网络结构 ├── options_setup.m % 函数文件设置训练选项 ├── digitDataset/ % 数据目录包含0~9共10个子文件夹 └── results/ ├── trained_net.mat % 保存训练好的网络 └── confusion_chart.png这个结构看着简单但有个很关键的设计逻辑把网络结构和训练参数独立成函数文件而不是全部堆在主脚本里。原因很简单你的数据集、迭代次数、学习率经常要来回调如果每次都要在主脚本里翻半天调试效率会很低。我见过太多人把所有东西写在一个文件里结果改一个MiniBatchSize都得滑动半天滚轮。分开写的好处是后面你要做调参实验或者换数据集只需要替换对应的函数文件即可。提示如果你是MATLAB新手建议刚开始不要碰实时脚本.mlx老老实实写.m脚本。实时脚本的调试输出格式虽然更像笔记本但它的单元格执行机制对很多新手反而容易造成混乱——变量在工作区里到底是哪个状态经常搞不清。2. 网络结构逐层拆解为什么是这几层顺序为什么不能乱2.1 一个基线结构的分层清单以一份标准的LeNet风格结构为例网络代码长这样layers [ imageInputLayer([28 28 1], Name, input) convolution2dLayer(5, 20, Padding, 0, Name, conv_1) batchNormalizationLayer(Name, bn_1) reluLayer(Name, relu_1) maxPooling2dLayer(2, Stride, 2, Name, pool_1) convolution2dLayer(5, 50, Padding, 0, Name, conv_2) batchNormalizationLayer(Name, bn_2) reluLayer(Name, relu_2) maxPooling2dLayer(2, Stride, 2, Name, pool_2) fullyConnectedLayer(500, Name, fc_1) reluLayer(Name, relu_3) fullyConnectedLayer(10, Name, fc_2) softmaxLayer(Name, softmax) classificationLayer(Name, output) ];这套结构你看着是不是特别眼熟它几乎是所有MATLAB CNN例程里最常见的样子。但眼熟归眼熟每一层为什么存在、参数为什么取5、20、50很多新手是说不清楚的。我逐个说一下。imageInputLayer([28 28 1])表示输入是28宽28高、单通道的灰度图。如果你用的是digitDataset它恰好就是28x28的灰度图所以这个输入尺寸直接匹配。这里有一个新手机很容易犯的错MNIST原始数据是28x28但有些第三方数据集的图片是32x32甚至更大如果直接喂给这个网络trainNetwork会在第一层就报维度不匹配错误报错信息往往很长但其实质就是输入尺寸问题。convolution2dLayer(5, 20)里的5指的是卷积核大小是5x520是输出通道数也可以理解为这一层提取了多少种特征图。为什么用5x5而不是3x3这是历史原因LeNet时代受限于计算能力5x5是验证过在数字识别上效果不错的尺寸。放到今天如果你想追求更高精度可以换成两个连续的3x3卷积层感受野等效于5x5但参数量更少、非线性更强。对新手来说我先劝你别瞎折腾结构5x5在数字识别这个任务上依然完全够用。batchNormalizationLayer在例程里很常见它的作用是让每一层的输入分布在训练过程中保持稳定。直觉上理解就是前一层输出经过激活函数后分布可能一直在飘加一个归一化能强制把它拉回均值为0、方差为1附近这样梯度下降走起来会平缓很多。加了BN层之后你可以把学习率设置得稍微大一点收敛速度会明显加快。maxPooling2dLayer(2, Stride, 2)下采样的逻辑就一句话对2x2的区域取最大值同时把特征图长宽各缩小一半。为什么要池化两个原因第一是降低计算量28x28经过一次池化变14x14再池化一次变7x7后面接全连接层时参数量会少很多第二是引入平移不变性数字稍微偏移一点池化后的特征仍然能保持较大的相似性。当然现代网络越来越不爱用池化了改用stride2的卷积来做下采样但对于这个小例程来说池化依然是省心可靠的选择。2.2 全连接层尺寸是怎么推算出来的很多新手在写fullyConnectedLayer时是懵的第一个全连接层的输入维度是多少答案是7x7x502450。你可以自己推一遍输入28x28经过一次5x5卷积无Padding后尺寸是24x2428-5124池化后变12x12再经过一次5x5卷积后是8x8池化后变4x4注意我这里用的是20和50个输出通道所以第二个池化层出来是4x4x50800不是2450。但在我上面给的代码里用的其实是Padding0输入28进过第一层卷积后是24x24池化成12x12第二个卷积后是8x8池化成4x4所以第一个全连接的输入维度是4x4x50800。这里不用你手动算trainNetwork在训练时会自动计算数据尺寸并完成适配但如果你自己写预测脚本或者要导出到其他平台这个数值就必须烂熟于心。我的建议是不管网络怎么改你都要习惯性地自己手推一遍特征图尺寸的变化这对排查报错极其有帮助。3. 训练选项里的门道准确率上不去先别急着怪网络3.1 trainingOptions的每个参数实际管什么事网络结构只是第一步真正决定模型能不能收敛的是trainingOptions里那串参数。看下面这段options trainingOptions(sgdm, ... InitialLearnRate, 0.01, ... MaxEpochs, 6, ... MiniBatchSize, 128, ... Shuffle, every-epoch, ... ValidationData, imdsValidation, ... ValidationFrequency, 30, ... Verbose, true, ... Plots, training-progress);逐行说。sgdm是带动量的随机梯度下降。动量这个概念你就想象成下山时手里推着一个铁球铁球一旦滚起来就有惯性即便当前这一小步方向有点偏整体趋势还是朝谷底走的。稳定性远好于纯SGD所以例程里默认用它。如果你电脑的内存比较富裕也可以试试adam在不少情况下收敛更快但最终精度未必比调好学习率的sgdm高。InitialLearnRate, 0.01这个值是新手最想调也最容易调崩的。0.01对这个小网络来说是一个比较中庸的起点。如果你把学习率调到0.1训练曲线大概率会剧烈震荡甚至直接发散如果调到0.001你会发现loss下降得像蜗牛爬6个epoch根本不够用。我的经验是当训练准确率一开始就在90%以上乱跳时优先怀疑学习率过大当loss曲线平缓得像一条直线时优先怀疑学习率过小。这个排查思路可以解决掉八成训练异常问题。MiniBatchSize, 128每次迭代用128张图算一次梯度。它的影响比较微妙批次越小梯度噪声越大但有时候这种噪声反而有助于跳出局部极小值批次越大梯度估计越准但内存占用也越大而且容易收敛到比较尖的极小值泛化能力可能略差。128是手写数字识别任务里一个非常均衡的选择。MaxEpochs, 6指的是把整个训练集完整过6遍。很多新手不理解epoch和iteration的关系这里一句话讲清楚一次epoch是看完全部训练数据一次iteration是看一个mini-batch的数据两者的关系是iteration数等于训练集样本数除以MiniBatchSize再乘以epoch数。6个epoch对MNIST级别的任务来说已经能将准确率推到99%附近再多会开始过拟合。3.2 训练过程中的实时监控别只盯着准确率那一栏用Plots, training-progress会在训练时弹出一个实时图窗上面同时画出损失和准确率的变化曲线。这里有个经验之谈刚开始训练时准确率曲线可能会先掉到很低再爬升甚至在前几百次迭代里只有10%左右的准确率这是正常的不要急着按CtrlC。因为网络权重是随机初始化的模型在一开始完全是在瞎猜。真正需要你警惕的是loss曲线在下降之后突然反弹并且一路走高不停那说明学习率太大了或者数据预处理有问题。还有一个非常实用的操作习惯不要等到训练完全结束才开始评估。我在自己跑实验时通常在训练跑到一半时就打开classify在验证集上测一把看看当前模型的错误样本长什么样。看错误样本比看准确率数字有用一百倍因为准确率只会告诉你好不好错误样本会告诉你为什么不好——是图片旋转太多是笔画太潦草还是某个数字类别和其他类别长得很像4. 从跑通到跑好数据增强、过拟合与常见报错排查4.1 为什么训练集准确率很高测试集却表现一般如果你照着前面的代码跑大概率训练集准确率能到99.5%以上但测试集或验证集只有98%左右。这不是bug这是过拟合。对于手写数字这种分辨率低、类别清晰的数据集98%和99.5%之间的差距其实已经算小的了如果你的数据更复杂这个差距会大得多。缓解过拟合的招数在MATLAB里最常用的是数据增强。用augmentedImageDatastore可以对训练图片做随机平移、旋转、缩放等变换相当于在有限的数据里变出更多样本。代码通常长这样imageAugmenter imageDataAugmenter( ... RandRotation, [-15 15], ... RandXTranslation, [-3 3], ... RandYTranslation, [-3 3]); augimds augmentedImageDatastore([28 28], imdsTrain, ... DataAugmentation, imageAugmenter, ... OutputSizeMode, centercrop);注意一点imageDataAugmenter只应该用在训练集上验证集和测试集做预处理时应当保持原样否则你评估出来的指标是不真实的——因为你在用和训练时一样的变形方式污染了测试样本。对数字识别这个具体任务来说我建议旋转范围不要超过15度平移不要超过3个像素。数字的语义方向性太强了一个9旋转180度就变成了6一个6旋转180度就变成了9旋转范围过大反而会误导模型。4.2 三个高频报错及排查链路这份例程虽然经典但新手跑的时候报错率一点都不低我总结三个最常见的你直接对照排查即可。报错一Undefined variable imds or function imds.这个报错几乎都是因为跳过了数据加载步骤直接复制了中间段代码。解决办法是确认你在运行主脚本之前已经把digitDataset赋值给imds并执行了splitEachLabel拆分。报错二Error using trainNetwork, Invalid training data. Predictors must be a cell array of image data or a single image datastore.这个报错说明训练输入不是一个ImageDatastore对象。最常见的场景是你手动加载了图片矩阵但矩阵维度是HxWxN而不是HxWxCxN四维输入要求最后一个是通道维。灰度图必须显式变成HxWx1xN才能作为四维数组输入。报错三Out of memory.训练深度学习模型时内存溢出十有八九是MiniBatchSize设得太大。先降成64甚至32试试如果还报错再检查你是不是同时开了一堆其他大变量在工作区。卡死前最好养成习惯训练前运行clear all; close all; clc;把工作区清干净。4.3 从MNIST到自己的数据集一个容易忽略的坑最后说一个我踩过的坑。MNIST的图片背景是黑色、数字是白色或者说近黑色背景、亮色笔迹。但很多人自己收集的数字图片往往是白底黑字也就是背景亮、笔画暗。如果你直接把这些图片喂给在网络里效果会断崖式下降。原因很简单卷积核提取的特征是基于梯度的黑白反演后同一位置的梯度方向完全反转模型学到的边缘特征就失效了。解决办法是在预处理阶段做一遍归一化或者反色。MATLAB里一句img 255 - img;就能把白底黑字变成黑底白字。我建议你在加载自定义数据时先随机挑几张图用montage或者imshow可视化确认一下方向和颜色是否正常再开始训练。这个习惯虽然不起眼但能帮你省掉大量无用功。提示如果你用的是自己组织的图片文件夹务必确认子文件夹名字是0、1、2这样因为imageDatastore会用文件夹名作为分类标签。一旦标签名和类别不对应训练准确率高得离谱也没有任何意义。5. 例程的进阶改造从演示到可用的四个方向基础例程跑通以后你会明显感觉到它能跑但不够用。这里给出四个我认为性价比最高的改造方向按难度从低到高排列。第一把训练好的模型导出成函数集成到GUI或者App Designer里。用genFunction(net, myModelFcn)可以直接生成一个独立的MATLAB函数文件输入图片返回预测类别和概率。这样别人用你的识别功能时不需要了解任何训练细节直接调用函数即可。第二引入交叉验证。默认的splitEachLabel按比例拆分一次训练集和验证集但单次拆分的运气成分很大。改成5折交叉验证后模型评估结果会更稳健。做法是用cvpartition生成交叉验证分组循环训练多次最后取平均准确率。第三加入对抗样本测试。手写数字识别模型虽然精度高但对微小的扰动非常敏感——比如给一个4的图片加上肉眼完全看不出来的噪声条纹模型可能就把它识别成9。用fastGradientMethod这类工具箱函数生成对抗样本测试一下模型的鲁棒性这是从会跑例程到理解模型短板的关键一步。第四导出到其它平台部署。如果你最终想把模型部署到嵌入式设备或者桌面应用上MATLAB提供了exportNetworkToTensorFlow或者coder工具链。这里有一个需要注意的地方MATLAB里的imageInputLayer默认做了归一化如果部署时外部环境的预处理方式和训练时不匹配精度会打折扣。所以预处理步骤一定要在训练前就封装成一个标准函数训练和部署共用同一套逻辑。这样后面移植到C或者Python环境时你能确保输入侧的变换方式完全一致。我本人更推荐先做第一和第三个方向。第一个能让你把模型真正用起来第三个则能帮你建立对深度学习本质的直觉——一个准确率99%的模型绝不等于一个可靠的模型。手写数字识别看似简单但把这里面的门道捋清楚了后面做图像分类、目标检测时你就能举一反三不会被例程牵着鼻子走。祝跑通顺利。本文还有配套的精品资源点击获取
返回列表