主要内容

本页采用了机器翻译。点击此处可查看英文原文。

训练视觉变换器网络用于图像分类

本示例演示了如何对预训练的视觉变换器 (ViT) 神经网络进行微调,以便对一组新的图像进行分类。

ViT [1] 是一种神经网络模型,它利用 Transformer 架构将图像输入编码为特征向量。该网络由两个主要组件构成:骨干网和主节点。骨干网负责网络中的编码步骤。该核心模块接收输入图像,并输出一个特征向量。负责人负责做出预测。该头部将编码后的特征向量映射到预测分数上。

在此示例中,预训练的 ViT 网络已经为图像学习到了强大的特征表示。您可以利用迁移学习,针对特定任务对模型进行微调。要迁移该特征表示并针对新数据集进行微调,请将网络的输出层替换为一个能对您的任务进行分类的新输出层,然后在新数据集上对网络进行微调。

该图概述了一个用于对 K 类进行预测的 ViT 网络的架构,以及如何修改该网络以实现针对包含 K* 类的新数据集的迁移学习。

例如,在此示例中,您将对基础规模的 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/

参考资料

  1. 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

  2. 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.

  3. TensorFlow.“Tf_flowers | TensorFlow Datasets.”Accessed June 16, 2023. https://www.tensorflow.org/datasets/catalog/tf_flowers.

另请参阅

| | (Deep Learning Toolbox) | (Deep Learning Toolbox) | (Deep Learning Toolbox)

主题