
💡💡💡本文核心贡献如下: 重写预测函数,支持单张图像预测和可视化(代码已开源)
预测结果如下:


💡💡💡本文核心贡献如下:

博主简介

AI小怪兽 | 计算机视觉布道者 | 视觉检测领域创新者
深耕计算机视觉与深度学习领域,专注于视觉检测前沿技术的探索与突破。长期致力于YOLO系列算法的结构性创新、性能极限优化与工业级落地实践,旨在打通从学术研究到产业应用的最后一公里。

论文:SubspaceAD: Training-Free Few-Shot Anomaly Detection via Subspace Modeling
摘要:工业检测中的视觉异常检测通常仅需每类别少量正常图像即可进行训练。近期的一些少样本方法利用基础模型的特征取得了优异的结果,但通常依赖于记忆库、辅助数据集或视觉语言模型的多模态微调。因此,我们质疑:鉴于视觉基础模型的特征表示能力,这种复杂性是否真的必要?为了回答这个问题,我们引入了 SubspaceAD,一种无需训练的方法,它通过两个简单阶段运行。首先,通过冻结的 DINOv2 骨干网络从少量正常图像中提取图像块级特征。其次,对这些特征拟合一个主成分分析模型,以估计正常变化的低维子空间。在推理时,通过计算相对于该子空间的重建残差来检测异常,从而产生可解释且具有统计基础的异常分数。尽管方法简单,SubspaceAD 在单样本和少样本设置下均实现了最先进的性能,无需训练、提示调优或记忆库。在单样本异常检测设置中,SubspaceAD 在 MVTec-AD 数据集上达到了 97.1% 的图像级 AUROC 和 97.5% 的像素级 AUROC,在 VisA 数据集上分别达到了 93.4% 和 98.2%,超越了先前的最先进结果。
检测图像中的视觉异常是计算机视觉领域一个长期存在的挑战 [7, 28]。在工业检测中,即使是与正常外观的细微偏差(如划痕、污染或缺失部件)也可能导致下游故障或安全风险。因此,开发能够自动检测此类缺陷的系统对于可靠且具有成本效益的生产至关重要。
工业异常检测的主要挑战在于数据稀缺:全样本方法每个类别需要数百张无缺陷图像来对正常性建模,这在实践中几乎不可行。另一个极端是零样本方法 [18, 40, 42],它们利用视觉语言模型和文本提示在没有任何正常样本的情况下检测异常。然而,这些方法通常难以检测那些仅靠语言难以捕捉的、细微的非语义缺陷。本文聚焦于实用且具有挑战性的少样本场景,其中仅有少量正常图像可用于定义给定物体类别的正常外观。
为了解决少样本挑战,最近的研究引入了日益复杂的深度学习技术,可分为三类。第一类是基于重建的方法 [3, 16, 34],它们学习仅复现正常样本,并将重建残差作为异常指标。第二类依赖于大型特征记忆库 [10, 11, 31],存储来自正常图像的大量图像块嵌入,并通过特征空间中的最近邻检索进行异常检测。最近,视觉语言模型方法通过提示调优 [18, 23, 25] 适配了CLIP [30] 等模型,以实现文本引导的异常检测。
尽管这三类方法在MVTec-AD [4] 和VisA [43] 等基准上取得了强劲的性能,但它们已变得日益复杂。这些方法通常需要大量的数据增强、仔细的超参数调优、多阶段训练、辅助学习目标或大容量的记忆库,这使得它们在实际工业环境中难以部署和维护。与此同时,表示学习已取得长足进步。像DINOv2这样的基础视觉模型能够产生密集且可迁移的特征,捕捉图像的语义和结构属性,即使对于它们从未训练过的领域也是如此 [6, 27, 36]。有了如此高质量的特征,人们不禁要问:我们是否仍然需要复杂的流程、大型记忆库和多阶段调优来检测异常?
本文认为答案是否定的。通过利用强大的基础特征,我们展示了一个更简单的替代方案不仅是可行的,而且更优越。具体而言,我们提出了一种基于主成分分析的纯统计方法。仅给定少量正常图像,PCA定义了一个低维子空间,用于捕捉正常外观的“主”要变化。相对于该子空间的偏差(通过重建残差量化)直接指示异常。该方法遵循了成熟的统计原理:异常表现为正常数据主PCA子空间的偏离。

