资讯详情

资讯详情

Anomalib 中的 UniNet:Teacher-Student 对比学习异常检测模型全解析

Anomalib 中的 UniNetTeacher-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/anomalibUniNet 是 anomalib 收录的一种统一对比学习异常检测框架Model Type 为 Classification 与 Segmentation同时支持监督与非监督两种训练模式并面向多类别异常检测场景设计。本文基于仓库文档页 uninet.md 所引用的六个核心模块结合 src/anomalib/models/image/uninet/ 下的源码实现完整讲解其 Lightning 入口、前向流程、对比损失、注意力瓶颈、域相关特征选择与推理阶段的加权决策机制并给出可直接复制的配置与调用方式。1. 模块布局与核心组件文档页 uninet.md 通过 Sphinxautomodule指令导入了 UniNet 的全部实现模块它们在仓库中的实际位置如下模块文件职责Lightning 入口lightning_model.pyUniNet类定义参数、训练/验证步骤与优化器PyTorch 核心网络torch_model.pyUniNetModel学生/教师/瓶颈前向与Teachers损失函数components/loss.pyUniNetLoss余弦 对比 margin 三元损失推理异常图components/anomaly_map.pyweighted_decision_mechanism加权决策机制注意力瓶颈components/attention_bottleneck.pyAttentionBottleneck与BottleneckLayer域相关特征选择components/dfs.pyDomainRelatedFeatureSelectionDFS从 模块 docstring 看UniNet 被描述为“面向多领域、适合监督与非监督异常检测、并聚焦多类别异常检测”的模型其设计源自 CVPR 2025 论文源码头部保留了原作者 MIT 许可与 Intel 修改后的 Apache-2.0 许可标注见 torch_model.py。2. Lightning 入口UniNet 的参数与训练配置UniNet继承自 anomalib 的AnomalibModule构造参数定义在 lightning_model.py#L38-L55参数类型默认值说明student_backbonestrwide_resnet50_2学生网络使用的 backbone作为 decoder 特征提取器teacher_backbonestrwide_resnet50_2教师网络使用的 backbone通过torchvision加载预训练权重temperaturefloat0.1对比损失的温度系数控制学生/教师相似度计算的锐度pre_processor/post_processor/evaluator/visualizer实例或boolTrueanomalib 标准前后处理、评估器、可视化组件几个值得注意的实现细节学习类型learning_type属性返回LearningType.ONE_CLASSlightning_model.py#L106-L109即模型按 one-class 学习框架管理。训练步骤training_step直接调用self.model(images..., masksbatch.gt_mask, labelsbatch.gt_label)并记录train_lossvalidation_step则要求batch.image必须是张量否则抛出ValueError。优化器配置lightning_model.py#L76-L99单个AdamW优化器分组管理四组参数——student、bottleneck、dfs 使用全局学习率5e-3而target_teacher单独以1e-6的低学习率缓慢更新weight_decay1e-5、amsgradTrue。调度器为MultiStepLR唯一 milestone 取max_steps或max_epochs的 80% 处gamma0.2。温度默认值的一个细节UniNetLoss自身的temperature默认值是2.0loss.py#L25但UniNet构造时会把外层temperature默认0.1显式传入损失因此实际生效的默认温度是0.1。3. 核心前向流程UniNetModelUniNetModeltorch_model.py#L30-L57由四部分组成self.teachers Teachers(teacher_backbone) # 双教师网络 self.student get_decoder(student_backbone) # 学生解码器 self.bottleneck BottleneckLayer(blockAttentionBottleneck, layers3) self.dfs DomainRelatedFeatureSelection() # 域相关特征选择3.1 双教师网络Teacherstorch_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 训练与推理两条路径forwardtorch_model.py#L59-L121的流程为教师抽特征 → BottleneckLayer → 学生 decoder。与学生 decoderde_resnet输出的处理有关的关键代码在 torch_model.py#L80-L95decoder 输出 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固定alpha0.01, beta3e-05融合出最终pred_score与anomaly_map封装为InferenceBatch返回。4. 训练损失UniNetLoss 的三元结构UniNetLossloss.py#L17-L108对 6 份学生/教师特征逐份累加损失每一份包含三项余弦损失特征展平并 L2 归一化后取1 - cosine_similarity的均值对比损失归一化特征矩阵相乘除以温度后做 exp 并逐行归一化取对角线元素diag_sum损失为-log(diag_sum)——本质是希望每个查询特征与其自身对应位置对角的相似度最大margin 损失margin1正常样本项为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-L102mask is None无监督分支按“仅正常样本”处理只用全量diag_sum计算对比损失与 margin 损失mask非空监督分支。若 mask 维数小于 3 视为图像级 label0/1否则视为像素级 mask 并interpolate到特征图分辨率再展平随后分别对正常/异常位置子集计算损失。5. 域相关特征选择DFSDomainRelatedFeatureSelectiondfs.py#L17-L93用于在训练时从教师特征中挑选“与当前域相关”的特征通道模式其机制定义三组可学习参数theta1/theta2/theta3对应 256/512/1024 通道初始为 0前向时theta clamp(sigmoid(theta_i) 0.5, max1)因此theta ∈ [0.5, 1]注释说明这是为避免局部权重丢失而保证的非零下界dfs.py#L70-L75权重计算对 target 特征沿空间维展平减去逐通道最大值maximizeTrue后做softmax得到空间权重再与“通道全局均值 theta”相乘最终逐元素乘以 source 特征完成加权选择dfs.py#L78-L92。6. 注意力瓶颈AttentionBottleneck 与 BottleneckLayerBottleneckLayer 是教师特征进入学生网络前的“过渡层”UniNetModel中以layers3实例化即 3 个残差块。其forwardattention_bottleneck.py#L399-L412将三个尺度输入分流处理尺度 0 经conv1→conv2、尺度 1 经conv3、尺度 2 直接透传三者 concat 后送入由 3 个AttentionBottleneck组成的bn_layer。AttentionBottleneckattention_bottleneck.py#L75-L148按 ResNet 惯例取channel_expansion4支持两种模式halve1标准瓶颈处理halve2双分支注意力把通道劈成两路分别用 3×3 与 7×7 卷积处理以捕获不同感受野的多尺度特征后再融合docstring 中给出了AttentionBottleneck(256, 64, halve1)与AttentionBottleneck(512, 128, halve2)的形状示例。此外模块还实现了fuse_bnattention_bottleneck.py#L54-L73可把 BatchNorm 参数折叠进卷积权重用于推理加速。7. 推理融合weighted_decision_mechanismweighted_decision_mechanismanomaly_map.py#L19-L103把 6 张尺度差异图融合为最终异常分与异常图流程为尺度权重对每张图的 batch 内最大值取softmax剔除低于均值的尺度取剩余最大值均值的alpha倍并与beta取较大者作为该样本的权重系数total_weights[i]alpha控制上限、beta控制下限异常图各尺度图bilinear插值到输入分辨率后直接累加得到anomaly_map图像分对累加图施加GaussianBlur2d(sigma4.0, kernel_size(5,5))展平后取top_k值其中top_k 分辨率像素数 × total_weights[i]至少 1以最大值作为pred_score。在UniNetModel.forward中该函数以alpha0.01, beta3e-05、output_sizeimages.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 方式源自 模块 READMEanomalib train --model UniNet --data MVTecAD --data.category categoryAPI 方式与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(modelmodel, datamoduledatamodule) engine.predict(modelmodel, datamoduledatamodule)模块 README 还记录了该实现在 MVTecAD 数据集上的基准seed 42供参考其相对水平指标AvgBottleCarpetGridPillScrewTransistorZipperImage-Level AUC0.9560.9990.8960.9960.8160.9190.9840.945Pixel-Level AUC0.9760.9890.9730.9920.9640.9920.9230.984Image F10.9570.9840.8830.9730.9210.9050.9610.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),仅供参考
觉得有用,分享给同行:

为您的企业打造数字门面

稳重轻奢商务风格,端正雅致视觉,长效耐看不易过时。

立即咨询 →