ARTICLE DETAIL

资讯详情

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

MATLAB迁移学习实战:用AlexNet搞定小样本图像分类与缺陷检测

MATLAB迁移学习实战:用AlexNet搞定小样本图像分类与缺陷检测 简介基于AlexNet的深度迁移学习小样本图像分类Matlab程序面向图像处理与计算机视觉开发者以及需要在工业质检中开展表面缺陷检测的工程师解决小样本数据下分类模型准确率低、易过拟合等问题。资源共76个文件包含75张jpg示例图像和1个Matlab脚本压缩包仅370KB轻量易用。示例图像覆盖螺丝刀、梅森罐、扑克牌、立方体等多类物体脚本完整实现了数据预处理、预训练模型加载、顶层微调、数据集划分、训练评估及混淆矩阵可视化等关键步骤。目前已有1437人学习下载。通过运行代码可直观理解迁移学习“冻结底层、更新顶层”的参数策略并掌握旋转、平移、缩放等数据增强方法在有限样本下显著提升模型泛化能力。这份代码将深度学习理论与实际任务结合是小样本图像分类和缺陷检测应用的良好实训参考。1. 深度迁移学习遇上小样本为什么 AlexNet 在 MATLAB 里还能打做图像缺陷检测的工程师应该都有过这种经历生产线拍了几百张图片标注好了想让 CNN 自己学特征结果从零训练一个网络验证准确率死活过不了 80%。不是模型不行是数据太少网络根本学不到泛化特征。这时候深度迁移学习的价值就体现出来了——拿 ImageNet 上训好的 AlexNet 当特征提取器冻结底层卷积只换掉最后三层重新训练几百张图就能把分类任务跑起来。这份 MATLAB 程序就是这么设计的用到的是MerchData小样本数据集五类小商品图片每类才 20 到 30 张配合预训练 AlexNet 微调代码走完一个完整的小样本图像分类流水线。适合两类人一是有 MATLAB 基础、想快速把深度学习落地到工业质检场景的工程师二是正在做课题、需要一份能跑通的迁移学习基线代码的研究生。核心就一句话——数据不够迁移来凑。2. 先把代码包拆开AlexNet 结构、替换层和那三个关键函数2.1 AlexNet.m 整个脚本到底做了什么打开AlexNet.m整个程序不长但结构非常紧凑。它做的事情可以用五步概括加载预训练网络、替换分类层、划分数据集、配置训练参数、跑trainNetwork。这个模式是所有 MATLAB 迁移学习任务的模板后续换 ResNet、GoogLeNet 都是同一套思路。% 加载预训练 AlexNet net alexnet; % 查看网络各层结构 analyzeNetwork(net); % 获取全连接层输入维度 fcLayer net.Layers(end-3); numClasses numel(categories(imds.Labels)); % 替换最后三层 layersTransfer [ net.Layers(1:end-3) fullyConnectedLayer(numClasses, Name, fc_new) softmaxLayer(Name, softmax_new) classificationLayer(Name, classoutput_new) ];这里需要注意net.Layers(1:end-3)截取了 AlexNet 从输入层到倒数第四层的全部层。AlexNet 的最后三层分别是fc81000 通道全连接、softmax和classification output这三层是和 ImageNet 的 1000 类绑死的必须拆掉换成自己数据集的类别数。fcLayer那行代码的作用是确认fc7的输出维度——2048这个数字后面配fullyConnectedLayer时会用到虽然本次替换的三层不需要它但在做特征提取方案时fc7的 2048 维特征就是你要的手工特征。analyzeNetwork这个函数值得单独说。它会把网络结构可视化成一张交互式图表每层的激活尺寸、可学习参数数量一目了然。新手最容易犯的错是不知道到底该替换哪些层跑一次analyzeNetwork就知道哪三层是分类头、哪几层是特征提取器。2.2 为什么只动最后三层而不是重新训练所有层这个问题几乎是每次分享必被问到的。答案是预训练模型的底层卷积学到的是通用特征——边缘、纹理、颜色块——这些特征在自然图像和工业缺陷图像之间是通用的。你换一个数据集底层的 Gabor-like 滤波器照样能激活。真正需要重新学的是高层语义组合方式对于 ImageNet 来说是狗嘴狗耳狗对于你的缺陷数据集来说是划痕凹坑次品。% 设置训练选项 options trainingOptions(sgdm, ... MiniBatchSize, 10, ... MaxEpochs, 6, ... InitialLearnRate, 1e-4, ... Shuffle, every-epoch, ... Verbose, false, ... Plots, training-progress);注意InitialLearnRate设成了1e-4比从零训练常用的1e-2低了两个数量级。原因很直接预训练权重已经处于一个较好的局部最优附近学习率太大会直接把权重踢出好区域导致灾难性遗忘——网络开始疯狂适配小样本数据里的噪声验证准确率不升反降。用1e-4这种小步长微调本质上是让网络在这个局部最优附近做精细搜索。MiniBatchSize设成 10 是因为小样本数据集本身不大批太大一步迭代就把整个训练集看完了梯度更新过于平滑收敛慢。Shuffle设为every-epoch保证每个 epoch 训练数据的顺序重新打乱避免模型学到样本顺序的假规律。MaxEpochs只设 6 轮这也是迁移学习的特点——不要训练太久小数据集上微调过头就是过拟合我见过有人把 epoch 设到 30 的最后测试集准确率反而掉了 15 个百分点。2.3 数据划分别把所有图像一股脑扔进去训练代码里用了splitEachLabel做分层抽样这是小样本分类的标配操作% 按标签分层划分数据集 [imdsTrain, imdsValidation, imdsTest] splitEachLabel(imds, 0.6, 0.2, 0.2, randomized); % 统计各类别样本数 trainCounts countEachLabel(imdsTrain);splitEachLabel的四个参数依次是训练集比例 0.6、验证集比例 0.2、测试集比例 0.2、抽取方式randomized。关键是randomized——如果不指定MATLAB 会按文件夹读取顺序从头开始切如果原始数据集是按类别排序存放的就会切出糟糕的分布。比如某个类别只有 20 张图按顺序切 60% 可能把其中 12 张全划到训练集验证集里这个类别一张都没有。countEachLabel这行强烈建议保留。它能发现一个隐蔽的坑如果某个类别样本数特别少比如只有 5 张那splitEachLabel划分后这个类别在验证集里可能只剩 1 张图准确率波动会极其剧烈。看到这种情况就该考虑对这个类别单独做数据增强而不是硬着头皮训练。3. 数据增强与图像预处理小样本能不能打就看这一步3.1 imageAugmenter给有限样本无中生有小样本图像分类最大的敌人不是模型复杂度不够而是过拟合。模型只有 20 张图可看它很快就会把背景是传送带银色反光当成次品的特征。数据增强是缓解这个问题的第一道防线。% 定义数据增强器 imageAugmenter imageDataAugmenter( ... RandRotation, [-10 10], ... RandXTranslation, [-5 5], ... RandYTranslation, [-5 5], ... RandXScale, [0.95 1.05], ... RandYScale, [0.95 1.05]); % 使用增强器创建增强图像数据存储 augimds augmentedImageDatastore([227 227 3], imdsTrain, ... DataAugmentation, imageAugmenter, ... OutputSizeMode, resize);augmentedImageDatastore有个容易被忽略的行为它不是在内存里先把增强后的图片全部生成好而是每个 iteration 随机抽取一批原始图像实时做增强变换后再喂给网络。这意味着同样一张原图在 6 个 epoch 里每轮看到的都是略微不同的版本——旋转角度不同、缩放比例不同、平移量不同。所以它不是把数据量翻了 N 倍而是让每次迭代看到的样本都不一样等效于无限多的训练样本。参数设计上RandRotation只给了 ±10 度不是 ±30 度。这是有讲究的工业质检场景中产品在传送带上的姿态通常是可控的不会出现倒置或 90 度翻转旋转范围太大反而会让模型学到旋转不变性这种在这个场景下不需要的冗余能力挤占本就有限的模型容量。RandXTranslation和RandYTranslation设置为 ±5 像素模拟产品在图像中的位置抖动这是缺陷检测中最常见的真实扰动。RandXScale的 0.95~1.05 模拟产品尺寸波动毕竟不同批次的产品大小会有细微差异。3.2 输入尺寸 227 与归一化一个像素都不能差AlexNet 的输入层要求 227×227×3 的 RGB 图像。这是历史遗留问题——AlexNet 论文里写的是 224×224但实际实现时因为 55×55 这个特征图尺寸计算出来正好是 227所以 Caffe 和 MATLAB 的实现都是按 227 来的。如果你把图片 resize 成 224 塞进去最后一层卷积的尺寸不匹配MATLAB 会直接报错。augmentedImageDatastore已经帮你做了 resize但要注意OutputSizeMode, resize的含义默认是resize即直接拉伸图片到 227×227不管原始宽高比。如果改成centercrop则是先按比例缩放再居中裁剪。这两种方式各有利弊resize会改变物体形状比例但对检测小目标更友好centercrop保持比例但可能裁掉边缘区域的缺陷。另外一个预处理细节是归一化。alexnet函数加载的模型自带inputLayer的归一化参数但augmentedImageDatastore不会自动应用它。标准做法是额外指定% 创建含归一化的增强数据存储 augimdsTrain augmentedImageDatastore([227 227], imdsTrain, ... DataAugmentation, imageAugmenter, ... OutputSizeMode, resize, ... Normalization, none);Normalization, none的含义是让数据保持原始像素范围因为trainNetwork在内部会自动处理。但如果你用的是predict函数而非trainNetwork就必须手动做imresize(double(img) / 255)的预处理否则预测结果会是一堆随机数。这个差异让很多人踩坑训练时效果好部署到独立脚本里预测就翻车原因就是训练时trainNetwork替你做了一部分预处理而部署时没人替你做了。4. 避坑指南五个让准确率莫名暴跌的元凶4.1 类别顺序颠倒导致的全乱套现象训练过程 loss 正常下降验证曲线跳动剧烈最终测试集准确率接近随机猜测。原因splitEachLabel按字母序或文件夹序排列类别但imds.Labels的类别顺序是 categorical 数组的枚举顺序。如果你在替换全连接层时用了numel(unique(imds.Labels))计算类数恰好unique返回的类别顺序和categories(imds.Labels)不一致后续标签映射就会错位。解决用categories(imds.Labels)获取类别列表并显式检查% 检查类别顺序 assert(isequal(categories(imdsTrain.Labels), categories(imdsValidation.Labels)), 类别顺序不一致);4.2 数据增强把标签也旋转没了现象训练集损失正常下降验证集准确率却极低且不随迭代提升。原因某些增强操作如RandXReflection水平翻转用在文本识别或左右不对称的缺陷上是致命的。比如检测左边缘划痕和右边缘划痕翻转后类别含义完全变了。解决查imageDataAugmenter的属性禁止使用RandXReflection和RandYReflection只保留旋转、平移、缩放这类几何变换。如果你不确定哪些变换安全最保守的组合是只留RandRotation和RandXTranslation。4.3 验证集没有 shuffle准确率曲线像心电图现象训练集准确率 100%验证集准确率在 60%~95% 之间剧烈震荡曲线呈锯齿状。原因trainingOptions中没有设置Shuffle, every-epoch。默认情况下验证数据在每个 epoch 内的顺序是固定的如果验证集里恰好某个类别的样本集中在末尾某一轮验证时模型对这批样本判断失误就会导致该轮验证准确率断崖式下跌。解决在trainingOptions中显式设置Shuffle, every-epoch。注意验证集 shuffle 是在每个 epoch 开始时执行不影响训练数据的随机性。4.4 图片读取失败扩展名大小写和损坏文件现象程序报错Unable to read file...或者训练到一半突然终止。原因数据集文件夹里混入了非图片文件如.txt说明文档、.db缩略图缓存或者图片扩展名大小写不一致.JPG和.jpg。MATLAB 的imageDatastore默认只识别常见扩展名且区分大小写。解决在创建imageDatastore前先清理文件夹% 删除非图片文件在数据集根目录执行 filelist dir(fullfile(rootPath, **, *.*)); for i 1:length(filelist) [~, ~, ext] fileparts(filelist(i).name); if ~ismember(lower(ext), {.jpg, .jpeg, .png, .bmp}) delete(fullfile(filelist(i).folder, filelist(i).name)); end end4.5 训练中途根本不收敛loss 直接 NaN现象训练开始后几个 iteration损失值变成NaN然后曲线消失。原因学习率设置过大迁移学习中常见或者MiniBatchSize太小导致 batch 内样本方差过大。还有一个冷门原因数据集里存在全黑或全白图片归一化后出现除零。解决先把学习率降到1e-5试跑 2 个 epoch确认 loss 正常递减后再调回1e-4。同时在做数据增强时排除低信息量图片用std(im, [1 2])筛掉标准差小于 10 的图片避免全黑全白图参与训练。这个操作也叫删垃圾样本在工业数据集里几乎总会遇到几张坏图。5. 训练与评估从混淆矩阵到缺陷检测场景落地5.1 classify 和 predict 的区别要准确率还是要置信度训练完成后AlexNet.m里用classify直接输出类别标签但实际项目里我更推荐用predict拿置信度向量% 对测试集做预测 YPred classify(netTransfer, augmentedTestImds); % 计算准确率 accuracy mean(YPred imdsTest.Labels); % 输出混淆矩阵 plotconfusion(imdsTest.Labels, YPred);% 用 predict 获取各类别概率 [YPredProb, scores] predict(netTransfer, augmentedTestImds); % scores 是 N×K 的矩阵N 为样本数K 为类别数 % 每行是该样本属于各个类别的概率行和为 1classify返回的是标签predict返回的是概率分布。在工业质检场景中概率分布的价值远大于标签。比如某个样本预测为次品的概率是 0.51说明这个样本处于边界状态应该送人工复检而不是直接判定。用predict可以设置双阈值大于 0.9 直接通过小于 0.7 直接拦截中间区域人工处理。scores矩阵的行数等于输入图片数列数等于类别数。找出每行的最大值和对应序号[maxScore, idx] max(scores, [], 2); % idx 是类别编号映射回类别名 predictedLabels imdsTest.Labels(idx);这里idx的数据类型是 double 的索引值imdsTest.Labels(idx)会找到索引对应位置的类别标签。这个映射关系要确保idx从 1 开始编号而 MATLAB 的 categorical 数组索引也从 1 开始直接调用不会越界。5.2 MerchData 数据集与工业缺陷检测的映射逻辑MerchData是 MathWorks 官方示例数据包含五类小商品图片Cap、MathWorks Tool、Screwdriver、Playing Cards、Cube每类约 20 到 30 张共 95 张左右。这个数据集的经典之处在于它与真实工业缺陷检测场景高度同构——类别间外观相似都是小型工业零件/工具、背景简单、光照可控。把它迁移到缺陷检测场景时只需要做一件事把五个类别改成两个即合格品和缺陷品。如果还想细分缺陷类型就按实际缺陷种类数替换numClasses。很多工程上的缺陷检测本质上是二分类问题OK/NGNext Goods / No Goods这时候把numClasses改成 2数据集按OK和NG两个子文件夹存放即可。但要注意缺陷检测的难点不在分类而在数据获取。缺陷样本往往比正常样本少一个数量级——工厂跑一天生产一万件产品缺陷可能只有五十件。这时候迁移学习的策略要做调整正负样本比例严重失衡时直接训练会倾向把所有样本预测为OK。常见做法是对缺陷类做更激进的数据增强旋转角度放宽到 ±30 度加入高斯噪声同时在 loss 函数中增加类别权重。MATLAB 里实现类别权重需要自定义损失层不太方便折中方案是把缺陷类样本做replication——复制 N 份到训练集里让模型看到更多的缺陷样本。这个方法简单有效代码也比较直观% 对缺陷类进行复制以平衡样本 ngImds imdsTrain.Labels NG; ngFiles imdsTrain.Files(ngImds); % 复制 NG 样本 4 次 imdsTrain.Files [imdsTrain.Files; repmat(ngFiles, 4, 1)]; % 注意复制后需要重建 imds 或调用 update 方法复制的本质是让每个 epoch 内模型看到更多的缺陷样本等效于提高了缺陷类在 loss 中的权重。但要注意复制操作是在内存层面展开的repmat会显著增加内存占用小数据集问题不大上万张图就要考虑在 datastore 层做采样控制。5.3 验证集准确率不是终点看看每类错在哪小样本分类最容易出现整体 90%某一类 60%的偏科现象。plotconfusion画出的混淆矩阵能一眼看出问题但要定位到具体是哪几张图被分错需要额外几行代码% 找出分错的样本 misclassifiedIdx find(YPred ~ imdsTest.Labels); % 逐张显示分错的图片 for i 1:min(numel(misclassifiedIdx), 9) idx misclassifiedIdx(i); subplot(3, 3, i); img readimage(imdsTest, idx); imshow(img); title(sprintf(真: %s, 预测: %s, ... char(imdsTest.Labels(idx)), char(YPred(idx)))); end这个可视化能帮你快速判断是图片本身就模糊难辨数据质量问题还是模型学偏了特征比如把螺丝刀的金属反光当成了易拉罐的银色瓶身。如果是后者说明数据增强里缺少了颜色扰动应该加入RandBrightness, [0.9 1.1]之类的亮度抖动让模型不再依赖绝对亮度值做判断。6. 一个收尾的实战技巧用 Grad-CAM 让网络说实话模型训练完准确率也达标了但在工业场景中直接上线风险很大——你怎么知道网络是学到了产品表面的划痕特征还是仅仅记住了缺陷样本都拍摄于传送带左侧这种环境伪影这时候需要用 Grad-CAM 做一次可解释性验证。% 计算 Grad-CAM 热力图 gradCAMMap gradCAM(netTransfer, img, LayerName, conv5); % 叠加显示 imshow(img); hold on; imagesc(gradCAMMap, AlphaData, 0.5); colormap jet; colorbar; hold off;gradCAM的输入参数分别是训练好的网络、单张测试图片、LayerName指定计算梯度的层名。AlexNet 的conv5层输出的是最后卷积特征图空间分辨率时 13×13能较好地定位注意力区域。如果换成fc7热力图会退化成 1×1没有空间信息。AlphaData设为 0.5 表示热力图 50% 透明度叠加在原图上注意叠加后要在imagesc之后设置colormap jet否则默认的 parula 色图在蓝色区域和原图背景容易混淆。判断标准很直接如果conv5热力图的高亮区域聚焦在产品本身的轮廓和纹理上说明模型学到了类别的真实特征如果高亮区域全部落在背景或某个固定角落那模型大概率是拿背景信息蒙混过关了。遇到后者处理办法不是在网络结构上纠结而是回数据端给训练集做背景替换或者把产品区域裁剪后再训练。分享一个亲身经历之前在做一个轴承表面缺陷检测项目训练准确率 98%Grad-CAM 一跑发现网络高亮区域全在轴承外围的齿环上——因为齿环区域有规则纹理划痕和齿环纹理在低层特征上容易混淆。后来在预处理环节加了背景掩膜把非产品区域像素置零准确率才真正稳定在 98% 线上。从那以后我每次训练完模型不管准确率多高都会强制过一遍 Grad-CAM确认网络注意力在该看的地方再谈上线。希望这个习惯也能帮到你。本文还有配套的精品资源点击获取
返回列表