训练视觉变换器网络用于图像分类
本示例演示了如何对预训练的视觉变换器 (ViT) 神经网络进行微调,以便对一组新的图像进行分类。
ViT [1] 是一种神经网络模型,它利用 Transformer 架构将图像输入编码为特征向量。该网络由两个主要组件构成:骨干网和主节点。骨干网负责网络中的编码步骤。该核心模块接收输入图像,并输出一个特征向量。负责人负责做出预测。该头部将编码后的特征向量映射到预测分数上。
在此示例中,预训练的 ViT 网络已经为图像学习到了强大的特征表示。您可以利用迁移学习,针对特定任务对模型进行微调。要迁移该特征表示并针对新数据集进行微调,请将网络的输出层替换为一个能对您的任务进行分类的新输出层,然后在新数据集上对网络进行微调。
该图概述了一个用于对 类进行预测的 ViT 网络的架构,以及如何修改该网络以实现针对包含 类的新数据集的迁移学习。

例如,在此示例中,您将对基础规模的 ViT 模型(8680 万个参数)进行微调,其补丁大小为 16,该模型使用分辨率为 384×384 的 ImageNet 2012 数据集进行微调。
加载预训练 ViT 网络
使用 visionTransformer 函数加载一个预训练的 ViT 网络。此函数需要 Deep Learning Toolbox™ 许可证和 Computer Vision Toolbox™ Model for Vision Transformer Network 支持包。您可以从附加功能资源管理器下载此支持包。如果您尚未安装该支持包,则该函数会提供一个下载链接。
net = visionTransformer
net =
dlnetwork with properties:
Layers: [143×1 nnet.cnn.layer.Layer]
Connections: [167×2 table]
Learnables: [200×3 table]
State: [0×3 table]
InputNames: {'imageinput'}
OutputNames: {'softmax'}
Initialized: 1
View summary with summary.
查看网络的输入规模。
inputSize = net.Layers(1).InputSize
inputSize = 1×3
384 384 3
要对 ViT 网络进行微调,通常只需对注意力层进行微调,并将其他可学习参数冻结即可 [2]。使用 freezeNetwork 函数冻结网络权重,该函数作为辅助文件附于本示例之后。要访问此函数,请以实时脚本形式打开此示例。
net = freezeNetwork(net,LayersToIgnore="SelfAttentionLayer");加载训练数据
下载并解压 Flowers 数据集 [3]。该数据集大小约为 218 MB,包含 3670 张花卉图像,分为以下五类:雏菊、蒲公英、玫瑰、向日葵以及郁金香。
url = "http://download.tensorflow.org/example_images/flower_photos.tgz"; downloadFolder = tempdir; filename = fullfile(downloadFolder,"flower_dataset.tgz"); imageFolder = fullfile(downloadFolder,"flower_photos"); if ~datasetExists(imageFolder) disp("Downloading Flowers data set (218 MB)...") websave(filename,url); untar(filename,downloadFolder) end
创建一个包含这些图像的图像数据存储。
imds = imageDatastore(imageFolder,IncludeSubfolders=true,LabelSource="foldernames");查看类数。
classNames = categories(imds.Labels); numClasses = numel(categories(imds.Labels))
numClasses = 5
使用 splitEachLabel 函数将数据集划分为训练集、验证集和测试集。将 80% 的图像用于训练,将 10% 留作验证集,另将 10% 留作测试集。
[imdsTrain,imdsValidation,imdsTest] = splitEachLabel(imds,0.8,0.1);
为了提高训练效果,应扩展训练数据,使其包含随机的旋转、缩放和水平翻转操作。将图像调整为与网络输入尺寸相匹配的大小。
augmenter = imageDataAugmenter( ... RandXReflection=true, ... RandRotation=[-90 90], ... RandScale=[1 2]); augimdsTrain = augmentedImageDatastore(inputSize(1:2),imdsTrain,DataAugmentation=augmenter);
创建增强型图像数据集,将验证集和测试集中的图像调整为与网络输入尺寸相匹配的大小。请勿对验证和测试数据进行任何额外的数据增强处理。
augimdsValidation = augmentedImageDatastore(inputSize(1:2),imdsValidation); augimdsTest = augmentedImageDatastore(inputSize(1:2),imdsTest);
替换网络分类头
ViT 网络主要由两个组件构成。网络的主体部分进行特征提取。分类头将提取的特征映射到概率向量上,这些概率向量代表了各类别的预测分数。要训练神经网络对一组新的类别进行图像分类,请将分类头替换为一个新的分类头,该分类头将提取的特征映射到针对这组新类别的预测分数上。
使用 analyzeNetwork 函数查看网络架构。找到网络末端将提取的特征映射为预测分数向量的层。在此情况下,名为 "head" 的全连接层将提取的特征映射为长度为 1000 的向量,1000 即为该网络经过训练可预测的类别数量。名为 "softmax" 的 softmax 层将这些向量映射为概率向量。
analyzeNetwork(net)

