ARTICLE DETAIL

资讯详情

深耕编程入门与网站建设的一线实战洞察。

Matlab基于CNN的手写数字识别系统:从MNIST到自定义数据集

Matlab基于CNN的手写数字识别系统:从MNIST到自定义数据集 手写数字识别是卷积神经网络入门最常见的实验但网上能找到的完整工程基本都是 Python 版。这次我们看的是一个 Matlab 版源码项目Matlab 基于卷积神经网络 CNN 的手写数字识别系统。它不只有 MNIST 标准数据集的支持还保留了普通数据集的训练入口也就是说你可以换成自己的手写数字图片重新训练。项目编号是源码 33 期下面按“环境准备 - 数据集处理 - 网络构建 - 训练测试 - 模型保存 - 重新训练”这条线完整拆一遍。先说这个项目最值得关注的点首先是双数据集支持MNIST 开箱即用普通数据集可以按文件夹整理后直接参与训练其次是可以重新训练不只是一个调好参数的模型训练脚本、验证脚本和预测脚本都是分开的再次是 Matlab 生态Precision 数据可视化、训练过程曲线、混淆矩阵、识别结果展示这些在 Matlab 里都很方便不需要额外搭前端。如果你正好在写课程设计或者毕业设计这套源码的改造空间很大可以把识别界面、批量测试、准确率统计都扩展进去。跑这个项目的门槛不高。Matlab 是主环境需要 Deep Learning Toolbox深度学习工具箱建议顺带装上 Parallel Computing Toolbox这样在支持 CUDA 的 N 卡上可以走 GPU 训练。MNIST 图像是 28x28 的灰度图数据量不大常规 CNN 结构下 CPU 也能在几分钟内完成一轮训练所以并不是必须要有高端显卡。更稳妥的判断是先按 CPU 流程跑通再根据本机的 Matlab 环境决定是否切 GPU。下面进入正题。1. 核心能力速览能力项说明项目类型Matlab 卷积神经网络CNN手写数字识别系统数据集支持MNIST 标准数据集 普通手写数字图片数据集核心功能CNN 模型训练、模型测试、手写数字识别、重新训练主要文件训练脚本、测试脚本、预测脚本、模型保存文件运行环境Matlab需要 Deep Learning Toolbox硬件要求CPU 可跑有 NVIDIA GPU CUDA 时可加速输入规格28x28 灰度图像单通道是否支持 API项目本身未明确提供 HTTP API可自行封装为 Matlab 函数是否支持批量任务可批量预测测试集图片需自行写循环脚本适合场景CNN 入门实验、Matlab 课程设计、数字识别演示系统需要说明的是以上能力项是基于项目标题和常见 Matlab CNN 工程结构得出的判断。不同源码包内的具体文件名和参数会略有差异实际使用时以你拿到的代码为准。从标题看这个源码的核心亮点是“MNIST 数据集和普通数据集都有”。这意味着你不必局限于官方 MNIST 的 0 到 9 数字可以把用画图工具写出来的数字、扫描的手写数字、自己采集的样本都放进数据集里参与训练。对普通数据集的处理关键是图片尺寸和灰度格式要统一。2. 适用场景与使用边界2.1 这个项目适合谁第一类是刚接触 CNN 的 Matlab 用户。相比 Python 里的 PyTorch 或 TensorFlowMatlab 的 Deep Learning Toolbox 把网络层封装成了可视化模块训练过程有曲线面板代码结构也更接近“搭积木”的思路很适合用来理解卷积、池化、全连接这些概念。第二类是正在做课程设计或毕业设计的学生。手写数字识别系统是一个经典选题这个源码可以直接作为基础版本后续可以扩展图形用户界面、实时摄像头识别、批量测试报告导出等功能。第三类是想快速验证“自己的数据集”能不能被 CNN 识别的工程师。如果手里有一批手写数字或类似小尺寸灰度图像完全可以借用这个工程的训练流程把自定义数据集整理好直接跑重新训练。2.2 能解决什么问题一套完整的 CNN 识别流程包含数据读取、网络定义、参数设置、训练、验证、保存模型、加载模型、单张预测等步骤。很多刚入门的人卡在“数据怎么喂给模型”这一步而 Matlab 的imageDatastore和augmentedImageDatastore能直接把文件夹结构转成训练集代码量比 Python 少很多。这个源码的实用价值就是把这些流程串起来给你一个能用、能改、能出图的结果。2.3 不适合什么场景如果要做高精度工业级 OCR 系统这个项目不太合适。手写数字识别只是一个 10 分类任务输入是 28x28 灰度图对复杂背景、倾斜角度、连笔数字的鲁棒性有限。它更适合作为教学演示和基础实验而不是生产级识别服务。另外如果完全没有 Matlab 授权也不想装 Matlab那这个项目就不合适。虽然 Matlab 有试用期但长时间使用需要正版授权这一点要提前确认。2.4 使用边界与合规提醒使用 MNIST 数据集本身是常见的科研和教学行为但如果你把自定义数据集用于论文、商用或公开发布需要确认数据来源和授权。MNIST 是公开数据集很多人写实验报告时直接引用没有问题。如果是自己采集的手写样本涉及他人笔迹时要注意隐私和授权。训练出来的人脸、签名、笔迹等生物特征相关模型不能用于未经同意的身份识别或验证场景。3. Matlab 环境准备与项目文件结构3.1 环境检查清单在打开源码之前先用命令窗口检查环境是否满足要求% 查看 Matlab 版本 version % 检查深度学习工具箱是否安装 matlab.addons.installedAddons % 使用 ver 命令定位工具箱 ver(deep)如果ver(deep)报错说明 Deep Learning Toolbox 没有安装。在 Matlab 的“附加功能”里可以搜索并安装或者重新运行安装程序勾选对应工具箱。建议同时安装 Parallel Computing Toolbox这样可以直接用trainingOptions中的ExecutionEnvironment,gpu参数。3.2 GPU 支持情况如果你有 NVIDIA 显卡可以在 Matlab 里执行% 检查 GPU 是否可用 gpuDevice能正常输出显卡型号和显存信息就说明 GPU 环境可用了。注意 Matlab 的 GPU 支持要求你的显卡驱动、CUDA 版本和 Matlab 版本匹配具体以 Matlab 官方系统要求为准。如果gpuDevice报错直接使用 CPU 训练即可MNIST 数据量小CPU 训练完全能接受。3.3 项目文件结构规划一个典型的 Matlab CNN 手写数字识别项目文件结构大致是这样的。你拿到源码后可以先按下面的结构梳理一遍CNN_HandwrittenDigit/ ├── data/ % 数据集目录 │ ├── mnist/ % MNIST 相关数据或脚本 │ └── custom/ % 普通数据集 │ ├── train/ │ │ ├── 0/ │ │ ├── 1/ │ │ ├── 2/ │ │ └── ... │ └── test/ │ ├── 0/ │ ├── 1/ │ ├── 2/ │ └── ... ├── models/ % 保存训练好的模型 │ └── trainedModel.mat ├── scripts/ % 主脚本 │ ├── trainCNN.m │ ├── testCNN.m │ └── predictDigit.m └── README.mdMNIST 数据在 Matlab 中通常不直接以图片形式存放而是通过内置函数或下载脚本加载。普通数据集则需要按数字类别分成子文件夹文件夹名就是标签。这种结构可以直接配合 Matlab 的imageDatastore读取。4. 数据集准备MNIST 与普通数据集4.1 MNIST 数据集加载方式MNIST 数据集包含 60000 张训练图和 10000 张测试图每张是 28x28 灰度图数字范围 0 到 9。在 Matlab 里加载 MNIST 有几种常见方式。如果使用 Matlab Deep Learning Toolbox 自带的演示数据可以用下面两个函数% 加载 MNIST 格式的训练集和测试集 [trainImages, trainLabels] digitTrain4DArrayData; [testImages, testLabels] digitTest4DArrayData;digitTrain4DArrayData返回的trainImages是一个 28x28x1xN 的四维数组最后一位是样本数trainLabels是对应的分类标签。这套数据是 Matlab 官方工具箱内置的不需要额外下载非常适合先跑通流程。如果项目代码里使用了loadMNISTImages和loadMNISTLabels这类函数说明是从 MNIST 原始文件加载。这种方式的调用形式通常是% 读取 MNIST 原始文件需要修改为实际文件路径 trainImages loadMNISTImages(data/mnist/train-images.idx3-ubyte); trainLabels loadMNISTLabels(data/mnist/train-labels.idx1-ubyte);不同源码包的 MNIST 加载函数可能封装在utils目录下拿到代码后先确认函数名和路径。4.2 普通数据集准备方式所谓“普通数据集”就是你自己准备的手写数字图片。准备过程分三步。第一步把图片按类别放到不同文件夹。比如要训练 0 到 9 十个数字就建 10 个文件夹分别命名为 0 到 9图片放在对应文件夹里。第二步统一图片格式。CNN 输入要求是固定尺寸的单通道灰度图MNIST 是 28x28所以普通图片最好也统一缩放到 28x28。如果原图是 JPG、PNG 彩色图需要先转灰度再缩放。可以用下面这段代码做批量预处理% 批量预处理普通数据集图片 srcFolder data/raw; % 原始图片目录 dstFolder data/custom/train; % 处理后目录 imageSize [28 28]; if ~exist(dstFolder, dir) mkdir(dstFolder); end fileList dir(fullfile(srcFolder, *.png)); for i 1:length(fileList) img imread(fullfile(srcFolder, fileList(i).name)); if size(img, 3) 3 img rgb2gray(img); end img imresize(img, imageSize); imwrite(img, fullfile(dstFolder, fileList(i).name)); end第三步用imageDatastore读取整个目录。imageDatastore会自动把一级子文件夹名作为标签不需要你手动生成标签文件。注意标签文件夹名称必须是数字字符串比如0、1、2这样classify返回的标签才是可读的数字类别。4.3 数据集划分建议普通数据集建议按 8:2 或 7:3 划分训练集和测试集训练集和测试集分别放在train和test两个根目录下。训练时用imageDatastore读取train测试时用imageDatastore读取test。如果每类样本太少可以先用最简单的图片增强平移、旋转、缩放扩充数据量。5. CNN 模型构建与训练流程5.1 网络结构设计针对 28x28 灰度图一个经典的简单 CNN 结构是卷积层 - 批归一化 - ReLU - 池化 - 卷积层 - 批归一化 - ReLU - 池化 - 卷积层 - 全连接层 - Softmax。这个结构在 MNIST 上通常能取得不错的效果同时训练速度很快。按 Matlab 语法网络层定义大致如下layers [ imageInputLayer([28 28 1], Name, input) convolution2dLayer(3, 8, Padding, same, Name, conv1) batchNormalizationLayer(Name, bn1) reluLayer(Name, relu1) maxPooling2dLayer(2, Stride, 2, Name, pool1) convolution2dLayer(3, 16, Padding, same, Name, conv2) batchNormalizationLayer(Name, bn2) reluLayer(Name, relu2) maxPooling2dLayer(2, Stride, 2, Name, pool2) convolution2dLayer(3, 32, Padding, same, Name, conv3) batchNormalizationLayer(Name, bn3) reluLayer(Name, relu3) fullyConnectedLayer(10, Name, fc) softmaxLayer(Name, softmax) classificationLayer(Name, classoutput)];这个结构输入是 28x28x1输出是 10 类。实际项目中可能用不同的卷积核数量比如 16、32、64这不是关键关键是fullyConnectedLayer的输出节点数必须等于类别数。如果你要识别 0 到 9就是 10如果只识别一部分数字就改成对应的类别数量。5.2 训练选项设置Matlab 中trainingOptions控制学习率、迭代轮数、批次大小、验证集等参数。一个常见的设置如下options trainingOptions(adam, ... InitialLearnRate, 0.001, ... MaxEpochs, 10, ... MiniBatchSize, 128, ... ValidationData, valImds, ... ValidationFrequency, 30, ... Shuffle, every-epoch, ... Verbose, true, ... Plots, training-progress);如果是第一次跑建议先设MaxEpochs为 5看训练曲线是否收敛再逐步增加到 10 或 15。MiniBatchSize在 CPU 上可以设小一点比如 64避免内存占用过高GPU 上可以设置为 128 或 256。InitialLearnRate一般从 0.001 开始如果损失下降慢可以试试 0.0005 或 0.0001。5.3 训练入口脚本训练脚本的核心逻辑是“读取数据集 - 构建网络 - 设置参数 - 调用 trainNetwork”。下面是一个整合模板% trainCNN.m % 使用 imageDatastore 读取普通数据集 trainImds imageDatastore(data/custom/train, ... IncludeSubfolders, true, ... LabelSource, foldernames); valImds imageDatastore(data/custom/test, ... IncludeSubfolders, true, ... LabelSource, foldernames); % 检查标签 disp(unique(trainImds.Labels)); % 归一化数据范围到 [0,1]使用 augmentedImageDatastore 调整尺寸 augTrain augmentedImageDatastore([28 28 1], trainImds); augVal augmentedImageDatastore([28 28 1], valImds); % 定义网络结构 layers [ ... ]; % 上面定义的 layers % 训练选项 options trainingOptions(adam, ... InitialLearnRate, 0.001, ... MaxEpochs, 10, ... MiniBatchSize, 128, ... ValidationData, augVal, ... ValidationFrequency, 30, ... Plots, training-progress, ... Verbose, true); % 训练 net trainNetwork(augTrain, layers, options); % 保存模型 save(models/trainedModel.mat, net);运行训练脚本后Matlab 会弹出训练进度窗口显示准确率和损失曲线。当验证准确率稳定在 95% 以上时说明训练基本成功。不同源码包里的训练脚本名可能是train.m或main_train.m要先看清楚入口文件名。5.4 CPU 与 GPU 训练差异在trainingOptions中可以通过ExecutionEnvironment指定运行环境% 自动选择 GPU 或 CPU options trainingOptions(adam, ... ExecutionEnvironment, auto, ... MaxEpochs, 10); % 强制使用 CPU options trainingOptions(adam, ... ExecutionEnvironment, cpu, ... MaxEpochs, 10);MNIST 训练数据量不算大CPU 和 GPU 的差异主要体现在每个 epoch 的耗时上。用 CPU 时训练 10 个 epoch 可能需要几分钟到十几分钟具体时间取决于处理器性能和MiniBatchSizeGPU 则通常能缩短到一两分钟以内。实际耗时以本机训练曲线上的时间为准。6. 功能测试与效果验证6.1 测试集准确率验证训练完成后第一步是计算测试集整体准确率。这个数字能反映模型在未见过的数据上的泛化能力。测试脚本的核心逻辑如下% testCNN.m load(models/trainedModel.mat, net); testImds imageDatastore(data/custom/test, ... IncludeSubfolders, true, ... LabelSource, foldernames); augTest augmentedImageDatastore([28 28 1], testImds); % 对测试集进行预测 YPred classify(net, augTest); YTest testImds.Labels; % 计算准确率 accuracy mean(YPred YTest); fprintf(测试集准确率%.2f%%\n, accuracy * 100);准确率超过 95% 可以认为模型有效。如果使用 MNIST 官方测试集且网络结构合理准确率通常在 98% 以上。这里要注意classify接收的输入必须与训练时一致augmentedImageDatastore的尺寸设置要保持不变。6.2 单张手写数字图片预测单张图片预测是最直观的验证方式。可以用画图工具写一个数字保存为 PNG然后在 Matlab 里读取并预测% predictDigit.m load(models/trainedModel.mat, net); % 读取并预处理图片 img imread(test_my_digit.png); if size(img, 3) 3 img rgb2gray(img); end img imresize(img, [28 28]); img im2double(img); % 使用网络预测 label classify(net, img); fprintf(识别结果%s\n, char(label));classify对单张图会输出对应的分类标签。如果你的普通数据集图片白底黑字而 MNIST 是黑底白字需要根据实际情况决定是否取反否则识别效果会受很大影响。可以用下面这行代码做反色处理img 1 - img;6.3 混淆矩阵与错分样本分析除了准确率还需要看模型在哪些数字上容易混淆。Matlab 的 Deep Learning Toolbox 可以直接绘制混淆矩阵% 绘制混淆矩阵 figure; confusionchart(YTest, YPred);confusionchart会显示每个类别的预测情况对角线越亮说明分类越准确非对角线上的亮点就是容易被混淆的样本。通常手写数字识别在4和9、3和8、7和1之间容易出现混淆遇到这种情况可以增加对应类别的训练样本或者调整网络结构增加特征提取能力。6.4 训练进度曲线观察训练时弹出的进度窗口包含两条曲线一条是准确率另一条是损失。正常情况下训练准确率应该逐步上升训练损失逐步下降验证准确率和损失会有轻微波动但整体趋势应当与训练一致。如果验证损失不断上升而训练损失下降说明模型过拟合了此时可以增加训练数据、使用数据增强或增大MiniBatchSize。7. 模型保存、加载与重新训练7.1 模型保存训练结束后用save命令把网络结构连同训练好的权重保存下来save(models/trainedModel.mat, net);保存后models目录下会出现trainedModel.mat文件。这个文件就是你的模型文件后续测试或部署时直接加载即可不需要重新训练。7.2 模型加载下次打开 Matlab 时只需要加载这个文件load(models/trainedModel.mat, net); whos netwhos net可以查看变量类型。确认net的类型是SeriesNetwork或DAGNetwork后就可以直接调用classify做预测了。7.3 重新训练重新训练有两种常见场景。第一种是在已有模型基础上继续训练也就是迁移学习思路。把加载出来的net的层提取出来替换最后的全连接层和分类层然后用新数据继续训练% 加载旧模型 load(models/trainedModel.mat, net); % 获取网络层去掉最后三层替换为新分类层 lgraph layerGraph(net); newLayers [ fullyConnectedLayer(10, Name, fc_new) softmaxLayer(Name, softmax_new) classificationLayer(Name, classoutput_new)]; lgraph replaceLayer(lgraph, fc, newLayers(1)); lgraph replaceLayer(lgraph, softmax, newLayers(2)); lgraph replaceLayer(lgraph, classoutput, newLayers(3)); % 使用新数据继续训练 netNew trainNetwork(augTrain, lgraph, options);第二种是完全从头训练也就是把trainNetwork的输入换成新的数据集网络定义保持不变。标题里的“可以重新训练”指的就是这个能力你把data/custom下的数据换成自己的图片直接运行训练脚本即可。重新训练时要特别注意标签数量一致。如果原来的模型是 10 分类新任务只有 5 个数字那么fullyConnectedLayer的输出节点数必须改成 5否则训练会报错。8. 常见问题与排查方法| 问题现象 | 可能原因 | 排查方式 | 解决方案 | | --- | --- | --- | --- | | ver(deep) 报错 | Deep Learning Toolbox 未安装 | 执行 ver 查看已安装工具箱 | 在 Matlab 附加功能中安装深度学习工具箱 | | imageDatastore 读取不到图片 | 路径错误或文件夹结构不符合规范 | 检查目录是否存在、子文件夹是否按类命名 | 使用绝对路径将图片按类别放入子文件夹 | | 标签数量不匹配 | fullyConnectedLayer 输出节点数不等于类别数 | 查看 unique(trainImds.Labels) 输出 | 修改全连接层输出节点数为实际类别数 | | 输入图像尺寸错误 | 图片不是 28x28 或通道数不是 1 | 在预处理前显示 size(img) | 统一缩放为 28x28 并转灰度 | | 训练时内存不足 | MiniBatchSize 过大或数据集过大 | 查看内存占用 | 调小 MiniBatchSize关闭其他程序 | | GPU 训练报错 | CUDA 版本或驱动不匹配 | 运行 gpuDevice 查看报错信息 | 改为 CPU 训练或更新显卡驱动/CUDA | | 白底黑字图片识别不准 | 图片与训练集颜色极性不一致 | 可视化预处理后的 img | 对图片取反使背景为黑色、笔迹为白色 | | classify 输入格式报错 | 输入不是 numeric 张量 | 检查图像数据类型 | 使用 im2double 转换为 double 类型 | | 验证准确率低于预期 | 网络结构过浅或训练不充分 | 观察损失曲线 | 增加 epoch、增加卷积核数量、加入数据增强 | | 训练数据太少导致过拟合 | 每类样本不足 | 查看训练集各类别数量 | 使用图片增强或采集更多样本 | | 重新训练时无法替换层 | 层名不存在或网络类型不支持 | 使用 analyzeNetwork 查看层结构 | 先用 layerGraph(net) 转为图结构再替换 | | 项目脚本路径报错 | 当前工作目录不是项目根目录 | 运行 pwd 查看当前路径 | 使用 cd 切换到项目根目录或添加路径到 Matlab path |表格里这些问题是手写数字识别项目里最常见的。如果你在启动训练脚本时遇到Undefined function or variable这类错误优先检查脚本名、函数路径和当前工作目录。Matlab 对路径比较敏感cd到项目根目录再运行脚本是一个很实用的习惯。9. 最佳实践与使用建议9.1 第一次运行先小参数测试不要一上来就跑 30 个 epoch。第一次先把MaxEpochs设为 2MiniBatchSize设为 64确认整个流程能跑通再逐步增加训练轮数。这样可以更快定位是代码问题、数据集问题还是参数问题避免把时间浪费在等待训练上。9.2 保留一套最小可运行配置把训练脚本、测试脚本和预测脚本复制到一个单独的“最小运行”目录只保留 MNIST 数据和必要的依赖文件。以后调试时直接用这套最小配置不会因为误删数据集或模型文件导致项目崩溃。9.3 数据集、模型、输出分目录管理建议按以下结构管理文件data/ % 原始数据和预处理脚本 models/ % 训练好的模型文件 results/ % 测试结果、混淆矩阵图、预测结果图 scripts/ % 训练、测试、预测、预处理脚本模型文件建议在命名中带上版本信息和准确率例如trainedModel_acc98.mat方便后续对比不同实验。9.4 图像增强提升泛化能力如果普通数据集样本少可以用imageDataAugmenter做随机平移、旋转、缩放imageAugmenter imageDataAugmenter( ... RandRotation, [-10 10], ... RandXTranslation, [-2 2], ... RandYTranslation, [-2 2]); augTrain augmentedImageDatastore([28 28 1], trainImds, ... DataAugmentation, imageAugmenter);增强后的数据能让模型对手写变形更鲁棒但要避免旋转角度过大导致数字语义改变。比如数字 6 旋转 180 度可能变成 9这类增强反而会引入噪声。9.5 批量预测与扩展方向如果需要对大量测试图片批量识别可以写一个循环读取目录文件并调用classify的脚本。把所有预测结果保存到表格里% 批量预测示例 fileList dir(fullfile(data/to_predict, *.png)); results table(); for i 1:length(fileList) img imread(fullfile(fileList(i).folder, fileList(i).name)); if size(img, 3) 3 img rgb2gray(img); end img imresize(img, [28 28]); img im2double(img); label classify(net, img); results [results; {fileList(i).name, char(label)}]; end writetable(results, results/prediction_results.xlsx);这就是一个简单的批量识别任务。如果需要对外提供接口可以考虑把预测函数封装成一个.m函数文件再通过 Matlab Web App Server 或编译成独立程序发布。10. 总结与下一步这个 Matlab CNN 手写数字识别项目最值得尝试的点是把“数据准备 - 网络搭建 - 训练 - 测试 - 保存模型 - 重新训练”这条链路完整跑通而且同时覆盖 MNIST 和普通数据集。对刚接触深度学习的人来说先用这个项目理解 CNN 怎么处理图像分类比直接啃复杂框架更容易上手。最先应该验证的是 MNIST 训练流程。把训练脚本跑一遍观察训练进度曲线和测试集准确率。如果准确率能达到 95% 以上说明环境、代码和数据都正常。接下来再尝试普通数据集准备一批自己的手写数字图片按文件夹分类跑重新训练重点看数据预处理是否到位尤其是尺寸、灰度、颜色极性这三个细节。最容易踩的坑有三个一是imageDatastore的目录结构不对导致读不到图片二是全连接层输出节点数与类别数不一致导致训练报错三是普通图片没有统一缩放到 28x28导致classify输入尺寸不匹配。先把这三个问题规避掉整个流程会顺畅很多。后续可以扩展的方向很多把网络结构换成 LeNet-5 或 ResNet对比不同结构在 MNIST 上的准确率加入图形用户界面做一个手写板工具用鼠标写数字实时识别把模型导出为 ONNX 格式接入其他推理框架把批量预测脚本扩展成目录级批量处理工具输出 Excel 报表。这套源码做底子往上加东西比从零开始写要省很多时间。建议收藏备用等真正跑起来的时候再对照这篇文章逐步排查。
返回列表