首页
学习
活动
专区
圈层
工具
发布
社区首页 >专栏 >产线部署耗时仅50ms!ArcAD斩获ECCV 2026:用3%的异常样本,跑出99.7%的检测神话!(源码已部署)

产线部署耗时仅50ms!ArcAD斩获ECCV 2026:用3%的异常样本,跑出99.7%的检测神话!(源码已部署)

原创
作者头像
AI小怪兽
发布2026-07-16 12:56:31
发布2026-07-16 12:56:31
2071
举报
文章被收录于专栏:毕业设计毕业设计YOLO大作战

💡💡💡ArcAD在产线部署耗时仅50ms,效果实测如下:

代码语言:javascript
复制
Image: ./datasets/mvtec/grid/test/bent/000.png
  Anomaly score (image-level): 0.785373
  Anomaly map shape: (392, 392)
  Min/Max pixel score: 0.031545 / 0.785373
  Prediction time: 71.70 ms
  原图
  原图
anomaly_map
anomaly_map
heatmap 
heatmap 
overlay
overlay

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

  1. 提出 ArcAD 校准框架:针对冷启动场景下正常样本不足、异常样本稀缺的工业异常检测挑战,提出一种即插即用的校准框架,可无缝集成至现有基于重建的检测模型,显著提升其性能。
  2. 实验验证与通用性:在 MVTec-AD、VisA、Real-IAD 和 MANTA 四个基准数据集上,ArcAD 显著提升了 Dinomaly、RD4AD 和 ReContrast 等多种基线模型的性能,在 Real-IAD 上,ArcAD 的 I-AUROC 达 92.5%,P-F1-max 达 49.8%,分别超越基线 3.7%3.1%;在 MVTec-AD 上 I-AUROC 达 99.7%,单类设置下达 100.0%。即使在异常比例仅 3% 时,仍能带来 0.9%~5.2% 的稳定增益。

博主简介

AI小怪兽 | 计算机视觉布道者 | 视觉检测领域创新者

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

0.原理介绍

Accepted to European Conference on Computer Vision (ECCV) 2026

论文:ArcAD: Anomaly-Rectified Calibration for Cold-Start Supervised Anomaly Detection

摘要:在现实制造业中部署工业异常检测常常遇到一个具有挑战性的冷启动瓶颈,即有限的正常样本无法代表完整的正常分布,且仅有少量异常样本可用。在这种机制下,现有方法难以形成紧凑的正常边界,也无法有效利用来自罕见缺陷的监督信号。为应对这一挑战,我们提出了异常修正冷启动异常检测,一个面向基于重建的工业异常检测基线的即插即用校准框架。ArcAD 遵循一种推拉学习范式,在数据稀缺条件下构建一个紧凑且有判别性的正常边界。一方面,ArcAD 将有限的正常样本投影到一个超球面上,并将它们拉入多个紧凑的聚类中,以最大化对正常流形的覆盖。另一方面,它在超球面上合成伪异常,并利用真实异常将边界向内推,增强异常判别能力。在 MVTec-AD、VisA、Real-IAD 和 MANTA 上的大量实验表明,ArcAD 在冷启动条件下的单类别和多类别设置中均显著优于最先进的有监督和无监督方法。

代码:https://github.com/LGC-AD/ArcAD

1 引言

工业异常检测旨在通过识别罕见缺陷来实现零缺陷制造。工业异常检测系统在实际应用中的部署经常遇到严重的冷启动瓶颈。与假设正常数据充足的标准无监督设置不同,新生产线初始爬坡阶段的冷启动场景呈现出不同的数据分布。在此阶段,可用的正常样本无法覆盖完整的正常模式。同时,在系统部署早期只能收集到少量异常样本。

最近的有监督方法在使用带标注的正常和异常样本训练时展现出有前景的性能。然而,在冷启动场景中,这些方法在仅用少量异常训练时容易过拟合,导致对未见缺陷模式的泛化能力较差(见图1)。相反,无监督方法,特别是基于重建的方法,在单类别和多类别设置中均占据主导地位,这主要归功于它们对多样化及未见异常模式的强大泛化能力。这些方法通常学习正常数据的分布,并将偏离所学正常流形的实例识别为异常。虽然充足的数据允许模型隐式学习一个连续且鲁棒的正常流形,但冷启动场景中的数据稀缺会导致碎片化且边界松散的潜在空间。此外,由于严格依赖正常数据,这些无监督范式本质上未能充分利用少数可用异常样本所提供的宝贵指导。因此,一个关键问题出现了:我们如何利用有限的正常样本和罕见的异常来构建一个紧凑的正常边界?

