Anomalib 中的 CFA 模型:耦合超球面特征适配实现目标导向异常定位实战指南
2026/9/17 1:59:48 网站建设 项目流程

Anomalib 中的 CFA 模型:耦合超球面特征适配实现目标导向异常定位实战指南

【免费下载链接】anomalibAn anomaly detection library comprising state-of-the-art algorithms and features such as experiment management, hyper-parameter optimization, and edge inference.项目地址: https://gitcode.com/GitHub_Trending/an/anomalib

CFA(Coupled-hypersphere-based Feature Adaptation,耦合超球面特征适配)是一种面向目标数据集做特征适配的异常分割模型,在 Anomalib 中以anomalib.models.Cfa提供完整实现。本文基于 模型文档 展开,结合仓库源码深入讲解其核心原理、配置参数、训练/推理流程与 MVTec AD 基准表现,帮助你快速上手并在自己的数据集上复现与调优。

模型概览与核心原理

CFA 由 Lee、Lee 与 Song 于 2022 年提出(论文 arXiv:2206.04325),模型类型为Segmentation(分割)。与直接使用通用预训练特征的方案不同,CFA 的核心思想是让特征适应目标数据集本身,从而提升异常的可区分度。其整体由两大组件构成(详见 架构图):

  1. 可学习的 Patch Descriptor(补丁描述器):从预训练 CNN 提取的多尺度特征中学习并嵌入目标导向特征;
  2. 与目标数据集规模无关的可扩展 Memory Bank(记忆库):存储正常样本的典型特征,规模不随数据集增大而线性膨胀。

配合预训练 CNN,CFA 采用迁移学习增大正常特征密度,使异常特征更易被区分。训练时,正常特征被约束在耦合超球面内(吸引力损失L_att),异常特征被推离超球面(排斥力损失L_rep);测试时则通过最近邻搜索计算每个 patch 的异常分数,生成像素级定位结果。

仓库实现结构

CFA 在 Anomalib 中的实现位于 src/anomalib/models/image/cfa/,由以下模块组成:

  • torch_model.py:底层 PyTorch 模型CfaModel,包含特征提取器、描述器、记忆库与距离计算;
  • lightning_model.py:PyTorch Lightning 模块Cfa,负责训练/验证流程、优化器与预处理器配置;
  • loss.py:CfaLoss损失函数,实现吸引力与排斥力双项损失;
  • anomaly_map.py:AnomalyMapGenerator,将距离张量转换为平滑的异常热力图。

特征提取与多尺度融合

CfaModel使用torchvisioncreate_feature_extractor从预训练骨干网络提取特征。根据 get_return_nodes,各骨干网络返回以下层:

