基于Deep SVDD的图像异常检测模型效果不佳,求技术排查
图像异常检测(Deep SVDD)效果差的问题排查与优化
我是异常检测领域新手,目前用Python实现基于Deep SVDD的图像异常检测模型,目标是识别图像集中的异常样本。数据集包含1000张尺寸各异的猫图像和1000张尺寸各异的熊猫图像,计划用800张猫图像训练模型,测试集采用200张猫图像+2张熊猫图像。预处理步骤为加载图像、统一尺寸、转数组并归一化至[0,1]范围,但模型测试结果极差,以下是我的代码,恳请帮忙排查问题:
import numpy as np import os from PIL import Image from sklearn.metrics import accuracy_score, precision_score, recall_score, roc_auc_score import torch import pandas as pd from deepod.models.tabular import DeepSVDD def load_images_from_folder(folder, size=(128, 128)): images = [] for filename in os.listdir(folder): img = Image.open(os.path.join(folder, filename)) if img is not None: img = img.resize(size, Image.ANTIALIAS) img_array = np.array(img).flatten() # Flatten the image img_array = img_array / 255.0 # Normalize to [0, 1] images.append(img_array) return images # Load and resize images from the 'cat' folder cat_images_array = load_images_from_folder("cat_folder_path") # Load and resize images from the 'panda' folder panda_images_array = load_images_from_folder("panda_folder_path") # split cat images into training and validation sets train_size = int(0.8 * len(cat_images_array)) # 80% of the cat images for training cat_images_train = cat_images_array[:train_size] cat_images_val = cat_images_array[train_size:] panda_images_test = panda_images_array[:10] # combine cat and dog images to create the validation set val_images = cat_images_val + panda_images_test cat_images_train = torch.tensor(cat_images_train) val_images = torch.tensor(val_images) # Initialize DeepSVDD model deep_svdd = DeepSVDD(epochs = 100, batch_size = 30, hidden_dims= '100',lr=0.0001) # Train model deep_svdd.fit(cat_images_train , y = None) # Test model scores = deep_svdd.decision_function(val_images) print(scores) threshold = 0.6 # Compute binary labels for your validation set # Assume that cat_images_val is a list of your actual cat images in the validation set y_true = [1 if i < len(cat_images_val) else 0 for i in range(len(val_images))] # Compute predicted labels based on the anomaly scores and threshold y_pred = [1 if score < threshold else 0 for score in scores] # Calculate metrics accuracy = accuracy_score(y_true, y_pred) precision = precision_score(y_true, y_pred) recall = recall_score(y_true, y_pred) roc_auc = roc_auc_score(y_true, scores) # Note: roc_auc_score uses scores, not labels print(f'Accuracy: {accuracy}') print(f'Precision: {precision}') print(f'Recall: {recall}') print(f'ROC AUC: {roc_auc}') from deepod.metrics import tabular_metrics auc, ap, f1 = tabular_metrics(val_images, scores) from sklearn.metrics import roc_curve
核心问题排查与优化方案
1. 数据预处理缺陷
- 图像通道未统一:代码未处理图像的通道差异(RGB/灰度),直接flatten会导致维度不一致(RGB图是128×128×3=49152维,灰度图是16384维),破坏数据结构。需统一转为RGB:
# 修改load_images_from_folder函数 img = img.convert('RGB') # 强制转为三通道RGB - 数据拆分无随机性:直接按文件夹顺序拆分训练集,可能存在分布偏差(比如某类猫图集中在尾部),需先打乱数据再拆分:
import random random.shuffle(cat_images_array) # 拆分前打乱猫图顺序 train_size = int(0.8 * len(cat_images_array)) cat_images_train = cat_images_array[:train_size] cat_images_val = cat_images_array[train_size:] - 测试集异常样本不符预期:代码取了10张熊猫图,但你描述中计划用2张,需修正为
panda_images_test = panda_images_array[:2],避免样本分布失衡。
2. 模型适配错误
- 误用表格数据模型:你使用的
deepod.models.tabular.DeepSVDD是针对低维表格数据设计的,无法处理高维图像数据的特征提取。需改用基于CNN的DeepSVDD实现,或自定义CNN特征提取器:import torch.nn as nn class CNNEncoder(nn.Module): def __init__(self): super().__init__() self.conv_layers = nn.Sequential( nn.Conv2d(3, 32, kernel_size=3, padding=1), nn.ReLU(), nn.MaxPool2d(2), nn.Conv2d(32, 64, kernel_size=3, padding=1), nn.ReLU(), nn.MaxPool2d(2), nn.Conv2d(64, 128, kernel_size=3, padding=1), nn.ReLU(), nn.MaxPool2d(2), ) self.fc = nn.Linear(128*16*16, 128) # 128×128经三次池化后为16×16 def forward(self, x): x = x.view(-1, 3, 128, 128) # 恢复图像形状:(batch, channels, h, w) x = self.conv_layers(x) x = x.flatten(1) x = self.fc(x) return x - 网络结构过于简单:原模型仅用
hidden_dims='100'的单层网络,无法从49152维的图像数据中提取有效特征,需搭配上述CNN编码器,或使用更深的网络结构。
3. 训练与评估参数问题
- 训练参数不合理:
- 学习率
0.0001过小,可尝试调整为0.001; - 加入早停机制,监控验证集得分,避免过拟合或训练不足;
- 学习率
- 阈值选择不科学:手动设置
threshold=0.6无依据,需根据ROC曲线自动选择最优阈值:fpr, tpr, thresholds = roc_curve(y_true, scores) optimal_idx = np.argmax(tpr - fpr) # 找tpr-fpr最大的点 optimal_threshold = thresholds[optimal_idx] - 评估函数参数错误:
tabular_metrics需传入真实标签而非数据,修正为:auc, ap, f1 = tabular_metrics(y_true, scores)
4. 数据类型规范
转换为torch张量时需指定浮点类型,避免默认整数类型导致计算错误:
cat_images_train = torch.tensor(cat_images_train, dtype=torch.float32) val_images = torch.tensor(val_images, dtype=torch.float32)
内容的提问来源于stack exchange,提问作者Ailisha
相关产品推荐
相关产品推荐