这种极简方法称为 SubspaceAD,它无需训练、参数轻量且可解释。如图1所示,即使每个类别仅提供一个正常参考样本,这种简单的公式也足以定位多样化的缺陷模式。大量实验表明,SubspaceAD超越了最近提出的基于重建、记忆库和视觉语言模型的方法,这表明当特征表达足够丰富时,经典的统计建模可以再次成为视觉异常检测的强大基础。总结而言,本文提供了以下贡献:
基于重建的方法通过学习仅复现正常样本,并通过重建误差识别偏差来检测异常。早期方法依赖于自编码器或变分自编码器来重建正常外观,假设异常无法被准确恢复 [3, 34]。生成模型通过将重建与学习到的正常数据流形对齐来扩展这一思想 [1, 33]。最近的发展引入了感知损失、基于扩散的先验或特征回归策略,以避免过度泛化(即模型无意中学会了重建异常模式)的常见陷阱 [13, 16]。最新的工作之一,FastRecon [15],通过带有分布正则化的回归,从少量正常样本中学习一个变换矩阵,将特征重建为正常状态。尽管这些方法已在工业缺陷基准上取得成功,但它们需要显式训练、超参数调优以及在重建质量和异常敏感性之间仔细平衡。
异常检测的另一个主要方向涉及将正常样本的代表性图像块特征存储在记忆库中,并通过最近邻匹配来识别异常。例如,SPADE [10] 使用受k-NN启发的多分辨率特征对应,以无需训练的方式运行,适用于少样本环境下的异常检测。PatchCore [31] 是另一种无需训练的异常检测方法,通过选择一个紧凑的核心嵌入集来减少内存冗余,提高检索效率。该方法已展示出处理少样本异常检测的能力。相关方法包括估计空间位置的特征分布 [12]、使用基于流的变换进行密度建模 [41],或蒸馏预训练教师网络以压缩正常性先验 [13]。最近的方法如AnomalyDINO [11] 利用视觉基础模型的特征来提升鲁棒性和定位质量。尽管性能强劲,但基于记忆库的方法通常需要存储数千到数百万个图像块描述符,并在推理时执行最近邻搜索,这在少样本或多类别部署场景中可能变得计算繁重。
大规模基础模型 [8, 17, 21, 24, 30],包括纯视觉方法如DINO [6, 27],已显著影响了视觉表示学习。随着CLIP [30] 等大规模视觉语言模型的成功,近期工作开始探索利用文本提示进行异常检测。例如,WinCLIP [18] 是首批采用CLIP进行异常检测的工作之一。它利用手动设计的文本提示在预定义的多尺度窗口上检测异常,同时在少样本设置中构建多尺度记忆库进行特征匹配。后续方法,如AnoVL [14] 和 PromptAD [23],自动化了提示创建或学习提示适配器,而其他方法则尝试使用辅助数据集学习跨类别的通用正常性和异常性提示 [5, 42]。具体而言,PromptAD [23] 提出了语义串联来反转提示的语义,并直接优化一组可学习的上下文向量。IIPAD [25] 则直接从可用的正常实例生成提示,而不是学习特定类别的提示。这使得能够使用一个跨类别泛化的单一共享提示空间,从而在没有额外训练数据的情况下提高少样本异常检测的效率。尽管这些方法提高了灵活性,但它们遵循“每类一个提示”的范式,并且通常依赖于额外的正常/异常数据、提示调优或特定领域的文本先验。
少样本异常检测方法在如何表征正常变化方面各不相同。免训练的纯视觉方法,如DN2 [2]、SPADE [10] 和 PatchCore [32],通常存储正常图像块特征,并通过最近邻检索检测异常。需要微调的方法,包括PaDiM [12] 和 GraphCore [39],则学习特征分布的参数化模型。除了纯视觉流程,视觉-语言方法如ADP [19]、WinCLIP [18] 以及基于GPT-4V的异常推理 [40],使用文本提示或语言对齐来引导少样本检测,而零样本或少样本方法如APRIL-GAN [9] 和 AnomalyCLIP [42] 旨在无需额外训练即可跨类别泛化。批处理零样本框架 MuSc [22] 和 ACR [20] 进一步利用测试集的集合统计信息,而不是独立评估样本。总的来说,这些方法展示了一种减少监督和消除训练开销的趋势,同时保持强大的异常判别能力。我们的工作遵循这一方向,但通过基于简单的PCA子空间公式对正常变化进行建模,摆脱了对记忆库或提示调优的依赖。
提出的 SubspaceAD 方法对正常图像块特征的线性子空间进行建模,消除了对记忆库、提示调优或外部数据的需求。通过相对于该子空间的重建误差来计算异常分数,形成了一种无需训练、紧凑且可解释的方法。SubspaceAD的概述如图2所示,它在两个简单的阶段中运行。首先,通过冻结的DINOv2-G骨干网络从一小部分(k张)正常图像中提取图像块级特征。其次,对这些特征拟合一个主成分分析模型,以估计正常变化的低维子空间。在推理时,通过测试特征相对于该子空间的重建残差来检测和定位异常。

