Overall positive prototype for few-shot open-set recognition
abstract
- 少样本开集识别(FSOR)是用有限数量的标注实例识别已知类别中的样本,同时检测不属于任何已知类别的样本的任务。这是一个具有挑战性的问题,因为模型必须学会从少量带标签的样本中进行归纳,并将它们与无限数量的潜在负面样本区分开来。在本文中,我们提出了一种新的方法称为整体positive 原型,以有效地提高性能。从概念上讲,负样本会分布在整个特征空间,很难描述。从相反的观点来看,我们建议构建一个整体positive 原型,作为分布在相对较小的邻域中的阳性样本的内聚表示。通过测量查询样本和总体positive 原型之间的距离,我们可以有效地将其分类为positive or negative。我们证明了这种简单而创新的方法在精度和AUROC方面提供了最先进的FSOR性能。GitHub - liangyu-git/FSOR-OPP。
- 本文针对少样本开放集识别(FSOR)问题,提出了整体正原型(OPP)方法 ,通过构建一个能凝聚正样本信息的整体正原型,来区分已知类和未知类样本。该方法利用 Transformer 生成整体正原型,并结合分类损失和正原型损失优化模型,在 MiniImageNet、TieredImageNet 等多个数据集上实现了准确率和 AUROC 的 state-of-the-art 性能 ,验证了通过建模正样本分布比生成负原型更有效的思路。少样本开放集识别(FSOR)挑战 :需在少量标注样本下识别已知类,并检测未知类。传统方法因负样本在特征空间分布广泛,生成有效负原型困难。从相反角度建模正样本分布,提出整体正原型(OPP),通过聚合正类原型形成紧凑表示,以距离判断样本是否为未知类。
- 采用 K-way N-shot 设置,通过 ResNet-12 提取特征,计算类原型均值。自注意力模块处理类原型,增强任务适应性。用 Transformer 编码器处理类原型和任务 token,生成整体正原型 P O P_O PO。FCN 网络优化 P O P_O PO,不同配置(如 640-2048-640)影响性能。总损失 L t o t a l = L C E + α L O P L_{total}=L_{CE}+\alpha L_{OP} Ltotal=LCE+αLOP,其中 L O P L_{OP} LOP 推动正样本与 P O P_O PO 相似度趋近 1,负样本趋近 - 1。Transformer 层数:5 层时性能最佳,过多层导致过拟合。OPP 通过建模正样本分布提升 FSOR 性能,在多数据集达 SOTA。
- OPP 的核心创新在于从相反角度建模正样本分布,通过 Transformer 生成一个能凝聚所有正类原型信息的整体正原型,避免了传统方法生成负原型的困难。实验表明,正样本在特征空间分布更紧凑,通过计算样本与整体正原型的距离,能有效区分已知类和未知类。该方案在训练过程中需要开集数据,用于学习模型对未知类别的区分能力。
Introduction
-
大型标注数据集和强大的卷积神经网络模型的可用性导致了图像识别研究的重大进展。然而,在医疗或军事图像等场景中,数据采集仍然是一项具有挑战性的任务,标记大型数据集需要大量的人力。为了应对这些挑战,少样本学习被引入作为解决有限标记数据问题的有前途的方法。当前的 few-shot 学习方法依赖于计算查询样本和已知类原型之间的相似性,这将任务限制为闭集识别。尽管如此,在现实世界中,并不是所有的实例都标有类,并且有无限多的实例类型。
-
因此,正在进行的研究工作是调查开集识别,其目的是赋予模型有效处理分布外测试数据的能力。闭集识别只需要模型将输入数据分类到预定义的已知类别中,与之相反,开集识别需要额外的任务来识别不能被分类到任何已知类别中的数据实例。因此,该模型必须能够检测这样的实例,并将它们分配给"未知"类。因此,这种进行少样本开集识别(FSOR)的能力在现实世界中至关重要。
-
许多以前的FSOR方法试图构造一个基于阈值的分类器,该分类器接受并识别一个肯定查询到一个少样本类,并检测从看不见的类中采样的否定查询。许多方法采用元学习方案来学习基于阈值的检测器。另一方面,黄等人提出了一种任务自适应方法,即从正原型估计负原型。如果查询更接近负原型,则将其分类为负的,如果查询更接近某个正原型,则将其分类为已知类别之一。
-
在Task-adaptive negative envision for fewshot open-set recognition中,负原型是由类似全连接网络的神经网络或基于正原型的 Transformer 生成的。然而,我们发现很难有效地生成负样本,从而很好地描述负样本的分布。与正样本相比,负样本在特征空间中分布更广。即使有多个负原型,FSOR性能的提高也是有限的。
-
在本文中,我们从相反的观点出发。不如我们用一个 overall positive prototype 的正面原型来模拟"相对较窄的邻域中的正面原型"的分布?对这种分布进行建模会相对容易一些,从而可以实现有效的FSOR 。可以将一个查询与 overall positive prototype 进行比较,它们之间的距离可以被视为该查询是来自未知类的负样本的可能性。如果距离低于阈值,则将该查询与正原型进行比较,并将其分类到具有最小距离的类别中。图1说明了在FSOR使用 overall positive prototype 的概念。通过设置不同的阈值来检测阴性样本,我们将表明这种简单的方法总体上产生更高的平均准确度和AUROC。
-

