ARTICLE DETAIL

资讯详情

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

Matlab实现CNN手写数字识别:从原理到实战的完整教程

Matlab实现CNN手写数字识别:从原理到实战的完整教程 简介深度学习作为人工智能的核心技术正在图像识别领域展现巨大价值。卷积神经网络CNN通过局部连接与权值共享机制能够自动提取图像特征彻底改变了传统手工设计特征的模式。在图像分类任务中CNN具备强大的特征学习能力与泛化性能广泛应用于手写识别、医学影像分析、工业缺陷检测等场景。手写数字识别是理解CNN原理的经典入门项目其数据规模适中、任务清晰非常适合工程实践。本文基于Matlab环境系统讲解了如何使用Deep Learning Toolbox搭建并训练卷积神经网络涵盖数据集预处理、网络结构设计、训练调参与模型评估全流程。通过实际代码演示展示了Matlab在快速原型验证与可视化调试方面的独特优势帮助读者高效掌握CNN的核心思想与工程实现方法。 手写数字识别是深度学习里公认的入门项目地位跟编程界的 Hello World 差不多。我做这个项目倒不是因为没别的可做而是课程设计恰好需要一个既能讲清楚原理、又能跑出漂亮结果的题目。最终我选了 Matlab 来做卷积神经网络CNN手写数字识别这个组合在工程效率上确实有它的独到之处。这篇文章把我从环境准备、数据处理、网络搭建到训练调参的完整过程记录下来代码都能直接跑适合正在做课设、毕设或者想快速上手深度学习的同学参考。先说结论Matlab 做这个任务代码量比 Python 的 PyTorch 少不少而且训练过程可视化做得非常直观调试起来很舒服。如果你手头已经有正版 Matlab或者学校提供了授权完全没有必要为了一个手写数字识别去装 Python 环境、配 CUDA、折腾一堆依赖。下面进入正题。1. 为什么用 Matlab 做 CNN 手写数字识别项目定位与整体思路1.1 手写数字识别到底在解决什么问题手写数字识别本质上是一个图像分类问题给出一张 28×28 像素的灰度图里面是一个人手写的 0 到 9 的数字模型需要判断这张图属于哪个类别。这个任务看起来简单但实际做起来并不 trivial——因为手写体的变形太丰富了倾斜、粗细不均、断笔、连笔、位置偏移同一个数字能写出无数种样子。这也是为什么经典的 MNIST 数据集能在深度学习历史上占据这么重要的位置。它虽然只有 10 个类别、图像尺寸也小但它完整地包含了一个图像分类任务的所有要素数据预处理、模型设计、训练调参、评估与错误分析。把 MNIST 上的流程跑通了迁移到猫狗分类、医学影像分类、工业缺陷检测这些更复杂的任务上思路是完全一致的。从应用场景来看手写数字识别也不是玩具项目。邮政编码自动分拣系统需要在信封上识别手写邮编银行需要自动读取支票上的手写金额和账号表单识别系统要把纸质调查表里的手写数字录入电脑。这些场景早期靠传统的机器视觉方法特征提取分类器实现后来基本都被 CNN 取代了。CNN 之所以能胜任是因为它不需要人工设计特征而是让网络自己从大量样本中学习出什么样的像素组合代表一个数字。1.2 为什么选择 Matlab 而不是 Python很多同学会下意识觉得深度学习就该用 Python这其实是被 PyTorch 和 TensorFlow 的生态带出来的惯性思维。我做这个项目选择 Matlab有几个非常实际的理由。第一Deep Learning Toolbox 已经封装得相当完善。在 Matlab 里定义 CNN你不需要自己写反向传播、不需要手动管理张量的 shape、不需要关心梯度是怎么回传的。trainNetwork这一个函数就把前向计算、损失计算、反向传播、参数更新全部处理完了。而在 PyTorch 里虽然也有高层 API但新手很容易在使用 DataLoader、transform、device 这些概念时卡住。第二Matlab 的可视化是天然优势。训练过程中损失曲线和准确率曲线实时绘制你一眼就能看出模型是在正常收敛还是已经在过拟合。网络结构可以用analyzeNetwork直接可视化每一层的输出尺寸、参数量一目了然。数据预处理阶段的图像查看、卷积核可视化、特征图可视化这些在 Matlab 里都只需要几行代码。做课程设计或毕设时这些图表直接就能放进论文里省去了用 Python 画图再导出的功夫。第三对于 MNIST 这个规模的任务Matlab 的运行效率完全够用。CPU 上训练一个 LeNet 级别的网络几分钟到十几分钟就能收敛到 98% 以上的准确率。不需要 GPU不需要 CUDA不需要配置深度学习环境。对很多人来说光是省掉装环境这一步就已经赢了一大半。当然如果你要做的是大规模图像分类、目标检测、Transformer 这类前沿方向那该用 PyTorch 还是得用 PyTorch。但在手写数字识别这个任务上Matlab 是完全够用且更省心的选择。1.3 整体技术路线与可行性评估我的项目技术路线可以概括为四步数据准备使用 MNIST 数据集或 Matlab 自带的 DigitDataset将数据划分为训练集和测试集并转为网络需要的格式。网络设计搭建一个类似 LeNet-5 的卷积神经网络包含卷积层、池化层、全连接层和分类输出层。训练配置设置优化器、学习率、批大小、训练轮数等超参数调用trainNetwork完成训练。评估分析在测试集上计算准确率绘制混淆矩阵分析错分样本观察哪些数字之间容易混淆。这个方案在技术上是完全成熟的。MNIST 分类准确率超过 99% 的模型有很多一个基础版 LeNet 就能达到 98% 到 99%足以满足课程设计和毕设的要求。难点不在于模型本身而在于你是否能理解每一层为什么要这么设计、每一个参数为什么取这个值。这篇文章后续章节会把这些为什么都讲清楚。2. 环境准备与数据集处理把原始数据变成网络能吃的格式2.1 环境配置与工具箱检查开始之前先确认 Matlab 环境。手写数字识别需要 Deep Learning Toolbox这个工具箱从 R2017b 之后基本就非常稳定了我这边的测试环境是 R2022b。如果你的版本比较老比如 R2016a 及更早可能没有trainNetwork这个函数建议直接换新版本。另外如果想要用 GPU 加速训练需要 Parallel Computing Toolbox并且电脑要有支持 CUDA 的 NVIDIA 显卡。不过就像前面说的MNIST 这种规模用 CPU 也完全能跑没有 GPU 也不用焦虑。打开 Matlab 后先检查工具箱是否可用% 检查 Deep Learning Toolbox 是否安装 ver(deep)如果能正常显示版本信息说明工具箱可用。如果提示未安装在 Matlab 的附加功能里搜索 Deep Learning Toolbox 安装即可。每次版本更新后Matlab 偶发 error 9 或者启动缓慢的问题多半是许可文件或路径缓存出了毛病重启软件后基本能自愈。2.2 数据集获取Matlab 自带 DigitDataset 与完整 MNIST 两种方案我建议的方案是先用 Matlab 自带的 DigitDataset 跑通整个流程再用完整 MNIST 验证模型的真实性能。两条路都走一遍你会对数据规模对模型的影响有更直观的感受。DigitDataset 是 Deep Learning Toolbox 自带的一个小型手写数字数据集路径通常在这里digitDatasetPath fullfile(matlabroot, toolbox, nnet, nndemos, nndatasets, DigitDataset);这个数据集的结构很友好根目录下有 0 到 9 共 10 个子文件夹每个文件夹里是该数字的图像总共约 10000 张已经是按类别分好文件夹的结构直接用imageDatastore就能加载不需要手写任何文件读取代码。完整 MNIST 则需要手动下载和读取。MNIST 官网提供的是 IDX 格式的二进制文件Matlab 不能直接识别需要自己写读取函数。完整 MNIST 包含 60000 张训练图和 10000 张测试图每张都是 28×28 灰度图。这里给出一个经过验证的读取函数function images loadMNISTImages(filename) % 读取 MNIST 图像文件idx3-ubyte 格式 fid fopen(filename, rb); % 跳过 magic number magic fread(fid, 1, uint32, 0, ieee-be); numImages fread(fid, 1, uint32, 0, ieee-be); numRows fread(fid, 1, uint32, 0, ieee-be); numCols fread(fid, 1, uint32, 0, ieee-be); % 读取全部像素并 reshape images fread(fid, inf, unsigned char); images reshape(images, numCols, numRows, numImages); images permute(images, [2 1 3]); fclose(fid); % 归一化到 [0,1] 区间 images images ./ 255; end标签文件idx1-ubyte 格式的读取函数更简单function labels loadMNISTLabels(filename) fid fopen(filename, rb); magic fread(fid, 1, uint32, 0, ieee-be); numLabels fread(fid, 1, uint32, 0, ieee-be); labels fread(fid, inf, unsigned char); fclose(fid); end这段代码里有两个细节值得注意。一是打开文件时用了ieee-be因为 IDX 格式是大端存储而 Matlab 在常见平台上默认按小端读取不指定的话读出来的数据全是错的。二是permute(images, [2 1 3])这个转置因为 MNIST 原始存储是行优先读出来后需要转置一下才是正常的图像方向。这两个坑我在第一次写的时候都踩过数据读出来全是横着的数字或者乱码排查了半天才发现是大端问题。2.3 数据预处理与训练集/测试集划分拿到图像数据之后还需要做两件事给数据打标签以及把数据划分成训练集和测试集。如果你用的是 DigitDataset用imageDatastore加载后标签是自动从文件夹名生成的然后通过splitEachLabel按比例划分即可imds imageDatastore(digitDatasetPath, IncludeSubfolders, true, LabelSource, foldernames); % 每个类别随机取 800 张作为训练集剩下作为测试集 [imdsTrain, imdsTest] splitEachLabel(imds, 800, randomized);如果你的数据是完整 MNIST 的数组格式则需要把数组转成arrayDatastore或augmentedImageDatastore的形式。这里我推荐直接构造augmentedImageDatastore因为它在训练时可以做实时数据增强XTrain loadMNISTImages(train-images.idx3-ubyte); YTrain categorical(loadMNISTLabels(train-labels.idx1-ubyte)); XTest loadMNISTImages(t10k-images.idx3-ubyte); YTest categorical(loadMNISTLabels(t10k-labels.idx1-ubyte)); % 将 28×28×60000 转为 28×28×1×60000 的四维数组 XTrain reshape(XTrain, 28, 28, 1, []); XTest reshape(XTest, 28, 28, 1, []); % 构造数据存储这里可以顺便做随机平移的数据增强 imageSize [28 28 1]; augTrain augmentedImageDatastore(imageSize, XTrain, YTrain, ... DataAugmentation, imageDataAugmenter( ... RandXTranslation, [-2 2], ... RandYTranslation, [-2 2]));这里有个概念值得讲清楚augmentedImageDatastore不是一次性把数据全部加载到内存再处理而是每次迭代时按需生成一个 batch 的数据。这样即使原始数据量大也不会撑爆内存。而且它允许在不改写原始数据的情况下做随机裁剪、随机平移、随机翻转等增强操作相当于免费扩大了训练集。手写数字识别里最常用的增强就是 ±2 像素的随机平移因为手写数字的位置偏移是最常见的变形方式。数据预处理的最后一步是确认输入数据范围。MNIST 图像像素值范围在 0 到 255如果不做归一化直接喂给网络数值过大容易导致梯度爆炸。我在读取函数里已经统一除以 255像素值落到了 [0,1] 区间。如果你自己构造数据务必记得做这一步。补充一点Matlab 的imageInputLayer默认自带一个 z-score 归一化操作对每个通道计算均值方差并标准化。如果你已经手动归一化过了可以在imageInputLayer中设置Normalization, none避免重复处理。但如果数据像素值只落在 [0,1] 区间而非标准正态分布保留默认的归一化也没问题具体可以对比测试。3. 卷积神经网络结构设计从原理到 Matlab 实现3.1 CNN 核心原理的直观理解在写代码之前先花点时间理解 CNN 到底在做什么。卷积神经网络跟传统的全连接网络最大的区别在于两点局部连接和权值共享。拿图像识别来说一个数字的7之所以是7并不需要看整张图片的所有像素点只要看到中间有一条斜线、顶部有一条横线基本就能判断了。CNN 的卷积层用一个小的卷积核比如 5×5在图像上滑动每次只看一个小区域提取局部特征。这就是局部连接——每个神经元只跟输入的一个小区域相连而不是跟整张图相连。权值共享则更进一步同一个卷积核在整个图像上滑动时参数是共享的。也就是说不管这个 5×5 的卷积核在图像的左上角还是右下角它用的都是同一组权重。这样的好处是大大减少了参数量也使得模型对特征的位置不那么敏感——同一个特征不管出现在图片的哪个位置都能被同一个卷积核捕捉到。池化层的作用是下采样最常用的是最大池化max pooling。它把一个小区域内的最大值取出来作为输出这样做的好处有两个一是缩小了特征图的尺寸减少了后续计算量二是带来了一定的平移不变性。比如数字稍微偏了几个像素经过池化后特征图的变化很小模型依然能正确识别。用生活化的类比来说卷积层就是在图像上放很多不同形状的滤镜每个滤镜负责检测一种模式横线、竖线、圆圈、弧线。浅层的卷积核检测的是边缘和纹理深层的卷积核会把浅层特征组合成更加抽象的结构比如两个圆圈叠加就是 8。这就是 CNN 能够端到端学习图像特征的原因不需要人工设计特征。3.2 网络结构搭建从 LeNet-5 到自己的网络我使用的网络结构参考了经典的 LeNet-5但做了轻微的现代化调整。LeNet-5 是 Yann LeCun 在 1998 年提出的专门用于手写数字识别可以说是 CNN 的开山之作。它的结构简洁而有效非常适合用来教学和做项目。完整网络定义如下layers [ imageInputLayer([28 28 1], Name, input) convolution2dLayer(5, 20, Padding, 0, Name, conv1) reluLayer(Name, relu1) maxPooling2dLayer(2, Stride, 2, Name, pool1) convolution2dLayer(5, 50, Padding, 0, Name, conv2) reluLayer(Name, relu2) maxPooling2dLayer(2, Stride, 2, Name, pool2) fullyConnectedLayer(500, Name, fc1) reluLayer(Name, relu3) fullyConnectedLayer(10, Name, fc2) softmaxLayer(Name, softmax) classificationLayer(Name, output) ];逐层解释一下设计意图第一层imageInputLayer([28 28 1])指定了输入图像的尺寸高 28、宽 28、通道数 1灰度图。如果输入是 RGB 彩色图通道数就是 3。第一层卷积convolution2dLayer(5, 20)表示使用 5×5 的卷积核输出 20 个特征图。5×5 的卷积核大小是 LeNet 系列的传统选择在 28×28 的小图像上感受野大小足够捕捉数字的笔画特征。输出 20 个特征图意味着学习 20 种不同的局部特征。跟着一个 ReLU 激活函数。ReLUmax(0, x)是目前 CNN 最常用的激活函数计算简单、能有效缓解梯度消失问题。早期 LeNet 用的是 sigmoid但 ReLU 在实践中的收敛速度和效果都更好。第一层池化maxPooling2dLayer(2, Stride, 2)使用 2×2 窗口、步长为 2 的最大池化。经过这层特征图尺寸从 24×24 减半到 12×12。第二层卷积输出 50 个特征图相当于在第一层特征的基础上组合出更复杂的模式。参数量比第一层大因为输入通道从 1 变成了 20这也是网络表示能力逐步增强的关键。最后是两个全连接层第一层 500 个神经元相当于把卷积层提取的所有特征展平后做一次非线性组合第二层 10 个神经元对应 0 到 9 十个数字类别。最后一层接softmaxLayer把 10 个输出值变成和为 1 的概率分布哪个类别的概率最大模型就预测哪个数字。这里有个细节值得注意convolution2dLayer的默认Padding是 0所以 28×28 的输入经过 5×5 卷积后输出是 24×2428 - 5 1 24。如果你的网络结构里有维度变化可以用analyzeNetwork(layers)查看每一层的输出尺寸避免在训练时才发现维度对不上。3.3 参数数量计算与内存估算理解网络的参数量可以帮助你判断模型是否过参数化以及训练时需要多少内存。以我上面定义的这个网络为例来算一笔账。第一层卷积输入 1 个通道输出 20 个通道卷积核 5×5。每个输出通道有一个偏置项所以参数量是 5×5×1×20 20 520。第二层卷积输入 20 个通道输出 50 个通道。参数量是 5×5×20×50 50 25050。第一个全连接层输入是第二层池化后的特征图展平尺寸为 4×4×50 800 个值输出 500 个神经元。参数量是 800×500 500 400500。第二个全连接层输入 500输出 10。参数量是 500×10 10 5010。总参数量520 25050 400500 5010 431080约 43 万个参数。以单精度浮点数存储每个参数占 4 字节模型本身约 1.7MB这个规模非常小完全不需要 GPUCPU 训练也毫无压力。对比一下现在稍微大一点的图像分类模型参数量都是千万级甚至亿级这个网络可以说是非常轻量了。如果训练时遇到内存不足的问题可以考虑减小MiniBatchSize这个参数直接影响训练时占用的内存。3.4 深度、宽度与过拟合的权衡很多同学在搭网络时容易陷入一个误区层数越多、通道数越大效果就一定越好。在 MNIST 这种小数据集上这个想法往往适得其反。MNIST 总共只有 60000 张训练图图像还很简单。如果网络过于复杂比如堆 10 层卷积每层 256 个通道模型容量远超数据量非常容易过拟合——训练集准确率冲到 99.9%测试集准确率反而只有 97%。过拟合的本质是模型把训练集里的噪声也当成特征记下来了遇到没见过的样本就露馅。所以选择网络结构的原则应该是在能完成任务的前提下模型越简单越好。LeNet-5 这个级别的网络对于 MNIST 来说是刚刚好的容量。如果你想要更高的准确率更值得做的不是一味加深网络而是做好数据增强、添加 Dropout 层、使用更好的优化器。Matlab 中添加 Dropout 很简单在fullyConnectedLayer之间插入dropoutLayer(0.5)即可。Dropout 在训练时随机让一半神经元失活迫使网络学到更加鲁棒的特征。我在实验中发现加了 Dropout 之后测试准确率通常能提升 0.2 到 0.5 个百分点而且训练曲线更平滑。4. 训练配置、模型训练与性能评估4.1 训练选项设置与优化器选择网络定义好之后接下来的重点就是训练配置。trainingOptions这个函数有非常多的可选参数但真正关键的只有几个。我使用的配置如下options trainingOptions(sgdm, ... MiniBatchSize, 128, ... MaxEpochs, 10, ... InitialLearnRate, 0.01, ... Shuffle, every-epoch, ... Verbose, true, ... Plots, training-progress);优化器选的是sgdm也就是带动量的随机梯度下降。在 MNIST 这种小规模问题上SGD 配合合适的动量已经能取得非常好的效果而且比 Adam 更容易调参。Adam 的优势在于自适应学习率对初始学习率不敏感但在这个任务上两者最终准确率差别很小。我个人习惯先试 SGD如果收敛太慢再换 Adam。MiniBatchSize设置为 128。批大小的选择有一个实际考量批太小比如 8 或 16梯度估计的噪声大训练不稳定批太大比如 512 或 1024单次迭代的计算时间变长而且需要更多内存。128 是一个在稳定性和计算效率之间均衡的选择。对于 DigitDataset 的 8000 张训练图128 的批大小意味着每轮迭代 62.5 次也就是最后一个 batch 只有 64 张。这没问题trainNetwork会自动处理不完整的批次。InitialLearnRate设为 0.01。学习率是 CNN 训练中最关键的超参数。0.01 配合 SGD 动量在本任务上是一个比较稳妥的起始值。如果训练曲线显示损失下降太慢可以适当调大到 0.05 或 0.1如果损失在训练初期就出现震荡说明学习率过大应调小到 0.001。判断学习率是否合适的快速方法是观察第一个 epoch 的损失如果第一个 epoch 后损失下降超过一半说明学习率偏大如果下降不明显说明学习率偏小。Shuffle设置为every-epoch意思是每轮训练前都打乱数据顺序。这能防止模型学习到数据排列中的顺序特征是提升泛化能力的简单有效手段。Plots设置为training-progress训练时会自动弹出进度窗口实时显示损失和准确率曲线。这是 Matlab 特别方便的功能训练过程中随时可以观察是否收敛、是否过拟合。4.2 训练过程监控与判断收敛训练代码只有一行net trainNetwork(imdsTrain, layers, options);如果数据集是augmentedImageDatastore写法完全一样net trainNetwork(augTrain, layers, options);训练过程中进度窗口会显示两个核心指标训练集准确率和训练集损失。随着迭代进行准确率应该逐步上升损失逐步下降。正常情况下10 个 epoch 后训练准确率能到 99% 以上损失降到 0.01 以下。看训练曲线时有一个重要的判断技巧如果训练集准确率一直在涨、但涨到 98% 左右就卡住不动了同时损失下降非常慢不一定是模型有问题很可能是学习率太小导致的收敛速度过慢。这时可以加大学习率或者换用 Adam 优化器。如果训练集准确率很快到了 100%但测试集准确率只有 96%这就是明显的过拟合信号应该加 Dropout、加数据增强、或者减小网络规模。如果用的是完整 MNIST训练完成后可以在测试集上做正式评估。但如果你跟我一样之前构建了augmentedImageDatastore作为训练数据测试集还是普通的数组形式直接调用classify和confusionchart即可不需要额外包装数据。4.3 测试集评估与混淆矩阵分析训练完成后在测试集上评估模型的真实表现YPred classify(net, XTest); YTest categorical(YTest); % 如果你之前用的是 categorical 数组 accuracy sum(YPred YTest) / numel(YTest); fprintf(测试集准确率: %.4f\n, accuracy);在我的实验环境中训练 10 个 epoch 后使用完整 MNIST测试集准确率稳定在 98.5% 到 99.0% 之间。如果使用数据增强并训练 15 到 20 个 epoch准确率可以超过 99%。单一准确率数字不能反映全部问题混淆矩阵才是分析模型错误模式的关键工具figure; confusionchart(YTest, YPred);运行后会生成一个 10×10 的热力图矩阵行代表真实标签列代表预测标签。对角线上的数字是正确分类的样本数非对角线上的数字是错分样本数。通过观察矩阵我发现最常被混淆的数字对是 4 和 9、7 和 9、3 和 8。这个现象背后的原因是手写体的笔画歧义手写 4 时如果不封口上半部分看着就像 9手写 7 时如果带横杠跟 9 也容易混淆。还有一个意想不到的发现4 和 9 的混淆是双向的但 3 和 8 的混淆几乎都是3 被识别成 8反向的情况很少。因为 3 下面如果写得太饱满确实像 8 的左边半个。如果你对某个误分类样本感兴趣可以用find找出索引然后显示原始图像idx find(YPred ~ YTest, 10); % 找出前 10 个预测错误的样本 for i 1:length(idx) subplot(2, 5, i); imshow(XTest(:, :, 1, idx(i))); title(sprintf(真值: %d, 预测: %d, double(YTest(idx(i))), double(YPred(idx(i))))); end看错误样本时大多数错分样本其实是人眼也无法轻易判断的模糊手写体。如果错误的样本明显是清晰的数字却被分错那就说明模型可能对某种特定风格的数据有偏差需要检查训练数据是否覆盖了这种风格。5. 常见问题与排查技巧实录5.1 经典报错与解决办法做这个项目过程中我整理了最常遇到的几个报错场景做成表格方便查阅报错信息可能原因解决办法未定义函数或变量 trainNetworkDeep Learning Toolbox 未安装检查工具箱版本确认已安装 Deep Learning Toolbox输入数据大小不一致数据集中存在尺寸不同的图像在imageInputLayer中设置统一尺寸或使用augmentedImageDatastore统一图像大小GPU 内存不足单卡显存不够减小MiniBatchSize或改用 CPU 训练索引超出数组边界MNIST 文件读取时字节序错误检查是否使用ieee-be读取 IDX 文件错误使用 trainNetwork无效的训练数据imds 或 datastore 格式不对确保使用imageDatastore或augmentedImageDatastore不能直接传普通数组数据集中没有有效的图像文件路径配置错误或文件夹结构不对检查IncludeSubfolders是否为 true文件夹名是否是标签这里面最值得专门说一句的是输入数据大小不一致这个报错。DigitDataset 里的图像虽然大部分是 28×28但偶尔有少数图像尺寸不同。解决方法是使用augmentedImageDatastore它会自动把不同尺寸的图像缩放到网络指定的输入尺寸。这也是我第一次做的时候没有注意到的坑直接用imageDatastore就报错包装一层就没事了。5.2 准确率上不去的调试思路如果你按照上面的代码跑准确率不到 95%不要急着改动网络结构先按以下顺序排查。第一步检查数据预处理。打开几张训练图像看看确认图像方向、像素范围都正常。如果图像是反的或者像素值超过 [0,1]模型是不可能收敛的。我遇到过图像方向反了导致训练损失一直不下降的情况那个问题排查了很久才意识到是转置惹的祸。第二步检查学习率。把训练曲线打开看损失曲线有没有下降趋势。如果损失几乎不变学习率可能太小如果损失上下剧烈震荡学习率可能太大。把学习率设为 0.001、0.01、0.1 各跑 3 个 epoch画在同一个坐标轴里对比这是最快确定合适学习率的方法。第三步检查数据规模。如果你只用了每个类别几十张图来训练准确率低是正常的。深度学习模型需要足够的数据才能学到泛化特征。试试把训练数据增加到每类至少几百张。第四步检查网络结构。确认卷积层输出尺寸没有在中间变成非常小的值。比如 28×28 的输入经过两次 2×2 池化后变成 7×7如果后面卷积核再减掉一些尺寸最后可能就只剩 1×1 了信息损失严重。用analyzeNetwork(layers)检查每一层的尺寸确保在进入全连接层之前特征图不是过于微小。5.3 用它做自己的手写数字图片测试模型训练好后还有个很酷的玩法自己手写一个数字在 Photoshop 或 Windows 画图里保存为 28×28 的灰度 PNG然后测试模型能不能识别。关键是把图片处理的格式跟训练数据保持一致。MNIST 的图像是黑色背景、白色笔迹的灰度图而平时我们用画图软件写的字是白色背景、黑色笔迹直接测试的话模型预测会非常混乱。需要先做反色处理。img imread(my_digit.png); % 如果是彩色图转成灰度 if size(img, 3) 3 img rgb2gray(img); end img imresize(img, [28 28]); % 反色MNIST 是黑底白字 img 255 - img; % 归一化并转为四维 img double(img) ./ 255; img reshape(img, 28, 28, 1, 1); % 预测 YPred classify(net, img); disp([识别结果: , char(YPred)]);注意imresize默认使用双三次插值缩放后数字的笔画粗细可能会改变。如果识别失败可以尝试不同的插值方法比如最近邻插值nearest或者调整阈值把灰度变成纯黑白。这个环节里预处理方式与训练数据的吻合程度比模型本身更影响结果。5.4 项目扩展方向如果这个项目做完了还想继续深入我列几个性价比很高的扩展方向。第一个方向是做数据增强的对比实验。分别用不增强、随机平移、随机旋转缩放三种方式训练模型在测试集上对比准确率。这个实验能直观地展示数据增强对模型泛化能力的影响写进论文里也是很有说服力的图表。第二个方向是使用预训练模型做迁移学习。Matlab 的 Deep Learning Toolbox 提供了 GoogLeNet、ResNet 等预训练模型可以直接用net googlenet加载然后替换最后的全连接层在 MNIST 上做 fine-tune。虽然对于 MNIST 来说这种大模型有点杀鸡用牛刀但对拓展视野有帮助。第三个方向是把模型部署成一个小应用。Matlab App Designer 可以做一个简单的 GUI画一个数字点击识别显示结果。这个做出来非常适合课程设计的实物演示环节比干巴巴的图表展示效果好得多。最后一个值得一提的方向是特征图可视化。提取训练好的网络第一层卷积核将其可视化出来你会发现 20 个卷积核分别学到了不同的边缘方向把某个测试样本经过第一层卷积后的特征图画出来也能直观看到模型看到了什么。这个可视化过程可以加深对 CNN 的理解。一点操作体会最后说一点个人的操作体会。我在跑这个项目的过程中最有价值的经验不是模型的准确率有多高而是养成了一个习惯每改一个参数就记录一次训练曲线和测试准确率。学习率、批大小、卷积核数量、是否加 Dropout、是否做数据增强这些变量单独看影响不大组合起来效果差异却很明显。建议你也建一个简单的表格做记录比凭感觉调参要靠谱得多。另外Matlab 做这个项目虽然省力但也要当心工具箱版本差异。不同版本的trainingOptions参数略有不同网上的代码不一定能直接跑报错了不要慌优先看 Matlab 自带的文档。祝各位都能顺利跑出自己的手写数字识别模型。本文还有配套的精品资源点击获取
返回列表