给定一组少量的无异常训练图像 I_train = {I_1, ..., I_k} 和一张测试图像 I_test,目标是定义一个异常评分函数 A,用于预测 I_test 中每个空间位置 p 的异常可能性:

SubspaceAD 的核心是从预训练视觉模型中提取的密集特征表示。采用冻结的 DINOv2-G 模型 [27] 作为特征提取模型,以获得图像块级特征。给定输入图像,模型生成一系列图像块标记,每个标记对应图像中的一个 14×14 块。
关键在于,我们不仅使用最后一个 Transformer 块的标记,而是从多个中间层聚合标记,以获得更鲁棒的表示,从而平衡高层语义与低层空间细节。这种多层融合提高了对细微异常的敏感性,同时保留了全局上下文线索,这一设计选择在消融研究中得到了支持。
令 f_l(p) ∈ R^D 表示来自 Transformer 块 l 的空间位置 p 处的图像块标记,其中 D 是模型的特征维度(例如,DINOv2-G 为 1536)。从一组层 L 中提取标记。位置 p 的最终特征向量 x_p ∈ R^D 定义为多层的平均池化表示:

对于基于 PCA 的建模,在多个中间层上平均特征特别有益。由于异常分数源自正交于主子空间的残差方差,其可靠性取决于特征分布是否能捕捉有意义的结构而非噪声。DINOv2 的中间层包含语义和结构信息的混合,而最深的层倾向于将局部细节坍缩为类别级的抽象。因此,在几个中间层上平均特征可以稳定协方差估计,减少特定层的方差,并确保主成分代表稳定的正常外观模式。
为了仅从 k 张正常图像构建具有代表性的协方差矩阵,我们应用了数据增强。对于每张 k 张正常图像,通过应用 0° 到 345° 之间的随机旋转生成 N_a=30 个增强视图(因为旋转变异在工业检测中很常见)。从所有 k×(1+N_a) 张图像中提取特征,形成所有图像块特征的集合 X_normal。这确保了估计的子空间能够捕捉常见的几何变化,并且不会因单一视图而产生偏差。该方法包括一个拟合阶段和一个应用于每张测试图像的推理阶段,如图2所示。
通过主成分分析对正常图像块特征进行建模。该模型拟合到所有正常特征 X_normal 的集合,该集合包含从 k 张原始和增强的正常图像中收集的所有图像块向量 x_p。从这个集合中,计算经验均值 μ ∈ R^D 和协方差矩阵 Σ ∈ R^{D×D}。
PCA 提供了数据主要线性子空间的封闭形式、无参数的估计。我们使用确定性 PCA 是为了简单和数值稳定性,使其非常适合必须避免过拟合的少样本场景。每个图像块特征 x ∈ R^D 被建模为:

,其中 C ∈ R^{D×r} 包含 Σ 的前 r 个特征向量,z ∈ R^r 是潜变量,ε 是各向同性噪声项。矩阵 C 构成了正常变化子空间的标准正交基。在此概率公式下,平方重建残差对应于正交于子空间的负对数似然分量,从而定义了异常分数。
保留的主成分数量 r 的选择使得解释方差超过预定义的阈值 τ:

其中 λ_i 是 Σ 的第 i 个特征值。选择这个高阈值是为了确保子空间捕捉正常变化的绝大部分,同时丢弃微小的噪声分量。得到的模型完全由均值向量 μ 和基矩阵 C 描述。
对于测试图像,按照第3.2节所述提取其对应的图像块特征图 X_test ∈ R^{h×w×D}。该图包含图像的所有图像块特征向量 x_p。
图像块级评分。每个图像块特征向量 x_p 被投影到正常子空间上:

并分配一个基于残差的异常分数:

该分数测量每个特征向量偏离正常变化主子空间的程度,产生一个低分辨率的异常图 M ∈ R^{h×w}。
图像级聚合。为了将图像块级分数聚合成单个图像级预测,我们采用了一个尾部鲁棒的统计量:经验尾部在险价值,它取异常图 M 中前 ρ% 的图像块分数的平均值。令 H_ρ(M) 表示 M 中处于或高于第 (100-ρ) 百分位数的分数集合。图像级分数 s_img 则计算为该集合的平均值:

我们设置 ρ = 1%,遵循先前的工作 [11],这平衡了对细微缺陷的敏感性与对稀疏假阳性的鲁棒性。
像素级定位。为了可视化和像素级评估,将图像块级异常图 M 双线性上采样到原始图像分辨率,并使用 σ=4 的高斯滤波器进行平滑,以抑制高频噪声,同时保持定位精度。最后,归一化的异常分数函数定义为:

