Anomalib 中的 UniNet:Teacher-Student 对比学习异常检测模型全解析
【免费下载链接】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
UniNet 是 anomalib 收录的一种统一对比学习异常检测框架(Model Type 为 Classification 与 Segmentation),同时支持监督与非监督两种训练模式,并面向多类别异常检测场景设计。本文基于仓库文档页 uninet.md 所引用的六个核心模块,结合 src/anomalib/models/image/uninet/ 下的源码实现,完整讲解其 Lightning 入口、前向流程、对比损失、注意力瓶颈、域相关特征选择与推理阶段的加权决策机制,并给出可直接复制的配置与调用方式。
1. 模块布局与核心组件
文档页 uninet.md 通过 Sphinxautomodule指令导入了 UniNet 的全部实现模块,它们在仓库中的实际位置如下:
| 模块 | 文件 | 职责 |
|---|---|---|
| Lightning 入口 | lightning_model.py | UniNet类,定义参数、训练/验证步骤与优化器 |
| PyTorch 核心网络 | torch_model.py | UniNetModel(学生/教师/瓶颈前向)与Teachers |
| 损失函数 | components/loss.py | UniNetLoss:余弦 + 对比 + margin 三元损失 |
| 推理异常图 | components/anomaly_map.py | weighted_decision_mechanism加权决策机制 |
| 注意力瓶颈 | components/attention_bottleneck.py | AttentionBottleneck与BottleneckLayer |
| 域相关特征选择 | components/dfs.py | DomainRelatedFeatureSelection(DFS) |
从 模块 docstring 看,UniNet 被描述为“面向多领域、适合监督与非监督异常检测、并聚焦多类别异常检测”的模型;其设计源自 CVPR 2025 论文,源码头部保留了原作者 MIT 许可与 Intel 修改后的 Apache-2.0 许可标注(见 torch_model.py)。
2. Lightning 入口:UniNet 的参数与训练配置
UniNet继承自 anomalib 的AnomalibModule,构造参数定义在 lightning_model.py#L38-L55:
| 参数 | 类型 | 默认值 | 说明 |
|---|---|---|---|
student_backbone | str | "wide_resnet50_2" | 学生网络使用的 backbone(作为 decoder 特征提取器) |
teacher_backbone | str | "wide_resnet50_2" | 教师网络使用的 backbone,通过torchvision加载预训练权重 |
temperature | float | 0.1 | 对比损失的温度系数,控制学生/教师相似度计算的锐度 |
pre_processor/post_processor/evaluator/visualizer | 实例或bool | True | anomalib 标准前后处理、评估器、可视化组件 |
几个值得注意的实现细节:
- 学习类型:
learning_type属性返回LearningType.ONE_CLASS(lightning_model.py#L106-L109),即模型按 one-class 学习框架管理。 - 训练步骤:
training_step直接调用self.model(images=..., masks=batch.gt_mask, labels=batch.gt_label)并记录train_loss;validation_step则要求batch.image必须是张量,否则抛出ValueError。 - 优化器配置(lightning_model.py#L76-L99):单个
AdamW优化器分组管理四组参数——student、bottleneck、dfs 使用全局学习率5e-3,而target_teacher单独以1e-6的低学习率缓慢更新;weight_decay=1e-5、amsgrad=True。调度器为MultiStepLR,唯一 milestone 取max_steps(或max_epochs)的 80% 处,gamma=0.2。 - 温度默认值的一个细节:
UniNetLoss自身的temperature默认值是2.0(loss.py#L25),但UniNet构造时会把外层temperature(默认0.1)显式传入损失,因此实际生效的默认温度是0.1。
3. 核心前向流程:UniNetModel
UniNetModel(torch_model.py#L30-L57)由四部分组成:
self.teachers = Teachers(teacher_backbone) # 双教师网络 self.student = get_decoder(student_backbone) # 学生解码器 self.bottleneck = BottleneckLayer(block=AttentionBottleneck, layers=3) self.dfs = DomainRelatedFeatureSelection() # 域相关特征选择3.1 双教师网络
Teachers(torch_model.py#L174-L222)用同一个teacher_backbone构造两个教师:
source_teacher:加载后立即.eval(),前向时全程包裹在torch.no_grad()中,提供冻结的参考特征;target_teacher:保持可训练,参与梯度回传。
两者均通过torchvision.models.feature_extraction.create_feature_extractor从预训练模型中抽取layer1/layer2/layer3三个尺度特征。前向时按尺度把 source 与 target 特征在 batch 维拼接,作为瓶颈层的输入(注释标明宽 ResNet 对应 512/1024/2048 通道,见 torch_model.py#L218-L221)。
3.2 训练与推理两条路径
forward(torch_model.py#L59-L121)的流程为:教师抽特征 → BottleneckLayer → 学生 decoder。与学生 decoder(de_resnet)输出的处理有关的关键代码在 torch_model.py#L80-L95:decoder 输出 3 个尺度特征,每个尺度再 chunk 成两半,重排为共6 份多尺度学生特征。
训练模式:先用DomainRelatedFeatureSelection对学生特征做域相关选择,再调用_compute_loss计算总损失——除了UniNetLoss之外,若提供了predictions与label,还会叠加两个BCEWithLogitsLoss分类损失(torch_model.py#L123-L155),分类预测头为AdaptiveAvgPool2d(1,1) + Linear(256,1)并 chunk 成两路。
推理模式:对每一对(教师特征,学生特征)计算1 - cosine_similarity得到 6 张逐像素相似度差异图,随后交给weighted_decision_mechanism(固定alpha=0.01, beta=3e-05)融合出最终pred_score与anomaly_map,封装为InferenceBatch返回。
4. 训练损失:UniNetLoss 的三元结构
UniNetLoss(loss.py#L17-L108)对 6 份学生/教师特征逐份累加损失,每一份包含三项:
- 余弦损失:特征展平并 L2 归一化后,取
1 - cosine_similarity的均值; - 对比损失:归一化特征矩阵相乘除以温度后做 exp 并逐行归一化,取对角线元素
diag_sum,损失为-log(diag_sum)——本质是希望每个查询特征与其自身对应位置(对角)的相似度最大; - margin 损失(
margin=1):正常样本项为relu(margin - diag_sum);当提供异常标注时,额外加入异常项relu(diag_sum - margin/2),把异常样本的自相似度往下压。
最终按cosine_loss * lambda_weight + contrastive_loss * (1 - lambda_weight) + margin_loss组合,lambda_weight默认0.7。
监督与非监督的分支切换逻辑(loss.py#L76-L102):
mask is None:无监督分支,按“仅正常样本”处理,只用全量diag_sum计算对比损失与 margin 损失;mask非空:监督分支。若 mask 维数小于 3 视为图像级 label(0/1),否则视为像素级 mask 并interpolate到特征图分辨率再展平;随后分别对正常/异常位置子集计算损失。
5. 域相关特征选择(DFS)
DomainRelatedFeatureSelection(dfs.py#L17-L93)用于在训练时从教师特征中挑选“与当前域相关”的特征通道模式,其机制:
- 定义三组可学习参数
theta1/theta2/theta3(对应 256/512/1024 通道),初始为 0; - 前向时
theta = clamp(sigmoid(theta_i) + 0.5, max=1),因此theta ∈ [0.5, 1],注释说明这是为避免局部权重丢失而保证的非零下界(dfs.py#L70-L75); - 权重计算:对 target 特征沿空间维展平,减去逐通道最大值(
maximize=True)后做softmax得到空间权重,再与“通道全局均值 + theta”相乘,最终逐元素乘以 source 特征完成加权选择(dfs.py#L78-L92)。
6. 注意力瓶颈:AttentionBottleneck 与 BottleneckLayer
BottleneckLayer 是教师特征进入学生网络前的“过渡层”,UniNetModel中以layers=3实例化(即 3 个残差块)。其forward(attention_bottleneck.py#L399-L412)将三个尺度输入分流处理:尺度 0 经conv1→conv2、尺度 1 经conv3、尺度 2 直接透传,三者 concat 后送入由 3 个AttentionBottleneck组成的bn_layer。
AttentionBottleneck(attention_bottleneck.py#L75-L148)按 ResNet 惯例取channel_expansion=4,支持两种模式:
halve=1:标准瓶颈处理;halve=2:双分支注意力,把通道劈成两路分别用 3×3 与 7×7 卷积处理,以捕获不同感受野的多尺度特征后再融合(docstring 中给出了AttentionBottleneck(256, 64, halve=1)与AttentionBottleneck(512, 128, halve=2)的形状示例)。
此外模块还实现了fuse_bn(attention_bottleneck.py#L54-L73),可把 BatchNorm 参数折叠进卷积权重,用于推理加速。
7. 推理融合:weighted_decision_mechanism
weighted_decision_mechanism(anomaly_map.py#L19-L103)把 6 张尺度差异图融合为最终异常分与异常图,流程为:
- 尺度权重:对每张图的 batch 内最大值取
softmax,剔除低于均值的尺度,取剩余最大值均值的alpha倍并与beta取较大者,作为该样本的权重系数total_weights[i](alpha控制上限、beta控制下限); - 异常图:各尺度图
bilinear插值到输入分辨率后直接累加得到anomaly_map; - 图像分:对累加图施加
GaussianBlur2d(sigma=4.0, kernel_size=(5,5)),展平后取top_k值,其中top_k = 分辨率像素数 × total_weights[i](至少 1),以最大值作为pred_score。
在UniNetModel.forward中该函数以alpha=0.01, beta=3e-05、output_size=images.shape[-2:]调用(torch_model.py#L111-L117)。
8. 训练与推理:配置与调用方式
仓库提供了现成 YAML 配置 examples/configs/model/uninet.yaml,完整内容如下:
model: class_path: anomalib.models.UniNet init_args: student_backbone: wide_resnet50_2 teacher_backbone: wide_resnet50_2 temperature: 0.1 trainer: max_epochs: 100 callbacks: - class_path: lightning.pytorch.callbacks.EarlyStopping init_args: patience: 20 monitor: image_AUROC mode: max即默认训练 100 epoch,并用image_AUROC(越大越好)做 EarlyStopping、patience 20。
CLI 方式(源自 模块 README):
anomalib train --model UniNet --data MVTecAD --data.category <category>API 方式(与init.py 中的 docstring 示例一致):
from anomalib.models import UniNet from anomalib.data import MVTecAD from anomalib.engine import Engine datamodule = MVTecAD() model = UniNet() engine = Engine() engine.train(model=model, datamodule=datamodule) engine.predict(model=model, datamodule=datamodule)模块 README 还记录了该实现在 MVTecAD 数据集上的基准(seed 42),供参考其相对水平:
| 指标 | Avg | Bottle | Carpet | Grid | Pill | Screw | Transistor | Zipper |
|---|---|---|---|---|---|---|---|---|
| Image-Level AUC | 0.956 | 0.999 | 0.896 | 0.996 | 0.816 | 0.919 | 0.984 | 0.945 |
| Pixel-Level AUC | 0.976 | 0.989 | 0.973 | 0.992 | 0.964 | 0.992 | 0.923 | 0.984 |
| Image F1 | 0.957 | 0.984 | 0.883 | 0.973 | 0.921 | 0.905 | 0.961 | 0.959 |
(完整 16 类数值见 README.md。)
9. 实现要点小结
- 从源码结构看,UniNet 的“对比”体现在两处:训练阶段
UniNetLoss的对角相似度对比损失 + margin 损失;推理阶段以教师-学生逐像素余弦距离作为异常证据,再由加权决策机制融合多尺度结果。 - 监督/非监督统一由
mask/label是否存在驱动:无 mask 时损失只走“仅正常样本”分支;有像素 mask 或图像 label 时自动拆分为正常/异常两段子损失,并叠加 BCE 分类损失。 source_teacher全程no_grad + eval冻结,target_teacher以1e-6学习率更新,这种“冻结参照 + 慢速跟随”的双教师设计是理解该模型参数分组优化器配置的关键。- 文档页 uninet.md 列出的六个
automodule与上文六个文件一一对应,是继续深入 API 级细节(各参数 docstring、继承关系)的入口。
【免费下载链接】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),仅供参考