为应对这一挑战,我们提出了异常修正冷启动异常检测,一个旨在增强基于重建模型的即插即用框架。ArcAD从两个互补的角度构建一个紧凑且有判别性的正常边界:显式地将正常样本组织成紧凑的聚类,并利用异常信号校准边界。首先,我们将图像块特征投影到一个超球面上,将方向信息与幅度变化分离,以建立一个有界的几何表示。然后,使用von Mises-Fisher分布对这些超球面特征进行建模。为确保有限的正常数据能充分覆盖潜在流形,我们引入了基于Sinkhorn的原型建模。通过将特征聚类表述为一个最优传输问题,SPM减轻了有偏的特征聚合,并将正常嵌入划分为紧凑且均匀分布的聚类。

除了对正常流形建模之外,ArcAD进一步利用异常信号来显式地修正正常边界。为缓解异常稀缺的问题,我们设计了一种原型约束的异常合成策略,该策略通过根据正常原型过滤候选样本来直接在超球面上生成合成异常。这些合成样本与少量真实缺陷一起,驱动缺陷引导校准。该模块采用一个对比目标,将异常聚集在一起,同时将它们推离最近的正常原型。这个过程显式地增强了所学正常流形的紧凑性和判别力。我们的贡献总结如下:

  • 我们提出了ArcAD,一个面向冷启动场景下基于重建模型的通用即插即用校准框架。
  • 我们引入了基于Sinkhorn的原型建模,以在超球面上将有限的正常样本组织成紧凑且均匀的聚类。
  • 我们提出了缺陷引导校准,它引入了一种原型约束的合成策略来生成伪异常,从而利用真实和合成缺陷来显式修正潜在空间中的正常边界。
  • 在四个数据集上的大量实验表明,ArcAD持续增强了最先进的基线模型。具体来说,在具有挑战性的Real-IAD数据集的多类别设置下,ArcAD分别为这些模型实现了+2.2%、+8.9%和+3.7%的图像级AUROC增益。

1.实战篇

1.1 环境安装

代码语言:javascript
复制
conda create -n arcad python=3.10
conda activate arcad
pip install -r requirements.txt

1.2 数据准备篇

nnh1012/ArcAD_Cold-start_Data_Splits at main

ps:本文以mvtec为案列进行展开,将上述所有json下载下来

1.3如何训练&预测

Step 1 — Generate prototypes (once per dataset)

代码语言:javascript
复制
python gen_protos.py --dataset mvtec

本文案列只训练grid数据集

gen_protos.py源码如下:

代码语言:javascript
复制
import torch
import torch.nn as nn
import torch.nn.functional as F
import numpy as np
import os
import argparse
import time
from sklearn.cluster import MiniBatchKMeans
from torch.utils.data import DataLoader, ConcatDataset
from functools import partial

from dataset import (
    RealIADDataset, MANTADataset, MVTecDataset, MVTecJSONDataset, VisADataset,
    AnomalyDataset, get_data_transforms,
)
from models import vit_encoder
from models.uad import ViTill
from models.vision_transformer import Block as VitBlock, bMlp, LinearAttention2


def get_feature_extractor(device, embed_dim=768, num_heads=12):
    encoder_name = 'dinov2reg_vit_base_14'
    target_layers = [2, 3, 4, 5, 6, 7, 8, 9]
    fuse_layer_encoder = [[0, 1, 2, 3], [4, 5, 6, 7]]
    fuse_layer_decoder = [[0, 1, 2, 3], [4, 5, 6, 7]]

    print(f"Loading Encoder: {encoder_name}...")
    encoder = vit_encoder.load(encoder_name)

    bottleneck = nn.ModuleList([bMlp(embed_dim, embed_dim * 4, embed_dim, drop=0.4)])
    decoder = nn.ModuleList([VitBlock(dim=embed_dim, num_heads=num_heads, mlp_ratio=4.,
                                      qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-8),
                                      attn_drop=0., attn=LinearAttention2) for _ in range(8)])

    model = ViTill(encoder=encoder, bottleneck=bottleneck, decoder=decoder, target_layers=target_layers,
                   mask_neighbor_size=0, fuse_layer_encoder=fuse_layer_encoder, fuse_layer_decoder=fuse_layer_decoder)
    return model.to(device)