令 n = k × (1 + N_a) × (h × w) 表示正常图像块特征的总数,D 为特征维度。PCA 拟合需要 O(n D²) 时间用于协方差计算和 O(D³) 用于特征分解,这在少样本场景中都是可忽略的。得到的模型仅包含 μ ∈ R^D 和 C ∈ R^{D×r},每个类别通常需要小于 1 MB 的存储空间。在单张 NVIDIA H100 GPU 上,对 672×672 图像进行推理大约需要 300 ms,其中大部分时间由 DINOv2-G 前向传播主导,而子空间投影和评分仅需约 74 ms。
我们将 SubspaceAD 的性能与最近最先进的方法在 1-shot、2-shot 和 4-shot 设置下进行了比较,并报告了图像级和像素级的结果。此外,我们还在批处理零样本设置下评估了其泛化能力,该设置下不使用任何参考图像,而是对整个未标记的测试集进行联合建模。最后,通过消融研究验证了模型的设计选择,包括基础模型主干、输入图像分辨率、层聚合策略以及 PCA 解释方差阈值 τ。
我们在两个广泛使用的工业异常检测基准上评估 SubspaceAD:MVTec-AD [4] 和 VisA [43]。两个数据集都包含多个不同的物体和纹理类别子集。MVTec-AD 包含 15 个类别,图像分辨率从 700×700 到 1024×1024 像素不等,而 VisA 包含更高分辨率的图像以及更广泛复杂的真实世界异常类型。由于异常检测被表述为一类问题,每个类别的训练集仅包含正常样本,而测试集包含正常和异常实例。测试集中的异常实例在图像级和像素级都有真实标签标注。
遵循标准做法 [11, 31],我们在图像级和像素级评估性能。图像级异常检测通过接收者操作特征曲线下面积和平均精度来衡量。像素级定位通过像素级 AUROC 和每区域重叠进行评估,后者考虑了异常区域的空间范围。
我们采用冻结的 DINOv2-G 模型 [27] 进行特征提取,特征在 22-28 层之间进行平均,如第 3.2 节所述。对于少样本拟合,随机选择 k ∈ {1,2,4} 张正常图像,并对每张图像应用 N_a = 30 次随机旋转(0° 到 345°),但对方向敏感的晶体管类别除外。为完整起见,附录 G 中提供了完全旋转无关的评估。PCA 方差阈值设置为 τ = 0.99,TVaR 聚合使用 ρ = 1% [11]。重要的是,对于 MVTec-AD 和 VisA 数据集的所有类别和 shot 设置,我们采用了统一的固定图像分辨率 672 px。每个少样本配置在 5 次独立运行(不同随机种子)上评估,报告均值和标准差。所有实验均在单张 NVIDIA H100 GPU 上执行。
我们将 SubspaceAD 与三种主流范式下的代表性少样本方法进行了比较:(1) 基于记忆库的方法,包括 SPADE [10]、PatchCore [31] 和 AnomalyDINO [11];(2) 基于重建的方法,如 FastRecon [15];(3) 基于视觉语言模型的方法,包括 WinCLIP [18]、PromptAD [23] 和 IIPAD [25]。
表 1 将 SubspaceAD 与最近的少样本方法在 MVTec-AD 和 VisA 上进行了比较。在 1-shot、2-shot 和 4-shot 设置下,SubspaceAD 在几乎所有图像级和像素级指标上持续取得了新的最先进性能。

在 MVTec-AD 上,SubspaceAD 在 1-shot 设置下达到了 97.1% 的图像级 AUROC 和 97.5% 的像素级 AUROC,优于所有先前的方法。在更具挑战性的 VisA 基准上,SubspaceAD 达到了 93.4% 的图像级 AUROC 和 93.5% 的 PRO,分别比先前的最先进方法高出 6.0% 和 1.0%。
随着参考样本数量的增加,我们的方法在几乎所有类别中保持领先。在 4-shot 设置下,SubspaceAD 在 MVTec-AD 上达到了最高的 PRO,并且在 VisA 上仍然具有很强的竞争力,PRO 达到 93.8%,仅次于 AnomalyDINO。这些结果验证了我们的核心主张:利用强大的基础特征和无参数的统计模型,可以达到并超越更复杂的最先进方法的性能。

此外,图 3 的定性结果表明,我们的方法在两个基准上都能产生更干净、更清晰、空间更精确的异常图。每个类别的详细结果包含在附录 A 中,代表性失败模式在附录 F 中讨论。
SubspaceAD 还在批处理零样本设置下进行了评估,这与基于提示的零样本或少样本范式有根本不同。遵循 AnomalyDINO [11] 和 MuSc [22] 的协议,使用该类别的整个测试集来构建模型,假设大多数图像块是无异常的。与存储跨图像所有图像块的基于记忆库的方法不同,SubspaceAD 对从该类别的未标记测试集中提取的所有图像块标记拟合单个 PCA 子空间,并基于重建残差计算异常分数。

