当我们面对大量没有打标签的图片(无监督学习的情况),又要对图像进行目标检测时,可以采用带伪标签的方案。

本文采用PadiM对经典数据集MTVec AD进行异常区域检测,显示出异常坐标,可转换成YOLO格式。

1.PadiM介绍

(1)PadiM(Patch Distribution Modeling)是一种用于无监督异常检测与定位(Anomaly Detection and Localization)的深度学习方法,最初由 Defard 等人在 2021 年提出(论文标题:PaDiM: A Patch Distribution Modeling Framework for Anomaly Detection and Localization)。该方法特别适用于工业视觉检测场景,比如在只有正常样本(无缺陷样本)的情况下训练模型,然后在测试阶段检测并定位图像中的异常区域(如划痕、污渍、缺失部件等)。

(2)PadiM 的核心思想是:
利用预训练的卷积神经网络(如 ResNet、EfficientNet 等)提取图像不同尺度下的特征图;

对每个空间位置(即“patch”)的特征向量建模其多元高斯分布(Multivariate Gaussian Distribution);

在推理阶段,通过计算测试图像 patch 特征与训练分布之间的马氏距离(Mahalanobis Distance)来判断是否为异常,并生成异常热力图。

2.数据集MTVec AD介绍

(1)MVTec AD(MVTec Anomaly Detection Dataset)是由德国机器视觉软件公司 MVTec Software GmbH 于 2019 年在 CVPR(IEEE/CVF Conference on Computer Vision and Pattern Recognition)上发布的工业级无监督异常检测基准数据集。

(2)主要架构:

MVTec AD 包含 15 个类别,在每个类别下都包含train,test。在train中只包含所有该类别下合格的物品的图片;test中既包含了合格图片,又有异常物品的图片,及异常类别。

我们用PadiM模型对train进行特征建模,模型只学习“好”的物品特征;在test上,去对每个(”好“与”坏“)图像进行推理。同时,我们只能出判断这个物品的好坏性,其异常的类别暂时不能(应该可以实现,大家可尝试尝试)。

3.代码部分

(1)我们先撰写数据加载。这是对训练集上的数据加载......

import cv2
from torch.utils.data import Dataset
from pathlib import Path

class MVTecTrainDataset(Dataset):
    def __init__(self, root, category, transform=None):
        self.transform = transform
        good_dir = Path(root) / category / "train" / "good"
        self.img_paths = sorted(good_dir.glob("*.png"))

    def __len__(self):
        return len(self.img_paths)

    def __getitem__(self, idx):
        img = cv2.imread(str(self.img_paths[idx]))
        if img is None:
            raise ValueError(f"Failed to load image: {self.img_paths[idx]}")
        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
        if self.transform:
            img = self.transform(img)
        return img

(2)模型PadiM的建立。使用 ResNet-18 作为特征提取骨干网络,在正常样本上学习每个图像块(patch)的多维高斯分布,测试时通过 Mahalanobis 距离判断是否异常

import torch
import torch.nn.functional as F
from torchvision import models

class PadimModel:
    def __init__(self, device, layers=[0, 1, 2]):
        self.device = device
        self.layers = layers
        self.model = models.resnet18(weights=models.ResNet18_Weights.IMAGENET1K_V1).to(device).eval()
        self.feature_maps = {}
        self._register_hooks()

    def _register_hooks(self):
        layer_names = ['layer1', 'layer2', 'layer3', 'layer4']
        for i in self.layers:
            name = layer_names[i]
            getattr(self.model, name).register_forward_hook(self._hook_fn(name))

    def _hook_fn(self, name):
        def hook(module, input, output):
            self.feature_maps[name] = output
        return hook

    def embed(self, x):
        with torch.no_grad():
            _ = self.model(x)
        features = [self.feature_maps[name] for name in ['layer1', 'layer2', 'layer3'][:len(self.layers)]]
        min_h = min(f.shape[2] for f in features)
        min_w = min(f.shape[3] for f in features)
        resized = [
            F.interpolate(f, size=(min_h, min_w), mode='bilinear', align_corners=False)
            for f in features
        ]
        return torch.cat(resized, dim=1)  # [B, C, H, W]

layer表示resnet的4个残差阶段,我们选取前3块,因为随着层数的增加,在第4个残差时,空间分辨率太低(通常是原图的 1/32),图片通道数大增,会丢失定位精度;且 PaDiM 原论文实验发现前三层效果最好。我们给定一下ResNet的网络搭建:

import torch
from torch import nn
from torchsummary import summary