def process_and_save_cluster(model, dataset, device, args, save_filename):
    """Extract bottleneck features from normal samples, cluster, and save prototypes."""
    print(f"\nProcessing target: {save_filename}")
    print(f"Total labeled images: {len(dataset)}")

    if len(dataset) == 0:
        print(f"Warning: dataset is empty for {save_filename}! Skipping...")
        return

    BATCH_SIZE = 32
    loader = DataLoader(dataset, batch_size=BATCH_SIZE, shuffle=True, num_workers=8, pin_memory=True)

    features_buffer = []
    TARGET_PATCHES = 1000000  # early-stop cap

    print(f"Extracting features...")
    start_time = time.time()
    normal_img_count = 0

    model.eval()
    with torch.no_grad():
        for i, batch in enumerate(loader):
            if len(batch) == 4:
                img, mask, label, _ = batch
            elif len(batch) == 3:
                img, _, label = batch
            else:
                img, label = batch[0], batch[1]

            img = img.to(device)
            label = label.to(device)

            normal_mask = (label == 0)
            if not normal_mask.any():
                continue

            img = img[normal_mask]
            normal_img_count += img.shape[0]

            _, _, blk = model(img)
            blk_spatial = blk[:, 5:, :]            # drop special tokens
            feats_flat = blk_spatial.reshape(-1, blk_spatial.shape[-1])

            num_keep = int(feats_flat.shape[0] * 0.3)   # keep 30% of patches
            if num_keep > 0:
                idx = torch.randperm(feats_flat.shape[0])[:num_keep]
                feats_sampled = feats_flat[idx]
                feats_norm = F.normalize(feats_sampled, dim=1).cpu().numpy()
                features_buffer.append(feats_norm)

            current_count = sum(f.shape[0] for f in features_buffer)
            if i % 20 == 0:
                print(f"  Batch {i}: collected {current_count} patches...")
            if current_count >= TARGET_PATCHES:
                break

    extract_time = time.time() - start_time
    print(f"Feature extraction finished in {extract_time:.2f}s.")

    if len(features_buffer) == 0:
        print(f"Warning: No features collected for {save_filename}! Skipping...")
        return

    all_feats = np.concatenate(features_buffer, axis=0)
    if all_feats.shape[0] > TARGET_PATCHES:
        all_feats = all_feats[:TARGET_PATCHES]

    print(f"Final feature shape for Clustering: {all_feats.shape}")

    print(f"Running MiniBatchKMeans (K={args.num_prototypes})...")
    cluster_start_time = time.time()
    kmeans = MiniBatchKMeans(
        n_clusters=args.num_prototypes,
        batch_size=16384,
        n_init=1,
        max_no_improvement=20,
        random_state=42,
        verbose=0,
    )
    kmeans.fit(all_feats)
    print(f"Clustering finished in {time.time() - cluster_start_time:.2f}s.")

    centers = torch.tensor(kmeans.cluster_centers_, dtype=torch.float)
    centers = F.normalize(centers, dim=1)

    os.makedirs(args.save_dir, exist_ok=True)
    full_save_path = os.path.join(args.save_dir, save_filename)
    torch.save(centers, full_save_path)
    print(f"Prototypes saved to: {full_save_path}")
    print("-" * 50)


def build_datasets(args, data_transform, gt_transform):
    """Build the per-class labeled dataset list for the selected dataset.

    Driven by --dataset when set; otherwise falls back to the active
    commented block below (manual comment-switch mode)."""
    datasets = []

    if args.dataset == 'manta':
        for item in args.item_list:
            datasets.append(MANTADataset(root=args.data_path, category=item,
                                         transform=data_transform, gt_transform=gt_transform, phase='labeled'))
        return datasets

    if args.dataset == 'realiad':
        for item in args.item_list:
            datasets.append(RealIADDataset(root=args.data_path, category=item,
                                           transform=data_transform, gt_transform=gt_transform, phase='labeled'))
        return datasets

    if args.dataset == 'mvtec':
        for item in args.item_list:
            label_dir = os.path.join(args.data_path, item, 'train', 'label')
            json_path = os.path.join(os.path.dirname(args.data_path), f'{item}.json')
            if os.path.isdir(label_dir):
                # cold-start layout: train/label/{good,bad}
                datasets.append(AnomalyDataset(root_dir=label_dir, transform=data_transform, mask_transform=gt_transform))
            elif os.path.isfile(json_path):
                # standard MVTec + JSON split file (project-provided cold-start splits)
                datasets.append(MVTecJSONDataset(root=args.data_path, category=item,
                                                 transform=data_transform, gt_transform=gt_transform, phase='labeled'))
            else:
                # plain standard MVTec layout: train/good, test/..., ground_truth/...
                datasets.append(MVTecDataset(root=os.path.join(args.data_path, item),
                                             transform=data_transform, gt_transform=gt_transform, phase='train'))
        return datasets

    if args.dataset == 'visa':
        for item in args.item_list:
            label_dir = os.path.join(args.data_path, item, 'train')
            datasets.append(AnomalyDataset(root_dir=label_dir, transform=data_transform, mask_transform=gt_transform))
        return datasets

    # ---- Manual comment-switch fallback (uncomment exactly one block) ----
    # for item in args.item_list:                                   # MANTA
    #     datasets.append(MANTADataset(root=args.data_path, category=item,
    #                                  transform=data_transform, gt_transform=gt_transform, phase='labeled'))
    for item in args.item_list:                                   # MVTec (standard or cold-start)
        label_dir = os.path.join(args.data_path, item, 'train', 'label')
        json_path = os.path.join(os.path.dirname(args.data_path), f'{item}.json')
        if os.path.isdir(label_dir):
            datasets.append(AnomalyDataset(root_dir=label_dir, transform=data_transform, mask_transform=gt_transform))
        elif os.path.isfile(json_path):
            datasets.append(MVTecJSONDataset(root=args.data_path, category=item,
                                             transform=data_transform, gt_transform=gt_transform, phase='labeled'))
        else:
            datasets.append(MVTecDataset(root=os.path.join(args.data_path, item),
                                         transform=data_transform, gt_transform=gt_transform, phase='train'))
    # for item in args.item_list:                                   # VisA (cold-start)
    #     label_dir = os.path.join(args.data_path, item, 'train')
    #     datasets.append(AnomalyDataset(root_dir=label_dir, transform=data_transform, mask_transform=gt_transform))
    return datasets


