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 的核心思想是让特征适应目标数据集本身,从而提升异常的可区分度。其整体由两大组件构成(详见 架构图):
- 可学习的 Patch Descriptor(补丁描述器):从预训练 CNN 提取的多尺度特征中学习并嵌入目标导向特征;
- 与目标数据集规模无关的可扩展 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使用torchvision的create_feature_extractor从预训练骨干网络提取特征。根据 get_return_nodes,各骨干网络返回以下层:
| 骨干网络 | 返回节点 |
|---|---|
| resnet18 / wide_resnet50_2 | layer1,layer2,layer3 |
| vgg19_bn | features.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类将CfaModel与CfaLoss组装为 Lightning 模块:
- on_train_start:调用
initialize_centroid初始化记忆库质心; - training_step:前向得到距离张量后计算
CfaLoss; - backward:由于计算图需求使用
loss.backward(retain_graph=True); - configure_optimizers:使用 AdamW 优化器,学习率
1e-3、权重衰减5e-4、amsgrad=True; - trainer_arguments:
gradient_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 数据集的类别名(如bottle、cable、screw等)。该命令会自动完成数据下载、预处理、训练与评估。
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封装了论文中的预处理方式:
- 默认将图像
Resize至256x256(antialias=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 中的构造函数一致):
| 参数 | 默认值 | 说明 |
|---|---|---|
backbone | wide_resnet50_2 | 骨干网络,可选resnet18、wide_resnet50_2、vgg19_bn(efficientnet_b5未实现) |
gamma_c | 1 | 质心(记忆库)损失权重参数;大于 1 时启用 K-Means 压缩记忆库 |
gamma_d | 1 | 距离损失权重参数,同时决定描述器输出通道数(dim // gamma_d) |
num_nearest_neighbors | 3 | 异常分数计算与吸引力损失所用的最近邻数量 |
num_hard_negative_features | 3 | 排斥力损失使用的硬负样本特征数量 |
radius | 1e-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-18与Wide ResNet50两种骨干。
图像级指标(Image-Level AUC / F1)
| 类别 | AUC (ResNet-18) | AUC (WRN50) | F1 (ResNet-18) | F1 (WRN50) |
|---|---|---|---|---|
| Bottle | 0.991 | 0.998 | 0.983 | 0.984 |
| Cable | 0.947 | 0.979 | 0.907 | 0.962 |
| Capsule | 0.858 | 0.872 | 0.938 | 0.946 |
| Carpet | 0.953 | 0.978 | 0.956 | 0.961 |
| Grid | 0.947 | 0.961 | 0.946 | 0.957 |
| Hazelnut | 0.995 | 1.000 | 0.996 | 1.000 |
| Leather | 0.999 | 0.990 | 0.995 | 0.973 |
| Metal_nut | 0.932 | 0.995 | 0.958 | 0.984 |
| Pill | 0.887 | 0.946 | 0.920 | 0.952 |
| Screw | 0.625 | 0.703 | 0.858 | 0.855 |
| Tile | 1.000 | 0.999 | 1.000 | 0.994 |
| Toothbrush | 0.994 | 1.000 | 0.984 | 1.000 |
| Transistor | 0.895 | 0.957 | 0.795 | 0.907 |
| Wood | 1.000 | 0.994 | 1.000 | 0.983 |
| Zipper | 0.919 | 0.967 | 0.949 | 0.975 |
| Average | 0.930 | 0.956 | 0.946 | 0.962 |
像素级指标(Pixel AUC / AUPRO / F1)
| 类别 | AUC (R18) | AUC (WRN50) | AUPRO (R18) | AUPRO (WRN50) | F1 (R18) | F1 (WRN50) |
|---|---|---|---|---|---|---|
| Bottle | 0.986 | 0.989 | 0.940 | 0.947 | 0.751 | 0.789 |
| Cable | 0.984 | 0.988 | 0.902 | 0.940 | 0.661 | 0.674 |
| Capsule | 0.987 | 0.989 | 0.946 | 0.939 | 0.507 | 0.500 |
| Carpet | 0.970 | 0.980 | 0.910 | 0.919 | 0.549 | 0.578 |
| Grid | 0.973 | 0.954 | 0.911 | 0.862 | 0.316 | 0.280 |
| Hazelnut | 0.987 | 0.985 | 0.931 | 0.930 | 0.598 | 0.561 |
| Leather | 0.992 | 0.989 | 0.974 | 0.955 | 0.461 | 0.378 |
| Metal_nut | 0.981 | 0.992 | 0.912 | 0.931 | 0.819 | 0.874 |
| Pill | 0.981 | 0.988 | 0.935 | 0.947 | 0.689 | 0.679 |
| Screw | 0.973 | 0.979 | 0.884 | 0.906 | 0.212 | 0.301 |
| Tile | 0.978 | 0.985 | 0.892 | 0.906 | 0.740 | 0.768 |
| Toothbrush | 0.990 | 0.991 | 0.895 | 0.899 | 0.609 | 0.627 |
| Transistor | 0.964 | 0.977 | 0.895 | 0.930 | 0.570 | 0.666 |
| Wood | 0.964 | 0.974 | 0.898 | 0.893 | 0.564 | 0.627 |
| Zipper | 0.978 | 0.990 | 0.925 | 0.958 | 0.561 | 0.668 |
| Average | 0.979 | 0.983 | 0.917 | 0.924 | 0.574 | 0.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),而Tile、Wood、Hazelnut、Toothbrush等类别接近或达到满分,这与目标纹理复杂度、缺陷形态有关; - 早停影响: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),仅供参考