class Residual(nn.Module):
    def __init__(self,input_channels,out_channels,use_1conv = False,strides=1):
        super(Residual,self).__init__()

        self.ReLu = nn.ReLU()
        self.conv1 = nn.Conv2d(in_channels=input_channels,out_channels=out_channels,kernel_size=3,padding=1,stride=strides)
        self.conv2 = nn.Conv2d(in_channels=out_channels,out_channels=out_channels,kernel_size=3,padding=1)
        self.bn1 = nn.BatchNorm2d(out_channels)
        self.bn2 = nn.BatchNorm2d(out_channels)

        if use_1conv:
            self.conv3 = nn.Conv2d(in_channels=input_channels,out_channels=out_channels,kernel_size=1,padding=0,stride=strides)
        else:
            self.conv3 = None


    def forward(self,x):

        y = self.ReLu(self.bn1(self.conv1(x)))
        y = self.bn2(self.conv2(y))

        if self.conv3 :
            x = self.conv3(x)

        y = self.ReLu(y+x)

        return y

class ResNet18(nn.Module):
    def __init__(self,Residual):
        super(ResNet18,self).__init__()

        self.b1 = nn.Sequential(
            nn.Conv2d(in_channels=1,out_channels=64,kernel_size=7,padding=3,stride=2),
            nn.ReLU(),
            nn.BatchNorm2d(64),
            nn.MaxPool2d(kernel_size=3,stride=2,padding=1)
        )

        self.b2 = nn.Sequential(
            Residual(64,64,use_1conv = False,strides=1),
            Residual(64,64,use_1conv = False,strides=1)
        )

        self.b3 = nn.Sequential(
            Residual(64,128,use_1conv = True,strides=2),
            Residual(128,128,use_1conv = False,strides=1)
        )

        self.b4 = nn.Sequential(
            Residual(128,256, use_1conv=True, strides=2),
            Residual(256,256, use_1conv=False, strides=1)
        )

        self.b5 = nn.Sequential(
            Residual(256,512, use_1conv=True, strides=2),
            Residual(512,512, use_1conv=False, strides=1)
        )

        self.b6 = nn.Sequential(
            nn.AdaptiveAvgPool2d((1,1)),
            nn.Flatten(),
            nn.Linear(1*1*512,10)
        )

    def forward(self,x):
        x = self.b1(x)
        x = self.b2(x)
        x = self.b3(x)
        x = self.b4(x)
        x = self.b5(x)
        x = self.b6(x)

        return x

if __name__ == "__main__":
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

    model = ResNet18(Residual).to(device)
    print(summary(model,(1,224,224)))

同时,我们加载 ImageNet 预训练权重并设为 .eval() 模式,通过迁移学习,避免从零训练(无监督场景下无法有效训练深层网络);提取更丰富、更稳定的特征。在PadiM中,我们只是前向推理,不进行反向传播或参数更新。

(3)主代码部分。基于 PaDiM方法的无监督工业异常检测系统,并将其输出转换为 YOLO 格式的伪标签数据集。

import os
import cv2
import numpy as np
import torch
from torch.utils.data import DataLoader
from torchvision import transforms
from sklearn.covariance import EmpiricalCovariance
from pathlib import Path
from tqdm import tqdm
from def_model import PadimModel
from def_dataset import MVTecTrainDataset

# ================== 配置 ==================
CATEGORIES = [
    "bottle", "cable", "capsule", "carpet", "grid",
    "hazelnut", "leather", "metal_nut", "pill", "screw",
    "tile", "toothbrush", "transistor", "wood", "zipper"
]

DATASET_PATH = "./MVTec AD"  # MVTec数据集路径
OUTPUT_DIR = "./output_yolo" # 输出保存yolo格式路径

THRESHOLD = 0.5  # anomaly map 阈值(可调)
MIN_AREA = 100  # 最小缺陷面积
IMAGE_SIZE = 256  # 输入尺寸

os.makedirs(OUTPUT_DIR, exist_ok=True)