DATASET_CONFIG = {
    'manta':   dict(data_path='./datasets/MANTA_CD',      save_dir='./MANTA_init_files800',     num_prototypes=800),
    'mvtec':   dict(data_path='./datasets/mvtec',      save_dir='./MVTec_init_files800',    num_prototypes=800),
    'visa':    dict(data_path='./datasets/VisA_CD',       save_dir='./VisA_S3_init_files500',  num_prototypes=500),
    'realiad': dict(data_path='./datasets/Real-IAD_CD',   save_dir='./RealIAD_init_files500', num_prototypes=500),
}

ITEM_LISTS = {
    'mvtec': ['grid'],
    #'mvtec': ['carpet', 'grid', 'leather', 'tile', 'wood', 'bottle', 'cable', 'capsule',
    #          'hazelnut', 'metal_nut', 'pill', 'screw', 'toothbrush', 'transistor', 'zipper'],
    'visa': ['candle', 'capsules', 'cashew', 'chewinggum', 'fryum', 'macaroni1', 'macaroni2',
             'pcb1', 'pcb2', 'pcb3', 'pcb4', 'pipe_fryum'],
    'manta': [
        'block_inductor', 'button', 'capsule', 'coated_tablet', 'coffee_beans',
        'copper_standoff', 'embossed_tablet', 'flat_nut', 'gear', 'goji_berries',
        'led', 'led_pad', 'lettered_tablet', 'long_button', 'maize',
        'nut', 'nut_cap', 'oblong_tablet', 'paddy', 'pink_tablet',
        'pistachios', 'power_inductor', 'red_tablet', 'red_washer', 'round_button_cap',
        'screw', 'short_button', 'soybean', 'square_button_cap', 'terminal',
        'thin_resistor', 'type_c', 'wafer_resistor', 'wheat', 'white_tablet',
        'wire_cap', 'yellow_green_washer', 'yellow_tablet',
    ],
    'realiad': ['audiojack', 'bottle_cap', 'button_battery', 'end_cap', 'eraser', 'fire_hood',
                'mint', 'mounts', 'pcb', 'phone_battery', 'plastic_nut', 'plastic_plug',
                'porcelain_doll', 'regulator', 'rolled_strip_base', 'sim_card_set', 'switch', 'tape',
                'terminalblock', 'toothbrush', 'toy', 'toy_brick', 'transistor1', 'usb',
                'usb_adaptor', 'u_block', 'vcpill', 'wooden_beads', 'woodstick', 'zipper'],
}


def run(args):
    device = os.environ.get('GEN_DEV', 'cuda:0' if torch.cuda.is_available() else 'cpu')
    print(f"Using device: {device}")

    model = get_feature_extractor(device)
    data_transform, gt_transform = get_data_transforms(448, 392)

    datasets = build_datasets(args, data_transform, gt_transform)

    if args.separate_classes:
        print(">>> Mode: Separate Classes (one prototype file per class)")
        for item, ds in zip(args.item_list, datasets):
            print(f"\nCurrently processing class: {item}")
            process_and_save_cluster(model, ds, device, args, f'prototypes_init_{item}.pth')
    else:
        print(">>> Mode: Concatenated (one global prototype file)")
        combined_dataset = ConcatDataset(datasets)
        process_and_save_cluster(model, combined_dataset, device, args, 'prototypes_init.pth')


