Matlab提供了功能强大的深度学习工具箱(Deep Learning Toolbox),支持从数据准备、网络构建、训练到部署的全流程。无论是直接调用成熟的预训练模型(如AlexNet、VGG、ResNet等),还是自定义网络结构,Matlab都能以简洁的代码实现。本章将围绕AlexNet预训练模型、自定义网络构建、数据采集和训练设置展开,帮助读者快速上手Matlab深度学习。
预训练模型是在大规模数据集(如ImageNet)上训练好的网络,可直接用于图像分类、特征提取或迁移学习。Matlab通过alexnet函数提供了AlexNet模型的便捷访问。
15.2.1 加载预训练模型
matlab
代码块
PlainText
net = alexnet
复制成功
执行后,net变量包含了AlexNet的网络结构和权重。可通过analyzeNetwork(net)查看网络各层的详细信息。
15.2.2 了解AlexNet结构
AlexNet包含8层:5个卷积层和3个全连接层,最后是一个softmax层和分类输出层。输入图像尺寸为227×227×3,输出1000个类别。可通过net.Layers查看各层属性。
15.2.3 迁移学习:修改网络以适配新任务
在实际应用中,通常需要将预训练模型调整到自己的分类任务(如分类数不同)。迁移学习的常见做法是保留网络前几层的特征提取能力,替换最后几层以适应新类别。
matlab
代码块
PlainText
% 加载预训练网络
net = alexnet;
% 获取输入层大小
inputSize = net.Layers(1).InputSize;
% 替换最后三层:全连接层、softmax层、分类层
layers = net.Layers;
layers(end-2) = fullyConnectedLayer(5, 'Name', 'fc_new'); % 假设新任务有5类
layers(end-1) = softmaxLayer('Name', 'softmax_new');
layers(end) = classificationLayer('Name', 'classoutput_new');
复制成功
如果需要冻结前几层(防止过拟合,加速训练),可将它们的学习率设置为0:
matlab
代码块
PlainText
% 冻结前7层(假设前7层为特征提取部分)
for i = 1:7
layers(i) = freezeWeights(layers(i));
end
复制成功
15.2.4 使用预训练网络进行特征提取
若不需重新训练,可直接将预训练网络作为特征提取器:
matlab
代码块
PlainText
% 去除最后三层,保留特征向量
featureLayer = 'fc7'; % 通常取全连接层输出
features = activations(net, img, featureLayer);
复制成功
除了调用现成模型,Matlab还支持从零开始构建网络。使用layer函数或层数组定义网络结构。
15.3.1 创建层数组
以一个简单的卷积神经网络为例:
matlab
代码块
PlainText
layers = [
imageInputLayer([28 28 1]) % 输入层,28x28灰度图
convolution2dLayer(3,8,'Padding','same') % 卷积层,8个3x3滤波器
batchNormalizationLayer % 批归一化层
reluLayer % ReLU激活层
maxPooling2dLayer(2,'Stride',2) % 最大池化层
convolution2dLayer(3,16,'Padding','same')
batchNormalizationLayer
reluLayer
maxPooling2dLayer(2,'Stride',2)
fullyConnectedLayer(10) % 全连接层,10个神经元
softmaxLayer % Softmax层
classificationLayer % 分类层
];
复制成功
15.3.2 检查网络结构
使用analyzeNetwork(layers)可直观检查网络维度兼容性和各层连接情况。
15.3.3 添加自定义层
深度学习工具箱也支持创建自定义层,继承nnet.layer.Layer并实现相应方法,但本章不展开。
15.3.4 从网络图构建
对于更复杂的结构(如残差连接),可以使用dlnetwork或layerGraph构建有向无环图。
matlab
代码块
PlainText
lgraph = layerGraph(layers);
lgraph = addLayers(lgraph, additionalLayer);
lgraph = connectLayers(lgraph, 'layer1', 'layer2');
复制成功
深度学习需要大量标注数据。Matlab支持从摄像头实时采集图像,以及从视频文件中提取帧作为数据集。
15.4.1 从摄像头采集图像
使用webcam对象访问摄像头:
matlab
代码块
PlainText
cam = webcam; % 创建摄像头对象
preview(cam); % 预览
img = snapshot(cam); % 采集一帧
clear cam; % 释放摄像头
复制成功
如需连续采集并保存为数据集,可编写循环:
matlab
代码块
PlainText
cam = webcam;
numImages = 100;
for i = 1:numImages
img = snapshot(cam);
imwrite(img, sprintf('image_%03d.jpg', i));
pause(0.5); % 间隔0.5秒
end
clear cam;
复制成功
15.4.2 从视频中切割图像
使用VideoReader读取视频文件,逐帧保存:
matlab
代码块
PlainText
v = VideoReader('video.mp4');
frameCount = 0;
while hasFrame(v)
frame = readFrame(v);
frameCount = frameCount + 1;
% 可选:每隔几帧保存一次
if mod(frameCount, 5) == 0 % 每5帧保存一帧
imwrite(frame, sprintf('frame_%04d.jpg', frameCount));
end
end
复制成功
15.4.3 创建图像数据存储
将图像文件组织到文件夹中(每个类别一个子文件夹),然后创建imageDatastore:
matlab
代码块
PlainText
imds = imageDatastore('dataset_path', 'IncludeSubfolders', true, 'LabelSource', 'foldernames');
复制成功
15.4.4 划分训练集和验证集
matlab
代码块
PlainText
[imdsTrain, imdsValidation] = splitEachLabel(imds, 0.7, 'randomized');
复制成功
15.4.5 数据增强
训练时可使用augmentedImageDatastore进行实时数据增强(随机旋转、缩放、平移等):
matlab
代码块
PlainText
imageSize = [227 227 3];
augimds = augmentedImageDatastore(imageSize, imdsTrain, ...
'DataAugmentation', imageDataAugmenter(...
'RandRotation', [-10 10], ...
'RandXTranslation', [-5 5], ...
'RandYTranslation', [-5 5]));
复制成功
15.5.1 训练选项设置
使用trainingOptions函数设置训练参数:
matlab
代码块
PlainText
options = trainingOptions('sgdm', ... % 优化器
'InitialLearnRate', 0.001, ... % 初始学习率
'MaxEpochs', 20, ... % 最大迭代轮数
'MiniBatchSize', 64, ... % 小批量大小
'ValidationData', imdsValidation, ... % 验证数据
'ValidationFrequency', 30, ... % 验证频率
'Shuffle', 'every-epoch', ... % 每轮打乱数据
'Plots', 'training-progress', ... % 实时显示训练进度
'Verbose', true, ... % 命令行输出
'ExecutionEnvironment', 'auto'); % 自动选择CPU/GPU
复制成功
15.5.2 启动训练
matlab
代码块
PlainText
net = trainNetwork(augimds, layers, options);
复制成功
15.5.3 自动化训练:检查点保存与早停
为防止意外中断丢失进度,可启用检查点保存:
matlab
代码块
PlainText
options = trainingOptions(..., ...
'CheckpointPath', 'checkpoints'); % 每隔一定epoch保存模型
复制成功
还可设置早停(validation patience),当验证损失连续若干次不再下降时提前终止:
matlab
代码块
PlainText
options = trainingOptions(..., ...
'ValidationPatience', 5); % 5次验证不改善则停止
复制成功
15.5.4 并行训练
若有多GPU或集群,可设置'ExecutionEnvironment'为'multi-gpu'或'parallel'。
15.5.5 迁移学习的训练技巧
15.5.6 训练后评估
训练完成后,可使用classify对测试图像进行分类,并用confusionchart绘制混淆矩阵评估性能。
matlab
代码块
PlainText
YPred = classify(net, testImds);
YTest = testImds.Labels;
accuracy = sum(YPred == YTest) / numel(YTest);
confusionchart(YTest, YPred);
复制成功
下面是一个综合示例,演示从摄像头采集图像,使用预训练的AlexNet进行实时分类。
matlab
代码块
PlainText
% 加载预训练AlexNet
net = alexnet;
% 获取输入尺寸
inputSize = net.Layers(1).InputSize;
% 打开摄像头
cam = webcam;
% 创建图形窗口
figure;
while true
% 采集图像
img = snapshot(cam);
% 调整图像大小以匹配网络输入
imgResized = imresize(img, inputSize(1:2));
% 分类
label = classify(net, imgResized);
% 显示结果
imshow(img);
title(char(label));
drawnow;
% 按'q'退出
if waitforbuttonpress && strcmp(get(gcf,'CurrentCharacter'),'q')
break;
end
end
clear cam;
复制成功
本章介绍了在Matlab中进行深度学习的基本流程,包括:
Matlab深度学习工具箱以其简洁的语法和强大的可视化功能,极大降低了深度学习应用的入门门槛。结合本章内容,读者可以快速开展自己的图像识别项目。
免责声明:本文系网络转载或改编,未找到原创作者,版权归原作者所有。如涉及版权,请联系删