if __name__ == "__main__":
    device = "cuda" if torch.cuda.is_available() else "cpu"

    # ================== 主循环 ==================
    for CATEGORY in CATEGORIES:
        print(f"\n{'=' * 60}")
        print(f"Processing category: {CATEGORY}")
        print(f"{'=' * 60}")

        # 1. 数据预处理
        transform = transforms.Compose([
            transforms.ToPILImage(),
            transforms.Resize((IMAGE_SIZE, IMAGE_SIZE)),
            transforms.ToTensor(),
            transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
        ])

        # 2. 加载正常样本并提取特征
        dataset = MVTecTrainDataset(DATASET_PATH, CATEGORY, transform=transform)
        dataloader = DataLoader(dataset, batch_size=32, shuffle=False, num_workers=4)

        padim = PadimModel(device=device, layers=[0, 1, 2])
        all_embeddings = []

        for imgs in tqdm(dataloader, desc="Extracting features"):
            imgs = imgs.to(device)
            emb = padim.embed(imgs)  # [B, C, H, W]
            B, C, H, W = emb.shape
            emb = emb.permute(0, 2, 3, 1).reshape(-1, C)  # [B*H*W, C]
            all_embeddings.append(emb.cpu().numpy())

        all_embeddings = np.concatenate(all_embeddings, axis=0)
        print(f"Total normal features: {all_embeddings.shape}")

        # 3. 计算高斯分布参数(均值 + 协方差)
        mean = np.mean(all_embeddings, axis=0)
        cov = EmpiricalCovariance().fit(all_embeddings)
        inv_cov = torch.from_numpy(cov.precision_).float().to(device)
        mean_tensor = torch.from_numpy(mean).float().to(device)

        # 4. 处理所有测试图像
        test_root = Path(DATASET_PATH) / CATEGORY / "test"
        subdirs = [d for d in test_root.iterdir() if d.is_dir()]

        for subdir in subdirs:
            is_good = (subdir.name == "good")
            img_paths = sorted(subdir.glob("*.png"))

            for img_path in tqdm(img_paths, desc=f"{CATEGORY}/{subdir.name}"):
                # 读原始图
                orig_img = cv2.imread(str(img_path))
                if orig_img is None:
                    continue
                h_orig, w_orig = orig_img.shape[:2]

                # 预处理
                img_rgb = cv2.cvtColor(orig_img, cv2.COLOR_BGR2RGB)
                img_tensor = transform(img_rgb).unsqueeze(0).to(device)  # [1,3,256,256]

                # 提取特征
                emb = padim.embed(img_tensor)  # [1, C, H, W]
                C, H, W = emb.shape[1], emb.shape[2], emb.shape[3]
                emb = emb.squeeze(0).permute(1, 2, 0).reshape(-1, C)  # [H*W, C]

                # 马氏距离
                delta = emb - mean_tensor  # [HW, C]
                dist = torch.diag(delta @ inv_cov @ delta.T)  # [HW]
                anomaly_map = torch.sqrt(dist).reshape(H, W).cpu().numpy()

                # 调整回原始尺寸
                anomaly_map = cv2.resize(anomaly_map, (w_orig, h_orig), interpolation=cv2.INTER_LINEAR)

                # 生成伪标签(仅 defective)
                bboxes_yolo = []
                if not is_good:
                    binary = (anomaly_map > THRESHOLD).astype(np.uint8)
                    contours, _ = cv2.findContours(binary, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
                    for cnt in contours:
                        if cv2.contourArea(cnt) < MIN_AREA:
                            continue
                        x, y, w, h = cv2.boundingRect(cnt)
                        cx, cy = (x + w / 2) / w_orig, (y + h / 2) / h_orig
                        nw, nh = w / w_orig, h / h_orig
                        cx, cy, nw, nh = np.clip([cx, cy, nw, nh], 0, 1)
                        bboxes_yolo.append([0, cx, cy, nw, nh])

                # 保存图像
                rel_path = img_path.relative_to(Path(DATASET_PATH))
                out_img_path = Path(OUTPUT_DIR) / "images" / rel_path
                out_img_path.parent.mkdir(parents=True, exist_ok=True)
                cv2.imwrite(str(out_img_path), orig_img)

                # 保存标签
                label_path = Path(OUTPUT_DIR) / "labels" / f"{rel_path.with_suffix('.txt')}"
                label_path.parent.mkdir(parents=True, exist_ok=True)
                with open(label_path, "w") as f:
                    for box in bboxes_yolo:
                        f.write(" ".join(f"{x:.6f}" for x in box) + "\n")

        print(f" Finished {CATEGORY}")

    print(f"\n YOLO dataset saved to: {OUTPUT_DIR}")

只有非 "good" 目录才生成伪标签(因为 good 图像不应有缺陷)。计算马氏距离时,值越大,越可能是异常。

(4)YOLO格式画框

import cv2
from pathlib import Path

IMG_DIR = Path("output_yolo/images")  # 输出里的images
LABEL_DIR = Path("output_yolo/labels") # 异常的标签坐标(数值)
VIS_DIR = Path("output_yolo/vis_check") # 保存的画框文件
VIS_DIR.mkdir(exist_ok=True)

for img_path in IMG_DIR.rglob("*.png"):
    label_path = LABEL_DIR / img_path.relative_to(IMG_DIR).with_suffix(".txt")
    if not label_path.exists():
        continue

    img = cv2.imread(str(img_path))
    h, w = img.shape[:2]

    with open(label_path) as f:
        for line in f:
            parts = line.strip().split()
            if len(parts) < 5:
                continue
            cls, cx, cy, nw, nh = map(float, parts)
            x1 = int((cx - nw / 2) * w)
            y1 = int((cy - nh / 2) * h)
            x2 = int((cx + nw / 2) * w)
            y2 = int((cy + nh / 2) * h)
            cv2.rectangle(img, (x1, y1), (x2, y2), (0, 0, 255), 2)

    out_path = VIS_DIR / img_path.name
    cv2.imwrite(str(out_path), img)

后续就可以将标签图像进行YOLO训练了......

以上就是基于PadiM的无监督工业数据检测的方案,代码还有很多不足,大家共同学习,相互进步!

Logo

魔乐社区(Modelers.cn) 是一个中立、公益的人工智能社区,提供人工智能工具、模型、数据的托管、展示与应用协同服务,为人工智能开发及爱好者搭建开放的学习交流平台。社区通过理事会方式运作,由全产业链共同建设、共同运营、共同享有,推动国产AI生态繁荣发展。

更多推荐