if __name__ == '__main__':
    parser = argparse.ArgumentParser()

    # --dataset selects everything (paths, K, item list, dataset class) at once.
    # Leave unset to use the manual comment-switch fallback in build_datasets().
    parser.add_argument('--dataset', type=str, default='mvtec', choices=['mvtec', 'visa', 'manta', 'realiad'])
    parser.add_argument('--data_path', type=str, default='./datasets/mvtec')
    parser.add_argument('--save_dir', type=str, default='./datasets/mvtec_vis')
    parser.add_argument('--num_prototypes', type=int, default=500)

    # Concatenated mode (default) -> single prototypes_init.pth used by arcad training.
    # Add --separate_classes to emit one file per class instead.
    parser.add_argument('--separate_classes', action='store_true',
                        help='If true, generate prototypes for each class separately.')

    args = parser.parse_args()

    if args.dataset is not None:
        cfg = DATASET_CONFIG[args.dataset]
        args.data_path = cfg['data_path']
        args.save_dir = cfg['save_dir']
        args.num_prototypes = cfg['num_prototypes']
        args.item_list = ITEM_LISTS[args.dataset]

    run(args)

Step 2 — Train + evaluate

代码语言:javascript
复制
# MVTec-AD
python arcad_mvtec_uni.py \
    --data_path /path/to/mvtec_CD \
    --save_name arcad_mvtec

arcad_mvtec_uni.py部分源码如下:

代码语言:javascript
复制
import torch
import torch.nn as nn
import torch.nn.functional as F
import math
import numpy as np
import random
import os
import csv
import argparse
import logging
import warnings
from functools import partial
from sklearn.cluster import KMeans
from torch.utils.data import DataLoader, ConcatDataset
from torch.nn.init import trunc_normal_

from dataset import MVTecDataset, MVTecJSONDataset, AnomalyDataset, get_data_transforms
from models import vit_encoder
from models.uad import ViTill
from models.vision_transformer import Block as VitBlock, bMlp, LinearAttention2
from optimizers import StableAdamW
from utils import evaluation_batch, global_cosine_hm_percent, WarmCosineScheduler, evaluation_fusion

warnings.filterwarnings("ignore")


class BinaryFocalLoss(nn.Module):
    def __init__(self, alpha=0.25, gamma=2.0, reduction='mean'):
        super(BinaryFocalLoss, self).__init__()
        self.alpha = alpha
        self.gamma = gamma
        self.reduction = reduction

    def forward(self, logits, targets):
        bce_loss = F.binary_cross_entropy_with_logits(logits, targets, reduction='none')
        pt = torch.exp(-bce_loss)
        # label 1 gets alpha, label 0 gets (1 - alpha)
        alpha_t = self.alpha * targets + (1 - self.alpha) * (1 - targets)
        focal_loss = alpha_t * (1 - pt) ** self.gamma * bce_loss

        if self.reduction == 'mean':
            return focal_loss.mean()
        elif self.reduction == 'sum':
            return focal_loss.sum()
        else:
            return focal_loss