骨干网络返回节点
resnet18 / wide_resnet50_2layer1,layer2,layer3
vgg19_bnfeatures.25,features.38,features.52
efficientnet_b5未实现(触发NotImplementedError

Descriptor网络首先对每层特征做avg_pool2d池化,再将不同层特征通过双线性插值对齐到同一分辨率后沿通道拼接,最后经过一个CoordConv2d1×1 卷积层完成降维嵌入。CoordConv 额外注入归一化的 x/y 坐标通道(可选径向通道r),使描述器能感知 patch 的空间位置,这一点在细粒度异常定位中尤为关键。

记忆库初始化与压缩

记忆库在训练开始前通过 initialize_centroid 初始化:遍历训练集(仅正常样本)提取目标导向特征并计算均值作为初始质心,随后按gamma_c参数决定是否用 K-Means 压缩:

  • gamma_c = 1:不压缩,保留全部特征;
  • gamma_c > 1:以scale[0]*scale[1] // gamma_c为聚类数执行 K-Means,将记忆库压缩为聚类中心。

该机制保证了记忆库规模与数据集大小解耦,正是"可扩展记忆库"的源码级体现。若在调用forward时记忆库尚未初始化(维度为 0),模型会抛出ValueError提示先运行initialize_centroid

训练流程

lightning_model.py 中的Cfa类将CfaModelCfaLoss组装为 Lightning 模块:

  • on_train_start:调用initialize_centroid初始化记忆库质心;
  • training_step:前向得到距离张量后计算CfaLoss
  • backward:由于计算图需求使用loss.backward(retain_graph=True)
  • configure_optimizers:使用 AdamW 优化器,学习率1e-3、权重衰减5e-4amsgrad=True
  • trainer_argumentsgradient_clip_val=0(禁用梯度裁剪)、num_sanity_val_steps=0(跳过验证 sanity 检查);
  • learning_type:返回LearningType.ONE_CLASS,表明这是一类分类(one-class)任务。

损失函数与异常图生成

CfaLoss 由两项组成(最终乘以 1000 放大):

  • 吸引力损失l_att:取前num_nearest_neighbors个最近邻距离与radius²比较,超出半径的部分被惩罚,将正常特征拉入超球面;
  • 排斥力损失l_rep:取后num_hard_negative_features个"硬负样本"距离,小于radius² - 0.1的部分被惩罚,将困难异常特征推出超球面。

AnomalyMapGenerator 在推理阶段对距离张量开方后取最近邻距离,用softmin加权得到 patch 分数,重排为特征图尺度后上采样回原图尺寸,最后用GaussianBlur2d(sigma=4)平滑生成热力图;图像级分数取异常图最大值(torch.amax)。

快速开始:训练与推理

命令行方式

在 Anomalib 中训练 CFA 最直接的方式是 CLI:

anomalib train --model Cfa --data MVTecAD --data.category <category>

其中<category>为 MVTec AD 数据集的类别名(如bottlecablescrew等)。该命令会自动完成数据下载、预处理、训练与评估。

Python API 方式

lightning_model.py 给出了等价的 Python 写法:

from anomalib.data import MVTecAD from anomalib.models import Cfa from anomalib.engine import Engine # 初始化模型与数据 datamodule = MVTecAD() model = Cfa() # 使用 Engine 训练 engine = Engine() engine.fit(model=model, datamodule=datamodule) # 获取预测结果 predictions = engine.predict(model=model, datamodule=datamodule) # 按论文设置配置预处理器:先缩放到 256x256,再做 224x224 中心裁剪 pre_processor = Cfa.configure_pre_processor( image_size=(256, 256), center_crop_size=(224, 224) )

注意:CLI 方式执行的是文档给出的标准训练命令;configure_pre_processor用于在自定义流程中复现论文的预处理设置。

预处理配置细节

Cfa.configure_pre_processor封装了论文中的预处理方式:

  • 默认将图像Resize256x256antialias=True);
  • 若指定center_crop_size,则额外执行CenterCrop(如 224×224),并在裁剪尺寸超过图像尺寸时抛出ValueError
  • 随后使用 ImageNet 统计值mean=[0.485, 0.456, 0.406]std=[0.229, 0.224, 0.225]归一化。

配置参数详解

官方配置示例位于 examples/configs/model/cfa.yaml,可直接作为训练配置文件的模板:

model: class_path: anomalib.models.Cfa init_args: backbone: wide_resnet50_2 gamma_c: 1 gamma_d: 1 num_nearest_neighbors: 3 num_hard_negative_features: 3 radius: 1.0e-05 trainer: max_epochs: 30 callbacks: - class_path: lightning.pytorch.callbacks.EarlyStopping init_args: patience: 5 monitor: pixel_AUROC mode: max

各参数含义如下(默认值与 lightning_model.py 中的构造函数一致):