表 2 将 SubspaceAD 的性能与其他批处理零样本方法进行了比较。SubspaceAD 在 VisA 数据集上取得了最先进的性能,图像级 AUROC 达到 94.1%,与 MuSc 的性能相当,并优于 AnomalyDINO。在 MVTec-AD 上,我们的方法达到了有竞争力的 96.6% AUROC。
SubspaceAD 在对未标记测试集进行建模的方式上与先前的批处理零样本方法不同。AnomalyDINO 从所有测试图像块构建一个记忆库,使得异常区域能够检索到其他异常作为最近邻,从而抑制了它们的异常分数。MuSc [22] 通过互相似性过滤来缓解这个问题,但它需要密集的跨图像比较,并且仍然对受污染的类别敏感。相比之下,SubspaceAD 对所有测试标记拟合单个 PCA 模型,其中主成分捕捉正常数据的共享、高方差结构,而罕见且不相关的异常重建效果差,从而获得高异常分数。这种紧凑的分布级建模在 MVTec-AD 上取得了有竞争力的性能,并在 VisA 上取得了最先进的结果,表明即使在批处理零样本模式下,一个简单的子空间也足以实现强大的异常判别。
我们分析了 SubspaceAD 中设计选择的影响,包括:(1) 输入图像分辨率,(2) 层聚合策略,(3) DINOv2 骨干网络规模,以及 (4) PCA 解释方差阈值 τ。所有实验均在 MVTec-AD 和 VisA 数据集上进行。

本文介绍了 SubspaceAD,一个用于少样本视觉异常检测的免训练框架,它利用了视觉基础模型的表示能力。通过从冻结的 DINOv2-G 编码器中提取图像块级特征,并通过简单的 PCA 子空间对正常变化进行建模,该方法通过重建残差检测异常,无需记忆库、辅助数据集、提示调优或任何形式的训练。尽管公式简单,SubspaceAD 在 MVTec-AD 和 VisA 数据集上的单样本和少样本设置中均取得了最先进的性能,表明当具有表达力的特征表示可用时,复杂的架构和多阶段优化并非必需。
# 1. Create environment
conda create -n subspacead python=3.10
conda activate subspacead
# 2. Install dependencies and the package
pip install -r requirements.txt
pip install -e .
https://www.mvtec.com/research-teaching/datasets/mvtec-ad
放置路径:
datasets/mvtec-ad/
python main.py \
--dataset_name mvtec_ad \
--dataset_path datasets/mvtec-ad \
--categories bottle \
--model_ckpt facebook/dinov2-with-registers-giant \
--image_res 672 \
--k_shot 5 \
--aug_count 30 \
--pca_ev 0.99 \
--outdir results/debug_run

