ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

Matlab实战CNN代码理解:从CIFAR-10到LeNet-5全流程解析

2026/9/16 6:23:53 拓冰建站 浏览量
Matlab实战CNN代码理解:从CIFAR-10到LeNet-5全流程解析 这篇笔记拖了两周才动笔因为“Matlab代码理解”这件事远比我想象中麻烦。前两篇笔记已经把卷积神经网络CNN的基本概念、卷积层池化层的原理梳理了一遍当时觉得自己懂了可真把Matlab跑起来才发现“看得懂公式”和“写得通代码”之间还隔着一条挺宽的沟。这篇就记录我啃Matlab实现CNN的完整过程重点落在代码理解上包括数据怎么进网络、每一层在代码里长什么样、训练参数怎么配、以及那些最容易让人卡住的报错是怎么排查的。适合已经对卷积、池化有概念正准备拿Matlab练手或者一直在网上找demo却跑不通的同学。1. 思路先行在敲代码以前先把网络跑通这件事想明白1.1 一段CNN代码真正在做什么很多人第一次看CNN的Matlab代码会觉得莫名其妙因为demo代码动辄上百行函数一个套一个光数据预处理就绕得人头晕。我自己的经验是别急着逐行读先搞清楚一段完整的CNN训练代码在整体上只做了四件事——准备数据、搭网络、调训练参数、看结果。对应到Matlab里就是四段代码理解起来非常清晰。准备数据是把图片和标签转成工具箱能吃的格式这是最琐碎但也最关键的一步一半的报错都出在这里。搭网络是用一长串layer对象把卷积、池化、全连接拼起来在Matlab里就是写一个layers数组像搭积木一样。调训练参数是设置trainingOptions里的学习率、批次大小、迭代轮数这部分决定了模型能不能收敛、收敛多快。看结果是用classify做预测、算准确率、画混淆矩阵验证网络到底学到了什么。为什么选Matlab而不是直接上Python PyTorch不是说哪个更好而是Matlab在矩阵调试上确实省心。训练前你可以直接在命令行里disp(size(X))看每一维的长度中间层特征也可以随时取出来可视化对于“我到底想理解CNN在算什么”这个目标来说这种即时反馈比看Python的tensor维度要直观得多。1.2 为什么选LeNet-5结构做入门我这次用来练手的是LeNet-5的变体输入是CIFAR-10数据集里的32x32彩色图片网络结构是“卷积-激活-池化”重复两次再接两个全连接层。选择这个结构的原因很简单它是1998年提出的经典CNN不大不小刚好用来理解卷积神经网络的核心机制。网络如果太深比如直接上ResNet那种几十层的结构新手会被参数量和各种残差连接淹没根本分不清是哪一层出了问题。网络如果太浅比如只有一个卷积层分类精度会很差跑完看着百分之三四十的准确率一点正反馈都没有很容易劝退。LeNet-5这套结构的另一个好处是它和教材里的结构图几乎一一对应。你打开任何一篇讲CNN原理的文章画的无非就是“卷积-池化-卷积-池化-全连接”LeNet-5就是这个模板的祖师爷。把这段代码吃透以后再去看AlexNet、VGG那些更复杂的网络会发现它们本质上是把“卷积池化”这个基本单元重复更多次而已。2. 数据才是第一个门槛CIFAR-10读取与维度问题2.1 CIFAR-10数据集在Matlab里长什么样很多Matlab官方文档的demo会直接用helperCIFAR10Data.load(cifar10)一行代码就把训练数据和标签拿出来了。这个helper函数确实方便但我强烈建议初学者不要直接用。原因有两个第一它封装得太好你看不到数据原本的存储结构下次换成自己的数据集照样一头雾水第二你以后做实际项目时不会有人给你写好helper函数数据整理早晚要自己面对。我这次是手动从CIFAR官网下载matlab版本的数据包。解压后你会看到5个data_batch_1.mat到data_batch_5.mat外加一个test_batch.mat。每个batch文件里有两个变量images和labels。images的尺寸是3072x10000的uint8矩阵每一列是一张图10000列就是10000张图labels是这10000张图对应的类别编号范围是0到9。这里有个非常容易让人懵的点为什么是3072而不是32x32因为CIFAR-10的每张图是32x32像素的RGB彩色图32乘以32等于1024三个通道就是3072。数据排列方式是先存红色通道的1024个像素再存绿色通道的1024个最后存蓝色通道的1024个每个通道内部按行优先排列。也就是说列向量里的前1024个数是左上角到右下角的红色亮度值而不是我们习惯的“一个像素的RGB三个数挨在一起”。如果不搞清楚这个存储方式后面的reshape步骤一定会出错。2.2 把原始数据变成网络能吃的格式Matlab的Deep Learning Toolbox要求输入数据是[H W C N]的四维数组H是高度W是宽度C是通道数N是样本数。对应CIFAR-10就是[32 32 3 10000]。直接从3072x10000变成[32 32 3 10000]需要一个reshape加一个permute代码如下dataDir fullfile(tempdir, cifar10); % 假设你已经把解压后的文件放到了 dataDir % 读取一个batch看看结构 d load(fullfile(dataDir, data_batch_1.mat)); disp(size(d.images)); % 3072 x 10000 disp(d.labels(1:5)); % 标签是数字 % 手动把全部训练batch合并 XTrain []; YTrain []; for b 1:5 d load(fullfile(dataDir, sprintf(data_batch_%d.mat, b))); XTrain [XTrain, d.images]; YTrain [YTrain; d.labels]; end % 换成网络需要的排列顺序 [高 宽 通道 样本] XTrain reshape(XTrain, [], 32, 32, 3); XTrain permute(XTrain, [2, 3, 4, 1]); % 转成单精度并归一化到0~1 XTrain single(XTrain) / 255; % 标签转成categorical YTrain categorical(YTrain); % 测试集同理 d load(fullfile(dataDir, test_batch.mat)); XTest reshape(d.images, [], 32, 32, 3); XTest permute(XTest, [2, 3, 4, 1]); XTest single(XTest) / 255; YTest categorical(d.labels);这段代码有三个地方值得停下来仔细理解。第一个是reshape和permute的组合。CIFAR-10的每一列是按“行优先”顺序存储的图像数据而Matlab的reshape是按“列优先”填充的这就导致单纯reshape出来的图像是“躺着的”宽和高反了需要permute把前两个维度交换回来。这个坑非常隐蔽我第一次跑的时候就栽在这里后面会专门讲。第二个是除以255的归一化。原始图像像素值是0到255的整数网络训练时希望输入数值范围在0到1附近这样梯度更新更稳定收敛也更快。除以255之后数据变成单精度浮点数省的是一半内存训练速度也有提升。第三个是categorical(YTrain)。网络训练时的标签必须是categorical类型不能直接用double数组。如果直接拿数字标签喂进去trainNetwork会直接报错告诉你标签类型不对。2.3 训练集和验证集为什么要分开数据准备好之后还需要从训练集里切出一部分做验证集。这段代码逻辑很简单但背后的道理值得说清楚rng(0); % 固定随机种子保证每次跑结果可复现 idx randperm(size(XTrain, 4)); XTrain XTrain(:, :, :, idx); YTrain YTrain(idx); % 前45000训练后5000验证 XValidation XTrain(:, :, :, 45001:end); YValidation YTrain(45001:end); XTrain XTrain(:, :, :, 1:45000); YTrain YTrain(1:45000);训练集是模型用来学习参数的验证集不参与梯度更新只用来观察模型在没见过的数据上的表现。如果只盯着训练集准确率很容易出现过拟合——模型把训练数据背下来了换一批新图片就原形毕露。验证集就是一个“模拟考场”让你提前知道模型泛化能力大概是什么水平。这里的45000和5000不是拍脑袋定的。CIFAR-10训练集一共有50000张图我用了90%做训练、10%做验证这是比较常见的划分比例。数据量大的时候可以留更少比例做验证但5000张图对这个任务来说足够稳定地估计验证准确率了。3. 网络结构代码逐块拆解卷积、池化、全连接到底在写什么3.1 layers数组一行一层的“积木式”写法Matlab里搭建CNN网络不需要像TensorFlow那样写一堆类和回调只需要按顺序写一个layers数组就行了。每一层是一个layer对象从上到下依次排列数据从第一层流到最后一层。代码如下layers [ imageInputLayer([32 32 3], Name, input) convolution2dLayer(5, 20, Padding, 2, Name, conv_1) reluLayer(Name, relu1) maxPooling2dLayer(2, Stride, 2, Name, pool1) convolution2dLayer(5, 50, Padding, 2, Name, conv_2) reluLayer(Name, relu2) maxPooling2dLayer(2, Stride, 2, Name, pool2) fullyConnectedLayer(100, Name, fc1) reluLayer(Name, relu3) fullyConnectedLayer(10, Name, fc2) softmaxLayer(Name, softmax) classificationLayer(Name, output) ];这个写法看起来平淡无奇但里面的每个参数都值得琢磨。imageInputLayer([32 32 3])规定了输入图片的宽、高、通道数这个必须和前面数据准备的尺寸完全一致少一个维都会报错。convolution2dLayer(5, 20, Padding, 2)是二维卷积层第一个参数5是卷积核的尺寸这里是5x5第二个参数20是卷积核的个数也就是输出的特征图数量Padding为2表示在图像四周各补一圈0保持特征图尺寸不缩小。reluLayer是激活函数层对应公式max(0,x)作用是给网络引入非线性。如果没有激活函数多层卷积叠加起来还是一个线性变换网络再深也白搭。maxPooling2dLayer(2, Stride, 2)是最大池化层把2x2区域里的最大值挑出来步长为2输出尺寸直接减半。3.2 每一层的参数为什么这么配我第一次看到这些数字时最大的疑惑是为什么卷积核是5不是3为什么第一个卷积层用20个滤波器而不是100个这些参数有没有一个标准答案老实说卷积核大小、滤波器数量这些超参数没有唯一正确答案但有一堆经验法则可以参考。5x5卷积核是LeNet时代的经典选择能覆盖更大的感受野3x3卷积核在VGG之后成为主流因为多层3x3堆叠可以取得类似大卷积核的效果但参数更少。在CIFAR-10这样的小图片上5x5完全够用而且特征图尺寸变化符合直觉方便初学者计算。滤波器数量的设计遵循一个常见套路逐层递增。第一个卷积层用20个第二个用50个。理由是浅层网络负责提取边缘、颜色块这类低级特征数量不需要太多深层网络负责组合出更复杂的语义特征需要更多的特征图来容纳不同模式。空间尺寸逐层减半32-16-8通道数逐层增加20-50这是CNN设计里最常见的“压缩-扩展”结构。全连接层用100个神经元接10个输出节点10正好是CIFAR-10的类别数。你可能会问100怎么来的其实100也是一个拍脑袋的数字换成64或者128都可以只是全连接层的参数量很大神经元太多容易过拟合100在效果和复杂度之间比较平衡。3.3 特征图尺寸变化的计算Padding存在的意义为什么要给卷积层加Padding2这得从卷积输出尺寸的计算公式说起。输出尺寸等于(W - F 2P) / S 1其中W是输入尺寸F是卷积核大小P是Padding数量S是步长。CIFAR-10图片是32x32卷积核是5步长默认1不加Padding的话输出是(32 - 5 0)/1 1 28图片缩小了4个像素。第一层缩小4个像素没什么但连续几层叠加下去特征图尺寸会迅速缩水到最后可能连1x1都不剩网络根本没法工作。加了Padding2之后输出尺寸变成(32 - 5 4)/1 1 32特征图尺寸和输入保持一致。卷积层不改变空间尺寸池化层负责缩小尺寸这种“卷积提特征、池化降采样”的分工更容易理解和调试。池化层的输出尺寸计算更简单输入尺寸除以步长就行。第一层池化把32x32变成16x16第二层池化把16x16变成8x8。到第二个池化层结束时特征图是8x8x50展平后是3200个数值。这个3200就是第一个全连接层的输入特征数全连接层做的就是把这3200个数值通过矩阵乘法组合成100个更高级的特征。3.4 为什么softmaxLayer和classificationLayer缺一不可很多初学者看到全连接层后面还要接softmaxLayer和classificationLayer觉得多余直接把这两层删了。这个想法很危险。softmaxLayer把全连接层的10个输出值转换成10个概率所有概率加起来等于1。如果没有softmax层网络的输出只是一个10维向量数值可能是正也可能是负很难解释成“属于某个类别的可能性”。classificationLayer则是训练时计算损失函数的地方Matlab要求分类网络的最后一层必须是classificationLayer否则trainNetwork会直接报错。这两层的分工可以这样理解全连接层输出的是“证据”softmax把证据变成概率classificationLayer把概率和真实标签对比算出损失。训练过程就是不断调整网络参数让这个损失变小。少任何一层整个训练流程都跑不通。4. 训练代码与参数选择trainingOptions背后的门道4.1 trainingOptions参数逐项解析网络搭好后接下来是设置训练选项。这段代码是我认为整个CNN代码中最“反直觉”的部分因为参数不多但每个参数背后都牵扯到训练动力学值得一个一个说明白。options trainingOptions(sgdm, ... InitialLearnRate, 0.01, ... MaxEpochs, 20, ... MiniBatchSize, 128, ... Shuffle, every-epoch, ... ValidationData, {XValidation, YValidation}, ... ValidationFrequency, 10, ... Verbose, true, ... Plots, training-progress);第一个参数sgdm是优化器全称是带动量的随机梯度下降。它和普通SGD的区别在于更新参数时考虑了历史梯度的方向像一个小球在损失曲面上滚动不容易停在局部最优点。Matlab还支持adam自适应学习率优化器实际使用中adam通常收敛更快但sgdm在中小型CNN上表现也很稳定而且更容易观察到学习率设置是否合理。InitialLearnRate是初始学习率控制每次参数更新的步长。0.01是一个比较保守的起点对于CIFAR-10这种规模的数据集通常不会出大问题。学习率太大loss会震荡甚至变成NaN学习率太小loss下降慢得像蜗牛20轮跑完可能还没有明显收敛。MaxEpochs设为20意思是把全部训练数据过20遍。一个epoch就是把5万张图都喂给网络一次。对LeNet-5这种规模的网络来说20轮足够收敛如果你发现验证准确率还在上升可以加大到30轮试试。MiniBatchSize是每次梯度更新用的样本数128意味着网络每看128张图才更新一次参数。批次大小太大内存吃不消而且模型容易收敛到尖锐的极小值太小梯度噪声大训练不稳定。128在显存和稳定性之间是平衡点如果你的GPU显存只有2G建议降到64。ValidationData是验证集数据ValidationFrequency10表示每10次迭代在验证集上算一次准确率方便实时监控。Plots设置为training-progress会在训练时弹出一个实时更新的进度图直观显示loss和准确率的变化。4.2 训练过程trainNetwork到底在干什么设置好options之后训练就一行代码net trainNetwork(XTrain, YTrain, layers, options);这一行命令会自动完成整个训练流程初始化网络参数把训练数据按MiniBatchSize切块在每个批次上前向传播计算预测值再用反向传播计算梯度更新参数重复这个过程直到所有epoch结束。训练过程中你会看到命令行不断输出迭代信息包括当前迭代数、训练loss、验证准确率等。如果开了Plots还会实时画出一条loss下降曲线和一条accuracy上升曲线。第一次看到这个画面的人很容易犯一个错误盯着曲线看半天却不知道该关注什么信号。我的经验是盯三个地方loss是不是在平稳下降、验证准确率是不是在稳步上升、训练loss和验证loss的差距是不是在拉大。第一个是是否收敛的信号第二个是学习是否有效的信号第三个是是否过拟合的信号。4.3 数据增强一个小改动就能提升精度如果你跑完上面的代码验证准确率大概在70%左右。想进一步提升最简单有效的方法是加数据增强。所谓数据增强就是训练时对图片做一些随机变换比如水平翻转、平移几个像素这样做相当于免费生成更多训练样本能显著提升模型的泛化能力。Matlab里通过imageDataAugmenter和augmentedImageDatastore实现imageAugmenter imageDataAugmenter(... RandXTranslation, [-4, 4], ... RandYTranslation, [-4, 4], ... RandXReflection, true); augimdsTrain augmentedImageDatastore([32 32 3], XTrain, YTrain, ... DataAugmentation, imageAugmenter, ... OutputSizeMode, randcrop);然后训练时把XTrain, YTrain替换成augimdsTrainnet trainNetwork(augimdsTrain, layers, options);这里要注意数据增强只应该用在训练集验证集和测试集千万不能做随机变换。原因很简单验证集和测试集要模拟真实世界中“模型见过的图片”如果也随机翻转平移评估结果就不稳定每次跑出来的验证准确率都不一样。我自己刚开始不懂这个把验证集也做了增强结果每次验证准确率波动好几个百分点还以为网络训练有问题折腾了半天才发现是数据增强加错了地方。5. 测试与可视化怎么验证网络真的学到了东西5.1 classify与activations一个看结果一个看过程训练完成后用classify函数对测试集做预测计算最终准确率YPred classify(net, XTest); accuracy sum(YPred YTest) / numel(YTest); fprintf(测试集准确率%.2f%%\n, accuracy * 100);classify返回的是一个categorical向量和YTest直接比较就能算出预测正确的比例。我跑出来的结果大概在73%到78%之间具体数值会根据数据增强和随机种子有所波动。除了看最终分类结果还有一个特别值得玩味的函数是activations。它可以提取网络中间层的特征输出让你看到一张图片在网络内部到底被加工成了什么样。用法如下layerName conv_1; act activations(net, XTest(:, :, :, 1), layerName);这个act是一个四维数组第一维和第二维是特征图的高和宽第三维是通道数第四维是样本数。取一个通道出来用imshow显示你会看到原始图片经过第一个卷积层后变成了一张张“特征高亮图”有的通道特别响应边缘有的通道特别响应纹理。这种可视化比任何教材都更能让你理解卷积层到底在做什么。5.2 混淆矩阵看网络在哪类图片上最容易翻车准确率只是一个数字更细致的信息藏在混淆矩阵里figure; confusionchart(YTest, YPred);confusionchart会画出一张10x10的热力图每一行是真实类别每一列是预测类别对角线上的数字是预测正确的数量非对角线则是错误的分布。看这个矩阵你能快速发现网络在哪些类之间容易混淆。我跑出来的结果里最典型的是猫和狗互相误判的比例明显偏高。这其实是意料之中的事猫和狗在姿态、毛色、背景上有很多相似之处在32x32这么小的分辨率下人类自己都不一定能分清楚网络分不清很正常。另一个常见的混淆是飞机和船很多俯拍的飞机图像轮廓和船的轮廓非常接近也容易被搞混。发现这些混淆规律后你会对“模型学到了什么、没学到什么”有一个非常具体的感知这比单纯看一个准确率数字有价值得多。5.3 把预测错误的样本可视化除了混淆矩阵这种统计层面的分析把错误样本直接画出来看也很有用。我写了一个小脚本显示前几个预测错误的测试图片wrongIdx find(YPred ~ YTest); figure; for i 1:min(9, numel(wrongIdx)) idx wrongIdx(i); subplot(3, 3, i); imshow(XTest(:, :, :, idx)); title(sprintf(真实:%s, 预测:%s, char(YTest(idx)), char(YPred(idx)))); end看到这些错误样本你会发现很多“错误”其实情有可原。有些图片本身就很模糊有些主体只占据画面一小块还有些是背景干扰严重。这提醒我一件事一个模型就算测试准确率只有75%也不意味着它“很笨”在很多情况下它只是被32x32分辨率下的图像信息量限制住了。如果想要更高的准确率要么换更大的输入尺寸要么用更强的网络结构单纯调参的天花板很快就到了。5.4 analyzeNetwork与deepNetworkDesigner最后推荐一个调试工具analyzeNetwork()。训练之前或之后都可以调用analyzeNetwork(net);它会弹出一个交互界面把网络的每一层都列出来包括每层输出数据的尺寸、参数量。如果网络结构哪里不对比如某一层的输入尺寸和上一层输出对不上它会直接标红报错省去你自己一个维度一个维度去算的功夫。新版Matlab还提供了deepNetworkDesigner这个图形化工具可以像画流程图一样拖拽搭建网络点几下鼠标就能生成对应的代码。我自己平时还是会手写layers数组因为方便用git做版本管理但你如果刚入门用deepNetworkDesigner搭一遍网络结构对理解层的顺序和连接关系会有很大帮助。6. 我踩过的坑报错信息与排查思路6.1 四个高频报错速查表这一节直接上干货把我在学习过程中遇到过的典型报错整理成一张表方便你对照排查。报错信息片段常见原因排查方向Invalid training data. Predictors must be a N-by-1 cell array of images, or an imageDatastore数据维度不是[H W C N]或者标签类型不对用size()逐维检查数据形状确认YTrain是categorical类型Input size mismatch网络定义的输入尺寸和实际数据尺寸不一致对比imageInputLayer的参数和Xtrain前三维最好打印出来看Out of memory内存或显存不足MiniBatchSize调到32或64或设置ExecutionEnvironment为cpuLayer xxx does not existactivations里写的层名和网络里的Name不一致用analyzeNetwork查看实际层名这些报错有一个共同点它们的信息提示往往不会直接告诉你问题在哪一行而是说“数据不对”或“尺寸不匹配”。所以排查的第一原则是先把变量尺寸打印出来确认每一个维度的数值都符合预期再去看网络结构。6.2 让我卡了一下午的崩溃现场说一个真实的排查经历。第一次跑完整流程时我自信满满地执行trainNetwork结果第一行就报错提示训练数据尺寸不匹配。我用size()打印了XTrain显示是32x32x3x50000看起来完全正常。又打印了YTrain是50000x1的categorical也正常。网络输入层是[32 32 3]和XTrain的前三维一比对也没问题。卡了快一个小时后我才发现问题出在最初的reshape顺序上。CIFAR-10的数据是“行优先”存储的Matlab的reshape是“列优先”填充结果就是reshape出来的图像虽然尺寸是32x32但内容左右翻转了相当于每张图都被旋转了90度。网络拿这种“躺着的”图像训练准确率一路上不去我还一直以为是网络结构写错了。后来在代码里加了一行permute(XTrain, [2, 3, 4, 1])就解决了。想明白之后觉得特别简单但当时因为完全没有打印图像可视化检查光靠尺寸判断根本发现不了这个隐藏的维度顺序问题。从那以后我养成了一个习惯数据预处理完一定随机抽几张用imshow看一眼图像是正的、颜色通道是正确的再进入训练环节。6.3 排查问题的通用套路踩了足够多的坑之后我总结了一套自己的排查顺序。第一步缩小问题范围先判断问题出在数据、网络、还是训练参数。最简单的办法是看报错发生在哪个阶段trainNetwork之前报错几乎都是数据或网络定义的问题训练过程中loss出现NaN大概率是学习率太大或数据里有异常值训练正常结束但准确率低才是模型结构或超参数的问题。第二步把网络参数临时调小。把卷积核数量从20和50改成8和16把全连接层神经元从100改成32先跑通一个“最小可行版本”。如果小网络能跑通再逐步把规模加回去这样能快速定位是不是某个参数过大导致的内存或数值问题。第三步善用搜索。报错信息翻译成英文去搜索时一定要带上你的Matlab版本号因为不同版本的Deep Learning Toolbox在很多函数接口上有差异。比如早一点的版本里layerGraph是layerGraph新版本里连接层的写法就变了不带版本搜出来的老方法很可能已经失效。7. 学习笔记3的收尾体会把这段代码彻底看懂之后我最大的感受是Matlab里跑CNN没有想象中那么高深真正繁琐的地方在于让数据按照框架默认的格式摆放整齐。只要数据形状对了网络层写对了训练参数往保守里调模型跑通只是时间问题。这个“时间”大部分不是花在理解卷积公式上而是花在和维度、类型、报错信息作斗争上而这恰恰是代码理解笔记最需要记录的东西。做这个系列时我还发现理解LeNet级别的网络代码是通往更复杂CNN结构的捷径。后面再看ResNet时虽然多了残差连接但本质上还是在“卷积-激活-池化”这个基本单元上做文章。参数再多、结构再复杂数据流的方向永远是单向的每一层做的事情依然是“改变数据的形状和数值”。把这份代码吃透后续学习就有了一个可以随时对照的锚点。下一步我打算写一写怎么把训练好的模型导出对单张真实图片做预测以及怎么把网络加深一点看准确率如何变化。如果你也在啃Matlab里的CNN建议别急着抄大模型先把这个级别的代码每一行都弄明白跑通了再往前走。