# Discriminator (SimpleNet-style 2-layer MLP), used as a structural regularizer.
class Discriminator(nn.Module):
    def __init__(self, in_planes, n_layers=2, hidden=None):
        super(Discriminator, self).__init__()
        _hidden = in_planes if hidden is None else hidden
        self.body = nn.Sequential()
        for i in range(n_layers - 1):
            _in = in_planes if i == 0 else _hidden
            _hidden = int(_hidden // 1.5) if hidden is None else hidden
            self.body.add_module('block%d' % (i + 1),
                                 nn.Sequential(
                                     nn.Linear(_in, _hidden),
                                     nn.BatchNorm1d(_hidden),
                                     nn.LeakyReLU(0.2)
                                 ))
        self.tail = nn.Linear(_hidden, 1, bias=False)

    def forward(self, x):
        x = self.body(x)
        x = self.tail(x)
        return x


class BottleneckPrototypeLearner(nn.Module):
    """SPM: vMF prototype modeling on the hypersphere with Sinkhorn assignment."""
    def __init__(self, feature_dim, num_prototypes=50, num_special_tokens=5,
                 temperature=0.1, momentum=0.99, epsilon=0.05, sinkhorn_iterations=3,
                 noise_std=0.015, cluster_mode='sinkhorn'):
        super().__init__()
        self.num_special_tokens = num_special_tokens
        self.K = num_prototypes
        self.tau = temperature
        self.epsilon = epsilon
        self.sinkhorn_iterations = sinkhorn_iterations
        self.momentum = momentum
        self.noise_std = noise_std
        self.cluster_mode = cluster_mode

        self.register_buffer("prototypes", torch.zeros(num_prototypes, feature_dim))
        self.is_initialized = False

    def normalize_prototypes(self):
        self.prototypes.data = F.normalize(self.prototypes.data, dim=1)

    @torch.no_grad()
    def distributed_sinkhorn(self, out):
        Q = torch.exp(out / self.epsilon).t()
        B = Q.shape[1]
        K = Q.shape[0]

        sum_Q = torch.sum(Q)
        Q /= sum_Q

        for it in range(self.sinkhorn_iterations):
            # normalize rows: uniform prototype assignment
            sum_of_rows = torch.sum(Q, dim=1, keepdim=True)
            Q /= sum_of_rows
            Q /= K
            # normalize cols: each sample assigned to exactly one prototype
            sum_of_cols = torch.sum(Q, dim=0, keepdim=True)
            Q /= sum_of_cols
            Q /= B

        Q *= B
        return Q.t()

    def augment_features(self, x):
        """Prototype-restricted synthetic anomaly synthesis (DGC).
        Input:  x [N, C] normalized anchors (real normal features).
        Output: x_aug [N, C] low-likelihood synthetic anomalies, normalized.
        """
        K = 50  # candidates per anchor
        N, C = x.shape
        x_expanded = x.unsqueeze(1).expand(N, K, C)

        # sample v ~ N(z, sigma^2 I) around each anchor
        noise = torch.randn_like(x_expanded)
        x_candidates = x_expanded + self.noise_std * noise
        x_candidates = F.normalize(x_candidates, dim=2)  # re-project onto hypersphere

        # keep the candidate farthest from every prototype (lowest likelihood)
        prototypes = F.normalize(self.prototypes.data, dim=1)            # [M, C]
        candidates_flat = x_candidates.view(N * K, C)
        sim_matrix = torch.mm(candidates_flat, prototypes.t())           # [N*K, M]
        max_sim_per_candidate, _ = sim_matrix.max(dim=1)                 # nearest prototype sim
        max_sim_per_candidate = max_sim_per_candidate.view(N, K)
        _, best_indices = max_sim_per_candidate.min(dim=1)               # lowest nearest-sim

        best_indices = best_indices.view(N, 1, 1).expand(N, 1, C)
        x_aug = torch.gather(x_candidates, 1, best_indices).squeeze(1)
        return x_aug

    def forward(self, x):
        x_spatial = x[:, self.num_special_tokens:, :]
        B, N, C = x_spatial.shape
        x_flat = x_spatial.reshape(-1, C)

        if self.cluster_mode == 'sinkhorn':
            x_norm = F.normalize(x_flat, dim=1)
            p_norm = F.normalize(self.prototypes, dim=1)

            logits = torch.matmul(x_norm, p_norm.t())

            with torch.no_grad():
                if x_flat.shape[0] > self.K:
                    q_weights = self.distributed_sinkhorn(logits)
                else:
                    q_weights = F.softmax(logits / self.epsilon, dim=1)

            # cross-entropy between Sinkhorn target Q and logits
            log_probs = F.log_softmax(logits / self.tau, dim=1)
            loss = -torch.mean(torch.sum(q_weights * log_probs, dim=1))

            return loss, x_norm, q_weights

    @torch.no_grad()
    def update_prototypes(self, x_norm, q_weights):
        """EMA prototype update."""
        z_sum = torch.matmul(q_weights.t(), x_norm)
        count = torch.sum(q_weights, dim=0).unsqueeze(1)  # [K, 1]

        if self.cluster_mode == 'sinkhorn':
            z_sum = F.normalize(z_sum, dim=1)
            self.prototypes.data = self.prototypes.data * self.momentum + z_sum * (1 - self.momentum)

        self.prototypes.data = F.normalize(self.prototypes.data, dim=1)

    def calculate_contrastive_calibration_loss(self, x, mask, z_syn):
        """DGC contrastive calibration: push real anomalies away from normal
        prototypes, pull them toward synthetic anomalies."""
        # 1. real anomaly features as anchors
        x_spatial = x[:, self.num_special_tokens:, :]
        B, N, C = x_spatial.shape
        H = int(math.sqrt(N))

        if mask.shape[-1] != H:
            mask_down = F.adaptive_max_pool2d(mask, (H, H))
        else:
            mask_down = mask

        mask_flat = mask_down.reshape(B, -1)
        x_all = x_spatial.reshape(-1, C)
        mask_all = mask_flat.reshape(-1)
        anchor_feats = x_all[mask_all > 0]

        if anchor_feats.shape[0] == 0:
            return torch.tensor(0.0).to(x.device)

        anchor_norm = F.normalize(anchor_feats, dim=1)

        # 2. push loss: drive anchors off the normal boundary (nearest prototype)
        p_norm = F.normalize(self.prototypes, dim=1)
        sim_matrix_push = torch.matmul(anchor_norm, p_norm.t())
        max_sim_push, _ = torch.max(sim_matrix_push, dim=1)
        loss_push = F.relu(max_sim_push).mean()

        # 3. pull loss: attract anchors toward synthetic anomalies
        if z_syn is None or z_syn.shape[0] == 0:
            loss_pull = torch.tensor(0.0).to(x.device)
        else:
            z_syn_norm = F.normalize(z_syn, dim=1)
            sim_matrix_pull = torch.matmul(anchor_norm, z_syn_norm.t())
            avg_sim_pull = torch.mean(sim_matrix_pull, dim=1)
            loss_pull = (1.0 - avg_sim_pull).mean()

        return loss_push + loss_pull


def get_logger(name, save_path=None, level='INFO'):
    logger = logging.getLogger(name)
    logger.setLevel(getattr(logging, level))
    log_format = logging.Formatter('%(message)s')
    streamHandler = logging.StreamHandler()
    streamHandler.setFormatter(log_format)
    logger.addHandler(streamHandler)
    if not save_path is None:
        os.makedirs(save_path, exist_ok=True)
        fileHandler = logging.FileHandler(os.path.join(save_path, 'log.txt'))
        fileHandler.setFormatter(log_format)
        logger.addHandler(fileHandler)
    return logger


def setup_seed(seed):
    torch.manual_seed(seed)
    torch.cuda.manual_seed_all(seed)
    np.random.seed(seed)
    random.seed(seed)
    torch.backends.cudnn.deterministic = False
    torch.backends.cudnn.benchmark = True



if __name__ == '__main__':
    os.environ['CUDA_LAUNCH_BLOCKING'] = "1"

    parser = argparse.ArgumentParser(description='')
    parser.add_argument('--data_path', type=str, default='./datasets/mvtec')
    parser.add_argument('--save_dir', type=str, default='./saved_results')
    parser.add_argument('--save_name', type=str,
                        default='vitill_uni_simplesnet_discriminator_focal')
    parser.add_argument('--cluster_mode', type=str, default='sinkhorn')
    parser.add_argument('--proto_path', type=str, default='./MVTec_init_files800/prototypes_init.pth',
                        )
    args = parser.parse_args()

    item_list = ['grid']

    #item_list = ['carpet', 'grid', 'leather', 'tile', 'wood', 'bottle', 'cable', 'capsule',
    #             'hazelnut', 'metal_nut', 'pill', 'screw', 'toothbrush', 'transistor', 'zipper']

    logger = get_logger(args.save_name, os.path.join(args.save_dir, args.save_name))
    print_fn = logger.info

    device = 'cuda:0' if torch.cuda.is_available() else 'cpu'
    print_fn(device)

    train(item_list)

Step 3— predict

代码语言:javascript
复制
python predict.py --image ./datasets/mvtec/grid/test/bent/

predict.py 部分核心源码如下

代码语言:javascript
复制
import os
import argparse
import time
from functools import partial

import torch
import torch.nn as nn
import numpy as np
from PIL import Image
import cv2

from dataset import get_data_transforms
from models import vit_encoder
from models.uad import ViTill
from models.vision_transformer import Block as VitBlock, bMlp, LinearAttention2
from utils import cal_anomaly_maps, get_gaussian_kernel, min_max_norm, cvt2heatmap, show_cam_on_image


def build_model(device='cuda'):
    """Build the same ViTill architecture used in arcad_mvtec_uni.py."""
    encoder_name = 'dinov2reg_vit_base_14'
    target_layers = [2, 3, 4, 5, 6, 7, 8, 9]
    fuse_layer_encoder = [[0, 1, 2, 3], [4, 5, 6, 7]]
    fuse_layer_decoder = [[0, 1, 2, 3], [4, 5, 6, 7]]

    encoder = vit_encoder.load(encoder_name)

    if 'small' in encoder_name:
        embed_dim, num_heads = 384, 6
    elif 'base' in encoder_name:
        embed_dim, num_heads = 768, 12
    elif 'large' in encoder_name:
        embed_dim, num_heads = 1024, 16
        target_layers = [4, 6, 8, 10, 12, 14, 16, 18]
    else:
        raise ValueError("Architecture not in small, base, large.")

    bottleneck = nn.ModuleList([bMlp(embed_dim, embed_dim * 4, embed_dim, drop=0.4)])

    decoder = nn.ModuleList([
        VitBlock(dim=embed_dim, num_heads=num_heads, mlp_ratio=4.,
                 qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-8), attn_drop=0.,
                 attn=LinearAttention2)
        for _ in range(8)
    ])

    model = ViTill(
        encoder=encoder,
        bottleneck=bottleneck,
        decoder=decoder,
        target_layers=target_layers,
        mask_neighbor_size=0,
        fuse_layer_encoder=fuse_layer_encoder,
        fuse_layer_decoder=fuse_layer_decoder
    )
    model = model.to(device)
    return model


