基于Matlab的CNN图像分类实战:从原理到代码解析
这篇笔记是系列第三篇前两篇我分别梳理了CNN的基本原理和网络结构怎么画、怎么理解到了这一篇我猜很多人跟我当初一样卡在了同一个地方博客和PPT看了不少卷积、池化、步长、填充这些概念都能背了但一打开Matlab面对trainNetwork、convolution2dLayer、trainingOptions这些函数还是不知道从哪里下手更别说把代码和脑子里那套“特征提取”的图对应起来。先说结论用Matlab学CNN尤其是想快速验证想法或者把深度学习用在自己课题里的人真没有网上说的那么“不专业”。Deep Learning Toolbox里的高层API封装得相当干净写起来比PythonKeras更像在写数学公式而且数据预览、中间层可视化、断点调试这些体验对刚入门的人来说实在太友好了。这篇笔记会以一段能直接运行的CNN训练代码作为主线逐段拆解每行代码背后的原理和调用逻辑再把我实际训练时踩过的几个坑完整复盘一遍包括损失变成NaN、验证集准确率突然归零这类问题。适合正在用Matlab做图像识别、信号分类或者毕业论文里需要跑深度学习的读者参考。1. 从“背函数”到“看代码”我为什么用Matlab学CNN1.1 我的学习路线和Matlab的位置我在前两篇笔记里走过一条弯路先花大量时间啃数学推导再想一把梭直接上Python和PyTorch结果卡在环境配置和数据预处理上一个周末什么都没跑出来。后来回到Matlab反而两天就把一个手写数字识别的CNN跑通了。原因很简单Matlab里图像读取、矩阵操作、标签管理、绘图这些基础能力是天然自带的不需要像Python那样拼一堆库深度学习工具箱又把这些能力直接串成了流水线。所以我现在的建议是如果你已经有Matlab基础不要有“深度学习必须用Python”的心结。Matlab的Deep Learning Toolbox完全可以用来入门CNN而且由于它把底层计算封装得比较彻底你反而能更早地把注意力集中在“网络结构怎么搭、参数为什么这么设”这些更关键的问题上。1.2 从“会用函数”到“看懂代码”的转变很多教程会直接给你一段trainNetwork的完整代码然后说“你运行一下看看效果”。这种教法的问题在于你跑通了也不知道发生了什么改一个参数也不会调换个任务直接懵。真正有效的学习方式是把这段代码当成一条链路来读数据怎么进去的、每一层做了什么事、训练循环在优化什么、最后输出的是什么。Matlab的CNN代码之所以适合做这种“逐行拆解”是因为它的API设计逻辑和CNN的物理结构几乎是一一对应的。imageInputLayer对应输入图像convolution2dLayer就是卷积操作maxPooling2dLayer就是下采样fullyConnectedLayer就是把特征图展平做分类你写的代码和网络结构图是能互相印证的这一点比看抽象的框架源码要直观得多。2. 一个能直接运行的CNN代码数字识别全流程2.1 数据准备imageDatastore到底帮你做了什么我用的数据集是Matlab自带的DigitDataset路径为matlabroot/toolbox/nnet/nndemos/nndatasets/DigitDataset里面是10000张28x28的灰度手写数字图0到9各1000张文件按类别分文件夹存放。加载代码就三行digitDatasetPath fullfile(matlabroot, toolbox, nnet, nndemos, ... nndatasets, DigitDataset); imds imageDatastore(digitDatasetPath, ... IncludeSubfolders, true, ... LabelSource, foldernames);这段代码的核心是imageDatastore它的作用有两个一是自动遍历子文件夹把每张图片的路径读进来二是按文件夹名字自动生成标签也就是foldernames这个参数的效果。换句话说只要你的数据按类别放在不同文件夹里Matlab会帮你把“图片”和“标签”的对应关系建立好这是后面训练的基础。紧接着要把数据分成训练集和验证集[imdsTrain, imdsValidation] splitEachLabel(imds, 0.8, randomized);splitEachLabel(imds, 0.8)表示每个类别的样本按80%和20%的比例拆分randomized表示先随机打乱再拆避免同一个类别的图片连在一起造成训练分布偏差。再往后是数据增强。这一步对数字识别来说很有必要因为手写数字的写法千变万化如果不做任何变换网络很容易过拟合。我用了imageDataAugmenterimageAugmenter imageDataAugmenter( ... RandRotation, [-15 15], ... RandXTranslation, [-3 3], ... RandYTranslation, [-3 3]); augimdsTrain augmentedImageDatastore([28 28], imdsTrain, ... DataAugmentation, imageAugmenter); augimdsValidation augmentedImageDatastore([28 28], imdsValidation);augmentedImageDatastore的作用是在训练过程中每次取一个mini-batch时对图片做随机旋转和平移相当于免费扩充了训练样本量。RandRotation的范围我后来从±180度改到了±15度原因后面会讲。注意验证集没有加数据增强因为验证集要尽量反映真实分布。2.2 搭建网络从输入层到分类层的每一层含义我用的网络结构不算深但对理解CNN的流程已经足够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) fullyConnectedLayer(10, Name, fc) softmaxLayer(Name, softmax) classificationLayer(Name, output) ];逐个说关键层imageInputLayer([28 28 1])输入层28 28是图像宽高1是通道数。灰度图是1通道RGB彩色图是3通道这个必须和你的数据一致。convolution2dLayer(3, 8, Padding, same)3是卷积核大小3x38是卷积核个数也就是输出特征图的深度。Padding, same表示保持输出尺寸和输入一致等会儿会算给你看。batchNormalizationLayer批归一化层作用是把每一批数据的分布拉回均值为0、方差为1附近让训练更稳定、收敛更快。不少新手会默认“CNN就是卷积池化全连接”把BN漏掉实测下来加不加BN收敛速度和最终准确率差别还是挺大的。reluLayer激活函数把负值置零引入非线性。没有它多层卷积叠加起来还是一个线性变换网络表达能力会大打折扣。maxPooling2dLayer(2, Stride, 2)最大池化核是2x2步长也是2。作用是把特征图尺寸缩小一半同时保留局部最明显的特征。池化没有需要学习的参数它的作用是减少计算量、扩大感受野。fullyConnectedLayer(10)全连接层10对应10个数字类别。这一层会把前面得到的特征图“展平”成一维向量然后做线性变换输出每个类别的得分。softmaxLayer和classificationLayersoftmax把得分转成概率分类层再根据概率输出最终类别标签。这两个是配套使用的不能只留一个。2.3 训练选项配置这些参数别只抄答案训练网络之前还要设置trainingOptions我常用的配置如下options trainingOptions(sgdm, ... InitialLearnRate, 0.01, ... MiniBatchSize, 128, ... MaxEpochs, 12, ... Shuffle, every-epoch, ... ValidationData, augimdsValidation, ... ValidationFrequency, 30, ... Verbose, true, ... Plots, training-progress);逐个解释为什么这么设sgdm带动量的随机梯度下降算法。动量可以理解为给参数更新加了一个“惯性”能在一定程度上抑制震荡让训练更平滑。对初学者这个就是最稳妥的默认选择。InitialLearnRate, 0.01初始学习率。学习率决定了每次参数更新的步长太大了会震荡甚至发散太小了收敛极慢。0.01对这个小网络是个合理的起点后文会展示学习率设成0.1时损失直接变NaN的翻车现场。MiniBatchSize, 128每批送入网络训练的样本数。这个值受显存或内存限制调小了训练更稳但更慢调大了对梯度估计更准但更占内存。MaxEpochs, 12整个训练集被完整遍历12遍。对数字识别这种简单任务12轮已经足够了继续增大收益很小还可能过拟合。Shuffle, every-epoch每一轮训练前都把样本重新打乱避免网络学到样本顺序带来的假规律。ValidationData, augimdsValidation指定验证集训练过程中会自动评估验证准确率并在图上显示方便你观察有没有过拟合。ValidationFrequency, 30每30次迭代验证一次。Plots, training-progress画出训练进程图这个功能对初学者太重要了能直接看到损失曲线和准确率曲线的变化。然后执行训练net trainNetwork(augimdsTrain, layers, options);2.4 训练后评估光看损失下降还不够训练完成后用验证集评估模型[YPred, scores] classify(net, augimdsValidation); YValidation imdsValidation.Labels; accuracy mean(YPred YValidation); fprintf(验证集准确率: %.2f%%\n, accuracy * 100);这里有个很容易搞错的地方YValidation不能从splitEachLabel之前的旧变量里拿而是要直接用imdsValidation.Labels。因为imdsValidation是拆分后的新数据存储它维护着自己的标签顺序如果你用旧的YTrain或者别的标签数组去对顺序对不上准确率会掉得莫名其妙。这个问题我后面还会详细复盘。如果想看得更细可以用confusionchart(YValidation, YPred)画混淆矩阵能直观地看出哪些类别之间容易混淆。比如7和1、9和4这类手写体在28x28的分辨率下确实容易搞混这时你会意识到数据增强、网络深度、训练参数都不是拍脑袋定的而是要根据这些实际反馈来调。3. 逐步拆解训练过程代码背后的计算逻辑3.1 一张图怎么穿过卷积层卷积核、步长和填充很多人对照代码看网络图最困惑的就是“卷积到底做了什么”。我用一个简单例子解释。假设输入是5x5的灰度图卷积核是3x3步长stride1padding0那么输出尺寸用公式算[ H_{out} \lfloor \frac{H_{in} 2P - K}{S} \rfloor 1 ]其中H_in是输入高度K是卷积核尺寸P是填充像素数S是步长。代进去就是[ H_{out} \lfloor \frac{5 0 - 3}{1} \rfloor 1 3 ]所以输出是3x3的特征图。这个3x3的每个位置都是原图对应3x3区域和卷积核逐元素相乘再求和得到的。比如原图某局部区域是[1 0 1; 0 1 0; 1 0 1]卷积核是[1 0 -1; 1 0 -1; 1 0 -1]两者点积就是11001*(-1)01100*(-1)11001*(-1)0。这个核是一个典型的垂直边缘检测核如果图像局部左右对称响应是0如果一边亮一边暗响应会很大。所以卷积层提取到的特征本质上就是“图像局部和这个核的相似程度”。那Padding, same是干嘛的如果不加padding卷积后尺寸会缩小多次卷积后图像就变得特别小边缘信息也会快速丢失。same的意思是给原图四周补上足够多的0让输出尺寸和输入一致这样网络层数就可以堆得更深。步长Stride控制的是卷积核每次滑动多远。步长1意味着逐像素滑动步长2相当于跳着看输出尺寸减半。步长越大特征图越小计算量也越小但可能会丢失细节所以实际使用中步长通常设在1或2。3.2 我这段代码里的特征图尺寸是怎么变化的把前面网络的尺寸计算全部画出来就能完整看到一张28x28的图像是怎么“流动”的层输入尺寸参数输出尺寸imageInputLayer--28x28x1conv128x28x13x3, 8个核, paddingsame, stride128x28x8pool128x28x82x2, stride214x14x8conv214x14x83x3, 16个核, paddingsame, stride114x14x16pool214x14x162x2, stride27x7x16展平7x7x16-1x784fc1x78410个神经元1x10注意最后一个池化输出的7x7x16展平后是7×7×16784个值这784个数就是网络提取到的高层特征全连接层把它们映射到10个类别的得分。这个尺寸链条非常重要尤其是当你自己搭网络时最后一层全连接神经元的输入维度和上一层展平后的维度必须对得上否则Matlab会直接报维度不匹配的错误。这也是为什么我建议你搭网络时先在草稿纸上把每一层的输出尺寸算出来再写代码。3.3 训练循环在做什么从损失值到反向传播trainNetwork封装了前向传播、损失计算、反向传播和参数更新。很多初学者以为options里设了sgdm就完事了其实你应该知道这几件事在每次迭代里是怎么发生的前向传播一个mini-batch的图像经过所有层得到预测概率。分类层用的是交叉熵损失简单理解就是“预测概率分布和真实标签分布有多不一致”损失越大说明错得越离谱。反向传播根据损失对每个参数的梯度从最后一层往第一层逐层回传。Matlab内部用的是自动微分你不需要手推梯度公式但你要知道梯度计算是逐层的靠近输出层的层梯度信号强、更新快靠近输入层的层梯度信号弱、更新慢。这也是为什么网络深了以后需要BN、残差连接这些手段来帮助梯度传播。参数更新sgdm在每一步会用当前梯度更新参数同时考虑上一次更新的方向公式可以近似理解为[ v_{t1} m v_t \text{lr} \cdot \nabla L ][ w_{t1} w_t - v_{t1} ]其中m是动量系数默认0.9lr是学习率。v累积了过去梯度的指数衰减平均相当于给更新方向加了“惯性”。还有一个关键认知BN层里的均值和方差是在训练过程里逐步估计出来的这决定了它在训练和推理验证两种模式下的行为略有不同。训练时用当前batch的统计量验证时用训练阶段累积的移动平均。所以把验证集也做增强或者让模型在验证集上“见过”数据都会污染验证效果。3.4 验证集为什么不能参与训练这句话听起来像废话但在实操中太容易踩雷了。我之前见过有人把整个imds既当训练集又当验证集传进ValidationData结果训练图上的验证准确率接近100%一换到真实新数据上就拉胯。这就是信息泄露验证集如果参与了训练哪怕只是被“观察”过你调参时就已经在隐式地拟合它了。正确的做法就是splitEachLabel拆开之后训练和验证严格分开验证集只在每个ValidationFrequency周期被评估一次它的作用只是给你一个“模型目前泛化得怎么样”的读数辅助你判断该停还是该调参数。4. 实测中踩过的坑损失曲线异常和内存不足4.1 学习率设成0.1损失直接NaN我在一次实验里图快把InitialLearnRate从0.01改成了0.1心想反正网络小步子大一点没关系。结果训练不到10次迭代损失值直接变成了NaN训练过程图上出现一条断崖式的曲线准确率也跟着崩了。排查思路是这样的先看是不是数据里有NaN。我检查了输入图像用ismissing查了标签数据本身没问题。然后把学习率降到0.01问题立刻消失。原因其实不复杂学习率过大时参数更新步长超出损失曲面允许的范围梯度爆炸到数值溢出NaN一旦出现后续所有参数更新都会受污染基本救不回来。这个坑给我的教训是改动超参数时最好一次只动一个而且改完后先跑两三个epoch看损失走向。如果你发现损失在前几个迭代就剧烈震荡或直接变NaN首选操作是把学习率往小调一个数量级比如从0.01改到0.001而不是怀疑网络结构写错了。4.2 miniBatchSize和内存不足8GB显存也扛不住做数字识别时把MiniBatchSize设成1024之后训练开始没多久就报Out of memory。这是因为每个mini-batch都需要把全部中间特征图保存在GPU显存里用于反向传播批越大特征图占用的显存就越多。解决办法有两个方向一是把MiniBatchSize降到128或64通常能立刻解决问题二是如果必须用大批次考虑改用分布式训练或者减小输入图像尺寸。另外提醒一句Matlab用GPU训练需要Parallel Computing Toolbox并且对显卡型号和驱动版本有要求在命令行输入gpuDevice可以查看GPU是否可用。如果显存还是不够可以考虑trainNetwork的另一个选项ExecutionEnvironment,cpu但CPU训练会慢很多只建议在调试小模型时用。4.3 验证集准确率突然归零一次完整的排查链路这个坑非常典型。有一次训练过程显示训练准确率稳定上升最后接近99%但验证集准确率却变成了10%左右大约是随机猜的概率。当时第一反应是网络过拟合但过拟合也不至于验证准确率掉到随机水平。我按下面的顺序一步步排查先检查YPred和YValidation的长度是否一致。用size一查发现两者样本数一样排除长度不匹配问题。画混淆矩阵发现所有验证样本都被预测成同一个数字。这说明网络并没有“部分学会”而是输出已经完全偏了。回头看训练曲线训练准确率明明很高验证集却全偏这不符合过拟合的表现更像是数据标签错位。检查imdsValidation.Labels和augimdsValidation的对应关系。augmentedImageDatastore内部维护了一套自己的样本队列它和imdsValidation.Labels的顺序在理论上是一致的但如果我在拆分后重新shuffle过验证集或者把两个不同数据存储混在一起就会导致标签顺序对不上。最后发现问题出在我一次清理代码时使用了imdsValidation imdsValidation.shuffle()而YValidation还是从imdsValidation.Labels里取的理论上这样没问题。真正的问题是我在一个子函数里把imdsValidation又当成了全局变量做了二次拆分导致标签和图片错位。这个案例给我的教训是Matlab的数据存储对象是带内部状态的你调了shuffle之后它内部的顺序会变但Labels属性的顺序只和你最后一次操作有关。训练代码里最好保持一个原则imdsValidation和它的Labels必须在同一个作用域内一起取出、一起使用不要跨函数传完再回头对标签。4.4 数据增强过度数字识别准确率不升反降数据增强的初衷是好的但如果幅度过大反而会让任务变得不真实。我一开始把RandRotation设成±180度想着数字旋转也能识别网络应该能学会。结果验证准确率比不增强还低了好几个百分点。原因是6倒过来会像97倒过来完全不像7数据增强把类别之间的边界搞模糊了网络被迫去学那些现实中不太可能出现的变形学偏了。后来把旋转范围缩到±15度保留小幅平移验证准确率才恢复并超过基线。这个小实验告诉我数据增强的幅度一定要结合任务本身来判断不能只追求“数据多”更要追求“数据像”。4.5 工具箱和版本导致的隐藏错误trainNetwork报错的时候很多情况不是代码逻辑问题而是工具箱缺失。Matlab的深度学习功能分布在Deep Learning Toolbox里GPU训练需要Parallel Computing Toolbox如果你用的是老版本可能还不支持某些层比如batchNormalizationLayer。建议先用ver命令查看已安装工具箱尤其是这几个Deep Learning Toolbox、Parallel Computing Toolbox、Computer Vision Toolbox做图像处理时会用到。另外提醒一句尽量用正版授权或学校提供的正版版本。我见过有人在奇怪版本上装各种“密钥补丁”训练到一半报出一堆看不懂的底层错误其实只是版本破解不完全。为了跑通代码花一整晚在这种事情上非常不值得。5. 把CNN概念和代码一块块对应起来5.1 特征图可视化用activations看网络学到了什么训练结束后最直观的验证方式是看看每一层到底提取了什么。Matlab用activations函数就可以做到testImg imread(fullfile(digitDatasetPath, 0, img_1.jpg)); testImg imresize(testImg, [28 28]); if size(testImg, 3) 3 testImg rgb2gray(testImg); end act1 activations(net, testImg, conv1);act1的形状是28x28x8对应conv1层输出的8个特征图。把8个特征图分别画出来你会看到有的特征图高亮区域集中在笔画边缘有的集中在角落这说明第一层卷积核学到的是不同朝向的边缘、拐角这类局部纹理。再把conv2层的特征图也画出来会看到响应越来越抽象开始倾向于组合低级特征比如某些小结构或笔画模式。用montage函数可以把所有特征图拼在一起显示figure; for i 1:size(act1, 3) subplot(2,4,i); imshow(act1(:,:,i), []); title(sprintf(conv1 特征图 %d, i)); end这种可视化对理解CNN特别有帮助它能让你直观看到“卷积核是在找什么”、“第几层开始变得抽象”。很多教程用抽象的示意图代码一跑可视化你才算真正建立感觉。5.2 感受野、展平和全连接为什么最后一层是784到10很多人不理解全连接层为什么能把二维特征图变成一维分类得分。其实fullyConnectedLayer在Matlab内部会自动把输入展平也就是把7x7x16的多维数组reshape成一维的784个值再做矩阵乘法。这个展平过程就是fc层做的第一件事。感受野这个概念也能在这个例子里体会第一层卷积核是3x3所以每个输出像素只能看到原图像的一个3x3局部。经过第一次池化后第二层卷积的3x3区域映射回原图大约是7x7的区域。层数越深单个输出单元能“看到”的原始图像范围越大这就是感受野的扩大。理解了这一点你才会明白为什么CNN不需要每个人都看整幅图而是从局部到全局逐层抽象。5.3 这套代码怎么迁移到自己的项目数字识别代码跑通之后大多数人要做的第一件事就是换成自己的数据集。你只需要注意几个改动点如果是彩色图像imageInputLayer的第三个维度要改成3比如[224 224 3]。同时augmentedImageDatastore的第一个参数也要改成对应尺寸比如[224 224 3]它会自动把不同分辨率的图像resize到统一尺寸。如果你的分类问题不是10类fullyConnectedLayer的神经元数量要改成你的类别数。比如二分类就改成2softmaxLayer不需要改。如果你的数据量特别大不要一次性把所有图片加载进内存请继续使用imageDatastore的方式它本身是“惰性加载”的只在训练取batch时读取图片内存友好得多。另外如果你的任务和通用图像分类比较接近比如区分猫狗、花卉、工业零件缺陷还有个更快的方案用预训练网络做迁移学习。Matlab内置了resnet18、googlenet等模型把最后几层替换成你自己的分类层然后只训练后面几层。这种做法在小数据集上效果很好训练时间也短。我写另一篇笔记时会专门展开迁移学习这篇先把从零搭建的理解打扎实。5.4 老生常谈但真的有用先跑通再调参如果让我给刚接触Matlab CNN的人一个最浓缩的建议我会说先把这份数字识别代码原样跑通再去动任何参数。跑通的意思是你能看到训练进度图、能得到一个80%以上的验证准确率并且能画出几张预测结果图。这个流程走完你对数据接口、网络定义、训练流程的基本盘就有数了之后调参、换数据集才不会像无头苍蝇。调参时也遵循“一次只改一个”的原则改完记录结果。我建议准备一个小本子专门记下每个实验的学习率、miniBatchSize、网络层数、数据增强参数和对应准确率。别看这个习惯简单它比任何调参技巧都重要因为深度学习实验属于典型的“没有记录就没有经验积累”的领域。我在实际使用中还有个体会Matlab的调试能力确实是学CNN的一大利器。你可以直接在trainNetwork之前加断点查看layers里每一层是否按顺序排列可以随时disp(size(XXX))查看数据尺寸甚至在自定义训练循环里逐行检查梯度值。我后来遇到过很多网络维度不匹配的报错几乎全是靠断点和size命令定位的。所以不要嫌麻烦该打断点就打该打印就打印这个习惯能帮你省下大量排错时间。最后再分享一个我自己一直在用的小技巧训练完不要只盯着准确率花两分钟把预测错的样本挑出来看。Matlab里可以用find(YPred ~ YValidation)找到预测错误的样本索引然后逐个显示原图、真实标签和预测标签。很多时候你会发现模型犯的错误往往也是人眼容易混淆的情况比如潦草的7和1、边缘残缺的4和9。这时候你就知道下一步该补什么数据、加强哪种数据增强而不是盲目地加深网络或者加大学习率。