-
图一。用 overall positive prototype 做FSOR的概念。将查询与 overall positive prototype 进行比较,以衡量查询正面或负面的可能性。
-
-
请注意,正面或负面原型本身并不新奇,这一点之前已经被广泛研究过了。我们工作的新颖之处在于提出了"overall positive prototype"来模拟正原型在相对较窄的邻域中的分布。关键的贡献是我们验证它相对更容易,并提供更有效的FSOR。
-
本文的其余部分组织如下。第二部分介绍了少样本学习、开集识别和FSOR的文献综述。在第3节中,我们描述了 overall positive prototype 的概念和生成模型的细节。第4节介绍了性能评估和消融研究,随后是第5节的结束语。
-
整体正原型生成器的目标是通过聚合少样本类原型,生成一个能代表所有正样本分布的紧凑表示。其具体架构及工作流程如下:
- 输入处理 :将少样本类原型 p 1 , p 2 , ... , p K p_1, p_2, \dots, p_K p1,p2,...,pK 作为输入,每个原型先通过线性层投影为 token。此外,添加一个随机初始化的任务 token,其作用是总结整个任务的正类信息,类似 Transformer 图像分类中的类 token。
- Transformer 编码器:将任务 token 与正类原型 token 共同输入到包含 5 层、每层 6 个头的 Transformer 编码器中。通过自注意力机制,模型捕捉不同原型之间的关联,例如计算第i个 token 与所有其他 token 的注意力权重,通过加权求和生成上下文感知的表示。
- 全连接网络(FCN) :Transformer 处理后的任务 token 经 FCN 生成最终的整体正原型 P O P_O PO。FCN 的配置(如层数和维度)影响性能,例如 "640-2048-640" 表示输入 640 维,中间层 2048 维,输出 640 维。实验表明,FCN 配置对性能影响较小,模型具有稳定性。
-
整体正原型损失 L OP L_{\text{OP}} LOP 的设计逻辑 : L OP L_{\text{OP}} LOP 的目标是迫使正样本与整体正原型 P O P_O PO 的相似度趋近于 1,负样本相似度趋近于 - 1,数学表达式为: L OP = { 1 − cos ( P O , q ) , 若 q 为正样本 , 1 + cos ( P O , q ) , 若 q 为负样本 , L_{\text{OP}} = \begin{cases} 1 - \cos(P_O, q), & \text{若 } q \text{ 为正样本}, \\ 1 + \cos(P_O, q), & \text{若 } q \text{ 为负样本}, \end{cases} LOP={1−cos(PO,q),1+cos(PO,q),若 q 为正样本,若 q 为负样本,其中 cos ( ⋅ ) \cos(\cdot) cos(⋅)为余弦相似度,取值范围(-1, 1)。
- 正样本优化目标 :当q为正样本时, L OP L_{\text{OP}} LOP 随 cos ( P O , q ) \cos(P_O, q) cos(PO,q) 增大而减小,推动 P O P_O PO 与正样本特征靠近,确保已知类样本的表示紧凑。负样本优化目标 :当q为负样本时, L OP L_{\text{OP}} LOP 随 c o s ( P O , q ) cos(P_O, q) cos(PO,q)减小而减小,迫使 P O P_O PO 与负样本特征远离,增强对未知类的区分能力。
-
生成器为损失函数提供优化对象 :生成器输出的 P O P_O PO 作为的 L OP L_{\text{OP}} LOP 输入,通过优化 L OP L_{\text{OP}} LOP 迫使 P O P_O PO 成为正样本的紧凑表示。损失函数引导生成器学习 : L OP L_{\text{OP}} LOP 的优化方向(正样本靠近、负样本远离)直接指导生成器调整 P O P_O PO 的参数,使其更适合区分已知类与未知类。在元训练阶段,生成器参数与特征提取器、自注意力模块共同更新,确保 P O P_O PO 能适应不同少样本任务的正样本分布。
Related works
Few-shot learning
- 少量学习问题需要从有限的标记数据中学习。已经提出了大量的方法 来解决这个问题。其中,原型网络近年来变得流行起来。原型网络为每个类别计算一个原型,通常通过取该类别中几个例子的特征向量的平均值来实现。在推理过程中,测试示例被分配给具有最近原型的类。这种方法已经在包括图像识别和文本分类在内的各种少样本分类任务中显示出有希望的结果。我们的工作建立在原型网络的基础上,通过引入一个额外的类原型,即 overall positive prototype 。该原型有助于识别测试集中的新类,因此可以处理少样本学习中的开集情况。
Open-set recognition
- 开集识别旨在识别已知类别的实例,同时检测属于未知或新类别的实例。这种设置实际上更符合现实世界的情况。研究人员提出了各种方法来解决这一问题,包括通过生成模型生成负样本,使用包含占位符的自适应决策边界来预测新的类别分布,通过分类器明确建模未知类别,以及利用开放世界分类器来识别未知类别。尽管这些方法很有效,但大多数都严重依赖大量的训练数据来促进模型学习。将这些方法直接应用于少样本的情况通常会产生有限的性能。
Few-shot open-set recognition
-
结合少样本学习和开集识别,少样本开集识别涉及识别只有少数例子的正类,同时还将它们与训练数据中不存在的类区分开来。已经提出了几种方法来关注FSOR,例如采用元学习来使分类器适应新的类别,生成否定原型来构建无阈值分类器以拒绝否定查询,以及使用图像的背景区域作为参考来更恰当地描述否定样本。
-
刘等人Few-shot open-set recognition using meta-learning最早提出了设置。他们将元学习技术引入FSOR,以减少经典softmax分类器在few-shot开集情况下使用时通常出现的过拟合。他们还引入了一种基于马氏距离的度量学习方法。周等人提出了新类数据占位符的概念,将开集问题转化为闭集问题,并提出了分类器占位符的概念,以校正典型分类器导致的过度自信预测。与对未知样本进行采样以估计未知类别分布不同,Jeong等人基于测量转换后的原型和修改后的原型集之间的差异,将估计问题转换为相对特征转换问题。Song等人提出利用来自可见类的图像的背景区域作为伪不可见信息,并通过学习可见类和不可见类之间的决策边界来开发分类器。黄等人提出从正原型预测负原型,并开发了一个无阈值的分类器。
-
大多数FSOR方法都是在归纳学习的设定下,在训练过程中只看到有标签的训练数据,没有标签的测试数据只有在测试的时候才能看到。因为在少样本学习问题中,训练数据的数量非常有限,所以越来越多的方法在直推式设置中发展起来,以进行少样本学习。最近,Boudiaf等人Open-set likelihood maximization for few-shot learning已经将直推式方法扩展到少样本开集识别。
Method
Overview of meta learning
-
我们首先简单介绍一下少样本学习中常用的元学习策略。元学习的目标是构建一个模型,该模型最初是基于一系列任务训练的,但可以适应具有少量数据点的新任务,而不需要从头开始训练。该模型是基于一组"元训练"任务训练的,并且适用于具有少量数据点的新"元测试"任务。
-
形式上,让我们将数据集𝐷分成非重叠的元训练集 𝐷𝑚𝑒𝑡𝑎_𝑡𝑟𝑎𝑖𝑛和元测试集𝐷𝑚𝑒𝑡𝑎_𝑡𝑒𝑠𝑡.请注意,𝐷𝑚𝑒𝑡𝑎_𝑡𝑟𝑎𝑖𝑛和𝐷𝑚𝑒𝑡𝑎_𝑡𝑒𝑠𝑡的类是不相交的。𝐷𝑚𝑒𝑡𝑎_𝑡𝑟𝑎𝑖𝑛和𝐷𝑚𝑒𝑡𝑎_𝑡𝑒𝑠𝑡被进一步细分为"支持集"和"查询集"。支持集合和查询集合中的类可以重叠。让我们分别把它们命名为𝐷𝑚𝑒𝑡𝑎_𝑡𝑟𝑎𝑖𝑛={𝐷𝑠𝑢𝑝𝑝𝑜𝑟𝑡𝑚𝑒𝑡𝑎_𝑡𝑟𝑎𝑖𝑛,𝐷𝑞𝑢𝑒𝑟𝑦 𝑚𝑒𝑡𝑎_𝑡𝑟𝑎𝑖𝑛}和𝐷𝑚𝑒𝑡𝑎_𝑡𝑒𝑠𝑡 = {𝐷 𝑠𝑢𝑝𝑝𝑜𝑟𝑡 𝑚𝑒𝑡𝑎_𝑡𝑒𝑠𝑡,𝐷𝑞𝑢𝑒𝑟𝑦 𝑚𝑒𝑡𝑎_𝑡𝑒𝑠𝑡}。元学习模型 M 是以情节模式训练的。在每集中,一个元训练任务是通过从𝐷 𝑠𝑢𝑝𝑝𝑜𝑟𝑡 𝑚𝑒𝑡𝑎_𝑡𝑟𝑎𝑖𝑛.的 k 个类别中取样而形成的基于采样的训练数据来训练模型 M,然后将其用于对从𝐷 𝑞𝑢𝑒𝑟𝑦 𝑚𝑒𝑡𝑎_𝑡𝑟𝑎𝑖𝑛.采样的相同 𝐾 类的图像进行预测在下一集,可能会对不同的 𝐾 类进行采样,以训练和测试模型 M. M 的参数根据剧集的预测误差不断更新,这被称为元训练阶段。对于𝐾-way 𝑁-shot分类问题,𝐾类是从𝐷 𝑠𝑢𝑝𝑝𝑜𝑟𝑡 𝑚𝑒𝑡𝑎_𝑡𝑟𝑎𝑖𝑛采样的,每个类有 𝑁 支持图像。𝑁 的数字通常很小,比如 𝑁 = 1 或 𝑁 = 5。
-
在元训练阶段之后,模型 M 带着𝐷 𝑠𝑢𝑝𝑝𝑜𝑟𝑡 𝑚𝑒𝑡𝑎_𝑡𝑒𝑠𝑡建立一个分类器。具体来说,元测试任务是通过从𝐷 𝑠𝑢𝑝𝑝𝑜𝑟𝑡 𝑚𝑒𝑡𝑎_𝑡𝑒𝑠𝑡.抽取𝐾类来完成的模型 M 适于作为特定于采样的𝐾支持类的分类器。评估时,从𝐷 𝑞𝑢𝑒𝑟𝑦 𝑚𝑒𝑡𝑎_𝑡𝑒𝑠𝑡随机抽取查询图像,并使用适应模型 M 对查询图像进行预测。属于 𝐾 支持类的查询图像被称为"正查询",而不属于𝐾支持类的查询图像被称为"负查询"。在实际实现中,我们会进行平衡采样,即肯定查询的数量等于否定查询的数量。
The proposed method
-
我们采用元学习策略来构建一个少样本开集识别系统。图2说明了我们提出的方法的总体工作流程。最初,一个特征提取器,比如ResNet-12 ,被用来从输入的少样本训练数据中提取特征。当进行元训练时,来自同一类的 𝑁 样本的特征被馈送到原型生成器以获得正原型。我们尝试了几种方法,如多层感知器或 Transformer 来构建原型生成器,但最终我们选择了与原型相同类别的特征的均值向量,因为它简单有效。图2示出了在四个类别的每一个中只有一个样本,这意味着只有一个样本的特征直接是正面原型。如果有𝐾类,我们在下面把这些𝐾类的正面原型表示为𝒑1,...,𝒑𝐾。
-

-
图二。用 overall positive prototype 说明建议的方法。
-
-
受transformer 的少样本嵌入适应的启发,𝒑1、...、𝒑𝐾的少样本类原型分别用自注意模块处理,以增强适应不同元任务的嵌入。𝒑̂𝐾𝒑̂1号的最终处理结果被视为FSOR的最终原型。具体地,自我关注模块由标准 Transformer 层实现,该 Transformer 层将输入令牌变换成查询向量、key 向量和值向量。将𝑖th令牌的查询向量和𝑗th令牌的关键向量的内积作为𝑖th令牌对𝑗th令牌的关注值。将𝑖th令牌与所有其他令牌之间的关注值作为权重,用于对所有令牌的值向量进行加权求和,得到𝑖th令牌的处理结果。每个令牌,即𝑖 = 1,...,𝐾,都用相同的过程处理。在我们的例子中,每一个正面的原型𝒑𝑖都被视为一个象征。上述过程可以从不同的角度进行。每个视角都由一个"头"来实现,它包括将标记投影到查询、键和值向量中,计算内积和加权和。在图2的实际实现中,自关注模块包括一个具有六个头的 Transformer 层。
-
图3示出了基于5路1次设定的三元训练任务的视觉样本。例如,在第一个元训练任务中,采样的类是金毛寻回犬、家雀、伊比赞猎犬、水母和狮子。这些类别是随机抽样的。在不同的元训练任务中,我们看到样本类是高度多样化的。我们可以预料,这五个积极的阶层已经是五花八门,而剩下的大量消极阶层应该更加分散。为了完整地描述正类,我们希望通过一个基于transformer的生成器来构建一个完整的正原型。少样本类原型𝒑1、...、𝒑𝐾被视为表征,通过发现表征之间的注意力,一个概括的表征被生成为 overall positive prototype 𝒑𝑂.整个 prototype generator 的细节将在第3.3节中提供。
-

-
图3。三个元训练任务的样本图像。从上到下,每一行代表一个元训练任务。每个元训练任务中有五个样本班。这里我们展示的是5路1次触发的情况。
-
-
当在元训练任务上工作时,图2中所示的几个组件的参数被更新,包括特征提取器、自我注意模块和 overall positive prototype 生成器。当处理元测试任务时,这些组件的参数被冻结。特征提取器从支持集 𝐷 𝑠𝑢𝑝𝑝𝑜𝑟𝑡 𝑚𝑒𝑡𝑎_𝑡𝑒𝑠𝑡 中提取图像的特征,然后用正原型𝒑̂ 1,𝒑̂ 2,...、𝒑̂ 𝐾和𝒑𝑂构成了 overall positive prototype 。给定一个查询𝒒,它和𝒑̂ 1,𝒑̂ 2,...、𝒑̂ 𝐾、𝒑𝑂分别进行了计算。相似性cos(𝒒,𝒑𝑂)表示查询属于已知的少样本类别之一的可能性。在评估中,如果cos(𝒒,𝒑𝑂)低于预定义的阈值𝜏,则该查询被分类为负面的,并且被视为来自看不见的类的人。如果𝒑𝑂cos(𝒒)大于𝜏,则此查询被分类为𝑖类 如果 i*= argmax𝑖cos(𝒒,𝒑̂ 𝑖),𝑖 = 1,2,...,N。
Overall positive generator
-
我们利用一个 transformer 来生成一个基于类原型𝒑1,...,𝒑𝑁的 overall positive prototype ,如图4所示。每个正面原型𝒑𝑖首先被线性层投影,并被视为一个 token。除了这些 token 之外,还会随机初始化一个任务token。然后,任务token和肯定token一起被馈送到包括五个 transformer 层的 transformer ,每个 transformer 层具有六个头,以在它们之间找到自我注意力。请注意,任务令牌的作用类似于通常在基于transformer的图像分类中使用的类令牌。这个任务令牌总结了整个任务的正面类信息,这就是我们如此命名它的原因。
-

-
图4。 overall positive prototype 发生器的框架。
-
-
经过 Transformer 编码器处理的任务令牌然后被馈送到全连接网络(FCN ),以生成最终的 overall positive prototype 𝒑𝑂.我们称之为"整体积极的",因为它总结了少样本类原型的信息。为了获得不同的𝒑𝑂's.,FCN的配置可以有不同的设计。我们将在评估部分展示不同fcn产生的𝒑𝑂's如何影响FSOR性能。在下文中,我们将所提出的方法简称为OPP( overall positive prototype )。
-
生成器的 Transformer 架构,其每层的具体计算方式如下:
- 输入层:原型投影与任务 token 初始化 ,少样本类原型 p 1 , p 2 , ... , p K p_1, p_2, \dots, p_K p1,p2,...,pK(K 为类别数)和一个随机初始化的任务 token (用于聚合全局信息)。每个原型和任务 token 通过线性层投影为固定维度的 token。形成 token 序列 task_token , p 1 , p 2 , ... , p K \\text{task\\_token}, p_1, p_2, \\dots, p_K task_token,p1,p2,...,pK,作为 Transformer 编码器的输入。
- Transformer 编码器层(每层结构) ,
- 多头自注意力(Multi-Head Self-Attention),QKV 生成 :对每个 token(包括任务 token 和类原型 token)分别线性投影为 Query (Q)、Key (K)、Value (V),即: Q = token ⋅ W Q , K = token ⋅ W K , V = token ⋅ W V Q = \text{token} \cdot W^Q, \quad K = \text{token} \cdot W^K, \quad V = \text{token} \cdot W^V Q=token⋅WQ,K=token⋅WK,V=token⋅WV 其中 W Q , W K , W V W^Q, W^K, W^V WQ,WK,WV 为可学习参数,多头设置为 6 头。缩放点积注意力 :对每个头独立计算注意力分数并加权求和: Attention ( Q , K , V ) = softmax ( Q K T d k ) V \text{Attention}(Q, K, V) = \text{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right) V Attention(Q,K,V)=softmax(dk QKT)V 其中 d k d_k dk 为键向量维度(如 64),缩放避免梯度消失。多头拼接 :6 个头的输出拼接后通过线性层整合,得到多头注意力结果 MultiHead ( Q , K , V ) \text{MultiHead}(Q, K, V) MultiHead(Q,K,V)。残差连接与层归一化 : out_attn = LayerNorm ( MultiHead ( Q , K , V ) + input_token ) \text{out\_attn} = \text{LayerNorm}(\text{MultiHead}(Q, K, V) + \text{input\_token}) out_attn=LayerNorm(MultiHead(Q,K,V)+input_token)
- 前馈神经网络(Feed-Forward Network, FCN) ,对自注意力输出逐位置应用两层 FCN: FFN ( x ) = GELU ( x ⋅ W 1 + b 1 ) ⋅ W 2 + b 2 \text{FFN}(x) = \text{GELU}(x \cdot W_1 + b_1) \cdot W_2 + b_2 FFN(x)=GELU(x⋅W1+b1)⋅W2+b2 文档未明确激活函数,默认采用 Transformer 常用的 GELU。残差连接与层归一化 : out_ffn = LayerNorm ( FFN ( out_attn ) + out_attn ) \text{out\_ffn} = \text{LayerNorm}(\text{FFN}(\text{out\_attn}) + \text{out\_attn}) out_ffn=LayerNorm(FFN(out_attn)+out_attn)。
- 任务 token 的聚合作用 ,任务 token 作为全局总结 :在每层 Transformer 中,任务 token 与类原型 token 通过自注意力交互,逐步聚合正类原型的全局信息。例如,第一层 Transformer 中,任务 token 通过注意力权重吸收所有原型的信息;后续层进一步提炼,最终输出的任务 token 成为整体正原型的基础 。跨层传递 :每层的任务 token 输出作为下一层的输入,通过 5 层堆叠,任务 token 逐渐从 "分散原型信息" 收敛为 "紧凑的整体正表示"。
- 输出层:生成整体正原型 ,最后一层 Transformer 的任务 token task_token final \text{task\token}{\text{final}} task_tokenfinal。通过全连接网络(FCN)将任务 token 映射为最终的整体正原型 P O : P O = FCN ( task_token final ) P_O:P_O = \text{FCN}(\text{task\token}{\text{final}}) PO:PO=FCN(task_tokenfinal)。
-
生成器的 Transformer 每层计算遵循:token 投影 → 多头自注意力(含任务 token)→ 残差 + 层归一化 → FCN → 残差 + 层归一化 ,最终通过任务 token 聚合生成 P O P_O PO。其核心创新是通过任务 token 和多头自注意力,将离散的少样本原型转化为全局一致的正类表示,为 FSOR 提供判别依据。
Loss function
-
为了实现少样本分类,使用预测和 GT 之间的交叉熵𝐿𝐶𝐸。此外,我们设计了一个损失函数,专门用来指导整个正向原型生成器的训练。对于一个良好的 overall positive prototype 𝒑𝑂,我们希望𝒑𝑂和正面查询𝒒之间的相似度应该大一些,而𝒑𝑂和负面查询𝒒之间的相似度应该小一些。因此,第二个损失𝐿𝑂𝑃被定义为
-
L O P = { 1 − cos ( p 0 , q ) , if q is a positive query , 1 + cos ( p 0 , q ) , if q is a negative query , L_{OP} = \begin{cases} 1 - \cos(\mathbf{p}_0, \mathbf{q}), & \text{if } q \text{ is a positive query}, \\ 1 + \cos(\mathbf{p}_0, \mathbf{q}), & \text{if } q \text{ is a negative query}, \end{cases} LOP={1−cos(p0,q),1+cos(p0,q),if q is a positive query,if q is a negative query,
-
其中cos(·)指余弦相似度。余弦值介于1和+1之间。我们希望好的𝒑𝑂和肯定查询之间的余弦相似度应该接近+1,好的𝒑𝑂和否定查询之间的余弦相似度应该接近-1。总体上,当生成的 overall positive prototype p o p_o po更接近正查询时,损失𝐿𝑂𝑃更小,或者当生成的 overall positive prototype p o p_o po远离负查询时,损失 L O P L_{OP} LOP 更小。最后,总损失是:
-
L t o t a l = L C E + α L O P L_{total}=L_{CE}+\alpha L_{OP} Ltotal=LCE+αLOP
-
其中, α \alpha α 是 generator 损失的权重,在我们的实验中设置为1。
-
Inductive vs. transductive
- 许多以前的FSOR方法是在归纳方案中训练和测试的。也就是说,模型在训练过程中只取带标签的训练数据,测试数据只有在测试时才被看到。然而,由于训练样本在少样本方案中是稀缺的,直推式学习方法最近作为一种替代方法出现,以提高少样本分类的性能。在直推式设置中,在训练过程中既可以看到标记的训练数据,也可以看到未标记的测试数据。聚类或伪标记可以被设计成发现未标记数据的数据分布,因此在直推式设置中训练的模型通常比在归纳式设置中更好地工作 。Boudiaf等人最近将他们的直推方法扩展到FSOR。他们提出了最大似然原则来降低潜在离群值的影响,并将其命名为开集似然优化(OSLO)问题。在我们的工作中,除了归纳设置,我们还可以在我们学习的原型上加入 OSLO 。我们可以在最后阶段使用来自查询图像的信息来调整原型,并探索OSLO模块如何有利于所提出的框架。
Experiments
Datasets
- 为了评估我们的FSOR方法,我们采用了该领域广泛使用的四个基准,包括MiniImageNet 、TieredImageNet 、CIFAR-FS 和FC100 数据集。MiniImageNet数据集由60,000幅大小为84 × 84的彩色图像的子集组成,这些图像取自更大的ImageNet数据集的100个不同类别。每个类别包含600张图片。TieredImageNet数据集包含608个类的集合,每个类包含779幅84 × 84像素的彩色图像,这些图像是从ImageNet数据集提取的。这些类别分为34个高级别分组,每个分组有20个低级别类别。这些低级类别被进一步分为三层,每层分别包含64、16和4个类。
- CIFAR-FS数据集由100个不同的类组成,每个类包含600幅图像。数据集分为两部分:一个由64个类组成的基本集和一个由36个类组成的新集。基本集用于训练和验证,而新集用于测试。小说中的每一个类只包含20张图片,这使得这个少样本的学习任务更具挑战性。FC100数据集旨在评估模型从少量示例中识别新类的能力。它包含100个类,每个类包含600张84 × 84像素的图像。
Implementation details
- 在我们的研究中,我们主要采用 ResNet-12 作为主干,并利用Task-adaptive negative envision for fewshot open-set recognition提供的预训练权重进行特征提取。在进行元学习时,我们采用𝐾-way 𝑁-shot设置,我们随机抽取𝐾阳性类和每个类的𝑁图像作为元训练和元测试任务的支持图像。此外,我们在评估过程中包括𝐾随机选择的负面类。对于每个正面和负面类别,我们选择15个图像作为查询集来评估模型的性能。进行元训练,每集抽样600个元训练任务。我们总共运行80集来训练图2所示的可学习组件。为了进行元测试,我们还随机抽取了600个元测试任务。通过计算模型在这600个任务上的平均性能来报告最终性能。我们通过计算最高精度来评估FSOR的性能。我们还遵循并报告了接收器工作特性曲线(AUROC)下的面积。
Performance comparison
Comparison with inductive learning-based approaches
-
我们进行了实验,并与其他最先进的(SOTA)在归纳学习设置进行比较。表1显示了我们的方法和其他归纳学习方法分别在MiniImageNet数据集和TieredImageNet数据集上的性能比较。可以看出,我们提出的方法优于现有的方法,并实现了SOTA结果。我们特别想强调我们和【TANE】之间的性能差异。即使有多个生成的否定原型,否定样本在特征空间中的广泛分散也使得否定原型的生成非常具有挑战性,因此所获得的性能在某种程度上受到限制。另一方面,我们提出的方法生成了一个 overall positive prototype,这有效地给出了性能增益。
-

-
表1 OPP和其他归纳学习方法分别在MiniImageNet数据集和TieredImageNet数据集上的性能比较。
-
-
按照TANE中提到的协议,我们还在CIFAR-FS数据集和FC100数据集上评估了建议的方法,分别如表2所示。还可以通过所提出的方法实现新的技术状态。我们注意到,5路5次样本设置的性能增益明显更大。这意味着当更多的训练样本可用时,overall positive prototype 的有效性会得到合理的提高。
-

-
表2 OPP和其他归纳学习方法分别在CIFAR-FS数据集和FC100数据集上的性能比较。
-
Comparison with transductive learning-based approaches
- 基于直推式学习的方法的出现来自于在少样本学习领域中训练数据的有限可用性。最近,Boudiaf等人介绍了他们的直推方法FSOR。在这种设置中,可以利用查询图像的特征来提高性能。他们提出了开集似然最大化来更新少样本类质心。我们也利用这种方法来更新我们的原型,并在直推设置中制作一个版本。表3显示了我们的直推式版本和其他直推式学习方法分别在MiniImageNet数据集和TieredImageNet数据集上的性能比较。实验结果再次表明,我们的方法优于现有的方法。这种性能增益是由 overall positive prototype 的有效性带来的。
- !-

5.png&pos_id=img-sIJGRNab-1789039467494) - 表3 OPP和其他直推式学习方法在MiniImageNet数据集和TieredImageNet数据集上的性能比较。
- !-
How different backbones affect performance
-
在FSOR,ResNet-12是常用的特征提取主干。这部分是因为这个小模型在数据有限的情况下表现很好,而且它能够在不同的方法之间进行公平的比较。在这一小节中,我们将研究用于特征提取的不同主干是如何影响性能的。像ResNet-50和vision transformer (ViT) 这样的较大主干是根据相应数据集中定义的训练集(与数据标签相关联的训练样本)从头开始训练的。然后它们被用于在元测试过程中提取特征,给出图2所示的起始特征。
-
表4显示了当不同的主干用于特征提取时的性能变化,分别基于MiniImageNet数据集和TieredImageNet数据集。从第一行到第三行,我们看到性能随着主干尺寸的增加而降低。该结果验证了Pushing the limits of simple pipelines for few-shot learning: External data and fine-tuning make a difference中提到的事实,即由于few-shot学习中训练数据的小规模,像ResNet-12这样的小型网络是很好的选择。像ResNet-50和ViT-small这样的较大网络不能被很好地训练,并且可能过度适应有限的训练数据量。
- !外链图片转存失败,源站可能有防盗链机制,建议将图片保存下来直接上传(https://img-home.csdnimg.cn/images/20230724024159.png?origin_url=C%3A\Users\seuic\AppData\Roaming\Typora\typora-user-images\image-
.png&pos_id=img-uYwyyouH-1789039467494) - 表4 基于MiniImageNet数据集和TieredImageNet数据集,使用不同主干进行特征提取时的性能变化。
- !外链图片转存失败,源站可能有防盗链机制,建议将图片保存下来直接上传(https://img-home.csdnimg.cn/images/20230724024159.png?origin_url=C%3A\Users\seuic\AppData\Roaming\Typora\typora-user-images\image-
-
我们也想知道在外部数据上预先训练的骨干如何使FSOR受益。为了验证这一点,我们采用了一个名为DINO 的自我监督学习框架来学习基于ImageNet数据集(不使用数据标签)的ResNet12、ResNet50和ViT-small模型。以表4的上半部分为例,行"OPP (ViT-small)"显示了基于ViT-small模型获得的性能,该ViT-small模型基于MiniImageNet数据集的训练数据以受监督的方式进行预训练。最后一行"OPP (DINO ViT-small)"显示了基于ViT-small模型获得的性能,该模型基于ImageNet数据集的训练数据以无监督的方式进行预训练。从"OPP (ResNet12)"、"OPP (ResNet 50)"和"OPP (ViTsmall)"行中,我们观察到当数据不足时,较小的模型比较大的模型表现得更好。然而,当我们应用DINO时,这一现象发生了逆转,DINO是一种利用大量未标记数据来训练模型的技术 。DINO增强了大型模型的特征提取能力。另一方面,较小的模型被过多的数据淹没,无法有效学习。在TieredImageNet数据集中也可以看到类似的结果。这些结果鼓励未来研究通过无监督学习获得的特征提取主干如何在许多方面有益于FSOR。
Ablation study
Weighting of the overall positive loss
- 在Eq(2)中。将分类损失和总体正损失与加权𝛼.相结合这里我们研究使用不同权重时的性能变化,如表5所示。一般来说,不同的权重对准确度和AUROC产生不同的影响。因此,在评估中,我们主要将𝛼设置为1,以获得平衡的性能。𝛼 = 0的设置意味着 overall positive loss 被完全忽略。我们看到它给出了最差的性能,这证明了在FSOR拥有 overall positive prototype 的有效性。
- !外链图片转存失败,源站可能有防盗链机制,建议将图片保存下来直接上传(https://img-home.csdnimg.cn/images/20230724024159.png?origin_url=C%3A\Users\seuic\AppData\Roaming\Typora\typora-user-images\image-
.png&pos_id=img-XdLznR68-1789039467494) - 表5使用不同权重组合总体正损失和分类损失时的性能变化。
- !外链图片转存失败,源站可能有防盗链机制,建议将图片保存下来直接上传(https://img-home.csdnimg.cn/images/20230724024159.png?origin_url=C%3A\Users\seuic\AppData\Roaming\Typora\typora-user-images\image-
The influence of self-attention module
- 图2所示的 prototype generator 包括一个自我关注模块。如果不经过基于transformer的自关注模块的处理,简单地用平均向量作为类原型会怎么样?表6显示了性能差异。可以看出,自我注意模块将原始平均向量改进为更有效的原型,从而提供更好的性能。表6清楚地显示了自我关注模块的价值。
- !外链图片转存失败,源站可能有防盗链机制,建议将图片保存下来直接上传(https://img-home.csdnimg.cn/images/20230724024159.png?origin_url=C%3A\Users\seuic\AppData\Roaming\Typora\typora-user-images\image-
.png&pos_id=img-1TevAtvK-1789039467494) - 表6 prototype generator 在有或没有自关注模块的情况下的性能变化。
- !外链图片转存失败,源站可能有防盗链机制,建议将图片保存下来直接上传(https://img-home.csdnimg.cn/images/20230724024159.png?origin_url=C%3A\Users\seuic\AppData\Roaming\Typora\typora-user-images\image-
Different settings of FCN
- 我们想调查 overall positive prototype 生成的不同设置如何影响最终性能。我们可以改变图4所示的FCN的配置,以产生不同的 overall positive prototypes。
- 图5和图6示出了分别基于MiniImageNet数据集上的5路1次设置和5路5次设置获得的性能变化。术语"640-1024-640"表示第一层FCN将原始的640维向量映射成1024维向量,第二层将1024维向量映射成640维向量,这是最终的 overall positive prototype 。图5和图6中的结果表明,当fcn以不同的配置设置时,仅获得轻微的性能变化。这表明所提出的方法是稳定的。在评估中,我们主要采用"640-2048-640"FCN生成的 overall positive prototype ,因为它在5路5拍设置中表现出相对更清晰的性能提升。不同的配置如何影响数据集的性能,可以在未来进行更多的研究。
-
!外链图片转存失败,源站可能有防盗链机制,建议将图片保存下来直接上传(https://img-home.csdnimg.cn/images/20230724024159.png?origin_url=C%3A\Users\seuic\AppData\Roaming\Typora\typora-user-images\image-
.png&pos_id=img-lLWdDLZL-1789039467495)
-
图5。性能差异基于不同fcn生成的 overall positive prototype ,基于MiniImageNet数据集上的5路单触发设置。
-

-
图6。性能差异基于不同fcn生成的 overall positive prototype ,基于MiniImageNet数据集上的5路5次拍摄设置。
-
Different settings of transformer encoder
- 我们评估改变图4所示的 Transformer 编码器的层数的效果。表7显示了基于MiniImageNet数据集的5路1次和5路5次设置的性能差异。随着层数从1增加到5,我们观察到性能略有提高。然而,当使用更多的层时,性能降低。我们认为太多的 Transformer 编码器层会显著增加生成器的复杂性,这可能会导致 overall positive prototype 的泛化能力的损失。根据表7,在报道的实验中采用了五个 Transformer 编码器层。
-

-
表7 当 Transformer 编码器中的不同层数用于 overall positive prototype 提取时的性能变化。
-
Discussion
-
上述实验结果证明了所提出的方法的性能优势。在这里,我们想更深入地研究为什么提出的方法更有效。为了获得这种洞察力,我们从MiniImageNet数据集抽取了600个测试用例,其中一些是正查询(来自已知类),一些是负查询(来自未知类)。我们计算正面查询和 overall positive prototype 之间的距离,并将距离分布显示为图7中的蓝色条。还计算否定查询和 overall positive prototype 之间的距离,并且距离分布在图7中显示为红色条。
- !外链图片转存失败,源站可能有防盗链机制,建议将图片保存下来直接上传(https://img-home.csdnimg.cn/images/20230724024159.png?origin_url=C%3A\Users\seuic\AppData\Roaming\Typora\typora-user-images\image-
.png&pos_id=img-CbntYrQA-1789039467495) - 图7。正总体距离和负总体距离的分布。
- !外链图片转存失败,源站可能有防盗链机制,建议将图片保存下来直接上传(https://img-home.csdnimg.cn/images/20230724024159.png?origin_url=C%3A\Users\seuic\AppData\Roaming\Typora\typora-user-images\image-
-
该直方图显示了正查询和正原型之间的距离相对更集中在较低值,而负查询和整个正原型之间的距离广泛分布。这一观察表明,肯定查询倾向于更紧密地聚集在特征空间中,而否定查询则分散在整个特征空间中。这一特点解释了为什么提出的 overall positive prototype 是有效的,以促进FSOR。
-
尽管图7显示了为什么从实验的角度来看 overall positive prototype 给出了性能增益,我们需要更多的基于数学的理论分析来显示为什么提出的 overall positive prototype有益于FSOR。但是现在还不清楚如何进行系统的理论分析。这部分可以以后再研究。
-
这项工作的另一个限制是,为了进行公平的性能比较,我们在MiniImageNet数据集、TieredImageNet数据集、CIFAR-FS数据集和FC100数据集上进行实验。然而,FSOR的本质是"开集识别",这些方法应该在真实世界、无限的场景中进行实验。虽然我们可以将FSOR方法应用于更大规模甚至无界的数据集,但很难比较不同的方法。
Visual samples
- 表8显示了一轮元测试(归纳设置,ResNet-12作为主干)的视觉样本,基于5路单触发方案,来自MiniImageNet数据集。取样的五个可见纲是线虫、帝王蟹、斑点狗、蚂蚁和黑脚雪貂。第二列显示正确识别的查询。第一个查询是正查询,即Ant,并且它被正确地识别为Ant。第二个查询是负查询,即Green Mamba,并且它被正确地识别为负类。第三列显示了被错误识别的查询。第一个查询是肯定的,即达尔马提亚,但它被错误地识别为否定的。主要原因可能是斑点狗只出现在很小的区域。第二个查询是一个否定的查询,即Malamute,但它被错误地识别为达尔马提亚。这可能是因为雪橇犬和斑点狗都是犬种,它们有相似的视觉特征。
Conclusion
-
我们提出了一种简单而有效的少样本开集识别方法。核心思想是生成一个 overall positive prototype 来描述正原型的整体信息。给定一个查询,它和整个肯定原型之间的相似性被计算为该查询是肯定类之一的可能性。我们开发了一个基于transformer的overall positive prototype 生成器,并演示了这种方法在多个数据集和多个设置上实现了新的 SOTA。我们还讨论了不同主干对特征提取的影响,并从距离分布的角度说明了为什么提出的方法更有效。在未来,可以研究 overall positive prototype 生成的不同设计,生成多个 overall positive prototype ,不同的特征相似性度量,以提高性能。
-
代码关键部分解析
pythonimport torch.utils.data as data from torchvision import transforms import os import pickle from PIL import Image class customDataset(data.Dataset): def __init__(self, args, partition='test', mode='episode', is_training=False, fix_seed=False): super(YourDataset, self).__init__() self.mode = mode self.fix_seed = fix_seed self.n_ways = args['n_ways'] self.n_open_ways = args['n_open_ways'] self.n_shots = args['n_shots'] self.n_queries = args['n_queries'] self.n_episodes = args['n_test_runs'] if partition == 'test' else args['n_train_runs'] self.n_aug_support_samples = 2 if partition == 'train' else args['n_aug_support_samples'] self.partition = partition mean = [120.39586422 / 255.0, 115.59361427 / 255.0, 104.54012653 / 255.0] std = [70.68188272 / 255.0, 68.27635443 / 255.0, 72.54505529 / 255.0] normalize = transforms.Normalize(mean=mean, std=std) if is_training: self.train_transform = transforms.Compose([ transforms.RandomCrop(84, padding=8), transforms.RandomHorizontalFlip(), transforms.RandomApply([transforms.ColorJitter(0.4, 0.4, 0.4, 0.1)], p=0.8), transforms.RandomGrayscale(p=0.8), transforms.ToTensor(), normalize ]) else: self.train_transform = transforms.Compose([ transforms.RandomCrop(84, padding=8), transforms.RandomHorizontalFlip(), transforms.ToTensor(), normalize ]) self.test_transform = transforms.Compose([transforms.ToTensor(), normalize]) self.init_episode(args['data_root'], partition) def init_episode(self, data_root, partition): # 加载数据集的逻辑 suffix = partition if partition in ['val', 'test'] else 'train_phase_train' filename = 'dataset_category_split_{}.pickle'.format(suffix) self.data = {} with open(os.path.join(data_root, filename), 'rb') as f: pack = pickle.load(f, encoding='latin1') imgs = pack['data'].astype('uint8') labels = pack['labels'] self.imgs = [Image.fromarray(x) for x in imgs] min_label = min(labels) self.labels = [x - min_label for x in labels] print('Load {} Data of {} for your_dataset in Meta-Learning Stage'.format(len(self.imgs), partition)) self.data = {} for idx in range(len(self.imgs)): if self.labels[idx] not in self.data: self.data[self.labels[idx]] = [] self.data[self.labels[idx]].append(self.imgs[idx]) self.classes = list(self.data.keys()) def get_episode(self, item): if self.fix_seed: np.random.seed(item) cls_sampled = np.random.choice(self.classes, self.n_ways, False) support_xs = [] support_ys = [] suppopen_xs = [] suppopen_ys = [] query_xs = [] query_ys = [] openset_xs = [] openset_ys = [] manyshot_xs = [] manyshot_ys = [] # Close set preparation for idx, the_cls in enumerate(cls_sampled): imgs = self.data[the_cls] support_xs_ids_sampled = np.random.choice(range(len(imgs)), self.n_shots, False) support_xs.extend([imgs[the_id] for the_id in support_xs_ids_sampled]) support_ys.extend([idx] * self.n_shots) query_xs_ids = np.setxor1d(np.arange(len(imgs)), support_xs_ids_sampled) query_xs_ids = np.random.choice(query_xs_ids, self.n_queries, False) query_xs.extend([imgs[the_id] for the_id in query_xs_ids]) query_ys.extend([idx] * self.n_queries) # Open set preparation cls_open_ids = np.setxor1d(np.arange(len(self.classes)), cls_sampled) cls_open_ids = np.random.choice(cls_open_ids, self.n_open_ways, False) for idx, the_cls in enumerate(cls_open_ids): imgs = self.data[the_cls] suppopen_xs_ids_sampled = np.random.choice(range(len(imgs)), self.n_shots, False) suppopen_xs.extend([imgs[the_id] for the_id in suppopen_xs_ids_sampled]) suppopen_ys.extend([idx] * self.n_shots) openset_xs_ids = np.setxor1d(np.arange(len(imgs)), suppopen_xs_ids_sampled) openset_xs_ids_sampled = np.random.choice(range(len(imgs)), self.n_queries, False) openset_xs.extend([imgs[the_id] for the_id in openset_xs_ids_sampled]) openset_ys.extend([the_cls] * self.n_queries) if self.partition == 'train': base_ids = np.setxor1d(np.arange(len(self.classes)), np.concatenate([cls_sampled, cls_open_ids])) assert len(set(base_ids).union(set(cls_open_ids)).union(set(cls_sampled))) == 64 base_ids = np.array(sorted(base_ids)) if self.n_aug_support_samples > 1: support_xs_aug = [support_xs[i:i + self.n_shots] * self.n_aug_support_samples for i in range(0, len(support_xs), self.n_shots)] support_ys_aug = [support_ys[i:i + self.n_shots] * self.n_aug_support_samples for i in range(0, len(support_ys), self.n_shots)] support_xs, support_ys = support_xs_aug[0], support_ys_aug[0] for next_xs, next_ys in zip(support_xs_aug[1:], support_ys_aug[1:]): support_xs.extend(next_xs) support_ys.extend(next_ys) suppopen_xs_aug = [suppopen_xs[i:i + self.n_shots] * self.n_aug_support_samples for i in range(0, len(support_xs), self.n_shots)] suppopen_ys_aug = [suppopen_ys[i:i + self.n_shots] * self.n_aug_support_samples for i in range(0, len(support_ys), self.n_shots)] suppopen_xs, suppopen_ys = suppopen_xs_aug[0], suppopen_ys_aug[0] for next_xs, next_ys in zip(suppopen_xs_aug[1:], suppopen_ys_aug[1:]): suppopen_xs.extend(next_xs) suppopen_ys.extend(next_ys) support_xs = torch.stack(list(map(lambda x: self.train_transform(x), support_xs))) suppopen_xs = torch.stack(list(map(lambda x: self.train_transform(x), suppopen_xs))) query_xs = torch.stack(list(map(lambda x: self.test_transform(x), query_xs))) openset_xs = torch.stack(list(map(lambda x: self.test_transform(x), openset_xs))) support_ys, query_ys, openset_ys = np.array(support_ys), np.array(query_ys), np.array(openset_ys) suppopen_ys = np.array(suppopen_ys) cls_sampled, cls_open_ids = np.array(cls_sampled), np.array(cls_open_ids) if self.partition == 'train': return support_xs, support_ys, query_xs, query_ys, suppopen_xs, suppopen_ys, openset_xs, openset_ys, cls_sampled, cls_open_ids, base_ids else: return support_xs, support_ys, query_xs, query_ys, suppopen_xs, suppopen_ys, openset_xs, openset_ys, cls_sampled, cls_open_ids def __getitem__(self, item): return self.get_episode(item) def __len__(self): return self.n_episodes # 创建数据集实例 train_dataset = customDataset(args, partition='train', is_training=True) train_loader = data.DataLoader(train_dataset, batch_size=1, shuffle=True)-
推理循环:进行推理并计算准确率。
pythonmodel.eval() correct = 0 total = 0 with torch.no_grad(): for i, (support_xs, support_ys, query_xs, query_ys, suppopen_xs, suppopen_ys, openset_xs, openset_ys, cls_sampled, cls_open_ids) in enumerate(test_loader): support_xs, support_ys = support_xs.cuda(), support_ys.cuda() query_xs, query_ys = query_xs.cuda(), query_ys.cuda() suppopen_xs, suppopen_ys = suppopen_xs.cuda(), suppopen_ys.cuda() openset_xs, openset_ys = openset_xs.cuda(), openset_ys.cuda() cls_sampled, cls_open_ids = cls_sampled.cuda(), cls_open_ids.cuda() the_input = [support_xs, query_xs, suppopen_xs, openset_xs] labels = [support_ys, query_ys, suppopen_ys, openset_ys] conj_ids = [cls_sampled, cls_open_ids] test_feats, cls_protos, test_cls_probs = model.open_forward(the_input, labels, conj_ids, None, test=True) query_cls_probs, openset_cls_probs = test_cls_probs _, predicted = torch.max(query_cls_probs, dim=-1) total += query_ys.size(0) * query_ys.size(1) correct += (predicted == query_ys).sum().item() accuracy = correct / total print(f'Test Accuracy: {accuracy}')-
以
OpenTiered类为例:
pythonclass OpenTiered(Dataset): # ... 其他代码 ... def get_episode(self, item): # ... 其他代码 ... # 选择闭集类别 cls_sampled = np.random.choice(self.classes, self.n_ways, False) # ... 闭集数据处理 ... # 选择开集类别 cls_open_ids = np.setxor1d(np.arange(len(self.classes)), cls_sampled) cls_open_ids = np.random.choice(cls_open_ids, self.n_open_ways, False) # 处理开集支持集数据 for idx, the_cls in enumerate(cls_open_ids): imgs = self.data[the_cls] suppopen_xs_ids_sampled = np.random.choice(range(len(imgs)), self.n_shots, False) suppopen_xs.extend([imgs[the_id] for the_id in suppopen_xs_ids_sampled]) suppopen_ys.extend([idx] * self.n_shots) # 处理开集查询集数据 openset_xs_ids = np.setxor1d(np.arange(len(imgs)), suppopen_xs_ids_sampled) openset_xs_ids_sampled = np.random.choice(range(len(imgs)), self.n_queries, False) openset_xs.extend([imgs[the_id] for the_id in openset_xs_ids_sampled]) openset_ys.extend([the_cls] * self.n_queries) # ... 其他代码 ... return support_xs, support_ys, query_xs, query_ys, suppopen_xs, suppopen_ys, openset_xs, openset_ys, cls_sampled, cls_open_ids- 在
get_episode方法中,会从所有类别中选出一部分作为闭集类别(cls_sampled),然后剩余的类别中再选出一部分作为开集类别(cls_open_ids)。接着分别处理开集的支持集数据(suppopen_xs和suppopen_ys)和查询集数据(openset_xs和openset_ys)。forward方法接收的features元组中包含了开集特征openset_feat,并且会使用self.metric计算开集数据的分类得分openset_cls_scores。
-