创建一个新的全连接层,其输出尺寸与训练数据中的类别数量相匹配:
将输出维度设置为训练数据的类数。
将图层名称设置为
"head"。
layer = fullyConnectedLayer(numClasses,Name="head");使用 replaceLayer (Deep Learning Toolbox) 函数将全连接层替换为新层。您无需替换 softmax 层,因为它不包含任何可学习参数。
net = replaceLayer(net,"head",layer);指定训练选项
指定训练选项。在选项中进行选择需要经验分析。要通过运行试验探索不同训练选项配置,您可以使用Experiment Manager (Deep Learning Toolbox)。
使用 Adam 优化器进行训练。
进行微调时,将学习率降低至 0.0001。
进行四轮训练。
使用小批量大小 12。训练一个 ViT 网络通常需要大量内存。如果内存不足,请尝试使用较小的小批量大小。或者,您也可以尝试使用更小的模型,例如超小型 ViT 模型(570 万个参数),方法是在
visionTransformer函数中将模型名指定为"tiny-16-imagenet-384"。每个 epoch 验证一次网络,使用验证数据进行验证。
输出导致验证损失最低的网络。
在图中监控训练进度并监控准确度度量。
禁用详细输出。
miniBatchSize = 12; numObservationsTrain = numel(augimdsTrain.Files); numIterationsPerEpoch = floor(numObservationsTrain/miniBatchSize); options = trainingOptions("adam", ... MaxEpochs=4, ... InitialLearnRate=0.0001, ... MiniBatchSize=miniBatchSize, ... ValidationData=augimdsValidation, ... ValidationFrequency=numIterationsPerEpoch, ... OutputNetwork="best-validation", ... Plots="training-progress", ... Metrics="accuracy", ... Verbose=false);
训练神经网络
使用 trainnet (Deep Learning Toolbox) 函数训练神经网络。对于分类,使用交叉熵损失。默认情况下,trainnet 函数使用 GPU(如果有)。在 GPU 上进行训练需要 Parallel Computing Toolbox™ 许可证和受支持的 GPU 设备。有关受支持设备的信息,请参阅GPU 计算要求 (Parallel Computing Toolbox)。否则,trainnet 函数使用 CPU。要指定执行环境,请使用 ExecutionEnvironment 训练选项。
例如,本示例使用配备 24 GB 内存的 NVIDIA Titan RTX GPU 对网络进行训练。训练大约需要 37 分钟。
net = trainnet(augimdsTrain,net,"crossentropy",options);
测试神经网络
使用测试数据评估网络的准确性。
使用测试数据进行预测。要将预测分数转换为类别标签,请使用 onehotdecode (Deep Learning Toolbox) 函数。
YTest = minibatchpredict(net,augimdsTest); YTest = onehotdecode(YTest,classNames,2);
将测试分类结果以混淆矩阵的形式显示出来。
figure TTest = imdsTest.Labels; confusionchart(TTest,YTest)

评估测试的准确性。
accuracy = mean(YTest == TTest)
accuracy = 0.9564
利用新数据进行预测
使用已训练好的神经网络,根据测试数据中的第一个图像进行预测。
从测试数据的第一份文件中读取图像。
idx = 1;
testData = readByIndex(augimdsTest,idx);
I = testData.input{1};根据图像做出预测。
Y = minibatchpredict(net,single(I));
使用 onehotdecode 函数获取概率最高的标签。
label = onehotdecode(Y,classNames,2);
显示图像和预测标签。
imshow(I) title(label)

fprintf("Image Credit: %s\n",flowerCredit(augimdsValidation.Files(idx)))Image Credit: CC-BY by mikeyskatie - https://www.flickr.com/photos/mikeyskatie/5948835387/
参考资料
Dosovitskiy, Alexey, Lucas Beyer, Alexander Kolesnikov, Dirk Weissenborn, Xiaohua Zhai, Thomas Unterthiner, Mostafa Dehghani et al."An Image is Worth 16x16 Words:Transformers for Image Recognition at Scale."Preprint, submitted June 3, 2021. https://doi.org/10.48550/arXiv.2010.11929
Touvron, Hugo, Matthieu Cord, Alaaeldin El-Nouby, Jakob Verbeek, and Hervé Jégou."Three things everyone should know about vision transformers."In Computer Vision–ECCV 2022, edited by Shai Avidan, Gabriel Brostow, Moustapha Cissé, Giovanni Maria Farinella, and Tal Hassner, 13684:497-515.Cham:Springer Nature Switzerland, 2022. https://doi.org/10.1007/978-3-031-20053-3_29.
TensorFlow.“Tf_flowers | TensorFlow Datasets.”Accessed June 16, 2023. https://www.tensorflow.org/datasets/catalog/tf_flowers.
另请参阅
visionTransformer | patchEmbeddingLayer | trainnet (Deep Learning Toolbox) | trainingOptions (Deep Learning Toolbox) | dlnetwork (Deep Learning Toolbox)
主题
- 在 MATLAB 中进行深度学习 (Deep Learning Toolbox)
- 深度学习层列表 (Deep Learning Toolbox)
- Deep Learning Tips and Tricks (Deep Learning Toolbox)
- Data Sets for Deep Learning (Deep Learning Toolbox)