def load_model(save_name=None, save_dir='./saved_results', checkpoint=None, device='cuda'):
    """Load a trained model from checkpoint."""
    model = build_model(device)

    if checkpoint is not None:
        checkpoint_path = checkpoint
    elif save_name is not None:
        checkpoint_path = os.path.join(save_dir, save_name, 'model.pth')
    else:
        raise ValueError("Either --checkpoint or --save_name must be provided.")

    if not os.path.exists(checkpoint_path):
        raise FileNotFoundError(f"Checkpoint not found: {checkpoint_path}")

    state_dict = torch.load(checkpoint_path, map_location=device)
    model.load_state_dict(state_dict)
    model.eval()
    print(f"Loaded checkpoint from {checkpoint_path}")
    return model


def preprocess_image(image_path, image_size=448, crop_size=392):
    """Apply the same transforms used during training."""
    data_transform, _ = get_data_transforms(image_size, crop_size)
    img = Image.open(image_path).convert('RGB')
    img_tensor = data_transform(img)
    return img_tensor.unsqueeze(0), img  # add batch dim, keep original image



def main():
    parser = argparse.ArgumentParser(description='ArcAD single-image inference (MVTec-AD unified model)')
    parser.add_argument('--image', type=str, required=True,
                        help='Path to the input image (or image directory).')
    parser.add_argument('--save_name', type=str, default=None,
                        help='Name of the trained model folder under --save_dir (used if --checkpoint is not set).')
    parser.add_argument('--save_dir', type=str, default='./saved_results',
                        help='Root directory where trained models are saved.')
    parser.add_argument('--checkpoint', type=str,
                        default='saved_results/vitill_uni_simplesnet_discriminator_focal/model.pth',
                        help='Direct path to a model.pth checkpoint.')
    parser.add_argument('--out_dir', type=str, default='./predictions',
                        help='Directory to save visualization results.')
    parser.add_argument('--device', type=str, default='cuda:0',
                        help='Device to run inference on.')
    args = parser.parse_args()

    device = args.device if torch.cuda.is_available() else 'cpu'
    print(f"Using device: {device}")

    model = load_model(save_name=args.save_name, save_dir=args.save_dir,
                       checkpoint=args.checkpoint, device=device)

    # Single image or directory
    if os.path.isdir(args.image):
        image_paths = [
            os.path.join(args.image, f)
            for f in sorted(os.listdir(args.image))
            if f.lower().endswith(('.png', '.jpg', '.jpeg', '.bmp', '.tiff'))
        ]
    else:
        image_paths = [args.image]

    if len(image_paths) == 0:
        raise ValueError(f"No valid images found in {args.image}")

    for image_path in image_paths:
        score, anomaly_map = predict(
            model, image_path,
            device=device,
            save_dir=args.out_dir,
            save_name=None  # use image base name
        )
        print(f"Image: {image_path}")
        print(f"  Anomaly score (image-level): {score:.6f}")
        print(f"  Anomaly map shape: {anomaly_map.shape}")
        print(f"  Min/Max pixel score: {anomaly_map.min():.6f} / {anomaly_map.max():.6f}")


if __name__ == '__main__':
    main()

3060+笔记本预测结果:包括保存各个可视化结果图

可视化结果如下:

代码语言:javascript
复制
Image: ./datasets/mvtec/grid/test/bent/000.png
  Anomaly score (image-level): 0.785373
  Anomaly map shape: (392, 392)
  Min/Max pixel score: 0.031545 / 0.785373
  Prediction time: 71.70 ms

原创声明:本文系作者授权腾讯云开发者社区发表,未经许可,不得转载。

如有侵权,请联系 cloudcommunity@tencent.com 删除。

原创声明:本文系作者授权腾讯云开发者社区发表,未经许可,不得转载。

如有侵权,请联系 cloudcommunity@tencent.com 删除。

评论
登录后参与评论
0 条评论
热度
最新
推荐阅读
目录
  • 0.原理介绍
  • 1 引言
  • 1.实战篇
    • 1.1 环境安装
    • 1.2 数据准备篇
    • 1.3如何训练&预测
      • Step 1 — Generate prototypes (once per dataset)
      • Step 2 — Train + evaluate
      • Step 3— predict
领券
问题归档专栏文章快讯文章归档关键词归档开发者手册归档开发者手册 Section 归档