尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

MATLAB手写数字识别:从CNN原理到训练部署全流程解析

MATLAB手写数字识别:从CNN原理到训练部署全流程解析 手写数字识别是深度学习入门阶段最经典的小任务之一输入一张 28x28 的灰度图片输出 0 到 9 的类别。很多人以为它简单真正用 MATLAB 独立做一遍后才会发现数据集格式、网络层参数、训练选项、模型保存和普通图片的预处理全都需要对齐差一步都会导致运行报错或识别结果离谱。本文围绕一个常见需求展开在 MATLAB 中基于卷积神经网络 CNN 搭建手写数字识别系统既可以用 MNIST 数据集快速训练也可以把普通数据集整理成规定格式后重新训练并用准确率、混淆矩阵和单张图片预测验证识别效果。1. 手写数字识别为什么选择 CNN而不是传统图像特征1.1 手写数字识别真正难在哪里手写数字看起来是一个非常简单的分类问题但不同人写出来的 0 到 9 在笔画粗细、倾斜角度、断笔连笔、位置偏移上差异很大。同一个数字 7有人会加一横有人会写得像 1同一个数字 9在笔画闭合程度不同时很容易被误判成 4 或 7。传统方法通常会先手工设计特征比如统计图像的边缘方向直方图、提取连通域、计算矩特征再送入 SVM 或决策树分类。这类方法在数据量小、场景固定时有效但特征设计依赖经验换一批笔迹、换一种背景识别率就可能明显下降。CNN 的优势在于它不要求人工设计特征。网络通过卷积核自动从图像中学习“短横线、圆弧、交叉点”这类局部结构再在高层组合成数字的整体语义。对于笔画变形较大的手写数字这种自动学习特征的方式比固定模板更鲁棒这也是为什么手写数字识别几乎成了 CNN 入门的第一课。1.2 CNN 的卷积、池化和全连接如何配合完成分类可以把 CNN 理解成一个分层的特征提取器。卷积层用一组小尺寸卷积核在图像上滑动每个卷积核负责检测一种局部模式例如横向笔画、纵向笔画、右上弧线。激活函数 ReLU 把负值变成 0为网络引入非线性否则多层线性变换叠加起来仍然等价于一层线性变换。池化层对局部区域做最大值或平均值采样降低特征图尺寸同时让特征对微小平移更不敏感。全连接层把卷积部分提取到的空间特征展平映射到 10 个类别的得分。Softmax 层把得分转换为概率分布分类层再根据真实标签计算交叉熵损失。在 MATLAB 的 Deep Learning Toolbox 中这些层都可以用一行代码声明。训练时只需要提供图像数据和标签工具箱会自动完成前向传播、反向传播和参数更新。1.3 MATLAB 训练 CNN 需要的前提条件本文示例基于 R2020a 及以上版本的 MATLAB 编写。运行训练代码前需要确认以下几点MATLAB 版本不要太旧Deep Learning Toolbox 的层 API 和训练选项在不同版本之间有差异。安装 Deep Learning Toolbox这是训练网络的核心工具箱。如果使用 GPU 训练还需要安装 Parallel Computing Toolbox并在 MATLAB 中执行gpuDevice确认 GPU 可用。没有独立显卡也不需要担心MNIST 这类 28x28 小图使用 CPU 训练几轮就能跑完只是时间会比 GPU 慢一些。2. 数据集准备MNIST 和普通图片都要统一成网络能吃的格式2.1 MATLAB 内置 MNIST 示例数据是最快的跑通方式Deep Learning Toolbox 内置了 MNIST 的两个数据集函数分别是digitTrain4DArrayData和digitTest4DArrayData。执行下面代码后不需要去网上下载任何文件就能直接得到训练和测试所需的图像与标签。[trainImages, trainLabels] digitTrain4DArrayData; [testImages, testLabels] digitTest4DArrayData; fprintf(trainImages size: %s\n, mat2str(size(trainImages))); fprintf(trainLabels size: %s\n, mat2str(size(trainLabels))); fprintf(testImages size: %s\n, mat2str(size(testImages)));正常情况下会看到类似下面的输出trainImages size: [28 28 1 50000] trainLabels size: [50000 1] testImages size: [28 28 1 10000]这里有一个容易忽视的细节图像数组的维度顺序是 HWC也就是 高度 x 宽度 x 通道数 x 样本数。trainImages每一张图片都是 28x28、单通道灰度图像素值已经做了归一化。这套数据格式非常适合直接作为trainNetwork的输入适合先把训练流程跑通再去处理自制数据。2.2 手动读取 MNIST 原始文件明白数据格式再继续如果拿到的是 MNIST 原始二进制文件例如train-images-idx3-ubyte和train-labels-idx1-ubyte用 MATLAB 读取时要注意它们是大端存储。下面的函数演示了如何读取 MNIST 图像和标签文件。function [images, labels] readMNIST(imageFile, labelFile) fid fopen(imageFile, rb, b); magic fread(fid, 1, uint32); numImages fread(fid, 1, uint32); numRows fread(fid, 1, uint32); numCols fread(fid, 1, uint32); images fread(fid, inf, uint8); images reshape(images, numCols, numRows, numImages); images permute(images, [2 1 3]); images single(images) / 255; fclose(fid); fid fopen(labelFile, rb, b); magic fread(fid, 1, uint32); numLabels fread(fid, 1, uint32); labels fread(fid, inf, uint8); fclose(fid); end调用方式[trainImages, trainLabels] readMNIST( ... train-images-idx3-ubyte, ... train-labels-idx1-ubyte);读取完成后把trainImages转成 4 维数组标签转成 categorical就可以作为训练数据。手动读取的意义在于理解数据格式因为换到普通数据集时同样要把图片统一成 4 维数组或 imageDatastore否则训练会报错。2.3 普通数据集用 imageDatastore 组织标签来自文件夹名实际项目里更多是拿到一批扫描图片或截图这时候推荐按类别建文件夹用imageDatastore一次性读取。目录结构可以像下面这样组织digit_dataset/ 0/ 0_001.png 0_002.png 1/ 1_001.png 2/ 2_001.png ... 9/ 9_001.png读取代码imds imageDatastore(digit_dataset, ... IncludeSubfolders, true, ... LabelSource, foldernames);接下来用countEachLabel查看每个类别的样本数这一步非常重要能提前发现某个类别样本过少的问题。tbl countEachLabel(imds); disp(tbl);这里要注意imageDatastore会把文件夹名当作标签因此文件夹名一定要是字符形式的0、1直到9。如果文件夹命名为zero、one虽然也能读但后面分析混淆矩阵时标签顺序会变得不直观。2.4 图像预处理灰度、缩放、归一化必须和训练一致普通图片可能来自不同来源尺寸、通道数、背景色都不一样。在送入网络之前需要统一做三步处理彩色图转灰度图。缩放到 28x28。像素值归一化到[0,1]区间。可以写一个预处理函数function img preprocessDigit(img) if size(img, 3) 3 img rgb2gray(img); end img imresize(img, [28 28]); img im2double(img); end如果数据集是 imageDatastore可以用transform对整个数据集执行预处理imdsProcessed transform(imds, preprocessDigit);也可以使用augmentedImageDatastore直接指定输出尺寸它会在读取图片时自动缩放augimds augmentedImageDatastore([28 28], imds);这里要特别提醒训练时怎么预处理预测时就必须怎么预处理。很多人训练阶段用im2double预测阶段却直接imread后把 uint8 图片传入网络结果必然出错或识别率极低。3. 构建 CNN 网络每一层都清楚它在做什么3.1 一个可跑通 28x28 灰度图的基础网络结构下面这段代码定义了一个适合 28x28 灰度图的小型 CNN它包含三组“卷积 批归一化 ReLU 池化”结构最后接全连接层和分类层。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, output) ];各层的数据尺寸变化如下表方便检查维度是否匹配。层作用输出尺寸imageInputLayer([28 28 1])接收 28x28 灰度图像28x28x1convolution2dLayer(3,8,Padding,same)8 个 3x3 卷积核提取局部特征28x28x8batchNormalizationLayer对每个通道做归一化稳定训练28x28x8reluLayer非线性激活28x28x8maxPooling2dLayer(2,Stride,2)2x2 最大值池化尺寸减半14x14x8convolution2dLayer(3,16,Padding,same)16 个卷积核提取更高层特征14x14x16maxPooling2dLayer(2,Stride,2)再次池化7x7x16convolution2dLayer(3,32,Padding,same)32 个卷积核继续抽象特征7x7x32fullyConnectedLayer(10)展平并映射到 10 个类别得分10softmaxLayer转成概率分布10classificationLayer计算分类损失-3.2 卷积核大小、步长和 Padding 怎么选卷积核大小 3x3 是当前最常用的小卷积核参数少、层数可以更深。5x5 卷积核感受野更大但参数更多对于 28x28 的小图反而容易使边缘信息快速丢失。Padding, same表示在图像周围补零保证卷积前后空间尺寸不变。如果使用Padding, valid特征图尺寸会不断缩小对后续池化层输出尺寸计算会更敏感基础项目里直接使用same更省心。池化层这里使用 2x2、步长 2 的最大值池化一半尺寸的信息被保留降低计算量并增强平移鲁棒性。如果让池化步长大于卷积步长输出维度也会变化调整网络时建议从简单结构入手先跑通再加深。3.3 池化、批归一化、ReLU 和 Softmax 的作用池化层不是必须每一层都加但加了之后特征图尺寸会变小网络参数总量下降训练速度更快。批归一化让每一层的输入分布更稳定使学习率可以适当调大训练过程不会因为中间层数值波动而崩溃。ReLU 解决 Sigmoid 在深层网络中的梯度消失问题计算简单正向传播和反向传播都快。Softmax 把全连接层输出的实数得分压缩成 0 到 1 之间的概率且所有类别概率之和为 1这样输出可以被解释成模型对每个数字的置信度。3.4 用 analyzeNetwork 检查网络定义完网络后在训练前先执行analyzeNetwork(layers)检查层连接、参数数量和数据流是否正确。这个命令会打开网络分析窗口显示每一层的输入输出尺寸和可学习参数数量。analyzeNetwork(layers);如果出现尺寸不匹配分析窗口会直接标红并给出错误原因比盲目运行训练命令更容易定位问题。4. 训练流程和“重新训练”的正确打开方式4.1 数据划分训练集、验证集和测试集不能混用训练网络之前要先把数据分成三部分训练集用于更新网络权重。验证集用于观察训练过程中的准确率和损失帮助判断是否过拟合。测试集训练结束后评估最终模型测试集不能参与训练和验证。对于内置 MNIST 数组数据可以这样划分rng(0); idx randperm(size(trainImages, 4)); numVal round(0.15 * numel(idx)); valIdx idx(1:numVal); trainIdx idx(numVal 1:end); trainImagesFinal trainImages(:, :, :, trainIdx); trainLabelsFinal trainLabels(trainIdx); valImages trainImages(:, :, :, valIdx); valLabels trainLabels(valIdx);对于 imageDatastore可以用splitEachLabel按比例划分[imdsTrain, imdsVal, imdsTest] splitEachLabel(imdsProcessed, 0.7, 0.15, randomized);4.2 trainingOptions 关键参数说明训练选项直接决定训练是否收敛、训练速度多快、模型最终效果如何。下面的表格列出了最常用的几个参数。参数含义常用值注意点InitialLearnRate初始学习率0.001 到 0.01过大 loss 发散过小收敛很慢MiniBatchSize每次迭代使用的样本数32 到 128太大占用内存高太小梯度不稳定MaxEpochs遍历训练集次数5 到 20MNIST 小网络通常 10 轮以内可收敛ValidationData验证集数据{valImages, valLabels}用于实时观察验证准确率ValidationFrequency验证频率每隔若干迭代验证一次太频繁训练变慢Plots训练进度图training-progress图形化观察 loss 和准确率ExecutionEnvironment训练环境auto没有 GPU 时会自动退回到 CPUCheckpointPath检查点保存目录./checkpoints训练中断后可以从检查点恢复一个完整的训练选项配置如下options trainingOptions(sgdm, ... InitialLearnRate, 0.01, ... MaxEpochs, 10, ... MiniBatchSize, 128, ... ValidationData, {valImages, valLabels}, ... ValidationFrequency, 30, ... Plots, training-progress, ... ExecutionEnvironment, auto, ... Verbose, true);4.3 从普通数据集重新训练先改分类层再看效果如果网络结构不变只是换一批训练数据直接把数据和网络层传入trainNetwork即可net trainNetwork(augimdsTrain, layers, options);如果换成普通数据集后类别数量不是 10或者图像尺寸不是 28x28必须同步修改网络修改imageInputLayer的输入尺寸例如[64 64 1]。修改fullyConnectedLayer的神经元数量为类别数。确认标签是 categorical 类型。可以用下面的方式查看并修改网络中的全连接层lgraph layerGraph(layers); newFC fullyConnectedLayer(numClasses, Name, fc); lgraph replaceLayer(lgraph, fc, newFC); net trainNetwork(augimdsTrain, lgraph, options);替换层时新层的 Name 必须和原层 Name 保持一致否则层连接关系会断裂训练前可以再执行一次analyzeNetwork(lgraph)确认。4.4 训练日志和 loss 曲线怎么看训练开始后MATLAB 会弹出训练进度窗口左侧显示迭代次数、耗时、损失值右侧显示训练准确率和验证准确率曲线。判断训练是否正常主要看两点训练 loss 是否整体下降。如果 loss 一直波动不下降优先降低学习率。验证准确率是否随着训练上升。如果训练准确率接近 100% 而验证准确率长期停滞说明模型过拟合需要增加数据量或加入数据增强。MNIST 这类数据量较大的任务一个小型 CNN 在 CPU 上训练 10 轮也能看到明显收敛。需要注意即使训练 loss 很低也不能说明模型质量好必须回到测试集上独立评估。4.5 模型保存、加载和恢复训练训练结束后保存模型save(digit_cnn.mat, net);以后直接加载预测loaded load(digit_cnn.mat); net loaded.net;训练中断是所有训练任务都可能遇到的问题。通过在trainingOptions中设置CheckpointPathMATLAB 会定期把当前网络保存为net_checkpoint__*.mat文件。中断后可以加载最近的检查点继续训练。mkdir(./checkpoints); options trainingOptions(sgdm, ... CheckpointPath, ./checkpoints, ... ...);恢复时加载检查点文件中的网络对象转成 layerGraph 后重新调用trainNetwork继续训练。生产环境中建议训练脚本加一个判断如果检查点文件存在就加载不存在才从零开始。5. 识别验证准确率、混淆矩阵和单张图片预测5.1 在测试集上计算整体准确率训练完成后最重要的一步是在测试集上评估模型泛化能力。MNIST 的测试集样本没有参与训练用它评估能看到模型对未见手写数字的真实效果。YPred classify(net, testImages, MiniBatchSize, 512); accuracy mean(YPred testLabels); fprintf(Test accuracy: %.2f%%\n, accuracy * 100);正常情况下这种规模的小型 CNN 在 MNIST 上能获得 98% 以上的准确率具体数值会受随机初始化影响。如果准确率明显偏低例如低于 95%需要回到数据处理和训练参数部分排查。5.2 混淆矩阵找出最容易混淆的数字准确率只能反映整体效果混淆矩阵能告诉我们哪些数字互相之间容易混淆。figure; confusionchart(testLabels, YPred);在生成的混淆矩阵中对角线越亮越好。实际训练 MNIST 模型时4 和 9、3 和 8、7 和 9 是常见的混淆对因为它们在笔画结构上有相似部分。如果某个数字的测试样本被大量误判例如 8 经常被识别成 3通常说明该数字的训练样本不足或者数据增强过于激进改变了笔画结构。5.3 单张手写图片的预测流程项目上线时最终要能输入一张普通图片输出预测结果。预测代码分为三步读图、预处理、预测。img imread(my_digit.png); img preprocessDigit(img); [YPred, scores] classify(net, img); [maxScore, maxIdx] max(scores); fprintf(预测类别: %s, 概率: %.4f\n, char(YPred), maxScore);这里再次强调preprocessDigit必须和训练时使用的预处理完全一致。如果训练数据是黑白颠倒预测时也需要做同样的反色处理。5.4 错误样本可视化只看准确率数字很难发现问题把预测错误的样本显示出来能直观判断错误是“看起来很像”还是“预处理有 bug”。wrongIdx find(YPred ~ testLabels); fprintf(错误样本数量: %d\n, numel(wrongIdx)); figure; for i 1:min(12, numel(wrongIdx)) subplot(3, 4, i); imshow(testImages(:, :, :, wrongIdx(i))); title(sprintf(True %s / Pred %s, ... char(testLabels(wrongIdx(i))), char(YPred(wrongIdx(i))))); end如果错误样本显示出来的图像都严重倾斜或模糊那么问题可能出在数据标注上
返回列表