
💡💡💡ArcAD在产线部署耗时仅50ms,效果实测如下:
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



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

博主简介

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

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

为应对这一挑战,我们提出了异常修正冷启动异常检测,一个旨在增强基于重建模型的即插即用框架。ArcAD从两个互补的角度构建一个紧凑且有判别性的正常边界:显式地将正常样本组织成紧凑的聚类,并利用异常信号校准边界。首先,我们将图像块特征投影到一个超球面上,将方向信息与幅度变化分离,以建立一个有界的几何表示。然后,使用von Mises-Fisher分布对这些超球面特征进行建模。为确保有限的正常数据能充分覆盖潜在流形,我们引入了基于Sinkhorn的原型建模。通过将特征聚类表述为一个最优传输问题,SPM减轻了有偏的特征聚合,并将正常嵌入划分为紧凑且均匀分布的聚类。
除了对正常流形建模之外,ArcAD进一步利用异常信号来显式地修正正常边界。为缓解异常稀缺的问题,我们设计了一种原型约束的异常合成策略,该策略通过根据正常原型过滤候选样本来直接在超球面上生成合成异常。这些合成样本与少量真实缺陷一起,驱动缺陷引导校准。该模块采用一个对比目标,将异常聚集在一起,同时将它们推离最近的正常原型。这个过程显式地增强了所学正常流形的紧凑性和判别力。我们的贡献总结如下:
conda create -n arcad python=3.10
conda activate arcad
pip install -r requirements.txt
nnh1012/ArcAD_Cold-start_Data_Splits at main

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

python gen_protos.py --dataset mvtec
本文案列只训练grid数据集

gen_protos.py源码如下:
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)
# MVTec-AD
python arcad_mvtec_uni.py \
--data_path /path/to/mvtec_CD \
--save_name arcad_mvtec
arcad_mvtec_uni.py部分源码如下:
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)
python predict.py --image ./datasets/mvtec/grid/test/bent/
predict.py 部分核心源码如下
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+笔记本预测结果:包括保存各个可视化结果图

可视化结果如下:
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 删除。