参数默认值说明
backbonewide_resnet50_2骨干网络,可选resnet18wide_resnet50_2vgg19_bnefficientnet_b5未实现)
gamma_c1质心(记忆库)损失权重参数;大于 1 时启用 K-Means 压缩记忆库
gamma_d1距离损失权重参数,同时决定描述器输出通道数(dim // gamma_d
num_nearest_neighbors3异常分数计算与吸引力损失所用的最近邻数量
num_hard_negative_features3排斥力损失使用的硬负样本特征数量
radius1e-5超球面决策边界初始半径(可学习参数,torch.ones(1, requires_grad=True) * radius

配置文件的 trainer 部分还演示了早停回调:以pixel_AUROC为监控指标、mode: max最大化、patience: 5,这与 README 中"使用早停(patience=5)产出基准数据"的说明一致。

参数调优建议

  • radius:源码中半径是可学习参数,但初始值影响收敛起点,若训练初期损失异常可尝试调整;
  • num_nearest_neighbors / num_hard_negative_features:二者之和决定了损失中参与 top-k 采样的距离数量,值过小可能导致负样本挖掘不足;
  • gamma_c:在显存受限或数据集较大时,调大gamma_c启用 K-Means 压缩,可显著减小记忆库规模。

MVTec AD 基准表现

README 报告了 seed=0 下、使用早停(patience=5)在 MVTec AD 15 个类别上的完整结果,涵盖图像级 AUC、图像 F1、像素级 AUC、像素级 AUPRO 与像素 F1 五类指标,分别评测ResNet-18Wide ResNet50两种骨干。

图像级指标(Image-Level AUC / F1)

类别AUC (ResNet-18)AUC (WRN50)F1 (ResNet-18)F1 (WRN50)
Bottle0.9910.9980.9830.984
Cable0.9470.9790.9070.962
Capsule0.8580.8720.9380.946
Carpet0.9530.9780.9560.961
Grid0.9470.9610.9460.957
Hazelnut0.9951.0000.9961.000
Leather0.9990.9900.9950.973
Metal_nut0.9320.9950.9580.984
Pill0.8870.9460.9200.952
Screw0.6250.7030.8580.855
Tile1.0000.9991.0000.994
Toothbrush0.9941.0000.9841.000
Transistor0.8950.9570.7950.907
Wood1.0000.9941.0000.983
Zipper0.9190.9670.9490.975
Average0.9300.9560.9460.962

像素级指标(Pixel AUC / AUPRO / F1)

类别AUC (R18)AUC (WRN50)AUPRO (R18)AUPRO (WRN50)F1 (R18)F1 (WRN50)
Bottle0.9860.9890.9400.9470.7510.789
Cable0.9840.9880.9020.9400.6610.674
Capsule0.9870.9890.9460.9390.5070.500
Carpet0.9700.9800.9100.9190.5490.578
Grid0.9730.9540.9110.8620.3160.280
Hazelnut0.9870.9850.9310.9300.5980.561
Leather0.9920.9890.9740.9550.4610.378
Metal_nut0.9810.9920.9120.9310.8190.874
Pill0.9810.9880.9350.9470.6890.679
Screw0.9730.9790.8840.9060.2120.301
Tile0.9780.9850.8920.9060.7400.768
Toothbrush0.9900.9910.8950.8990.6090.627
Transistor0.9640.9770.8950.9300.5700.666
Wood0.9640.9740.8980.8930.5640.627
Zipper0.9780.9900.9250.9580.5610.668
Average0.9790.9830.9170.9240.5740.598

结果解读

  • Wide ResNet50 整体更优:图像级 AUC 平均 0.956(vs 0.930)、像素级 AUC 平均 0.983(vs 0.979)、像素级 AUPRO 平均 0.924(vs 0.917),说明更强的骨干能带来一致的定位精度收益;
  • 类别差异显著Screw的图像级 AUC 最低(0.625/0.703),而TileWoodHazelnutToothbrush等类别接近或达到满分,这与目标纹理复杂度、缺陷形态有关;
  • 早停影响:README 明确指出所有数字均在早停(patience=5)下产生,增大 patience 可能获得更高指标,复现或调优时可据此调整。

README 同时提供了三张样例可视化(结果图 1、结果图 2、结果图 3),每张图按列展示原始图像、真实掩码、预测热力图、预测掩码与分割叠加结果,直观验证了 CFA 的像素级定位能力。

引用

复现或使用 CFA 时,可引用论文:

@article{lee2022cfa, title={CFA: Coupled-hypersphere-based Feature Adaptation for Target-Oriented Anomaly Localization}, author={Lee, Sungwook and Lee, Seunghyun and Song, Byung Cheol}, journal={arXiv preprint arXiv:2206.04325}, year={2022} }

原始参考实现可查阅 sungwool/cfa_for_anomaly_localization(外部仓库,供对照研究使用)。

小结

本文从 Anomalib 的 CFA 实现出发,系统梳理了耦合超球面特征适配的原理、CfaModel/Cfa/CfaLoss/AnomalyMapGenerator的源码级分工、CLI 与 Python 两种训练方式、全部超参数含义以及 MVTec AD 上的五类基准指标。如果你想快速验证效果,直接运行anomalib train --model Cfa --data MVTecAD --data.category <category>;若要追求更高指标,可从增大早停 patience、更换wide_resnet50_2骨干或调整gamma_c/radius入手。

【免费下载链接】anomalibAn anomaly detection library comprising state-of-the-art algorithms and features such as experiment management, hyper-parameter optimization, and edge inference.项目地址: https://gitcode.com/GitHub_Trending/an/anomalib

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询