python predict_all_image.py --normal_dir ./datasets/mvtec_ad/bottle/train/good --query_image ./datasets/mvtec_ad/bottle/test/broken_large/ --model_ckpt facebook/dinov2-with-registers-giant --image_res 320 --patch_size 224 --patch_overlap 0.2 --pca_ev 0.99 --outdir ./prediction_output
显存足够的话 可以增加image_res,完整预测代码可以找我
#!/usr/bin/env python3
"""
predict_single_image.py
Standalone single-image anomaly detection script for SubspaceAD.
Given:
- a folder (or list) of normal reference images, and
- one query image path,
fit a PCA subspace on the normal DINOv2 patch features and predict:
- an image-level anomaly score
- a pixel-level anomaly heatmap
Usage example:
python predict_single_image.py \
--normal_dir ./datasets/mvtec-ad/bottle/train/good \
--query_image ./datasets/mvtec-ad/bottle/test/broken_large/000.png \
--model_ckpt facebook/dinov2-with-registers-giant \
--image_res 672 \
--pca_ev 0.99 \
--outdir ./prediction_output
"""
import argparse
import logging
import math
import os
import sys
from pathlib import Path
from types import SimpleNamespace
import cv2
import numpy as np
import torch
from PIL import Image
# ---------------------------------------------------------------------------
# Imports: try installed package first, otherwise use local src/ folder
# ---------------------------------------------------------------------------
try:
from subspacead.config import parse_layer_indices, parse_grouped_layers
from subspacead.core.extractor import FeatureExtractor
from subspacead.core.pca import PCAModel
from subspacead.core.patching import process_image_patched, get_patch_coords
from subspacead.post_process.scoring import (
aggregate_image_score,
calculate_anomaly_scores,
post_process_map,
)
except ImportError:
repo_root = Path(__file__).resolve().parent
src_path = str(repo_root / "src")
if src_path not in sys.path:
sys.path.insert(0, src_path)
from subspacead.config import parse_layer_indices, parse_grouped_layers
from subspacead.core.extractor import FeatureExtractor
from subspacead.core.pca import PCAModel
from subspacead.core.patching import process_image_patched, get_patch_coords
from subspacead.post_process.scoring import (
aggregate_image_score,
calculate_anomaly_scores,
post_process_map,
)
DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s - %(levelname)s - %(message)s",
)
def _collect_image_paths(path_source, extensions=(".png", ".jpg", ".jpeg", ".bmp", ".tiff")):
"""Return a list of image paths from a directory or a single file path."""
path_source = Path(path_source)
if path_source.is_file():
return [str(path_source)]
if not path_source.is_dir():
raise ValueError(f"Path must be an existing directory or file: {path_source}")
paths = sorted(
p for p in path_source.iterdir()
if p.is_file() and p.suffix.lower() in extensions
)
if not paths:
raise ValueError(f"No images found in {path_source} with extensions {extensions}")
return [str(p) for p in paths]
def _build_config(args):
"""Convert argparse Namespace to the attribute-style config used by core code."""
return SimpleNamespace(
model_ckpt=args.model_ckpt,
image_res=args.image_res,
layers=args.layers,
agg_method=args.agg_method,
grouped_layers=args.grouped_layers,
docrop=args.docrop,
use_clahe=args.use_clahe,
score_method=args.score_method,
drop_k=args.drop_k,
img_score_agg=args.img_score_agg,
bg_mask_method=args.bg_mask_method,
mask_threshold_method=args.mask_threshold_method,
percentile_threshold=args.percentile_threshold,
dino_saliency_layer=args.dino_saliency_layer,
patch_size=args.patch_size,
patch_overlap=args.patch_overlap,
batch_size=args.batch_size,
)
def _get_mask_background(saliency_map, cfg):
"""Return a boolean background mask from a DINO saliency map."""
background_mask = np.zeros_like(saliency_map, dtype=bool)
try:
if cfg.mask_threshold_method == "percentile":
threshold = np.percentile(saliency_map, cfg.percentile_threshold * 100)
background_mask = saliency_map < threshold
else: # otsu
norm_mask = cv2.normalize(
saliency_map, None, 0, 255, cv2.NORM_MINMAX, dtype=cv2.CV_8U
)
_, binary_mask = cv2.threshold(
norm_mask, 0, 255, cv2.THRESH_BINARY + cv2.THRESH_OTSU
)
background_mask = binary_mask == 0
except Exception as e:
logging.warning(f"Background masking failed: {e}. Returning empty mask.")
return background_mask
def fit_pca_model(normal_paths, extractor, cfg, pca_ev=0.99, pca_dim=None, whiten=False):
"""
Fit a PCAModel on the normal reference images.
Returns:
pca_params (dict): parameters required by calculate_anomaly_scores
feature_dim (int): feature dimensionality
h_p (int), w_p (int): token grid size
"""
logging.info(f"Fitting PCA on {len(normal_paths)} normal reference images...")
# Determine feature dimension and token grid from one reference image
temp_img = Image.open(normal_paths[0]).convert("RGB")
layers = parse_layer_indices(cfg.layers)
grouped_layers = parse_grouped_layers(cfg.grouped_layers) if cfg.agg_method == "group" else []
temp_tokens, (h_p, w_p), _ = extractor.extract_tokens(
[temp_img],
cfg.image_res,
layers,
cfg.agg_method,
grouped_layers,
cfg.docrop,
use_clahe=cfg.use_clahe,
dino_saliency_layer=cfg.dino_saliency_layer,
)
feature_dim = temp_tokens.shape[-1]
if cfg.patch_size:
# ---------------- Patch-based fitting ----------------
temp_patch = temp_img.crop((0, 0, cfg.patch_size, cfg.patch_size))
temp_tokens, (h_p, w_p), _ = extractor.extract_tokens(
[temp_patch],
cfg.image_res,
layers,
cfg.agg_method,
grouped_layers,
cfg.docrop,
use_clahe=cfg.use_clahe,
dino_saliency_layer=cfg.dino_saliency_layer,
)
feature_dim = temp_tokens.shape[-1]
tokens_per_patch = h_p * w_p
total_patches = 0
num_batches = 0
for path in normal_paths:
img = Image.open(path).convert("RGB")
patch_coords = get_patch_coords(img.height, img.width, cfg.patch_size, cfg.patch_overlap)
total_patches += len(patch_coords)
num_batches += math.ceil(len(patch_coords) / cfg.batch_size)
total_tokens = total_patches * tokens_per_patch
def feature_generator():
for path in normal_paths:
pil_img = Image.open(path).convert("RGB")
patch_coords = get_patch_coords(
pil_img.height, pil_img.width, cfg.patch_size, cfg.patch_overlap
)
for i in range(0, len(patch_coords), cfg.batch_size):
coord_batch = patch_coords[i : i + cfg.batch_size]
patch_batch = [pil_img.crop(c) for c in coord_batch]
tokens_batch, _, saliency_masks_batch = extractor.extract_tokens(
patch_batch,
cfg.image_res,
layers,
cfg.agg_method,
grouped_layers,
cfg.docrop,
use_clahe=cfg.use_clahe,
dino_saliency_layer=cfg.dino_saliency_layer,
)
tokens_flat = tokens_batch.reshape(-1, feature_dim)
if cfg.bg_mask_method == "dino_saliency":
masks_flat = saliency_masks_batch.reshape(-1)
foreground_tokens = tokens_flat[_get_mask_background(masks_flat, cfg) == False]
yield foreground_tokens if foreground_tokens.shape[0] > 0 else tokens_flat
else:
yield tokens_flat
else:
# ---------------- Full-image fitting ----------------
total_train_images = len(normal_paths)
total_tokens = total_train_images * h_p * w_p
num_batches = math.ceil(total_train_images / cfg.batch_size)
def feature_generator():
for i in range(0, len(normal_paths), cfg.batch_size):
path_batch = normal_paths[i : i + cfg.batch_size]
img_batch = [Image.open(p).convert("RGB") for p in path_batch]
tokens_batch, _, saliency_masks_batch = extractor.extract_tokens(
img_batch,
cfg.image_res,
layers,
cfg.agg_method,
grouped_layers,
cfg.docrop,
use_clahe=cfg.use_clahe,
dino_saliency_layer=cfg.dino_saliency_layer,
)
tokens_flat = tokens_batch.reshape(-1, feature_dim)
if cfg.bg_mask_method == "dino_saliency":
masks_flat = saliency_masks_batch.reshape(-1)
foreground_tokens = tokens_flat[_get_mask_background(masks_flat, cfg) == False]
yield foreground_tokens if foreground_tokens.shape[0] > 0 else tokens_flat
else:
yield tokens_flat
pca_model = PCAModel(k=pca_dim, ev=pca_ev, whiten=whiten)
pca_params = pca_model.fit(feature_generator, feature_dim, total_tokens, num_batches)
logging.info(f"PCA fitted: kept k={pca_params['k']} components.")
return pca_params, feature_dim, h_p, w_p
def save_visualization(pil_img, anomaly_map, out_path, colormap=cv2.COLORMAP_JET):
"""Save original image, anomaly heatmap, and an overlay."""
out_path = Path(out_path)
out_path.parent.mkdir(parents=True, exist_ok=True)
img_np = np.array(pil_img)
anomaly_norm = (anomaly_map - anomaly_map.min()) / (anomaly_map.max() - anomaly_map.min() + 1e-8)
anomaly_u8 = (anomaly_norm * 255).astype(np.uint8)
heatmap = cv2.applyColorMap(anomaly_u8, colormap)
heatmap = cv2.cvtColor(heatmap, cv2.COLOR_BGR2RGB)
# Resize heatmap to match image if needed
if img_np.shape[:2] != heatmap.shape[:2]:
heatmap = cv2.resize(heatmap, (img_np.shape[1], img_np.shape[0]), interpolation=cv2.INTER_LINEAR)
alpha = 0.5
overlay = (img_np.astype(np.float32) * (1 - alpha) + heatmap.astype(np.float32) * alpha).astype(np.uint8)
# Stack horizontally
vis = np.concatenate([img_np, heatmap, overlay], axis=1)
Image.fromarray(vis).save(out_path)
logging.info(f"Visualization saved to {out_path}")
def main():
parser = argparse.ArgumentParser(description="SubspaceAD single-image prediction")
# Data
parser.add_argument("--normal_dir", type=str, required=True,
help="Directory (or single image path) containing normal reference images.")
parser.add_argument("--query_image", type=str, required=True,
help="Path to a query image or a directory containing query images.")
parser.add_argument("--outdir", type=str, default="./prediction_output",
help="Directory where the output heatmap will be saved.")
parser.add_argument("--save_heatmap", action="store_true", default=True,
help="Save a visualization of the anomaly heatmap.")
# Model
parser.add_argument("--model_ckpt", type=str,
default="facebook/dinov2-with-registers-large",
help="HuggingFace DINOv2 checkpoint.")
parser.add_argument("--image_res", type=int, default=256,
help="Input resolution for the model.")
parser.add_argument("--batch_size", type=int, default=4,
help="Batch size for feature extraction.")
# Feature aggregation
parser.add_argument("--layers", type=str, default="-12,-13,-14,-15,-16,-17,-18",
help="Comma-separated layer indices to aggregate.")
parser.add_argument("--agg_method", type=str, default="mean",
choices=["concat", "mean", "group"],
help="How to aggregate selected layers.")
parser.add_argument("--grouped_layers", type=str, default=None,
help="Layer groups for 'group' agg. Format: '-1,-2:-3,-4'.")
parser.add_argument("--docrop", action="store_true",
help="Apply center crop during preprocessing.")
parser.add_argument("--use_clahe", action="store_true",
help="Apply CLAHE contrast enhancement.")
# PCA
parser.add_argument("--pca_ev", type=float, default=0.99,
help="Explained-variance ratio for PCA. Ignored if --pca_dim is set.")
parser.add_argument("--pca_dim", type=int, default=None,
help="Fixed number of PCA components.")
parser.add_argument("--whiten", action="store_true",
help="Whitening in PCA.")
# Scoring
parser.add_argument("--score_method", type=str, default="reconstruction",
choices=["reconstruction", "mahalanobis", "cosine", "euclidean"],
help="Anomaly scoring method.")
parser.add_argument("--drop_k", type=int, default=0,
help="Drop first k principal components during scoring.")
parser.add_argument("--img_score_agg", type=str, default="mtop1p",
choices=["max", "mean", "p99", "mtop5", "mtop1p"],
help="Aggregation method for the image-level score.")
# Background masking
parser.add_argument("--bg_mask_method", type=str, default=None,
choices=[None, "dino_saliency"],
help="Optional background removal. 'pca_normality' is omitted for simplicity.")
parser.add_argument("--mask_threshold_method", type=str, default="percentile",
choices=["percentile", "otsu"],
help="Binarization method for the saliency mask.")
parser.add_argument("--percentile_threshold", type=float, default=0.15,
help="Percentile threshold when using percentile masking.")
parser.add_argument("--dino_saliency_layer", type=int, default=6,
help="Transformer layer index for DINO saliency mask.")
# Patching (optional)
parser.add_argument("--patch_size", type=int, default=None,
help="If set, process images in overlapping patches of this size.")
parser.add_argument("--patch_overlap", type=float, default=0.0,
help="Overlap ratio between patches (0.0-1.0).")
args = parser.parse_args()
cfg = _build_config(args)
logging.info(f"Using device: {DEVICE}")
# Load data
normal_paths = _collect_image_paths(args.normal_dir)
logging.info(f"Found {len(normal_paths)} normal reference images.")
query_paths = _collect_image_paths(args.query_image)
logging.info(f"Found {len(query_paths)} query image(s).")
# Build feature extractor
extractor = FeatureExtractor(args.model_ckpt)
# Fit PCA model
pca_params, feature_dim, h_p, w_p = fit_pca_model(
normal_paths,
extractor,
cfg,
pca_ev=args.pca_ev,
pca_dim=args.pca_dim,
whiten=args.whiten,
)
# Predict all query images
results = []
for query_path in query_paths:
logging.info(f"Predicting anomaly for {query_path}...")
anomaly_map, img_score, pil_img = predict_single_image(
query_path, extractor, pca_params, cfg, feature_dim, h_p, w_p
)
print(f"\n[{Path(query_path).name}] Image-level anomaly score ({args.img_score_agg}): {img_score:.6f}")
results.append((Path(query_path).name, img_score))
query_stem = Path(query_path).stem
# Save heatmap visualization
if args.save_heatmap:
out_path = Path(args.outdir) / f"{query_stem}_anomaly.png"
save_visualization(pil_img, anomaly_map, out_path)
# Save raw anomaly map as .npy
raw_out_path = Path(args.outdir) / f"{query_stem}_anomaly_map.npy"
np.save(raw_out_path, anomaly_map)
logging.info(f"Raw anomaly map saved to {raw_out_path}")
if len(results) > 1:
print("\n" + "=" * 60)
print("Summary of image-level anomaly scores:")
for name, score in results:
print(f" {name:<40s} {score:.6f}")
print("=" * 60 + "\n")
if __name__ == "__main__":
main()
异常得分如下:

预测结果如下:


原创声明:本文系作者授权腾讯云开发者社区发表,未经许可,不得转载。
如有侵权,请联系 cloudcommunity@tencent.com 删除。
原创声明:本文系作者授权腾讯云开发者社区发表,未经许可,不得转载。
如有侵权,请联系 cloudcommunity@tencent.com 删除。