diff --git a/change_plan b/change_plan new file mode 100644 index 0000000..55eecf3 --- /dev/null +++ b/change_plan @@ -0,0 +1,5 @@ +ResADの場合、例えばmvtecでtestする場合、visaの全データを用いてtrainしている。このときtrainをnormal+dataaug anomalt or normal + anomaly + dataaug anomalyにする? +fewshotなので他のデータセットのクラスの異常にも対応できるようにしたいcutpasteよりはdream?dreamを物体があるに場所しか適応できないように変える。maskのデータはあるからperlinnoiseを生成してそこから絞り込み。 +NSA: cutpasteの応用、コピー元が正常画像、切り取った画像をノイズやブレンドなどを行い異常風に生成、それをほかの正常画像にはりつける。 +SPADEの設定を適応するにはどうするか? +新しいテスト用のデータセットを作る diff --git a/classes.py b/classes.py index 6ab46be..529730f 100644 --- a/classes.py +++ b/classes.py @@ -36,4 +36,19 @@ MVTEC_TO_BRATS = {'seen': ['bottle', 'cable', 'capsule', 'carpet', 'grid', 'hazelnut', 'leather', 'metal_nut', 'pill', 'screw', 'tile', 'toothbrush', 'transistor', 'wood', 'zipper'], - 'unseen': ['brain']} \ No newline at end of file + 'unseen': ['brain']} + +MVTEC_TO_MVTEC = {'seen': ['bottle', 'cable', 'capsule', 'carpet', 'grid', +'hazelnut', 'leather', 'metal_nut', 'pill', 'screw', +'tile', 'toothbrush', 'transistor', 'wood', 'zipper'], +'unseen': ['bottle', 'cable', 'capsule', 'carpet', 'grid', +'hazelnut', 'leather', 'metal_nut', 'pill', 'screw', +'tile', 'toothbrush', 'transistor', 'wood', 'zipper']} + +VISA_TO_VISA = {'seen': ['candle', 'capsules', 'cashew', 'chewinggum', 'fryum', +'macaroni1', 'macaroni2', 'pcb1', 'pcb2', 'pcb3', 'pcb4', 'pipe_fryum'], +'unseen': ['candle', 'capsules', 'cashew', 'chewinggum', 'fryum', +'macaroni1', 'macaroni2', 'pcb1', 'pcb2', 'pcb3', 'pcb4', 'pipe_fryum']} + +CAPSULES_TO_CAPSULES = {'seen': ['capsules'], + 'unseen': ['capsules']} diff --git a/datasets/capsules.py b/datasets/capsules.py new file mode 100644 index 0000000..c7111e3 --- /dev/null +++ b/datasets/capsules.py @@ -0,0 +1,309 @@ +import os +import pandas +import torch +import numpy as np +from PIL import Image +from typing import Callable, Optional +from torch.utils.data import Dataset +from torchvision import transforms as T +from torchvision.transforms.transforms import RandomHorizontalFlip + + +IMAGENET_MEAN = [0.485, 0.456, 0.406] +IMAGENET_STD = [0.229, 0.224, 0.225] + + +class CAPSULES(Dataset): + + CLASS_NAMES = ['capsules'] + + def __init__(self, + root: str, + class_name: str, + train: bool = True, + normalize: str = 'imagebind', + transform: Optional[Callable] = None, + target_transform: Optional[Callable] = None, + **kwargs): + + self.root = root + self.class_name = class_name + self.train = train + self.cropsize = [kwargs.get('crp_size'), kwargs.get('crp_size')] + + # load dataset + if isinstance(self.class_name, str): + self.image_paths, self.labels, self.mask_paths, self.class_names = self._load_data(self.class_name) + elif self.class_name is None: # load all classes + self.image_paths, self.labels, self.mask_paths, self.class_names = self._load_all_data() + else: + self.image_paths, self.labels, self.mask_paths, self.class_names = self._load_all_data(self.class_name) + + # set transforms + if normalize == "imagebind": + self.transform = T.Compose( # for imagebind + [ + T.Resize( + 224, interpolation=T.InterpolationMode.BICUBIC + ), + T.CenterCrop(224), + T.ToTensor(), + T.Normalize( + mean=(0.48145466, 0.4578275, 0.40821073), + std=(0.26862954, 0.26130258, 0.27577711), + ), + ] + ) + else: + self.transform = T.Compose([ + T.Resize(kwargs.get('img_size', 224), T.InterpolationMode.BICUBIC), + T.CenterCrop(kwargs.get('crp_size', 224)), + T.ToTensor(), + T.Compose([T.Normalize(IMAGENET_MEAN, IMAGENET_STD)])]) + + # mask + self.target_transform = T.Compose([ + T.Resize(kwargs.get('img_size'), Image.NEAREST), + T.CenterCrop(kwargs.get('crp_size')), + T.ToTensor()]) + + self.class_to_idx = {'capsules': 0 } + self.idx_to_class = {0:'capsules' } + + def __getitem__(self, idx): + image_path, label, mask, class_name = self.image_paths[idx], self.labels[idx], self.mask_paths[idx], self.class_names[idx] + + image = Image.open(image_path).convert('RGB') + image = self.transform(image) + + if label == 0: + mask = torch.zeros([1, self.cropsize[0], self.cropsize[1]]) + else: + mask = Image.open(mask) + mask = np.array(mask) + mask[mask != 0] = 255 + mask = Image.fromarray(mask) + mask = self.target_transform(mask) + + if self.train: + label = self.class_to_idx[class_name] + + return image, label, mask, class_name + + def __len__(self): + return len(self.image_paths) + + def _load_data(self, class_name): + split_csv_file = os.path.join(self.root, 'split_csv', '1cls.csv') + csv_data = pandas.read_csv(split_csv_file) + + class_data = csv_data.loc[csv_data['object'] == class_name] + + if self.train: + train_data = class_data.loc[class_data['split'] == 'train'] + image_paths = train_data['image'].to_list() + image_paths = [os.path.join(self.root, file_name) for file_name in image_paths] + labels = [0] * len(image_paths) + mask_paths = [None] * len(image_paths) + else: + image_paths, labels, mask_paths = [], [], [] + + test_data = class_data.loc[class_data['split'] == 'test'] + test_normal_data = test_data.loc[test_data['label'] == 'normal'] + test_anomaly_data = test_data.loc[test_data['label'] == 'anomaly'] + + normal_image_paths = test_normal_data['image'].to_list() + normal_image_paths = [os.path.join(self.root, file_name) for file_name in normal_image_paths] + image_paths.extend(normal_image_paths) + labels.extend([0] * len(normal_image_paths)) + mask_paths.extend([None] * len(normal_image_paths)) + + anomaly_image_paths = test_anomaly_data['image'].to_list() + anomaly_mask_paths = test_anomaly_data['mask'].to_list() + anomaly_image_paths = [os.path.join(self.root, file_name) for file_name in anomaly_image_paths] + anomaly_mask_paths = [os.path.join(self.root, file_name) for file_name in anomaly_mask_paths] + image_paths.extend(anomaly_image_paths) + labels.extend([1] * len(anomaly_image_paths)) + mask_paths.extend(anomaly_mask_paths) + + class_names = [class_name] * len(image_paths) + return image_paths, labels, mask_paths, class_names + + def _load_all_data(self, class_names=None): + all_image_paths = [] + all_labels = [] + all_mask_paths = [] + all_class_names = [] + CLASS_NAMES = class_names if class_names is not None else self.CLASS_NAMES + for class_name in CLASS_NAMES: + image_paths, labels, mask_paths, class_names = self._load_data(class_name) + all_image_paths.extend(image_paths) + all_labels.extend(labels) + all_mask_paths.extend(mask_paths) + all_class_names.extend(class_names) + return all_image_paths, all_labels, all_mask_paths, all_class_names + + def update_class_to_idx(self, class_to_idx): + for class_name in self.class_to_idx.keys(): + self.class_to_idx[class_name] = class_to_idx[class_name] + class_names = self.class_to_idx.keys() + idxs = self.class_to_idx.values() + self.idx_to_class = dict(zip(idxs, class_names)) + + +class CAPSULESANO(Dataset): + + CLASS_NAMES = ['capsules'] + + def __init__(self, + root: str, + class_name: str, + train: bool = True, + normalize: str = 'imagebind', + transform: Optional[Callable] = None, + target_transform: Optional[Callable] = None, + **kwargs): + + self.root = root + self.class_name = class_name + self.train = train + self.cropsize = [kwargs.get('crp_size'), kwargs.get('crp_size')] + + # load dataset + if isinstance(self.class_name, str): + self.image_paths, self.labels, self.mask_paths, self.class_names = self._load_data(self.class_name) + elif self.class_name is None: + self.image_paths, self.labels, self.mask_paths, self.class_names = self._load_all_data() + else: + self.image_paths, self.labels, self.mask_paths, self.class_names = self._load_all_data(self.class_name) + + if normalize == "imagebind": + self.transform = T.Compose( # for imagebind + [ + T.Resize( + 224, interpolation=T.InterpolationMode.BICUBIC + ), + T.CenterCrop(224), + T.ToTensor(), + T.Normalize( + mean=(0.48145466, 0.4578275, 0.40821073), + std=(0.26862954, 0.26130258, 0.27577711), + ), + ] + ) + else: + self.transform = T.Compose([ + T.Resize(kwargs.get('img_size', 224), T.InterpolationMode.BICUBIC), + T.CenterCrop(kwargs.get('crp_size', 224)), + T.ToTensor(), + T.Compose([T.Normalize(IMAGENET_MEAN, IMAGENET_STD)])]) + + # mask + self.target_transform = T.Compose([ + T.Resize(kwargs.get('img_size'), Image.NEAREST), + T.CenterCrop(kwargs.get('crp_size')), + T.ToTensor()]) + + self.class_to_idx = {'capsules':0} + self.idx_to_class = {0:'capsules'} + + def __getitem__(self, idx): + image_path, label, mask, class_name = self.image_paths[idx], self.labels[idx], self.mask_paths[idx], self.class_names[idx] + + image = Image.open(image_path).convert('RGB') + image = self.transform(image) + + if label == 0: + mask = torch.zeros([1, self.cropsize[0], self.cropsize[1]]) + else: + mask = Image.open(mask) + mask = np.array(mask) + mask[mask != 0] = 255 + mask = Image.fromarray(mask) + mask = self.target_transform(mask) + + if self.train: + label = self.class_to_idx[class_name] + + return image, label, mask, class_name + + def __len__(self): + return len(self.image_paths) + + def _load_data(self, class_name): + split_csv_file = os.path.join(self.root, 'split_csv', '1cls.csv') + csv_data = pandas.read_csv(split_csv_file) + + class_data = csv_data.loc[csv_data['object'] == class_name] + all_image_paths, all_labels, all_mask_paths = [], [], [] + + # train + train_data = class_data.loc[class_data['split'] == 'train'] + image_paths = train_data['image'].to_list() + image_paths = [os.path.join(self.root, file_name) for file_name in image_paths] + labels = [0] * len(image_paths) + mask_paths = [None] * len(image_paths) + all_image_paths.extend(image_paths) + all_labels.extend(labels) + all_mask_paths.extend(mask_paths) + + # test + image_paths, labels, mask_paths = [], [], [] + test_data = class_data.loc[class_data['split'] == 'test'] + test_normal_data = test_data.loc[test_data['label'] == 'normal'] + test_anomaly_data = test_data.loc[test_data['label'] == 'anomaly'] + + normal_image_paths = test_normal_data['image'].to_list() + normal_image_paths = [os.path.join(self.root, file_name) for file_name in normal_image_paths] + image_paths.extend(normal_image_paths) + labels.extend([0] * len(normal_image_paths)) + mask_paths.extend([None] * len(normal_image_paths)) + + anomaly_image_paths = test_anomaly_data['image'].to_list() + anomaly_mask_paths = test_anomaly_data['mask'].to_list() + anomaly_image_paths = [os.path.join(self.root, file_name) for file_name in anomaly_image_paths] + anomaly_mask_paths = [os.path.join(self.root, file_name) for file_name in anomaly_mask_paths] + image_paths.extend(anomaly_image_paths) + labels.extend([1] * len(anomaly_image_paths)) + mask_paths.extend(anomaly_mask_paths) + + all_image_paths.extend(image_paths) + all_labels.extend(labels) + all_mask_paths.extend(mask_paths) + + class_names = [class_name] * len(all_image_paths) + return all_image_paths, all_labels, all_mask_paths, class_names + + def _load_all_data(self, class_names=None): + all_image_paths = [] + all_labels = [] + all_mask_paths = [] + all_class_names = [] + CLASS_NAMES = class_names if class_names is not None else self.CLASS_NAMES + for class_name in CLASS_NAMES: + image_paths, labels, mask_paths, class_names = self._load_data(class_name) + all_image_paths.extend(image_paths) + all_labels.extend(labels) + all_mask_paths.extend(mask_paths) + all_class_names.extend(class_names) + return all_image_paths, all_labels, all_mask_paths, all_class_names + + def update_class_to_idx(self, class_to_idx): + for class_name in self.class_to_idx.keys(): + self.class_to_idx[class_name] = class_to_idx[class_name] + class_names = self.class_to_idx.keys() + idxs = self.class_to_idx.values() + self.idx_to_class = dict(zip(idxs, class_names)) + + +def get_normal_image_paths_visa(root, class_name): + split_csv_file = os.path.join(root, 'split_csv', '1cls.csv') + csv_data = pandas.read_csv(split_csv_file) + + class_data = csv_data.loc[csv_data['object'] == class_name] + + train_data = class_data.loc[class_data['split'] == 'train'] + image_paths = train_data['image'].to_list() + image_paths = [os.path.join(root, file_name) for file_name in image_paths] + + return image_paths diff --git a/datasets/mvtec.py b/datasets/mvtec.py index c52ef6b..495620b 100644 --- a/datasets/mvtec.py +++ b/datasets/mvtec.py @@ -66,11 +66,11 @@ def __init__( # load dataset if isinstance(self.class_name, str): - self.image_paths, self.labels, self.mask_paths, self.class_names = self._load_data(self.class_name) + self.image_paths, self.labels, self.mask_paths, self.class_names,self.anomaly_types = self._load_data(self.class_name) elif self.class_name is None: - self.image_paths, self.labels, self.mask_paths, self.class_names = self._load_all_data() + self.image_paths, self.labels, self.mask_paths, self.class_names,self.anomaly_types = self._load_all_data() else: - self.image_paths, self.labels, self.mask_paths, self.class_names = self._load_all_data(self.class_name) + self.image_paths, self.labels, self.mask_paths, self.class_names,self.anomaly_types = self._load_all_data(self.class_name) if normalize == "imagebind": self.transform = T.Compose( # for imagebind @@ -109,10 +109,10 @@ def __init__( 12: 'transistor', 13: 'wood', 14: 'zipper'} def __getitem__(self, idx): - image_path, label, mask_path, class_name = self.image_paths[idx], self.labels[idx], self.mask_paths[idx], self.class_names[idx] + image_path, label, mask_path, class_name,anomaly_type = self.image_paths[idx], self.labels[idx], self.mask_paths[idx], self.class_names[idx],self.anomaly_types[idx] img, label, mask = self._load_image_and_mask(image_path, label, mask_path) - return img, label, mask, class_name + return img, label, mask, class_name,anomaly_type def _load_image_and_mask(self, image_path, label, mask_path): img = Image.open(image_path).convert('RGB') @@ -136,7 +136,9 @@ def __len__(self): return len(self.image_paths) def _load_data(self, class_name): - image_paths, labels, mask_paths = [], [], [] + image_paths, labels, mask_paths,anomaly_types = [], [], [], [] + class_names_list = [] # Initialize class_names_list here + for phase in ['train', 'test']: image_dir = os.path.join(self.root, class_name, phase) @@ -157,6 +159,7 @@ def _load_data(self, class_name): if img_type == 'good': labels.extend([0] * len(img_fpath_list)) mask_paths.extend([None] * len(img_fpath_list)) + anomaly_types.extend(['good'] * len(img_fpath_list)) # 'good'を追加 else: labels.extend([1] * len(img_fpath_list)) gt_type_dir = os.path.join(mask_dir, img_type) @@ -164,23 +167,26 @@ def _load_data(self, class_name): gt_fpath_list = [os.path.join(gt_type_dir, img_fname + '_mask.png') for img_fname in img_fname_list] mask_paths.extend(gt_fpath_list) + anomaly_types.extend([img_type] * len(img_fpath_list)) # 異常タイプ名を追加 - class_names = [class_name] * len(image_paths) - return image_paths, labels, mask_paths, class_names + class_names_list = [class_name] * len(image_paths) # 変数名が衝突しないように変更 + return image_paths, labels, mask_paths, class_names_list, anomaly_types # anomaly_types も返す def _load_all_data(self, class_names=None): all_image_paths = [] all_labels = [] all_mask_paths = [] all_class_names = [] + all_anomaly_types = [] # anomaly_types を追加 CLASS_NAMES = class_names if class_names is not None else self.CLASS_NAMES for class_name in CLASS_NAMES: - image_paths, labels, mask_paths, class_names = self._load_data(class_name) + image_paths, labels, mask_paths, class_names,anomaly_types_from_load_data = self._load_data(class_name) all_image_paths.extend(image_paths) all_labels.extend(labels) all_mask_paths.extend(mask_paths) all_class_names.extend(class_names) - return all_image_paths, all_labels, all_mask_paths, all_class_names + all_anomaly_types.extend(anomaly_types_from_load_data) # anomaly_types を追加 + return all_image_paths, all_labels, all_mask_paths, all_class_names, all_anomaly_types # anomaly_types も返す class MVTEC(Dataset): @@ -201,11 +207,11 @@ def __init__(self, self.cropsize = [kwargs.get('msk_crp_size'), kwargs.get('msk_crp_size')] if isinstance(self.class_name, str): - self.image_paths, self.labels, self.mask_paths, self.class_names = self._load_data(self.class_name) + self.image_paths, self.labels, self.mask_paths, self.class_names,self.anomaly_types = self._load_data(self.class_name) elif self.class_name is None: - self.image_paths, self.labels, self.mask_paths, self.class_names = self._load_all_data() + self.image_paths, self.labels, self.mask_paths, self.class_names,self.anomaly_types = self._load_all_data() else: - self.image_paths, self.labels, self.mask_paths, self.class_names = self._load_all_data(self.class_name) + self.image_paths, self.labels, self.mask_paths, self.class_names,self.anomaly_types = self._load_all_data(self.class_name) # set transforms if normalize == "imagebind": @@ -239,10 +245,10 @@ def __len__(self): return len(self.image_paths) def __getitem__(self, idx): - image_path, label, mask_path, class_name = self.image_paths[idx], self.labels[idx], self.mask_paths[idx], self.class_names[idx] + image_path, label, mask_path, class_name,anomaly_type = self.image_paths[idx], self.labels[idx], self.mask_paths[idx], self.class_names[idx], self.anomaly_types[idx] # anomaly_type を追加 img, label, mask = self._load_image_and_mask(image_path, label, mask_path) - return img, label, mask, class_name + return img, label, mask, class_name,anomaly_type # anomaly_type を返す def _load_image_and_mask(self, image_path, label, mask_path): img = Image.open(image_path).convert('RGB') @@ -263,7 +269,9 @@ def _load_image_and_mask(self, image_path, label, mask_path): return img, label, mask def _load_data(self, class_name): - image_paths, labels, mask_paths = [], [], [] + image_paths, labels, mask_paths,anomaly_types = [], [], [], [] + class_names_list = [] # Initialize class_names_list here + phase = 'train' if self.train else 'test' image_dir = os.path.join(self.root, class_name, phase) @@ -284,6 +292,7 @@ def _load_data(self, class_name): if img_type == 'good': labels.extend([0] * len(img_fpath_list)) mask_paths.extend([None] * len(img_fpath_list)) + anomaly_types.extend(['good'] * len(img_fpath_list)) # 'good'を追加 else: labels.extend([1] * len(img_fpath_list)) gt_type_dir = os.path.join(mask_dir, img_type) @@ -291,23 +300,26 @@ def _load_data(self, class_name): gt_fpath_list = [os.path.join(gt_type_dir, img_fname + '_mask.png') for img_fname in img_fname_list] mask_paths.extend(gt_fpath_list) + anomaly_types.extend([img_type] * len(img_fpath_list)) # 異常タイプ名を追加 - class_names = [class_name] * len(image_paths) - return image_paths, labels, mask_paths, class_names + class_names_list = [class_name] * len(image_paths) # 変数名が衝突しないように変更 + return image_paths, labels, mask_paths, class_names_list, anomaly_types # anomaly_types も返す def _load_all_data(self, class_names=None): all_image_paths = [] all_labels = [] all_mask_paths = [] all_class_names = [] + all_anomaly_types = [] # anomaly_types を追加 CLASS_NAMES = class_names if class_names is not None else self.CLASS_NAMES for class_name in CLASS_NAMES: - image_paths, labels, mask_paths, class_names = self._load_data(class_name) + image_paths, labels, mask_paths, class_names, anomaly_types_from_load_data = self._load_data(class_name) all_image_paths.extend(image_paths) all_labels.extend(labels) all_mask_paths.extend(mask_paths) all_class_names.extend(class_names) - return all_image_paths, all_labels, all_mask_paths, all_class_names + all_anomaly_types.extend(anomaly_types_from_load_data) # anomaly_types を追加 + return all_image_paths, all_labels, all_mask_paths, all_class_names, all_anomaly_types # anomaly_types も返す def get_normal_image_paths_mvtec(root, class_name): @@ -326,4 +338,4 @@ def get_normal_image_paths_mvtec(root, class_name): for f in os.listdir(img_type_dir)]) image_paths.extend(img_fpath_list) - return image_paths \ No newline at end of file + return image_paths diff --git a/datasets/mvtec_fewclass.py b/datasets/mvtec_fewclass.py new file mode 100644 index 0000000..709f3d7 --- /dev/null +++ b/datasets/mvtec_fewclass.py @@ -0,0 +1,331 @@ +""" +The dataset defined in this script is only used for cross-class training, +where we use both normal and abnormal samples for training. And we use all +abnormal samples from the test set, as these abnormal samples will not be +tested in cross-class setting. +""" +import os +import random +from typing import Any, Callable, Optional, Tuple +import torch +import numpy as np +from PIL import Image +from torchvision.transforms.transforms import RandomHorizontalFlip +from tqdm import tqdm +from torch.utils.data import Dataset +from torchvision import transforms as T + +import cv2 +import glob +import imgaug.augmenters as iaa +import albumentations as A + + +IMAGENET_MEAN = [0.485, 0.456, 0.406] +IMAGENET_STD = [0.229, 0.224, 0.225] + + +class MVTECFEWANO(Dataset): + """This dataset is used for cross-class training, where we use all the normal and abnomal + samples for training. As we will not test on the training classes, using abnormal samples + in test set is actually reasonable. + + Args: + root (string): Root directory of dataset, i.e ``../../mvtec_anomaly_detection``. + train (bool, optional): If True, creates dataset for training, otherwise for testing. + download (bool, optional): If true, downloads the dataset from the internet and + puts it in root directory. If dataset is already downloaded, it is not + downloaded again. + transform (callable, optional): A function/transform that takes in an PIL image + and returns a transformed version. E.g, ``transforms.Resize`` + target_transform (callable, optional): A function/transform that takes in the + target and transforms it. + """ + + MVTEC_URL = 'ftp://guest:GU.205dldo@ftp.softronics.ch/mvtec_anomaly_detection/mvtec_anomaly_detection.tar.xz' + + CLASS_NAMES = ['capsule','screw','transistor'] + + def __init__( + self, + root: str, + class_name: str, + train: bool = True, + normalize: str = 'imagebind', + transform: Optional[Callable] = None, + target_transform: Optional[Callable] = None, + download: bool = False, + **kwargs): + + self.root = root + self.class_name = class_name + self.train = train + self.cropsize = [kwargs.get('msk_crp_size'), kwargs.get('msk_crp_size')] + + # load dataset + if isinstance(self.class_name, str): + self.image_paths, self.labels, self.mask_paths, self.class_names,self.anomaly_types = self._load_data(self.class_name) + elif self.class_name is None: + self.image_paths, self.labels, self.mask_paths, self.class_names,self.anomaly_types = self._load_all_data() + else: + self.image_paths, self.labels, self.mask_paths, self.class_names,self.anomaly_types = self._load_all_data(self.class_name) + + if normalize == "imagebind": + self.transform = T.Compose( # for imagebind + [ + T.Resize( + 224, interpolation=T.InterpolationMode.BICUBIC + ), + T.CenterCrop(224), + T.ToTensor(), + T.Normalize( + mean=(0.48145466, 0.4578275, 0.40821073), + std=(0.26862954, 0.26130258, 0.27577711), + ), + ] + ) + else: + self.transform = T.Compose([ + T.Resize(kwargs.get('img_size', 224), T.InterpolationMode.BICUBIC), + T.CenterCrop(kwargs.get('crp_size', 224)), + T.ToTensor(), + T.Compose([T.Normalize(IMAGENET_MEAN, IMAGENET_STD)])]) + + # mask + self.target_transform = T.Compose([ + T.Resize(kwargs.get('msk_size'), Image.NEAREST), + T.CenterCrop(kwargs.get('msk_crp_size')), + T.ToTensor()]) + + self.class_to_idx = {'capsule': 0,'screw': 1, 'toothbrush': 2} + self.idx_to_class = { 0: 'capsule', 1: 'screw', 2: 'transistor'} + + def __getitem__(self, idx): + image_path, label, mask_path, class_name,anomaly_type = self.image_paths[idx], self.labels[idx], self.mask_paths[idx], self.class_names[idx],self.anomaly_types[idx] + img, label, mask = self._load_image_and_mask(image_path, label, mask_path) + + return img, label, mask, class_name,anomaly_type + + def _load_image_and_mask(self, image_path, label, mask_path): + img = Image.open(image_path).convert('RGB') + class_name = image_path.split('/')[-4] + # if class_name in ['zipper', 'screw', 'grid']: # handle greyscale classes + # img = np.expand_dims(np.asarray(img), axis=2) + # img = np.concatenate([img, img, img], axis=2) + # img = Image.fromarray(img.astype('uint8')).convert('RGB') + # + img = self.transform(img) + # + if label == 0: + mask = torch.zeros([1, self.cropsize[0], self.cropsize[1]]) + else: + mask = Image.open(mask_path) + mask = self.target_transform(mask) + + return img, label, mask + + def __len__(self): + return len(self.image_paths) + + def _load_data(self, class_name): + image_paths, labels, mask_paths,anomaly_types = [], [], [], [] + class_names_list = [] # Initialize class_names_list here + + + for phase in ['train', 'test']: + image_dir = os.path.join(self.root, class_name, phase) + mask_dir = os.path.join(self.root, class_name, 'ground_truth') + + img_types = sorted(os.listdir(image_dir)) + for img_type in img_types: + # load images + img_type_dir = os.path.join(image_dir, img_type) + if not os.path.isdir(img_type_dir): + continue + img_fpath_list = sorted([os.path.join(img_type_dir, f) + for f in os.listdir(img_type_dir) + if f.endswith('.png')]) + image_paths.extend(img_fpath_list) + + # load gt labels + if img_type == 'good': + labels.extend([0] * len(img_fpath_list)) + mask_paths.extend([None] * len(img_fpath_list)) + anomaly_types.extend(['good'] * len(img_fpath_list)) # 'good'を追加 + else: + labels.extend([1] * len(img_fpath_list)) + gt_type_dir = os.path.join(mask_dir, img_type) + img_fname_list = [os.path.splitext(os.path.basename(f))[0] for f in img_fpath_list] + gt_fpath_list = [os.path.join(gt_type_dir, img_fname + '_mask.png') + for img_fname in img_fname_list] + mask_paths.extend(gt_fpath_list) + anomaly_types.extend([img_type] * len(img_fpath_list)) # 異常タイプ名を追加 + + class_names_list = [class_name] * len(image_paths) # 変数名が衝突しないように変更 + return image_paths, labels, mask_paths, class_names_list, anomaly_types # anomaly_types も返す + + def _load_all_data(self, class_names=None): + all_image_paths = [] + all_labels = [] + all_mask_paths = [] + all_class_names = [] + all_anomaly_types = [] # anomaly_types を追加 + CLASS_NAMES = class_names if class_names is not None else self.CLASS_NAMES + for class_name in CLASS_NAMES: + image_paths, labels, mask_paths, class_names,anomaly_types_from_load_data = self._load_data(class_name) + all_image_paths.extend(image_paths) + all_labels.extend(labels) + all_mask_paths.extend(mask_paths) + all_class_names.extend(class_names) + all_anomaly_types.extend(anomaly_types_from_load_data) # anomaly_types を追加 + return all_image_paths, all_labels, all_mask_paths, all_class_names, all_anomaly_types # anomaly_types も返す + + +class MVTECFEW(Dataset): + + CLASS_NAMES = ['capsule','screw','transistor'] + def __init__(self, + root: str, + class_name: str = 'capsule', + train: bool = True, + normalize: str = 'imagebind', + **kwargs) -> None: + + self.root = root + self.class_name = class_name + self.train = train + self.cropsize = [kwargs.get('msk_crp_size'), kwargs.get('msk_crp_size')] + + if isinstance(self.class_name, str): + self.image_paths, self.labels, self.mask_paths, self.class_names,self.anomaly_types = self._load_data(self.class_name) + elif self.class_name is None: + self.image_paths, self.labels, self.mask_paths, self.class_names,self.anomaly_types = self._load_all_data() + else: + self.image_paths, self.labels, self.mask_paths, self.class_names,self.anomaly_types = self._load_all_data(self.class_name) + + # set transforms + if normalize == "imagebind": + self.transform = T.Compose( # for imagebind + [ + T.Resize( + 224, interpolation=T.InterpolationMode.BICUBIC + ), + T.CenterCrop(224), + T.ToTensor(), + T.Normalize( + mean=(0.48145466, 0.4578275, 0.40821073), + std=(0.26862954, 0.26130258, 0.27577711), + ), + ] + ) + else: + self.transform = T.Compose([ + T.Resize(kwargs.get('img_size', 224), T.InterpolationMode.BICUBIC), + T.CenterCrop(kwargs.get('crp_size', 224)), + T.ToTensor(), + T.Compose([T.Normalize(IMAGENET_MEAN, IMAGENET_STD)])]) + + # mask + self.target_transform = T.Compose([ + T.Resize(kwargs.get('msk_size', 224), Image.NEAREST), + T.CenterCrop(kwargs.get('msk_crp_size', 224)), + T.ToTensor()]) + + def __len__(self): + return len(self.image_paths) + + def __getitem__(self, idx): + image_path, label, mask_path, class_name,anomaly_type = self.image_paths[idx], self.labels[idx], self.mask_paths[idx], self.class_names[idx], self.anomaly_types[idx] # anomaly_type を追加 + img, label, mask = self._load_image_and_mask(image_path, label, mask_path) + + return img, label, mask, class_name,anomaly_type # anomaly_type を返す + + def _load_image_and_mask(self, image_path, label, mask_path): + img = Image.open(image_path).convert('RGB') + class_name = image_path.split('/')[-4] + # if class_name in ['zipper', 'screw', 'grid']: # handle greyscale classes + # img = np.expand_dims(np.asarray(img), axis=2) + # img = np.concatenate([img, img, img], axis=2) + # img = Image.fromarray(img.astype('uint8')).convert('RGB') + # + img = self.transform(img) + # + if label == 0: + mask = torch.zeros([1, self.cropsize[0], self.cropsize[1]]) + else: + mask = Image.open(mask_path) + mask = self.target_transform(mask) + + return img, label, mask + + def _load_data(self, class_name): + image_paths, labels, mask_paths,anomaly_types = [], [], [], [] + class_names_list = [] # Initialize class_names_list here + + phase = 'train' if self.train else 'test' + + image_dir = os.path.join(self.root, class_name, phase) + mask_dir = os.path.join(self.root, class_name, 'ground_truth') + + img_types = sorted(os.listdir(image_dir)) + for img_type in img_types: + # load images + img_type_dir = os.path.join(image_dir, img_type) + if not os.path.isdir(img_type_dir): + continue + img_fpath_list = sorted([os.path.join(img_type_dir, f) + for f in os.listdir(img_type_dir) + if f.endswith('.png')]) + image_paths.extend(img_fpath_list) + + # load gt labels + if img_type == 'good': + labels.extend([0] * len(img_fpath_list)) + mask_paths.extend([None] * len(img_fpath_list)) + anomaly_types.extend(['good'] * len(img_fpath_list)) # 'good'を追加 + else: + labels.extend([1] * len(img_fpath_list)) + gt_type_dir = os.path.join(mask_dir, img_type) + img_fname_list = [os.path.splitext(os.path.basename(f))[0] for f in img_fpath_list] + gt_fpath_list = [os.path.join(gt_type_dir, img_fname + '_mask.png') + for img_fname in img_fname_list] + mask_paths.extend(gt_fpath_list) + anomaly_types.extend([img_type] * len(img_fpath_list)) # 異常タイプ名を追加 + + class_names_list = [class_name] * len(image_paths) # 変数名が衝突しないように変更 + return image_paths, labels, mask_paths, class_names_list, anomaly_types # anomaly_types も返す + + def _load_all_data(self, class_names=None): + all_image_paths = [] + all_labels = [] + all_mask_paths = [] + all_class_names = [] + all_anomaly_types = [] # anomaly_types を追加 + CLASS_NAMES = class_names if class_names is not None else self.CLASS_NAMES + for class_name in CLASS_NAMES: + image_paths, labels, mask_paths, class_names, anomaly_types_from_load_data = self._load_data(class_name) + all_image_paths.extend(image_paths) + all_labels.extend(labels) + all_mask_paths.extend(mask_paths) + all_class_names.extend(class_names) + all_anomaly_types.extend(anomaly_types_from_load_data) # anomaly_types を追加 + return all_image_paths, all_labels, all_mask_paths, all_class_names, all_anomaly_types # anomaly_types も返す + + +def get_normal_image_paths_mvtec(root, class_name): + phase = 'train' + image_paths = [] + + image_dir = os.path.join(root, class_name, phase) + + img_types = sorted(os.listdir(image_dir)) + for img_type in img_types: + # load images + img_type_dir = os.path.join(image_dir, img_type) + if not os.path.isdir(img_type_dir): + continue + img_fpath_list = sorted([os.path.join(img_type_dir, f) + for f in os.listdir(img_type_dir)]) + image_paths.extend(img_fpath_list) + + return image_paths diff --git a/extract_ref_features.py b/extract_ref_features.py index 384d475..fbf503e 100644 --- a/extract_ref_features.py +++ b/extract_ref_features.py @@ -9,6 +9,7 @@ import torch.nn as nn from torch.utils.data import DataLoader, Dataset import torchvision.transforms as T +from models.fc_flow import load_flow_model from datasets.mvtec import MVTEC from datasets.visa import VISA @@ -18,7 +19,7 @@ from datasets.mvtec_loco import MVTECLOCO from datasets.brats import BRATS from models.imagebind import ImageBindModel - +from utils import load_weights class FEWSHOTDATA(Dataset): @@ -113,10 +114,20 @@ def main(args): image_size = 224 device = 'cuda:0' root_dir = args.few_shot_dir - encoder = timm.create_model("wide_resnet50_2", features_only=True, - out_indices=(1, 2, 3), pretrained=True).eval() - encoder.to(device) + if args.backbone == 'wide_resnet50_2': + encoder = timm.create_model('wide_resnet50_2', features_only=True, + out_indices=(1, 2, 3), pretrained=True).eval() # the pretrained checkpoint will be in /home/.cache/torch/hub/checkpoints/ + encoder = encoder.to(device) + elif args.backbone == 'tf_efficientnet_b6':#10/26追加 + encoder = timm.create_model('tf_efficientnet_b6', features_only=True, + out_indices=(1, 2, 3), pretrained=True).eval() # the pretrained checkpoint will be in /home/.cache/torch/hub/checkpoints/ + encoder = encoder.to(device) + feat_dims = encoder.feature_info.channels() + decoders = [load_flow_model(args, feat_dim) for feat_dim in feat_dims] + decoders = [decoder.to(args.device) for decoder in decoders] + if args.bgadweight_dir: + load_weights(encoder, decoders, args.bgadweight_dir) if args.dataset in SETTINGS.keys(): CLASS_NAMES = SETTINGS[args.dataset] else: @@ -143,14 +154,21 @@ def main(args): print(layer1_features.shape) print(layer2_features.shape) print(layer3_features.shape) - - layer1_features = layer1_features.permute(0, 2, 3, 1).reshape(-1, 256) - layer2_features = layer2_features.permute(0, 2, 3, 1).reshape(-1, 512) - layer3_features = layer3_features.permute(0, 2, 3, 1).reshape(-1, 1024) - + #修正10/26 + layer1_channels = layer1_features.shape[1] + layer2_channels = layer2_features.shape[1] + layer3_channels = layer3_features.shape[1] + + layer1_features = layer1_features.permute(0, 2, 3, 1).reshape(-1, layer1_channels) + layer2_features = layer2_features.permute(0, 2, 3, 1).reshape(-1, layer2_channels) + layer3_features = layer3_features.permute(0, 2, 3, 1).reshape(-1, layer3_channels) + #修正終わり os.makedirs(os.path.join(args.save_dir, class_name), exist_ok=True) + print(f"Attempting to save layer1.npy for {class_name}...") np.save(os.path.join(args.save_dir, class_name, 'layer1.npy'), layer1_features.cpu().numpy()) + print(f"Successfully saved layer1.npy for {class_name}.") + np.save(os.path.join(args.save_dir, class_name, 'layer2.npy'), layer2_features.cpu().numpy()) np.save(os.path.join(args.save_dir, class_name, 'layer3.npy'), layer3_features.cpu().numpy()) @@ -223,7 +241,13 @@ def main2(args): parser = argparse.ArgumentParser() parser.add_argument('--dataset', type=str, default="mvtec") parser.add_argument('--few_shot_dir', type=str, default="./4shot/mvtec") + parser.add_argument('--flow_arch', type=str, default='conditional_flow_model') + parser.add_argument('--bgadweight_dir', type=str, default="")# 12/16追加 parser.add_argument('--save_dir', type=str, default="./ref_features/w50/mvtec_4shot") - + parser.add_argument('--backbone', type=str, default="wide_resnet50_2")#10/26追加 + parser.add_argument('--coupling_layers', type=int, default=10) + parser.add_argument('--clamp_alpha', type=float, default=1.9) + parser.add_argument('--pos_embed_dim', type=int, default=256) + parser.add_argument('--device', type=str, default="cuda:0") args = parser.parse_args() - main(args) \ No newline at end of file + main(args) diff --git a/extract_ref_features1.py b/extract_ref_features1.py new file mode 100644 index 0000000..22aff0d --- /dev/null +++ b/extract_ref_features1.py @@ -0,0 +1,283 @@ +#wideresnetの特徴抽出後に少し加工 +import os +import argparse +import numpy as np +from PIL import Image + +import torch +import tqdm +import timm +import torch.nn as nn +from torch.utils.data import DataLoader, Dataset +import torchvision.transforms as T +from models.fc_flow import load_flow_model + +from datasets.mvtec import MVTEC +from datasets.visa import VISA +from datasets.btad import BTAD +from datasets.mvtec_3d import MVTEC3D +from datasets.mpdd import MPDD +from datasets.mvtec_loco import MVTECLOCO +from datasets.brats import BRATS +from models.imagebind import ImageBindModel +from utils import load_weights +import torch.nn.functional as F + +class WideResNetFeatureExtractor(nn.Module): + def __init__(self, model_name='wide_resnet50_2', out_indices=(1, 2, 3)): + super().__init__() + # 元のモデルを読み込み + self.encoder = timm.create_model(model_name, features_only=True, out_indices=out_indices, pretrained=True) + self.embed_dims = self.encoder.feature_info.channels() + + # 重みを持たない平滑化レイヤー + self.layer1_pool = nn.AvgPool2d(kernel_size=3, stride=1, padding=1) # 浅い層用 + self.layer3_pool = nn.AvgPool2d(kernel_size=5, stride=1, padding=2) # 深い層用(より広範囲をぼかす) + + def forward(self, x): + features = self.encoder(x) + processed_features = [] + + for i, feat in enumerate(features): + if i == 0: + # 浅い層: 局所的な細かい変動ノイズを吸収 + feat = self.layer1_pool(feat) + elif i == 2: + # 深い層: 意味的な情報を少し広げてマッチングを安定させる + feat = self.layer3_pool(feat) + + # 全層共通: ベクトルのスケールを統一し、次元が大きくても距離計算を安定させる + feat = F.normalize(feat, p=2, dim=1) + processed_features.append(feat) + + return processed_features +class FEWSHOTDATA(Dataset): + + def __init__(self, + root: str, + class_name: str = 'bottle', + train: bool = True, + **kwargs) -> None: + + self.root = root + self.class_name = class_name + self.train = train + self.mask_size = [kwargs.get('msk_crp_size'), kwargs.get('msk_crp_size')] + + self.image_paths, self.labels, self.mask_paths, self.class_names = self._load_data(self.class_name) + + # set transforms + self.transform = T.Compose([ + T.Resize(kwargs.get('img_size', 224), T.InterpolationMode.BICUBIC), + T.CenterCrop(kwargs.get('crp_size', 224)), + T.ToTensor(), + T.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])]) + + # mask + self.target_transform = T.Compose([ + T.Resize(kwargs.get('msk_size', 256), T.InterpolationMode.NEAREST), + T.CenterCrop(kwargs.get('msk_crp_size', 256)), + T.ToTensor()]) + + def __len__(self): + return len(self.image_paths) + + def __getitem__(self, idx): + image_path, label, mask_path, class_name = self.image_paths[idx], self.labels[idx], self.mask_paths[idx], self.class_names[idx] + img, label, mask = self._load_image_and_mask(image_path, label, mask_path) + + return img, label, mask, class_name + + def _load_image_and_mask(self, image_path, label, mask_path): + img = Image.open(image_path).convert('RGB') + + img = self.transform(img) + + if label == 0: + mask = torch.zeros([1, self.mask_size[0], self.mask_size[1]]) + else: + mask = Image.open(mask_path) + mask = self.target_transform(mask) + + return img, label, mask + + def _load_data(self, class_name): + image_paths, labels, mask_paths = [], [], [] + phase = 'train' if self.train else 'test' + + image_dir = os.path.join(self.root, class_name, phase) + mask_dir = os.path.join(self.root, class_name, 'ground_truth') + + img_types = sorted(os.listdir(image_dir)) + for img_type in img_types: + # load images + img_type_dir = os.path.join(image_dir, img_type) + if not os.path.isdir(img_type_dir): + continue + img_fpath_list = sorted([os.path.join(img_type_dir, f) + for f in os.listdir(img_type_dir)]) + image_paths.extend(img_fpath_list) + + # load gt labels + if img_type == 'good': + labels.extend([0] * len(img_fpath_list)) + mask_paths.extend([None] * len(img_fpath_list)) + else: + labels.extend([1] * len(img_fpath_list)) + gt_type_dir = os.path.join(mask_dir, img_type) + img_fname_list = [os.path.splitext(os.path.basename(f))[0] for f in img_fpath_list] + gt_fpath_list = [os.path.join(gt_type_dir, img_fname + '_mask.png') + for img_fname in img_fname_list] + mask_paths.extend(gt_fpath_list) + + class_names = [class_name] * len(image_paths) + return image_paths, labels, mask_paths, class_names + + +SETTINGS = {'mvtec': MVTEC.CLASS_NAMES, 'visa': VISA.CLASS_NAMES, + 'btad': BTAD.CLASS_NAMES, 'mvtec3d': MVTEC3D.CLASS_NAMES, + 'mpdd': MPDD.CLASS_NAMES, 'mvtecloco': MVTECLOCO.CLASS_NAMES, + 'brats': BRATS.CLASS_NAMES} + + +def main(args): + image_size = 224 + device = 'cuda:0' + root_dir = args.few_shot_dir + if args.backbone == 'wide_resnet50_2': + encoder = WideResNetFeatureExtractor(model_name='wide_resnet50_2', out_indices=(1, 2, 3)).eval() + encoder = encoder.to(device) + feat_dims = encoder.embed_dims + elif args.backbone == 'tf_efficientnet_b6':#10/26追加 + encoder = timm.create_model('tf_efficientnet_b6', features_only=True, + out_indices=(1, 2, 3), pretrained=True).eval() # the pretrained checkpoint will be in /home/.cache/torch/hub/checkpoints/ + encoder = encoder.to(device) + #feat_dims = encoder.feature_info.channels() + decoders = [load_flow_model(args, feat_dim) for feat_dim in feat_dims] + decoders = [decoder.to(args.device) for decoder in decoders] + + if args.bgadweight_dir: + load_weights(encoder, decoders, args.bgadweight_dir) + if args.dataset in SETTINGS.keys(): + CLASS_NAMES = SETTINGS[args.dataset] + else: + raise ValueError(f"Dataset setting must be in {SETTINGS.keys()}, but got {args.dataset}.") + + for class_name in CLASS_NAMES: + train_dataset = FEWSHOTDATA(root_dir, class_name=class_name, train=True, img_size=image_size, crp_size=image_size, + msk_size=image_size, msk_crp_size=image_size) + train_loader = DataLoader( + train_dataset, batch_size=8, shuffle=False, num_workers=8, drop_last=False + ) + layer1_features, layer2_features, layer3_features = [], [], [] + + for batch in tqdm.tqdm(train_loader): + images, _, _, _ = batch + with torch.no_grad(): + patch_tokens = encoder(images.to(device)) + layer1_features.append(patch_tokens[0]) + layer2_features.append(patch_tokens[1]) + layer3_features.append(patch_tokens[2]) + layer1_features = torch.cat(layer1_features, dim=0) + layer2_features = torch.cat(layer2_features, dim=0) + layer3_features = torch.cat(layer3_features, dim=0) + print(layer1_features.shape) + print(layer2_features.shape) + print(layer3_features.shape) + #修正10/26 + layer1_channels = layer1_features.shape[1] + layer2_channels = layer2_features.shape[1] + layer3_channels = layer3_features.shape[1] + + layer1_features = layer1_features.permute(0, 2, 3, 1).reshape(-1, layer1_channels) + layer2_features = layer2_features.permute(0, 2, 3, 1).reshape(-1, layer2_channels) + layer3_features = layer3_features.permute(0, 2, 3, 1).reshape(-1, layer3_channels) + #修正終わり + os.makedirs(os.path.join(args.save_dir, class_name), exist_ok=True) + + print(f"Attempting to save layer1.npy for {class_name}...") + np.save(os.path.join(args.save_dir, class_name, 'layer1.npy'), layer1_features.cpu().numpy()) + print(f"Successfully saved layer1.npy for {class_name}.") + + np.save(os.path.join(args.save_dir, class_name, 'layer2.npy'), layer2_features.cpu().numpy()) + np.save(os.path.join(args.save_dir, class_name, 'layer3.npy'), layer3_features.cpu().numpy()) + + +def main2(args): + image_size = 224 + device = 'cuda:0' + root_dir = args.few_shot_dir + encoder = ImageBindModel(device=device) + encoder.to(device) + preprocess = T.Compose( # for imagebind + [ + T.Resize( + image_size, interpolation=T.InterpolationMode.BICUBIC + ), + T.CenterCrop(image_size), + T.ToTensor(), + T.Normalize( + mean=(0.48145466, 0.4578275, 0.40821073), + std=(0.26862954, 0.26130258, 0.27577711), + ), + ] + ) + + if args.dataset in SETTINGS.keys(): + CLASS_NAMES = SETTINGS[args.dataset] + else: + raise ValueError(f"Dataset setting must be in {SETTINGS.keys()}, but got {args.dataset}.") + + for class_name in CLASS_NAMES: + train_dataset = FEWSHOTDATA(root_dir, class_name=class_name, train=True, img_size=image_size, crp_size=image_size, + msk_size=image_size, msk_crp_size=image_size) + train_dataset.transform = preprocess + train_loader = DataLoader( + train_dataset, batch_size=4, shuffle=False, num_workers=8, drop_last=False + ) + layer1_features, layer2_features, layer3_features, layer4_features = [], [], [], [] + + for batch in tqdm.tqdm(train_loader): + images, _, _, _ = batch + with torch.no_grad(): + patch_features = encoder.encode_image_from_tensors(images.to(device)) + layer1_features.append(patch_features[0]) + layer2_features.append(patch_features[1]) + layer3_features.append(patch_features[2]) + layer4_features.append(patch_features[3]) + layer1_features = torch.cat(layer1_features, dim=0) + layer2_features = torch.cat(layer2_features, dim=0) + layer3_features = torch.cat(layer3_features, dim=0) + layer4_features = torch.cat(layer4_features, dim=0) + print(layer1_features.shape) + print(layer2_features.shape) + print(layer3_features.shape) + print(layer4_features.shape) + + layer1_features = layer1_features.reshape(-1, 1280) + layer2_features = layer2_features.reshape(-1, 1280) + layer3_features = layer3_features.reshape(-1, 1280) + layer4_features = layer4_features.reshape(-1, 1280) + + os.makedirs(os.path.join(args.save_dir, class_name), exist_ok=True) + + np.save(os.path.join(args.save_dir, class_name, 'layer1.npy'), layer1_features.cpu().numpy()) + np.save(os.path.join(args.save_dir, class_name, 'layer2.npy'), layer2_features.cpu().numpy()) + np.save(os.path.join(args.save_dir, class_name, 'layer3.npy'), layer3_features.cpu().numpy()) + np.save(os.path.join(args.save_dir, class_name, 'layer4.npy'), layer4_features.cpu().numpy()) + + +if __name__ == '__main__': + parser = argparse.ArgumentParser() + parser.add_argument('--dataset', type=str, default="mvtec") + parser.add_argument('--few_shot_dir', type=str, default="./4shot/mvtec") + parser.add_argument('--flow_arch', type=str, default='conditional_flow_model') + parser.add_argument('--bgadweight_dir', type=str, default="")# 12/16追加 + parser.add_argument('--save_dir', type=str, default="./ref_features/w50/mvtec_4shot") + parser.add_argument('--backbone', type=str, default="wide_resnet50_2")#10/26追加 + parser.add_argument('--coupling_layers', type=int, default=10) + parser.add_argument('--clamp_alpha', type=float, default=1.9) + parser.add_argument('--pos_embed_dim', type=int, default=256) + parser.add_argument('--device', type=str, default="cuda:0") + args = parser.parse_args() + main(args) diff --git a/extract_ref_features_filter.py b/extract_ref_features_filter.py new file mode 100644 index 0000000..d643fe7 --- /dev/null +++ b/extract_ref_features_filter.py @@ -0,0 +1,269 @@ +import os +import argparse +import numpy as np +from PIL import Image + +import torch +import tqdm +import timm +import torch.nn as nn +from torch.utils.data import DataLoader, Dataset +import torchvision.transforms as T +from models.fc_flow import load_flow_model +import torch.nn.functional as F +from datasets.mvtec import MVTEC +from datasets.visa import VISA +from datasets.btad import BTAD +from datasets.mvtec_3d import MVTEC3D +from datasets.mpdd import MPDD +from datasets.mvtec_loco import MVTECLOCO +from datasets.brats import BRATS +from models.imagebind import ImageBindModel +from utils import load_weights + +def apply_laplacian_filter(feature_maps): + filtered_maps = [] + # ラプラシアンカーネル (中心が4、周辺が-1) + kernel = torch.tensor([[0., -1., 0.], + [-1., 4., -1.], + [0., -1., 0.]]).view(1, 1, 3, 3) + + for f_map in feature_maps: + # f_map shape: [B, C, H, W] + B, C, H, W = f_map.shape + weight = kernel.expand(C, 1, 3, 3).to(f_map.device) + filtered = F.conv2d(f_map, weight, padding=1, groups=C) + filtered_maps.append(filtered) + + return filtered_maps +class FEWSHOTDATA(Dataset): + + def __init__(self, + root: str, + class_name: str = 'bottle', + train: bool = True, + **kwargs) -> None: + + self.root = root + self.class_name = class_name + self.train = train + self.mask_size = [kwargs.get('msk_crp_size'), kwargs.get('msk_crp_size')] + + self.image_paths, self.labels, self.mask_paths, self.class_names = self._load_data(self.class_name) + + # set transforms + self.transform = T.Compose([ + T.Resize(kwargs.get('img_size', 224), T.InterpolationMode.BICUBIC), + T.CenterCrop(kwargs.get('crp_size', 224)), + T.ToTensor(), + T.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])]) + + # mask + self.target_transform = T.Compose([ + T.Resize(kwargs.get('msk_size', 256), T.InterpolationMode.NEAREST), + T.CenterCrop(kwargs.get('msk_crp_size', 256)), + T.ToTensor()]) + + def __len__(self): + return len(self.image_paths) + + def __getitem__(self, idx): + image_path, label, mask_path, class_name = self.image_paths[idx], self.labels[idx], self.mask_paths[idx], self.class_names[idx] + img, label, mask = self._load_image_and_mask(image_path, label, mask_path) + + return img, label, mask, class_name + + def _load_image_and_mask(self, image_path, label, mask_path): + img = Image.open(image_path).convert('RGB') + + img = self.transform(img) + + if label == 0: + mask = torch.zeros([1, self.mask_size[0], self.mask_size[1]]) + else: + mask = Image.open(mask_path) + mask = self.target_transform(mask) + + return img, label, mask + + def _load_data(self, class_name): + image_paths, labels, mask_paths = [], [], [] + phase = 'train' if self.train else 'test' + + image_dir = os.path.join(self.root, class_name, phase) + mask_dir = os.path.join(self.root, class_name, 'ground_truth') + + img_types = sorted(os.listdir(image_dir)) + for img_type in img_types: + # load images + img_type_dir = os.path.join(image_dir, img_type) + if not os.path.isdir(img_type_dir): + continue + img_fpath_list = sorted([os.path.join(img_type_dir, f) + for f in os.listdir(img_type_dir)]) + image_paths.extend(img_fpath_list) + + # load gt labels + if img_type == 'good': + labels.extend([0] * len(img_fpath_list)) + mask_paths.extend([None] * len(img_fpath_list)) + else: + labels.extend([1] * len(img_fpath_list)) + gt_type_dir = os.path.join(mask_dir, img_type) + img_fname_list = [os.path.splitext(os.path.basename(f))[0] for f in img_fpath_list] + gt_fpath_list = [os.path.join(gt_type_dir, img_fname + '_mask.png') + for img_fname in img_fname_list] + mask_paths.extend(gt_fpath_list) + + class_names = [class_name] * len(image_paths) + return image_paths, labels, mask_paths, class_names + + +SETTINGS = {'mvtec': MVTEC.CLASS_NAMES, 'visa': VISA.CLASS_NAMES, + 'btad': BTAD.CLASS_NAMES, 'mvtec3d': MVTEC3D.CLASS_NAMES, + 'mpdd': MPDD.CLASS_NAMES, 'mvtecloco': MVTECLOCO.CLASS_NAMES, + 'brats': BRATS.CLASS_NAMES} + + +def main(args): + image_size = 224 + device = 'cuda:0' + root_dir = args.few_shot_dir + if args.backbone == 'wide_resnet50_2': + encoder = timm.create_model('wide_resnet50_2', features_only=True, + out_indices=(1, 2, 3), pretrained=True).eval() # the pretrained checkpoint will be in /home/.cache/torch/hub/checkpoints/ + encoder = encoder.to(device) + elif args.backbone == 'tf_efficientnet_b6':#10/26追加 + encoder = timm.create_model('tf_efficientnet_b6', features_only=True, + out_indices=(1, 2, 3), pretrained=True).eval() # the pretrained checkpoint will be in /home/.cache/torch/hub/checkpoints/ + encoder = encoder.to(device) + feat_dims = encoder.feature_info.channels() + decoders = [load_flow_model(args, feat_dim) for feat_dim in feat_dims] + decoders = [decoder.to(args.device) for decoder in decoders] + + if args.bgadweight_dir: + load_weights(encoder, decoders, args.bgadweight_dir) + if args.dataset in SETTINGS.keys(): + CLASS_NAMES = SETTINGS[args.dataset] + else: + raise ValueError(f"Dataset setting must be in {SETTINGS.keys()}, but got {args.dataset}.") + + for class_name in CLASS_NAMES: + train_dataset = FEWSHOTDATA(root_dir, class_name=class_name, train=True, img_size=image_size, crp_size=image_size, + msk_size=image_size, msk_crp_size=image_size) + train_loader = DataLoader( + train_dataset, batch_size=8, shuffle=False, num_workers=8, drop_last=False + ) + layer1_features, layer2_features, layer3_features = [], [], [] + + for batch in tqdm.tqdm(train_loader): + images, _, _, _ = batch + with torch.no_grad(): + patch_tokens = encoder(images.to(device)) + filtered_tokens = apply_laplacian_filter(patch_tokens) + layer1_features.append(filtered_tokens[0]) + layer2_features.append(filtered_tokens[1]) + layer3_features.append(filtered_tokens[2]) + layer1_features = torch.cat(layer1_features, dim=0) + layer2_features = torch.cat(layer2_features, dim=0) + layer3_features = torch.cat(layer3_features, dim=0) + print(layer1_features.shape) + print(layer2_features.shape) + print(layer3_features.shape) + #修正10/26 + layer1_channels = layer1_features.shape[1] + layer2_channels = layer2_features.shape[1] + layer3_channels = layer3_features.shape[1] + + layer1_features = layer1_features.permute(0, 2, 3, 1).reshape(-1, layer1_channels) + layer2_features = layer2_features.permute(0, 2, 3, 1).reshape(-1, layer2_channels) + layer3_features = layer3_features.permute(0, 2, 3, 1).reshape(-1, layer3_channels) + #修正終わり + os.makedirs(os.path.join(args.save_dir, class_name), exist_ok=True) + + print(f"Attempting to save layer1.npy for {class_name}...") + np.save(os.path.join(args.save_dir, class_name, 'layer1.npy'), layer1_features.cpu().numpy()) + print(f"Successfully saved layer1.npy for {class_name}.") + + np.save(os.path.join(args.save_dir, class_name, 'layer2.npy'), layer2_features.cpu().numpy()) + np.save(os.path.join(args.save_dir, class_name, 'layer3.npy'), layer3_features.cpu().numpy()) + + +def main2(args): + image_size = 224 + device = 'cuda:0' + root_dir = args.few_shot_dir + encoder = ImageBindModel(device=device) + encoder.to(device) + preprocess = T.Compose( # for imagebind + [ + T.Resize( + image_size, interpolation=T.InterpolationMode.BICUBIC + ), + T.CenterCrop(image_size), + T.ToTensor(), + T.Normalize( + mean=(0.48145466, 0.4578275, 0.40821073), + std=(0.26862954, 0.26130258, 0.27577711), + ), + ] + ) + + if args.dataset in SETTINGS.keys(): + CLASS_NAMES = SETTINGS[args.dataset] + else: + raise ValueError(f"Dataset setting must be in {SETTINGS.keys()}, but got {args.dataset}.") + + for class_name in CLASS_NAMES: + train_dataset = FEWSHOTDATA(root_dir, class_name=class_name, train=True, img_size=image_size, crp_size=image_size, + msk_size=image_size, msk_crp_size=image_size) + train_dataset.transform = preprocess + train_loader = DataLoader( + train_dataset, batch_size=4, shuffle=False, num_workers=8, drop_last=False + ) + layer1_features, layer2_features, layer3_features, layer4_features = [], [], [], [] + + for batch in tqdm.tqdm(train_loader): + images, _, _, _ = batch + with torch.no_grad(): + patch_features = encoder.encode_image_from_tensors(images.to(device)) + layer1_features.append(patch_features[0]) + layer2_features.append(patch_features[1]) + layer3_features.append(patch_features[2]) + layer4_features.append(patch_features[3]) + layer1_features = torch.cat(layer1_features, dim=0) + layer2_features = torch.cat(layer2_features, dim=0) + layer3_features = torch.cat(layer3_features, dim=0) + layer4_features = torch.cat(layer4_features, dim=0) + print(layer1_features.shape) + print(layer2_features.shape) + print(layer3_features.shape) + print(layer4_features.shape) + + layer1_features = layer1_features.reshape(-1, 1280) + layer2_features = layer2_features.reshape(-1, 1280) + layer3_features = layer3_features.reshape(-1, 1280) + layer4_features = layer4_features.reshape(-1, 1280) + + os.makedirs(os.path.join(args.save_dir, class_name), exist_ok=True) + + np.save(os.path.join(args.save_dir, class_name, 'layer1.npy'), layer1_features.cpu().numpy()) + np.save(os.path.join(args.save_dir, class_name, 'layer2.npy'), layer2_features.cpu().numpy()) + np.save(os.path.join(args.save_dir, class_name, 'layer3.npy'), layer3_features.cpu().numpy()) + np.save(os.path.join(args.save_dir, class_name, 'layer4.npy'), layer4_features.cpu().numpy()) + + +if __name__ == '__main__': + parser = argparse.ArgumentParser() + parser.add_argument('--dataset', type=str, default="mvtec") + parser.add_argument('--few_shot_dir', type=str, default="./4shot/mvtec") + parser.add_argument('--flow_arch', type=str, default='conditional_flow_model') + parser.add_argument('--bgadweight_dir', type=str, default="")# 12/16追加 + parser.add_argument('--save_dir', type=str, default="./ref_features/w50/mvtec_4shot") + parser.add_argument('--backbone', type=str, default="wide_resnet50_2")#10/26追加 + parser.add_argument('--coupling_layers', type=int, default=10) + parser.add_argument('--clamp_alpha', type=float, default=1.9) + parser.add_argument('--pos_embed_dim', type=int, default=256) + parser.add_argument('--device', type=str, default="cuda:0") + args = parser.parse_args() + main(args) diff --git a/extract_ref_features_vit.py b/extract_ref_features_vit.py new file mode 100644 index 0000000..644d618 --- /dev/null +++ b/extract_ref_features_vit.py @@ -0,0 +1,284 @@ +import os +import argparse +import numpy as np +from PIL import Image + +import torch +import tqdm +import timm +import torch.nn as nn +from torch.utils.data import DataLoader, Dataset +import torchvision.transforms as T +from models.fc_flow import load_flow_model + +from datasets.mvtec import MVTEC +from datasets.visa import VISA +from datasets.btad import BTAD +from datasets.mvtec_3d import MVTEC3D +from datasets.mpdd import MPDD +from datasets.mvtec_loco import MVTECLOCO +from datasets.brats import BRATS +from models.imagebind import ImageBindModel +from utils import load_weights + +class ViTFeatureExtractor(nn.Module): + def __init__(self, model_name="deit_base_patch16_224", out_indices=(3, 7, 11)): + super().__init__() + self.vit = timm.create_model(model_name, pretrained=True) + self.out_indices = out_indices + self.patch_size = self.vit.patch_embed.patch_size[0] + self.embed_dim = self.vit.embed_dim + + def forward(self, x): + B, C, H, W = x.shape + h_out, w_out = H // self.patch_size, W // self.patch_size + + x = self.vit.patch_embed(x) + x = self.vit._pos_embed(x) + x = self.vit.norm_pre(x) + + features = [] + for i, blk in enumerate(self.vit.blocks): + x = blk(x) + if i in self.out_indices: + num_prefix_tokens = self.vit.num_prefix_tokens + tokens = x[:, num_prefix_tokens:] + feat = tokens.transpose(1, 2).reshape(B, self.embed_dim, h_out, w_out) + features.append(feat) + + return features +class FEWSHOTDATA(Dataset): + + def __init__(self, + root: str, + class_name: str = 'bottle', + train: bool = True, + **kwargs) -> None: + + self.root = root + self.class_name = class_name + self.train = train + self.mask_size = [kwargs.get('msk_crp_size'), kwargs.get('msk_crp_size')] + + self.image_paths, self.labels, self.mask_paths, self.class_names = self._load_data(self.class_name) + + # set transforms + self.transform = T.Compose([ + T.Resize(kwargs.get('img_size', 224), T.InterpolationMode.BICUBIC), + T.CenterCrop(kwargs.get('crp_size', 224)), + T.ToTensor(), + T.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])]) + + # mask + self.target_transform = T.Compose([ + T.Resize(kwargs.get('msk_size', 256), T.InterpolationMode.NEAREST), + T.CenterCrop(kwargs.get('msk_crp_size', 256)), + T.ToTensor()]) + + def __len__(self): + return len(self.image_paths) + + def __getitem__(self, idx): + image_path, label, mask_path, class_name = self.image_paths[idx], self.labels[idx], self.mask_paths[idx], self.class_names[idx] + img, label, mask = self._load_image_and_mask(image_path, label, mask_path) + + return img, label, mask, class_name + + def _load_image_and_mask(self, image_path, label, mask_path): + img = Image.open(image_path).convert('RGB') + + img = self.transform(img) + + if label == 0: + mask = torch.zeros([1, self.mask_size[0], self.mask_size[1]]) + else: + mask = Image.open(mask_path) + mask = self.target_transform(mask) + + return img, label, mask + + def _load_data(self, class_name): + image_paths, labels, mask_paths = [], [], [] + phase = 'train' if self.train else 'test' + + image_dir = os.path.join(self.root, class_name, phase) + mask_dir = os.path.join(self.root, class_name, 'ground_truth') + + img_types = sorted(os.listdir(image_dir)) + for img_type in img_types: + # load images + img_type_dir = os.path.join(image_dir, img_type) + if not os.path.isdir(img_type_dir): + continue + img_fpath_list = sorted([os.path.join(img_type_dir, f) + for f in os.listdir(img_type_dir)]) + image_paths.extend(img_fpath_list) + + # load gt labels + if img_type == 'good': + labels.extend([0] * len(img_fpath_list)) + mask_paths.extend([None] * len(img_fpath_list)) + else: + labels.extend([1] * len(img_fpath_list)) + gt_type_dir = os.path.join(mask_dir, img_type) + img_fname_list = [os.path.splitext(os.path.basename(f))[0] for f in img_fpath_list] + gt_fpath_list = [os.path.join(gt_type_dir, img_fname + '_mask.png') + for img_fname in img_fname_list] + mask_paths.extend(gt_fpath_list) + + class_names = [class_name] * len(image_paths) + return image_paths, labels, mask_paths, class_names + + +SETTINGS = {'mvtec': MVTEC.CLASS_NAMES, 'visa': VISA.CLASS_NAMES, + 'btad': BTAD.CLASS_NAMES, 'mvtec3d': MVTEC3D.CLASS_NAMES, + 'mpdd': MPDD.CLASS_NAMES, 'mvtecloco': MVTECLOCO.CLASS_NAMES, + 'brats': BRATS.CLASS_NAMES} + + +def main(args): + image_size = 224 + device = 'cuda:0' + root_dir = args.few_shot_dir + if args.backbone == 'wide_resnet50_2': + encoder = timm.create_model('wide_resnet50_2', features_only=True, + out_indices=(1, 2, 3), pretrained=True).eval() # the pretrained checkpoint will be in /home/.cache/torch/hub/checkpoints/ + encoder = encoder.to(device) + feat_dims = encoder.feature_info.channels() + elif args.backbone == 'tf_efficientnet_b6':#10/26追加 + encoder = timm.create_model('tf_efficientnet_b6', features_only=True, + out_indices=(1, 2, 3), pretrained=True).eval() # the pretrained checkpoint will be in /home/.cache/torch/hub/checkpoints/ + encoder = encoder.to(device) + feat_dims = encoder.feature_info.channels() + elif args.backbone == 'vit_base_patch14': + encoder = ViTFeatureExtractor(model_name='vit_base_patch16_224_dino', out_indices=(3, 7, 11)).eval() + encoder = encoder.to(device) + feat_dims = [encoder.embed_dim] * len(encoder.out_indices) + decoders = [load_flow_model(args, feat_dim) for feat_dim in feat_dims] + decoders = [decoder.to(args.device) for decoder in decoders] + + if args.bgadweight_dir: + load_weights(encoder, decoders, args.bgadweight_dir) + if args.dataset in SETTINGS.keys(): + CLASS_NAMES = SETTINGS[args.dataset] + else: + raise ValueError(f"Dataset setting must be in {SETTINGS.keys()}, but got {args.dataset}.") + + for class_name in CLASS_NAMES: + train_dataset = FEWSHOTDATA(root_dir, class_name=class_name, train=True, img_size=image_size, crp_size=image_size, + msk_size=image_size, msk_crp_size=image_size) + train_loader = DataLoader( + train_dataset, batch_size=8, shuffle=False, num_workers=8, drop_last=False + ) + layer1_features, layer2_features, layer3_features = [], [], [] + + for batch in tqdm.tqdm(train_loader): + images, _, _, _ = batch + with torch.no_grad(): + patch_tokens = encoder(images.to(device)) + layer1_features.append(patch_tokens[0]) + layer2_features.append(patch_tokens[1]) + layer3_features.append(patch_tokens[2]) + layer1_features = torch.cat(layer1_features, dim=0) + layer2_features = torch.cat(layer2_features, dim=0) + layer3_features = torch.cat(layer3_features, dim=0) + print(layer1_features.shape) + print(layer2_features.shape) + print(layer3_features.shape) + #修正10/26 + layer1_channels = layer1_features.shape[1] + layer2_channels = layer2_features.shape[1] + layer3_channels = layer3_features.shape[1] + + layer1_features = layer1_features.permute(0, 2, 3, 1).reshape(-1, layer1_channels) + layer2_features = layer2_features.permute(0, 2, 3, 1).reshape(-1, layer2_channels) + layer3_features = layer3_features.permute(0, 2, 3, 1).reshape(-1, layer3_channels) + #修正終わり + os.makedirs(os.path.join(args.save_dir, class_name), exist_ok=True) + + print(f"Attempting to save layer1.npy for {class_name}...") + np.save(os.path.join(args.save_dir, class_name, 'layer1.npy'), layer1_features.cpu().numpy()) + print(f"Successfully saved layer1.npy for {class_name}.") + + np.save(os.path.join(args.save_dir, class_name, 'layer2.npy'), layer2_features.cpu().numpy()) + np.save(os.path.join(args.save_dir, class_name, 'layer3.npy'), layer3_features.cpu().numpy()) + + +def main2(args): + image_size = 224 + device = 'cuda:0' + root_dir = args.few_shot_dir + encoder = ImageBindModel(device=device) + encoder.to(device) + preprocess = T.Compose( # for imagebind + [ + T.Resize( + image_size, interpolation=T.InterpolationMode.BICUBIC + ), + T.CenterCrop(image_size), + T.ToTensor(), + T.Normalize( + mean=(0.48145466, 0.4578275, 0.40821073), + std=(0.26862954, 0.26130258, 0.27577711), + ), + ] + ) + + if args.dataset in SETTINGS.keys(): + CLASS_NAMES = SETTINGS[args.dataset] + else: + raise ValueError(f"Dataset setting must be in {SETTINGS.keys()}, but got {args.dataset}.") + + for class_name in CLASS_NAMES: + train_dataset = FEWSHOTDATA(root_dir, class_name=class_name, train=True, img_size=image_size, crp_size=image_size, + msk_size=image_size, msk_crp_size=image_size) + train_dataset.transform = preprocess + train_loader = DataLoader( + train_dataset, batch_size=4, shuffle=False, num_workers=8, drop_last=False + ) + layer1_features, layer2_features, layer3_features, layer4_features = [], [], [], [] + + for batch in tqdm.tqdm(train_loader): + images, _, _, _ = batch + with torch.no_grad(): + patch_features = encoder.encode_image_from_tensors(images.to(device)) + layer1_features.append(patch_features[0]) + layer2_features.append(patch_features[1]) + layer3_features.append(patch_features[2]) + layer4_features.append(patch_features[3]) + layer1_features = torch.cat(layer1_features, dim=0) + layer2_features = torch.cat(layer2_features, dim=0) + layer3_features = torch.cat(layer3_features, dim=0) + layer4_features = torch.cat(layer4_features, dim=0) + print(layer1_features.shape) + print(layer2_features.shape) + print(layer3_features.shape) + print(layer4_features.shape) + + layer1_features = layer1_features.reshape(-1, 1280) + layer2_features = layer2_features.reshape(-1, 1280) + layer3_features = layer3_features.reshape(-1, 1280) + layer4_features = layer4_features.reshape(-1, 1280) + + os.makedirs(os.path.join(args.save_dir, class_name), exist_ok=True) + + np.save(os.path.join(args.save_dir, class_name, 'layer1.npy'), layer1_features.cpu().numpy()) + np.save(os.path.join(args.save_dir, class_name, 'layer2.npy'), layer2_features.cpu().numpy()) + np.save(os.path.join(args.save_dir, class_name, 'layer3.npy'), layer3_features.cpu().numpy()) + np.save(os.path.join(args.save_dir, class_name, 'layer4.npy'), layer4_features.cpu().numpy()) + + +if __name__ == '__main__': + parser = argparse.ArgumentParser() + parser.add_argument('--dataset', type=str, default="mvtec") + parser.add_argument('--few_shot_dir', type=str, default="./4shot/mvtec") + parser.add_argument('--flow_arch', type=str, default='conditional_flow_model') + parser.add_argument('--bgadweight_dir', type=str, default="")# 12/16追加 + parser.add_argument('--save_dir', type=str, default="./ref_features/w50/mvtec_4shot") + parser.add_argument('--backbone', type=str, default="wide_resnet50_2")#10/26追加 + parser.add_argument('--coupling_layers', type=int, default=10) + parser.add_argument('--clamp_alpha', type=float, default=1.9) + parser.add_argument('--pos_embed_dim', type=int, default=256) + parser.add_argument('--device', type=str, default="cuda:0") + args = parser.parse_args() + main(args) diff --git a/extract_ref_features_wav.py b/extract_ref_features_wav.py new file mode 100644 index 0000000..9a41841 --- /dev/null +++ b/extract_ref_features_wav.py @@ -0,0 +1,299 @@ +import os +import argparse +import numpy as np +from PIL import Image + +import torch +import tqdm +import timm +import torch.nn as nn +import torch.nn.functional as F +from torch.utils.data import DataLoader, Dataset +import torchvision.transforms as T +from models.fc_flow import load_flow_model + +from datasets.mvtec import MVTEC +from datasets.visa import VISA +from datasets.btad import BTAD +from datasets.mvtec_3d import MVTEC3D +from datasets.mpdd import MPDD +from datasets.mvtec_loco import MVTECLOCO +from datasets.brats import BRATS +from models.imagebind import ImageBindModel +from utils import load_weights + +# ========================================== +# Haar Wavelet Filter の追加 +# ========================================== +class HaarWaveletFilter(nn.Module): + def __init__(self, low_freq_weight=0.1, high_freq_weight=1.2): + super().__init__() + self.lf_w = low_freq_weight + self.hf_w = high_freq_weight + + ll = torch.tensor([[0.5, 0.5], [0.5, 0.5]]) + hl = torch.tensor([[-0.5, -0.5], [0.5, 0.5]]) + lh = torch.tensor([[-0.5, 0.5], [-0.5, 0.5]]) + hh = torch.tensor([[0.5, -0.5], [-0.5, 0.5]]) + + self.register_buffer('k_ll', ll.view(1, 1, 2, 2)) + self.register_buffer('k_hl', hl.view(1, 1, 2, 2)) + self.register_buffer('k_lh', lh.view(1, 1, 2, 2)) + self.register_buffer('k_hh', hh.view(1, 1, 2, 2)) + + def forward(self, x): + B, C, H, W = x.shape + ll = F.conv2d(x, self.k_ll.expand(C, 1, 2, 2), stride=2, groups=C) + hl = F.conv2d(x, self.k_hl.expand(C, 1, 2, 2), stride=2, groups=C) + lh = F.conv2d(x, self.k_lh.expand(C, 1, 2, 2), stride=2, groups=C) + hh = F.conv2d(x, self.k_hh.expand(C, 1, 2, 2), stride=2, groups=C) + + ll = ll * self.lf_w + hl = hl * self.hf_w + lh = lh * self.hf_w + hh = hh * self.hf_w + + out = F.conv_transpose2d(ll, self.k_ll.expand(C, 1, 2, 2), stride=2, groups=C) + \ + F.conv_transpose2d(hl, self.k_hl.expand(C, 1, 2, 2), stride=2, groups=C) + \ + F.conv_transpose2d(lh, self.k_lh.expand(C, 1, 2, 2), stride=2, groups=C) + \ + F.conv_transpose2d(hh, self.k_hh.expand(C, 1, 2, 2), stride=2, groups=C) + return out +# ========================================== + +class FEWSHOTDATA(Dataset): + def __init__(self, + root: str, + class_name: str = 'bottle', + train: bool = True, + **kwargs) -> None: + + self.root = root + self.class_name = class_name + self.train = train + self.mask_size = [kwargs.get('msk_crp_size'), kwargs.get('msk_crp_size')] + + self.image_paths, self.labels, self.mask_paths, self.class_names = self._load_data(self.class_name) + + # set transforms + self.transform = T.Compose([ + T.Resize(kwargs.get('img_size', 224), T.InterpolationMode.BICUBIC), + T.CenterCrop(kwargs.get('crp_size', 224)), + T.ToTensor(), + T.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])]) + + # mask + self.target_transform = T.Compose([ + T.Resize(kwargs.get('msk_size', 256), T.InterpolationMode.NEAREST), + T.CenterCrop(kwargs.get('msk_crp_size', 256)), + T.ToTensor()]) + + def __len__(self): + return len(self.image_paths) + + def __getitem__(self, idx): + image_path, label, mask_path, class_name = self.image_paths[idx], self.labels[idx], self.mask_paths[idx], self.class_names[idx] + img, label, mask = self._load_image_and_mask(image_path, label, mask_path) + return img, label, mask, class_name + + def _load_image_and_mask(self, image_path, label, mask_path): + img = Image.open(image_path).convert('RGB') + img = self.transform(img) + if label == 0: + mask = torch.zeros([1, self.mask_size[0], self.mask_size[1]]) + else: + mask = Image.open(mask_path) + mask = self.target_transform(mask) + return img, label, mask + + def _load_data(self, class_name): + image_paths, labels, mask_paths = [], [], [] + phase = 'train' if self.train else 'test' + + image_dir = os.path.join(self.root, class_name, phase) + mask_dir = os.path.join(self.root, class_name, 'ground_truth') + + img_types = sorted(os.listdir(image_dir)) + for img_type in img_types: + img_type_dir = os.path.join(image_dir, img_type) + if not os.path.isdir(img_type_dir): + continue + img_fpath_list = sorted([os.path.join(img_type_dir, f) + for f in os.listdir(img_type_dir)]) + image_paths.extend(img_fpath_list) + + if img_type == 'good': + labels.extend([0] * len(img_fpath_list)) + mask_paths.extend([None] * len(img_fpath_list)) + else: + labels.extend([1] * len(img_fpath_list)) + gt_type_dir = os.path.join(mask_dir, img_type) + img_fname_list = [os.path.splitext(os.path.basename(f))[0] for f in img_fpath_list] + gt_fpath_list = [os.path.join(gt_type_dir, img_fname + '_mask.png') + for img_fname in img_fname_list] + mask_paths.extend(gt_fpath_list) + + class_names = [class_name] * len(image_paths) + return image_paths, labels, mask_paths, class_names + + +SETTINGS = {'mvtec': MVTEC.CLASS_NAMES, 'visa': VISA.CLASS_NAMES, + 'btad': BTAD.CLASS_NAMES, 'mvtec3d': MVTEC3D.CLASS_NAMES, + 'mpdd': MPDD.CLASS_NAMES, 'mvtecloco': MVTECLOCO.CLASS_NAMES, + 'brats': BRATS.CLASS_NAMES} + + +def main(args): + image_size = 224 + device = args.device + root_dir = args.few_shot_dir + if args.backbone == 'wide_resnet50_2': + encoder = timm.create_model('wide_resnet50_2', features_only=True, + out_indices=(1, 2, 3), pretrained=True).eval() + encoder = encoder.to(device) + elif args.backbone == 'tf_efficientnet_b6': + encoder = timm.create_model('tf_efficientnet_b6', features_only=True, + out_indices=(1, 2, 3), pretrained=True).eval() + encoder = encoder.to(device) + + # ウェーブレットフィルタの初期化 + wav_filter = HaarWaveletFilter(low_freq_weight=args.lf_weight, high_freq_weight=args.hf_weight).to(device) + wav_filter.eval() + + feat_dims = encoder.feature_info.channels() + decoders = [load_flow_model(args, feat_dim) for feat_dim in feat_dims] + decoders = [decoder.to(args.device) for decoder in decoders] + + if args.bgadweight_dir: + load_weights(encoder, decoders, args.bgadweight_dir) + if args.dataset in SETTINGS.keys(): + CLASS_NAMES = SETTINGS[args.dataset] + else: + raise ValueError(f"Dataset setting must be in {SETTINGS.keys()}, but got {args.dataset}.") + + for class_name in CLASS_NAMES: + train_dataset = FEWSHOTDATA(root_dir, class_name=class_name, train=True, img_size=image_size, crp_size=image_size, + msk_size=image_size, msk_crp_size=image_size) + train_loader = DataLoader( + train_dataset, batch_size=8, shuffle=False, num_workers=8, drop_last=False + ) + layer1_features, layer2_features, layer3_features = [], [], [] + + for batch in tqdm.tqdm(train_loader): + images, _, _, _ = batch + with torch.no_grad(): + patch_tokens = encoder(images.to(device)) + # 【追加】保存用に変換する前に、ウェーブレット変換をかける + patch_tokens = [wav_filter(f) for f in patch_tokens] + + layer1_features.append(patch_tokens[0]) + layer2_features.append(patch_tokens[1]) + layer3_features.append(patch_tokens[2]) + + layer1_features = torch.cat(layer1_features, dim=0) + layer2_features = torch.cat(layer2_features, dim=0) + layer3_features = torch.cat(layer3_features, dim=0) + print(layer1_features.shape) + print(layer2_features.shape) + print(layer3_features.shape) + + layer1_channels = layer1_features.shape[1] + layer2_channels = layer2_features.shape[1] + layer3_channels = layer3_features.shape[1] + + layer1_features = layer1_features.permute(0, 2, 3, 1).reshape(-1, layer1_channels) + layer2_features = layer2_features.permute(0, 2, 3, 1).reshape(-1, layer2_channels) + layer3_features = layer3_features.permute(0, 2, 3, 1).reshape(-1, layer3_channels) + + os.makedirs(os.path.join(args.save_dir, class_name), exist_ok=True) + + print(f"Attempting to save layer1.npy for {class_name}...") + np.save(os.path.join(args.save_dir, class_name, 'layer1.npy'), layer1_features.cpu().numpy()) + print(f"Successfully saved layer1.npy for {class_name}.") + + np.save(os.path.join(args.save_dir, class_name, 'layer2.npy'), layer2_features.cpu().numpy()) + np.save(os.path.join(args.save_dir, class_name, 'layer3.npy'), layer3_features.cpu().numpy()) + + +def main2(args): + # (既存のImageBind用の処理はそのまま残しています) + image_size = 224 + device = args.device + root_dir = args.few_shot_dir + encoder = ImageBindModel(device=device) + encoder.to(device) + preprocess = T.Compose( + [ + T.Resize(image_size, interpolation=T.InterpolationMode.BICUBIC), + T.CenterCrop(image_size), + T.ToTensor(), + T.Normalize( + mean=(0.48145466, 0.4578275, 0.40821073), + std=(0.26862954, 0.26130258, 0.27577711), + ), + ] + ) + + if args.dataset in SETTINGS.keys(): + CLASS_NAMES = SETTINGS[args.dataset] + else: + raise ValueError(f"Dataset setting must be in {SETTINGS.keys()}, but got {args.dataset}.") + + for class_name in CLASS_NAMES: + train_dataset = FEWSHOTDATA(root_dir, class_name=class_name, train=True, img_size=image_size, crp_size=image_size, + msk_size=image_size, msk_crp_size=image_size) + train_dataset.transform = preprocess + train_loader = DataLoader( + train_dataset, batch_size=4, shuffle=False, num_workers=8, drop_last=False + ) + layer1_features, layer2_features, layer3_features, layer4_features = [], [], [], [] + + for batch in tqdm.tqdm(train_loader): + images, _, _, _ = batch + with torch.no_grad(): + patch_features = encoder.encode_image_from_tensors(images.to(device)) + layer1_features.append(patch_features[0]) + layer2_features.append(patch_features[1]) + layer3_features.append(patch_features[2]) + layer4_features.append(patch_features[3]) + + layer1_features = torch.cat(layer1_features, dim=0) + layer2_features = torch.cat(layer2_features, dim=0) + layer3_features = torch.cat(layer3_features, dim=0) + layer4_features = torch.cat(layer4_features, dim=0) + print(layer1_features.shape) + print(layer2_features.shape) + print(layer3_features.shape) + print(layer4_features.shape) + + layer1_features = layer1_features.reshape(-1, 1280) + layer2_features = layer2_features.reshape(-1, 1280) + layer3_features = layer3_features.reshape(-1, 1280) + layer4_features = layer4_features.reshape(-1, 1280) + + os.makedirs(os.path.join(args.save_dir, class_name), exist_ok=True) + + np.save(os.path.join(args.save_dir, class_name, 'layer1.npy'), layer1_features.cpu().numpy()) + np.save(os.path.join(args.save_dir, class_name, 'layer2.npy'), layer2_features.cpu().numpy()) + np.save(os.path.join(args.save_dir, class_name, 'layer3.npy'), layer3_features.cpu().numpy()) + np.save(os.path.join(args.save_dir, class_name, 'layer4.npy'), layer4_features.cpu().numpy()) + + +if __name__ == '__main__': + parser = argparse.ArgumentParser() + parser.add_argument('--dataset', type=str, default="mvtec") + parser.add_argument('--few_shot_dir', type=str, default="./4shot/mvtec") + parser.add_argument('--flow_arch', type=str, default='conditional_flow_model') + parser.add_argument('--bgadweight_dir', type=str, default="") + parser.add_argument('--save_dir', type=str, default="./ref_features/w50/mvtec_4shot_wav") + parser.add_argument('--backbone', type=str, default="wide_resnet50_2") + parser.add_argument('--coupling_layers', type=int, default=10) + parser.add_argument('--clamp_alpha', type=float, default=1.9) + parser.add_argument('--pos_embed_dim', type=int, default=256) + parser.add_argument('--device', type=str, default="cuda:0") + + # 追加: ウェーブレット変換のパラメータ + parser.add_argument("--lf_weight", type=float, default=0.1, help="Weight for low frequency (LL) components") + parser.add_argument("--hf_weight", type=float, default=1.2, help="Weight for high frequency (LH, HL, HH) components") + + args = parser.parse_args() + main(args) diff --git a/main.py b/main.py index 3468d6f..f366589 100644 --- a/main.py +++ b/main.py @@ -17,6 +17,7 @@ from datasets.mpdd import MPDD from datasets.mvtec_loco import MVTECLOCO from datasets.brats import BRATS +from datasets.capsules import CAPSULES, CAPSULESANO from models.fc_flow import load_flow_model from models.modules import MultiScaleConv @@ -26,6 +27,8 @@ from losses.loss import calculate_log_barrier_bi_occ_loss from classes import VISA_TO_MVTEC, MVTEC_TO_VISA, MVTEC_TO_BTAD, MVTEC_TO_MVTEC3D from classes import MVTEC_TO_MPDD, MVTEC_TO_MVTECLOCO, MVTEC_TO_BRATS +from classes import MVTEC_TO_MVTEC, VISA_TO_VISA +from classes import CAPSULES_TO_CAPSULES warnings.filterwarnings('ignore') @@ -34,7 +37,7 @@ SETTINGS = {'visa_to_mvtec': VISA_TO_MVTEC, 'mvtec_to_visa': MVTEC_TO_VISA, 'mvtec_to_btad': MVTEC_TO_BTAD, 'mvtec_to_mvtec3d': MVTEC_TO_MVTEC3D, 'mvtec_to_mpdd': MVTEC_TO_MPDD, 'mvtec_to_mvtecloco': MVTEC_TO_MVTECLOCO, - 'mvtec_to_brats': MVTEC_TO_BRATS} + 'mvtec_to_brats': MVTEC_TO_BRATS,'mvtec_to_mvtec':MVTEC_TO_MVTEC, 'visa_to_visa':VISA_TO_VISA, 'capsules_to_capsules': CAPSULES_TO_CAPSULES} def main(args): @@ -42,8 +45,22 @@ def main(args): CLASSES = SETTINGS[args.setting] else: raise ValueError(f"Dataset setting must be in {SETTINGS.keys()}, but got {args.setting}.") + + if args.classes == 'capsules': # from mvtec to other datasets # from mvtec to other datasets + train_dataset1 = CAPSULES(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, + normalize="w50", + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader1 = DataLoader( + train_dataset1, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True + ) + train_dataset2 = CAPSULESANO(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader2 = DataLoader( + train_dataset2, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True + ) - if CLASSES['seen'][0] in MVTEC.CLASS_NAMES: # from mvtec to other datasets + elif CLASSES['seen'][0] in MVTEC.CLASS_NAMES: # from mvtec to other datasets train_dataset1 = MVTEC(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, normalize="w50", img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) @@ -69,12 +86,16 @@ def main(args): train_loader2 = DataLoader( train_dataset2, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True ) - - encoder = timm.create_model('wide_resnet50_2', features_only=True, - out_indices=(1, 2, 3), pretrained=True).eval() # the pretrained checkpoint will be in /home/.cache/torch/hub/checkpoints/ - encoder = encoder.to(args.device) - feat_dims = encoder.feature_info.channels() - + if args.backbone == 'wide_resnet50_2': + encoder = timm.create_model('wide_resnet50_2', features_only=True, + out_indices=(1, 2, 3), pretrained=True).eval() # the pretrained checkpoint will be in /home/.cache/torch/hub/checkpoints/ + encoder = encoder.to(args.device) + feat_dims = encoder.feature_info.channels() + elif args.backbone == 'tf_efficientnet_b6':#10/26追加 + encoder = timm.create_model('tf_efficientnet_b6', features_only=True, + out_indices=(1, 2, 3), pretrained=True).eval() # the pretrained checkpoint will be in /home/.cache/torch/hub/checkpoints/ + encoder = encoder.to(args.device) + feat_dims = encoder.feature_info.channels() boundary_ops = BoundaryAverager(num_levels=args.feature_levels) vq_ops = MultiScaleVQ(num_embeddings=args.num_embeddings, channels=feat_dims).to(args.device) optimizer_vq = torch.optim.Adam(vq_ops.parameters(), lr=args.lr, weight_decay=0.0005) @@ -94,7 +115,7 @@ def main(args): optimizer1 = torch.optim.Adam(params, lr=args.lr, weight_decay=0.0005) scheduler1 = torch.optim.lr_scheduler.MultiStepLR(optimizer1, milestones=[70, 90], gamma=0.1) - best_pro = 0 + best_img_auc = 0 N_batch = 8192 for epoch in range(args.epochs): vq_ops.train() @@ -174,7 +195,11 @@ def main(args): s1_res, s2_res, s_res = [], [], [] test_ref_features = load_mc_reference_features(args.test_ref_feature_dir, CLASSES['unseen'], args.device, args.num_ref_shot) for class_name in CLASSES['unseen']: - if class_name in MVTEC.CLASS_NAMES: + if args.classes == 'capsules': + test_dataset = CAPSULES(args.test_dataset_dir, class_name=class_name, train=False, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in MVTEC.CLASS_NAMES: test_dataset = MVTEC(args.test_dataset_dir, class_name=class_name, train=False, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) @@ -229,13 +254,14 @@ def main(args): print('(Merged) Average Image AUC | AP | F1_Score: {:.3f} | {:.3f} | {:.3f}, Average Pixel AUC | AP | F1_Score | AUPRO: {:.3f} | {:.3f} | {:.3f} | {:.3f}'.format( img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro)) - if pix_aupro > best_pro: + if img_auc > best_img_auc: os.makedirs(args.checkpoint_path, exist_ok=True) - best_pro = pix_aupro + best_img_auc = img_auc state_dict = {'vq_ops': vq_ops.state_dict(), 'constraintor': constraintor.state_dict(), 'estimators': [estimator.state_dict() for estimator in estimators]} - torch.save(state_dict, os.path.join(args.checkpoint_path, f'{args.setting}_checkpoints.pth')) + torch.save(state_dict, os.path.join(args.checkpoint_path, f'{args.setting}_epoch_{epoch}_checkpoints.pth')) + #torch.save(state_dict, os.path.join(args.checkpoint_path, f'{args.setting}_checkpoints.pth')) def load_mc_reference_features(root_dir: str, class_names, device: torch.device, num_shot=4): @@ -264,10 +290,11 @@ def load_mc_reference_features(root_dir: str, class_names, device: torch.device, if __name__ == "__main__": parser = argparse.ArgumentParser() parser.add_argument('--setting', type=str, default="visa_to_mvtec") + parser.add_argument('--classes', type=str, default="none") parser.add_argument('--train_dataset_dir', type=str, default="") parser.add_argument('--test_dataset_dir', type=str, default="") parser.add_argument('--test_ref_feature_dir', type=str, default="./ref_features/w50/mvtec_4shot") - + parser.add_argument('--bgadweight_dir', type=str, default="none")# 12/16追加 parser.add_argument('--batch_size', type=int, default=32) parser.add_argument('--lr', type=float, default=1e-5) parser.add_argument('--epochs', type=int, default=100) @@ -275,6 +302,7 @@ def load_mc_reference_features(root_dir: str, class_names, device: torch.device, parser.add_argument('--checkpoint_path', type=str, default="./checkpoints/") parser.add_argument('--eval_freq', type=int, default=1) parser.add_argument('--backbone', type=str, default="wide_resnet50_2") + parser.add_argument('--rank', type=int, default="0") # flow parameters parser.add_argument('--flow_arch', type=str, default='conditional_flow_model') @@ -295,8 +323,4 @@ def load_mc_reference_features(root_dir: str, class_names, device: torch.device, init_seeds(42) main(args) - - - - - \ No newline at end of file + diff --git a/main_1.py b/main_1.py new file mode 100644 index 0000000..45922dd --- /dev/null +++ b/main_1.py @@ -0,0 +1,357 @@ +#widerenの特徴抽出後に少し加工してみたver +import os +import warnings +import argparse +from tqdm import tqdm +import numpy as np +import torch +import timm +import torch.nn.functional as F +from torch.utils.data import DataLoader +import torch.nn as nn +from train import train +from validate import validate +from datasets.mvtec import MVTEC, MVTECANO +from datasets.visa import VISA, VISAANO +from datasets.btad import BTAD +from datasets.mvtec_3d import MVTEC3D +from datasets.mpdd import MPDD +from datasets.mvtec_loco import MVTECLOCO +from datasets.brats import BRATS +from datasets.capsules import CAPSULES, CAPSULESANO + +from models.fc_flow import load_flow_model +from models.modules import MultiScaleConv +from models.vq import MultiScaleVQ +from utils import init_seeds, get_residual_features, get_mc_matched_ref_features, get_mc_reference_features +from utils import BoundaryAverager +from losses.loss import calculate_log_barrier_bi_occ_loss +from classes import VISA_TO_MVTEC, MVTEC_TO_VISA, MVTEC_TO_BTAD, MVTEC_TO_MVTEC3D +from classes import MVTEC_TO_MPDD, MVTEC_TO_MVTECLOCO, MVTEC_TO_BRATS +from classes import MVTEC_TO_MVTEC, VISA_TO_VISA +from classes import CAPSULES_TO_CAPSULES + +warnings.filterwarnings('ignore') + +TOTAL_SHOT = 4 # total few-shot reference samples +FIRST_STAGE_EPOCH = 10 +SETTINGS = {'visa_to_mvtec': VISA_TO_MVTEC, 'mvtec_to_visa': MVTEC_TO_VISA, + 'mvtec_to_btad': MVTEC_TO_BTAD, 'mvtec_to_mvtec3d': MVTEC_TO_MVTEC3D, + 'mvtec_to_mpdd': MVTEC_TO_MPDD, 'mvtec_to_mvtecloco': MVTEC_TO_MVTECLOCO, + 'mvtec_to_brats': MVTEC_TO_BRATS,'mvtec_to_mvtec':MVTEC_TO_MVTEC, 'visa_to_visa':VISA_TO_VISA, 'capsules_to_capsules': CAPSULES_TO_CAPSULES} + +import torch.nn.functional as F + +class WideResNetFeatureExtractor(nn.Module): + def __init__(self, model_name='wide_resnet50_2', out_indices=(1, 2, 3)): + super().__init__() + # 元のモデルを読み込み + self.encoder = timm.create_model(model_name, features_only=True, out_indices=out_indices, pretrained=True) + self.embed_dims = self.encoder.feature_info.channels() + + # 重みを持たない平滑化レイヤー + self.layer1_pool = nn.AvgPool2d(kernel_size=3, stride=1, padding=1) # 浅い層用 + self.layer3_pool = nn.AvgPool2d(kernel_size=5, stride=1, padding=2) # 深い層用(より広範囲をぼかす) + + def forward(self, x): + features = self.encoder(x) + processed_features = [] + + for i, feat in enumerate(features): + if i == 0: + # 浅い層: 局所的な細かい変動ノイズを吸収 + feat = self.layer1_pool(feat) + elif i == 2: + # 深い層: 意味的な情報を少し広げてマッチングを安定させる + feat = self.layer3_pool(feat) + + # 全層共通: ベクトルのスケールを統一し、次元が大きくても距離計算を安定させる + feat = F.normalize(feat, p=2, dim=1) + processed_features.append(feat) + + return processed_features +def main(args): + if args.setting in SETTINGS.keys(): + CLASSES = SETTINGS[args.setting] + else: + raise ValueError(f"Dataset setting must be in {SETTINGS.keys()}, but got {args.setting}.") + + if args.classes == 'capsules': # from mvtec to other datasets # from mvtec to other datasets + train_dataset1 = CAPSULES(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, + normalize="w50", + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader1 = DataLoader( + train_dataset1, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True + ) + train_dataset2 = CAPSULESANO(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader2 = DataLoader( + train_dataset2, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True + ) + + elif CLASSES['seen'][0] in MVTEC.CLASS_NAMES: # from mvtec to other datasets + train_dataset1 = MVTEC(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, + normalize="w50", + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader1 = DataLoader( + train_dataset1, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True + ) + train_dataset2 = MVTECANO(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader2 = DataLoader( + train_dataset2, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True + ) + else: # from visa to mvtec + train_dataset1 = VISA(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, + normalize="w50", + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader1 = DataLoader( + train_dataset1, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True + ) + train_dataset2 = VISAANO(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, + normalize="w50", + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader2 = DataLoader( + train_dataset2, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True + ) + if args.backbone == 'wide_resnet50_2': + + encoder = WideResNetFeatureExtractor(model_name='wide_resnet50_2', out_indices=(1, 2, 3)).eval() + encoder = encoder.to(args.device) + feat_dims = encoder.embed_dims + elif args.backbone == 'tf_efficientnet_b6':#10/26追加 + encoder = timm.create_model('tf_efficientnet_b6', features_only=True, + out_indices=(1, 2, 3), pretrained=True).eval() # the pretrained checkpoint will be in /home/.cache/torch/hub/checkpoints/ + encoder = encoder.to(args.device) + feat_dims = encoder.feature_info.channels() + boundary_ops = BoundaryAverager(num_levels=args.feature_levels) + vq_ops = MultiScaleVQ(num_embeddings=args.num_embeddings, channels=feat_dims).to(args.device) + optimizer_vq = torch.optim.Adam(vq_ops.parameters(), lr=args.lr, weight_decay=0.0005) + scheduler_vq = torch.optim.lr_scheduler.MultiStepLR(optimizer_vq, milestones=[70, 90], gamma=0.1) + + constraintor = MultiScaleConv(feat_dims).to(args.device) + # weight_decay is the l2 weight penalty lambda, weight_decay = lambda / 2 + optimizer0 = torch.optim.Adam(constraintor.parameters(), lr=args.lr, weight_decay=0.0005) + scheduler0 = torch.optim.lr_scheduler.MultiStepLR(optimizer0, milestones=[70, 90], gamma=0.1) + + # Normflow decoder + estimators = [load_flow_model(args, feat_dim) for feat_dim in feat_dims] + estimators = [decoder.to(args.device) for decoder in estimators] + params = list(estimators[0].parameters()) + for l in range(1, args.feature_levels): + params += list(estimators[l].parameters()) + optimizer1 = torch.optim.Adam(params, lr=args.lr, weight_decay=0.0005) + scheduler1 = torch.optim.lr_scheduler.MultiStepLR(optimizer1, milestones=[70, 90], gamma=0.1) + + best_img_auc = 0 + N_batch = 8192 + for epoch in range(args.epochs): + vq_ops.train() + constraintor.train() + for estimator in estimators: + estimator.train() + if epoch < FIRST_STAGE_EPOCH: + train_loader = train_loader1 + else: + train_loader = train_loader2 + train_loss_total, total_num = 0, 0 + progress_bar = tqdm(total=len(train_loader)) + progress_bar.set_description(f"Epoch[{epoch}/{args.epochs}]") + for step, batch in enumerate(train_loader): + progress_bar.update(1) + #images, _, masks, class_names = batch + images, _, masks, class_names = batch[0], batch[1], batch[2], batch[3] + + images = images.to(args.device) + masks = masks.to(args.device) + + with torch.no_grad(): + features = encoder(images) + + ref_features = get_mc_reference_features(encoder, args.train_dataset_dir, class_names, images.device, args.train_ref_shot) + mfeatures = get_mc_matched_ref_features(features, class_names, ref_features) + rfeatures = get_residual_features(features, mfeatures, pos_flag=True) + + lvl_masks = [] + for l in range(args.feature_levels): + _, _, h, w = rfeatures[l].size() + m = F.interpolate(masks, size=(h, w), mode='nearest').squeeze(1) + lvl_masks.append(m) + rfeatures_t = [rfeature.detach().clone() for rfeature in rfeatures] + + loss_vq = vq_ops(rfeatures, lvl_masks, train=True) + train_loss_total += loss_vq.item() + total_num += 1 + optimizer_vq.zero_grad() + loss_vq.backward() + optimizer_vq.step() + + rfeatures = constraintor(*rfeatures) + loss = 0 + for l in range(args.feature_levels): # backward svdd loss + e = rfeatures[l] + t = rfeatures_t[l] + bs, dim, h, w = e.size() + e = e.permute(0, 2, 3, 1).reshape(-1, dim) + t = t.permute(0, 2, 3, 1).reshape(-1, dim) + m = lvl_masks[l] + m = m.reshape(-1) + + loss_i, _, _ = calculate_log_barrier_bi_occ_loss(e, m, t) + loss += loss_i + optimizer0.zero_grad() + loss.backward() + optimizer0.step() + + train_loss_total += loss.item() + total_num += 1 + + # detach the rfeatures for flow optimization + rfeatures = [rfeature.detach().clone() for rfeature in rfeatures] + # train flow corresponding to with neck + loss, num = train(args, rfeatures, estimators, optimizer1, masks, boundary_ops, epoch, N_batch=N_batch, FIRST_STAGE_EPOCH=FIRST_STAGE_EPOCH) + train_loss_total += loss + total_num += num + + scheduler_vq.step() + scheduler0.step() + scheduler1.step() + + progress_bar.close() + print(f"Epoch[{epoch}/{args.epochs}]: train_loss: {train_loss_total / total_num}") + + if (epoch + 1) % args.eval_freq == 0: + s1_res, s2_res, s_res = [], [], [] + test_ref_features = load_mc_reference_features(args.test_ref_feature_dir, CLASSES['unseen'], args.device, args.num_ref_shot) + for class_name in CLASSES['unseen']: + if args.classes == 'capsules': + test_dataset = CAPSULES(args.test_dataset_dir, class_name=class_name, train=False, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in MVTEC.CLASS_NAMES: + test_dataset = MVTEC(args.test_dataset_dir, class_name=class_name, train=False, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in VISA.CLASS_NAMES: + test_dataset = VISA(args.test_dataset_dir, class_name=class_name, train=False, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in BTAD.CLASS_NAMES: + test_dataset = BTAD(args.test_dataset_dir, class_name=class_name, train=False, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in MVTEC3D.CLASS_NAMES: + test_dataset = MVTEC3D(args.test_dataset_dir, class_name=class_name, train=False, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in MPDD.CLASS_NAMES: + test_dataset = MPDD(args.test_dataset_dir, class_name=class_name, train=False, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in MVTECLOCO.CLASS_NAMES: + test_dataset = MVTECLOCO(args.test_dataset_dir, class_name=class_name, train=False, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in BRATS.CLASS_NAMES: + test_dataset = BRATS(args.test_dataset_dir, class_name=class_name, train=False, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + else: + raise ValueError('Unrecognized class name: {}'.format(class_name)) + test_loader = DataLoader( + test_dataset, batch_size=1, shuffle=False, num_workers=8, drop_last=False + ) + metrics = validate(args, encoder, vq_ops, constraintor, estimators, test_loader, test_ref_features[class_name], args.device, class_name) + img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro = metrics['scores'] + + print("Epoch: {}, Class Name: {}, Image AUC | AP | F1_Score: {} | {} | {}, Pixel AUC | AP | F1_Score | AUPRO: {} | {} | {} | {}".format( + epoch, class_name, img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro)) + s1_res.append(metrics['scores1']) + s2_res.append(metrics['scores2']) + s_res.append(metrics['scores']) + + s1_res = np.array(s1_res) + s2_res = np.array(s2_res) + s_res = np.array(s_res) + img_auc1, img_ap1, img_f1_score1, pix_auc1, pix_ap1, pix_f1_score1, pix_aupro1 = np.mean(s1_res, axis=0) + img_auc2, img_ap2, img_f1_score2, pix_auc2, pix_ap2, pix_f1_score2, pix_aupro2 = np.mean(s2_res, axis=0) + img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro = np.mean(s_res, axis=0) + print('(Logps) Average Image AUC | AP | F1_Score: {:.3f} | {:.3f} | {:.3f}, Average Pixel AUC | AP | F1_Score | AUPRO: {:.3f} | {:.3f} | {:.3f} | {:.3f}'.format( + img_auc1, img_ap1, img_f1_score1, pix_auc1, pix_ap1, pix_f1_score1, pix_aupro1)) + print('(BScores) Average Image AUC | AP | F1_Score: {:.3f} | {:.3f} | {:.3f}, Average Pixel AUC | AP | F1_Score | AUPRO: {:.3f} | {:.3f} | {:.3f} | {:.3f}'.format( + img_auc2, img_ap2, img_f1_score2, pix_auc2, pix_ap2, pix_f1_score2, pix_aupro2)) + print('(Merged) Average Image AUC | AP | F1_Score: {:.3f} | {:.3f} | {:.3f}, Average Pixel AUC | AP | F1_Score | AUPRO: {:.3f} | {:.3f} | {:.3f} | {:.3f}'.format( + img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro)) + + if img_auc > best_img_auc: + os.makedirs(args.checkpoint_path, exist_ok=True) + best_img_auc = img_auc + state_dict = {'vq_ops': vq_ops.state_dict(), + 'constraintor': constraintor.state_dict(), + 'estimators': [estimator.state_dict() for estimator in estimators]} + torch.save(state_dict, os.path.join(args.checkpoint_path, f'{args.setting}_epoch_{epoch}_checkpoints.pth')) + #torch.save(state_dict, os.path.join(args.checkpoint_path, f'{args.setting}_checkpoints.pth')) + + +def load_mc_reference_features(root_dir: str, class_names, device: torch.device, num_shot=4): + refs = {} + for class_name in class_names: + layer1_refs = np.load(os.path.join(root_dir, class_name, 'layer1.npy')) + layer2_refs = np.load(os.path.join(root_dir, class_name, 'layer2.npy')) + layer3_refs = np.load(os.path.join(root_dir, class_name, 'layer3.npy')) + + layer1_refs = torch.from_numpy(layer1_refs).to(device) + layer2_refs = torch.from_numpy(layer2_refs).to(device) + layer3_refs = torch.from_numpy(layer3_refs).to(device) + + K1 = (layer1_refs.shape[0] // TOTAL_SHOT) * num_shot + layer1_refs = layer1_refs[:K1, :] + K2 = (layer2_refs.shape[0] // TOTAL_SHOT) * num_shot + layer2_refs = layer2_refs[:K2, :] + K3 = (layer3_refs.shape[0] // TOTAL_SHOT) * num_shot + layer3_refs = layer3_refs[:K3, :] + + refs[class_name] = (layer1_refs, layer2_refs, layer3_refs) + + return refs + + +if __name__ == "__main__": + parser = argparse.ArgumentParser() + parser.add_argument('--setting', type=str, default="visa_to_mvtec") + parser.add_argument('--classes', type=str, default="none") + parser.add_argument('--train_dataset_dir', type=str, default="") + parser.add_argument('--test_dataset_dir', type=str, default="") + parser.add_argument('--test_ref_feature_dir', type=str, default="./ref_features/w50/mvtec_4shot") + parser.add_argument('--bgadweight_dir', type=str, default="none")# 12/16追加 + parser.add_argument('--batch_size', type=int, default=32) + parser.add_argument('--lr', type=float, default=1e-5) + parser.add_argument('--epochs', type=int, default=100) + parser.add_argument('--device', type=str, default="cuda:0") + parser.add_argument('--checkpoint_path', type=str, default="./checkpoints/") + parser.add_argument('--eval_freq', type=int, default=1) + parser.add_argument('--backbone', type=str, default="wide_resnet50_2") + parser.add_argument('--rank', type=int, default="0") + + # flow parameters + parser.add_argument('--flow_arch', type=str, default='conditional_flow_model') + parser.add_argument('--feature_levels', default=3, type=int) + parser.add_argument('--coupling_layers', type=int, default=10) + parser.add_argument('--clamp_alpha', type=float, default=1.9) + parser.add_argument('--pos_embed_dim', type=int, default=256) + parser.add_argument('--pos_beta', type=float, default=0.05) + parser.add_argument('--margin_tau', type=float, default=0.1) + parser.add_argument('--bgspp_lambda', type=float, default=1) + + parser.add_argument('--fdm_alpha', type=float, default=0.4) # low value, more training distribution + parser.add_argument('--num_embeddings', type=int, default=1536) # VQ embeddings + parser.add_argument("--train_ref_shot", type=int, default=4) + parser.add_argument("--num_ref_shot", type=int, default=4) + + args = parser.parse_args() + init_seeds(42) + + main(args) + diff --git a/main_Fourier.py b/main_Fourier.py new file mode 100644 index 0000000..7eed5d2 --- /dev/null +++ b/main_Fourier.py @@ -0,0 +1,328 @@ +import os +import warnings +import argparse +from tqdm import tqdm +import numpy as np +import torch +import timm +import torch.nn.functional as F +from torch.utils.data import DataLoader + +from train import train +from validate_Fourier import validate +from datasets.mvtec import MVTEC, MVTECANO +from datasets.visa import VISA, VISAANO +from datasets.btad import BTAD +from datasets.mvtec_3d import MVTEC3D +from datasets.mpdd import MPDD +from datasets.mvtec_loco import MVTECLOCO +from datasets.brats import BRATS +from datasets.capsules import CAPSULES, CAPSULESANO + +from models.fc_flow import load_flow_model +from models.modules import MultiScaleConv +from models.vq import MultiScaleVQ +from utils import init_seeds, get_residual_features, get_mc_matched_ref_features, get_mc_reference_features,get_fourier_residual_features +from utils import init_seeds, get_residual_features, get_mc_image_level_matched_features, get_mc_reference_features,get_fourier_residual_features +from utils import BoundaryAverager +from losses.loss import calculate_log_barrier_bi_occ_loss +from classes import VISA_TO_MVTEC, MVTEC_TO_VISA, MVTEC_TO_BTAD, MVTEC_TO_MVTEC3D +from classes import MVTEC_TO_MPDD, MVTEC_TO_MVTECLOCO, MVTEC_TO_BRATS +from classes import MVTEC_TO_MVTEC, VISA_TO_VISA +from classes import CAPSULES_TO_CAPSULES + +warnings.filterwarnings('ignore') + +TOTAL_SHOT = 4 # total few-shot reference samples +FIRST_STAGE_EPOCH = 10 +SETTINGS = {'visa_to_mvtec': VISA_TO_MVTEC, 'mvtec_to_visa': MVTEC_TO_VISA, + 'mvtec_to_btad': MVTEC_TO_BTAD, 'mvtec_to_mvtec3d': MVTEC_TO_MVTEC3D, + 'mvtec_to_mpdd': MVTEC_TO_MPDD, 'mvtec_to_mvtecloco': MVTEC_TO_MVTECLOCO, + 'mvtec_to_brats': MVTEC_TO_BRATS,'mvtec_to_mvtec':MVTEC_TO_MVTEC, 'visa_to_visa':VISA_TO_VISA, 'capsules_to_capsules': CAPSULES_TO_CAPSULES} + + +def main(args): + if args.setting in SETTINGS.keys(): + CLASSES = SETTINGS[args.setting] + else: + raise ValueError(f"Dataset setting must be in {SETTINGS.keys()}, but got {args.setting}.") + + if args.classes == 'capsules': # from mvtec to other datasets # from mvtec to other datasets + train_dataset1 = CAPSULES(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, + normalize="w50", + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader1 = DataLoader( + train_dataset1, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True + ) + train_dataset2 = CAPSULESANO(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader2 = DataLoader( + train_dataset2, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True + ) + + elif CLASSES['seen'][0] in MVTEC.CLASS_NAMES: # from mvtec to other datasets + train_dataset1 = MVTEC(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, + normalize="w50", + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader1 = DataLoader( + train_dataset1, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True + ) + train_dataset2 = MVTECANO(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader2 = DataLoader( + train_dataset2, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True + ) + else: # from visa to mvtec + train_dataset1 = VISA(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, + normalize="w50", + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader1 = DataLoader( + train_dataset1, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True + ) + train_dataset2 = VISAANO(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, + normalize="w50", + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader2 = DataLoader( + train_dataset2, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True + ) + if args.backbone == 'wide_resnet50_2': + encoder = timm.create_model('wide_resnet50_2', features_only=True, + out_indices=(1, 2, 3), pretrained=True).eval() # the pretrained checkpoint will be in /home/.cache/torch/hub/checkpoints/ + encoder = encoder.to(args.device) + feat_dims = encoder.feature_info.channels() + elif args.backbone == 'tf_efficientnet_b6':#10/26追加 + encoder = timm.create_model('tf_efficientnet_b6', features_only=True, + out_indices=(1, 2, 3), pretrained=True).eval() # the pretrained checkpoint will be in /home/.cache/torch/hub/checkpoints/ + encoder = encoder.to(args.device) + feat_dims = encoder.feature_info.channels() + boundary_ops = BoundaryAverager(num_levels=args.feature_levels) + vq_ops = MultiScaleVQ(num_embeddings=args.num_embeddings, channels=feat_dims).to(args.device) + optimizer_vq = torch.optim.Adam(vq_ops.parameters(), lr=args.lr, weight_decay=0.0005) + scheduler_vq = torch.optim.lr_scheduler.MultiStepLR(optimizer_vq, milestones=[70, 90], gamma=0.1) + + constraintor = MultiScaleConv(feat_dims).to(args.device) + # weight_decay is the l2 weight penalty lambda, weight_decay = lambda / 2 + optimizer0 = torch.optim.Adam(constraintor.parameters(), lr=args.lr, weight_decay=0.0005) + scheduler0 = torch.optim.lr_scheduler.MultiStepLR(optimizer0, milestones=[70, 90], gamma=0.1) + + # Normflow decoder + estimators = [load_flow_model(args, feat_dim) for feat_dim in feat_dims] + estimators = [decoder.to(args.device) for decoder in estimators] + params = list(estimators[0].parameters()) + for l in range(1, args.feature_levels): + params += list(estimators[l].parameters()) + optimizer1 = torch.optim.Adam(params, lr=args.lr, weight_decay=0.0005) + scheduler1 = torch.optim.lr_scheduler.MultiStepLR(optimizer1, milestones=[70, 90], gamma=0.1) + + best_img_auc = 0 + N_batch = 8192 + for epoch in range(args.epochs): + vq_ops.train() + constraintor.train() + for estimator in estimators: + estimator.train() + if epoch < FIRST_STAGE_EPOCH: + train_loader = train_loader1 + else: + train_loader = train_loader2 + train_loss_total, total_num = 0, 0 + progress_bar = tqdm(total=len(train_loader)) + progress_bar.set_description(f"Epoch[{epoch}/{args.epochs}]") + for step, batch in enumerate(train_loader): + progress_bar.update(1) + images, _, masks, class_names = batch + + images = images.to(args.device) + masks = masks.to(args.device) + + with torch.no_grad(): + features = encoder(images) + + ref_features = get_mc_reference_features(encoder, args.train_dataset_dir, class_names, images.device, args.train_ref_shot) + #mfeatures = get_mc_matched_ref_features(features, class_names, ref_features) + mfeatures = get_mc_image_level_matched_features(features, class_names, ref_features) + rfeatures = get_fourier_residual_features(features, mfeatures, pos_flag=True) + + lvl_masks = [] + for l in range(args.feature_levels): + _, _, h, w = rfeatures[l].size() + m = F.interpolate(masks, size=(h, w), mode='nearest').squeeze(1) + lvl_masks.append(m) + rfeatures_t = [rfeature.detach().clone() for rfeature in rfeatures] + + loss_vq = vq_ops(rfeatures, lvl_masks, train=True) + train_loss_total += loss_vq.item() + total_num += 1 + optimizer_vq.zero_grad() + loss_vq.backward() + optimizer_vq.step() + + rfeatures = constraintor(*rfeatures) + loss = 0 + for l in range(args.feature_levels): # backward svdd loss + e = rfeatures[l] + t = rfeatures_t[l] + bs, dim, h, w = e.size() + e = e.permute(0, 2, 3, 1).reshape(-1, dim) + t = t.permute(0, 2, 3, 1).reshape(-1, dim) + m = lvl_masks[l] + m = m.reshape(-1) + + loss_i, _, _ = calculate_log_barrier_bi_occ_loss(e, m, t) + loss += loss_i + optimizer0.zero_grad() + loss.backward() + optimizer0.step() + + train_loss_total += loss.item() + total_num += 1 + + # detach the rfeatures for flow optimization + rfeatures = [rfeature.detach().clone() for rfeature in rfeatures] + # train flow corresponding to with neck + loss, num = train(args, rfeatures, estimators, optimizer1, masks, boundary_ops, epoch, N_batch=N_batch, FIRST_STAGE_EPOCH=FIRST_STAGE_EPOCH) + train_loss_total += loss + total_num += num + + scheduler_vq.step() + scheduler0.step() + scheduler1.step() + + progress_bar.close() + print(f"Epoch[{epoch}/{args.epochs}]: train_loss: {train_loss_total / total_num}") + + if (epoch + 1) % args.eval_freq == 0: + s1_res, s2_res, s_res = [], [], [] + test_ref_features = load_mc_reference_features(args.test_ref_feature_dir, CLASSES['unseen'], args.device, args.num_ref_shot) + for class_name in CLASSES['unseen']: + if args.classes == 'capsules': + test_dataset = CAPSULES(args.test_dataset_dir, class_name=class_name, train=False, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in MVTEC.CLASS_NAMES: + test_dataset = MVTEC(args.test_dataset_dir, class_name=class_name, train=False, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in VISA.CLASS_NAMES: + test_dataset = VISA(args.test_dataset_dir, class_name=class_name, train=False, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in BTAD.CLASS_NAMES: + test_dataset = BTAD(args.test_dataset_dir, class_name=class_name, train=False, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in MVTEC3D.CLASS_NAMES: + test_dataset = MVTEC3D(args.test_dataset_dir, class_name=class_name, train=False, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in MPDD.CLASS_NAMES: + test_dataset = MPDD(args.test_dataset_dir, class_name=class_name, train=False, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in MVTECLOCO.CLASS_NAMES: + test_dataset = MVTECLOCO(args.test_dataset_dir, class_name=class_name, train=False, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in BRATS.CLASS_NAMES: + test_dataset = BRATS(args.test_dataset_dir, class_name=class_name, train=False, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + else: + raise ValueError('Unrecognized class name: {}'.format(class_name)) + test_loader = DataLoader( + test_dataset, batch_size=1, shuffle=False, num_workers=8, drop_last=False + ) + metrics = validate(args, encoder, vq_ops, constraintor, estimators, test_loader, test_ref_features[class_name], args.device, class_name) + img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro = metrics['scores'] + + print("Epoch: {}, Class Name: {}, Image AUC | AP | F1_Score: {} | {} | {}, Pixel AUC | AP | F1_Score | AUPRO: {} | {} | {} | {}".format( + epoch, class_name, img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro)) + s1_res.append(metrics['scores1']) + s2_res.append(metrics['scores2']) + s_res.append(metrics['scores']) + + s1_res = np.array(s1_res) + s2_res = np.array(s2_res) + s_res = np.array(s_res) + img_auc1, img_ap1, img_f1_score1, pix_auc1, pix_ap1, pix_f1_score1, pix_aupro1 = np.mean(s1_res, axis=0) + img_auc2, img_ap2, img_f1_score2, pix_auc2, pix_ap2, pix_f1_score2, pix_aupro2 = np.mean(s2_res, axis=0) + img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro = np.mean(s_res, axis=0) + print('(Logps) Average Image AUC | AP | F1_Score: {:.3f} | {:.3f} | {:.3f}, Average Pixel AUC | AP | F1_Score | AUPRO: {:.3f} | {:.3f} | {:.3f} | {:.3f}'.format( + img_auc1, img_ap1, img_f1_score1, pix_auc1, pix_ap1, pix_f1_score1, pix_aupro1)) + print('(BScores) Average Image AUC | AP | F1_Score: {:.3f} | {:.3f} | {:.3f}, Average Pixel AUC | AP | F1_Score | AUPRO: {:.3f} | {:.3f} | {:.3f} | {:.3f}'.format( + img_auc2, img_ap2, img_f1_score2, pix_auc2, pix_ap2, pix_f1_score2, pix_aupro2)) + print('(Merged) Average Image AUC | AP | F1_Score: {:.3f} | {:.3f} | {:.3f}, Average Pixel AUC | AP | F1_Score | AUPRO: {:.3f} | {:.3f} | {:.3f} | {:.3f}'.format( + img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro)) + + if img_auc > best_img_auc: + os.makedirs(args.checkpoint_path, exist_ok=True) + best_img_auc = img_auc + state_dict = {'vq_ops': vq_ops.state_dict(), + 'constraintor': constraintor.state_dict(), + 'estimators': [estimator.state_dict() for estimator in estimators]} + torch.save(state_dict, os.path.join(args.checkpoint_path, f'{args.setting}_epoch_{epoch}_checkpoints.pth')) + #torch.save(state_dict, os.path.join(args.checkpoint_path, f'{args.setting}_checkpoints.pth')) + + +def load_mc_reference_features(root_dir: str, class_names, device: torch.device, num_shot=4): + refs = {} + for class_name in class_names: + layer1_refs = np.load(os.path.join(root_dir, class_name, 'layer1.npy')) + layer2_refs = np.load(os.path.join(root_dir, class_name, 'layer2.npy')) + layer3_refs = np.load(os.path.join(root_dir, class_name, 'layer3.npy')) + + layer1_refs = torch.from_numpy(layer1_refs).to(device) + layer2_refs = torch.from_numpy(layer2_refs).to(device) + layer3_refs = torch.from_numpy(layer3_refs).to(device) + + K1 = (layer1_refs.shape[0] // TOTAL_SHOT) * num_shot + layer1_refs = layer1_refs[:K1, :] + K2 = (layer2_refs.shape[0] // TOTAL_SHOT) * num_shot + layer2_refs = layer2_refs[:K2, :] + K3 = (layer3_refs.shape[0] // TOTAL_SHOT) * num_shot + layer3_refs = layer3_refs[:K3, :] + + refs[class_name] = (layer1_refs, layer2_refs, layer3_refs) + + return refs + + +if __name__ == "__main__": + parser = argparse.ArgumentParser() + parser.add_argument('--setting', type=str, default="visa_to_mvtec") + parser.add_argument('--classes', type=str, default="none") + parser.add_argument('--train_dataset_dir', type=str, default="") + parser.add_argument('--test_dataset_dir', type=str, default="") + parser.add_argument('--test_ref_feature_dir', type=str, default="./ref_features/w50/mvtec_4shot") + parser.add_argument('--bgadweight_dir', type=str, default="none")# 12/16追加 + parser.add_argument('--batch_size', type=int, default=32) + parser.add_argument('--lr', type=float, default=1e-5) + parser.add_argument('--epochs', type=int, default=100) + parser.add_argument('--device', type=str, default="cuda:0") + parser.add_argument('--checkpoint_path', type=str, default="./checkpoints/") + parser.add_argument('--eval_freq', type=int, default=1) + parser.add_argument('--backbone', type=str, default="wide_resnet50_2") + parser.add_argument('--rank', type=int, default="0") + + # flow parameters + parser.add_argument('--flow_arch', type=str, default='conditional_flow_model') + parser.add_argument('--feature_levels', default=3, type=int) + parser.add_argument('--coupling_layers', type=int, default=10) + parser.add_argument('--clamp_alpha', type=float, default=1.9) + parser.add_argument('--pos_embed_dim', type=int, default=256) + parser.add_argument('--pos_beta', type=float, default=0.05) + parser.add_argument('--margin_tau', type=float, default=0.1) + parser.add_argument('--bgspp_lambda', type=float, default=1) + + parser.add_argument('--fdm_alpha', type=float, default=0.4) # low value, more training distribution + parser.add_argument('--num_embeddings', type=int, default=1536) # VQ embeddings + parser.add_argument("--train_ref_shot", type=int, default=4) + parser.add_argument("--num_ref_shot", type=int, default=4) + + args = parser.parse_args() + init_seeds(42) + + main(args) + diff --git a/main_ad.py b/main_ad.py new file mode 100644 index 0000000..2843a2a --- /dev/null +++ b/main_ad.py @@ -0,0 +1,353 @@ +import os +import warnings +import argparse +from tqdm import tqdm +import numpy as np +import torch +import timm +import torch.nn.functional as F +from torch.utils.data import DataLoader +import torch.nn as nn + +from train import train +from validate import validate +from datasets.mvtec import MVTEC, MVTECANO +from datasets.visa import VISA, VISAANO +from datasets.btad import BTAD +from datasets.mvtec_3d import MVTEC3D +from datasets.mpdd import MPDD +from datasets.mvtec_loco import MVTECLOCO +from datasets.brats import BRATS +from datasets.capsules import CAPSULES, CAPSULESANO + +from models.fc_flow import load_flow_model +from models.modules import MultiScaleConv +from models.vq import MultiScaleVQ +from utils import init_seeds, get_residual_features, get_mc_matched_ref_features, get_mc_reference_features +from utils import load_weights_ada +from utils import BoundaryAverager +from losses.loss import calculate_log_barrier_bi_occ_loss +from classes import VISA_TO_MVTEC, MVTEC_TO_VISA, MVTEC_TO_BTAD, MVTEC_TO_MVTEC3D +from classes import MVTEC_TO_MPDD, MVTEC_TO_MVTECLOCO, MVTEC_TO_BRATS +from classes import MVTEC_TO_MVTEC, VISA_TO_VISA +from classes import CAPSULES_TO_CAPSULES + +warnings.filterwarnings('ignore') + +TOTAL_SHOT = 4 # total few-shot reference samples +FIRST_STAGE_EPOCH = 10 +SETTINGS = {'visa_to_mvtec': VISA_TO_MVTEC, 'mvtec_to_visa': MVTEC_TO_VISA, + 'mvtec_to_btad': MVTEC_TO_BTAD, 'mvtec_to_mvtec3d': MVTEC_TO_MVTEC3D, + 'mvtec_to_mpdd': MVTEC_TO_MPDD, 'mvtec_to_mvtecloco': MVTEC_TO_MVTECLOCO, + 'mvtec_to_brats': MVTEC_TO_BRATS,'mvtec_to_mvtec':MVTEC_TO_MVTEC, 'visa_to_visa':VISA_TO_VISA, 'capsules_to_capsules': CAPSULES_TO_CAPSULES} + + +def main(args): + if args.setting in SETTINGS.keys(): + CLASSES = SETTINGS[args.setting] + else: + raise ValueError(f"Dataset setting must be in {SETTINGS.keys()}, but got {args.setting}.") + + if args.classes == 'capsules': # from mvtec to other datasets # from mvtec to other datasets + train_dataset1 = CAPSULES(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, + normalize="w50", + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader1 = DataLoader( + train_dataset1, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True + ) + train_dataset2 = CAPSULESANO(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader2 = DataLoader( + train_dataset2, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True + ) + + elif CLASSES['seen'][0] in MVTEC.CLASS_NAMES: # from mvtec to other datasets + train_dataset1 = MVTEC(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, + normalize="w50", + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader1 = DataLoader( + train_dataset1, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True + ) + train_dataset2 = MVTECANO(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader2 = DataLoader( + train_dataset2, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True + ) + else: # from visa to mvtec + train_dataset1 = VISA(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, + normalize="w50", + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader1 = DataLoader( + train_dataset1, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True + ) + train_dataset2 = VISAANO(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, + normalize="w50", + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader2 = DataLoader( + train_dataset2, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True + ) + if args.backbone == 'wide_resnet50_2': + encoder = timm.create_model('wide_resnet50_2', features_only=True, + out_indices=(1, 2, 3), pretrained=True).eval() # the pretrained checkpoint will be in /home/.cache/torch/hub/checkpoints/ + encoder = encoder.to(args.device) + feat_dims = encoder.feature_info.channels() + elif args.backbone == 'tf_efficientnet_b6':#10/26追加 + encoder = timm.create_model('tf_efficientnet_b6', features_only=True, + out_indices=(1, 2, 3), pretrained=True).eval() # the pretrained checkpoint will be in /home/.cache/torch/hub/checkpoints/ + encoder = encoder.to(args.device) + feat_dims = encoder.feature_info.channels() +#feat_fims [40, 72, 200] + adapters = nn.ModuleList([ + nn.Conv2d(in_channels=feat_dim, out_channels=feat_dim, kernel_size=1, stride=1) + for feat_dim in feat_dims + ]).to(args.device) #追加1/8 + params_ada = list(adapters[0].parameters()) #追加1/8 + optimizer_ada = torch.optim.Adam(params_ada, lr=args.lr, weight_decay=0.0005) #追加1/8 + scheduler_ada = torch.optim.lr_scheduler.MultiStepLR(optimizer_ada, milestones=[70, 90], gamma=0.1) #追加1/8 + + boundary_ops = BoundaryAverager(num_levels=args.feature_levels) + vq_ops = MultiScaleVQ(num_embeddings=args.num_embeddings, channels=feat_dims).to(args.device) + optimizer_vq = torch.optim.Adam(vq_ops.parameters(), lr=args.lr, weight_decay=0.0005) + scheduler_vq = torch.optim.lr_scheduler.MultiStepLR(optimizer_vq, milestones=[70, 90], gamma=0.1) + + constraintor = MultiScaleConv(feat_dims).to(args.device) + # weight_decay is the l2 weight penalty lambda, weight_decay = lambda / 2 + optimizer0 = torch.optim.Adam(constraintor.parameters(), lr=args.lr, weight_decay=0.0005) + scheduler0 = torch.optim.lr_scheduler.MultiStepLR(optimizer0, milestones=[70, 90], gamma=0.1) + + # Normflow decoder + estimators = [load_flow_model(args, feat_dim) for feat_dim in feat_dims] + estimators = [decoder.to(args.device) for decoder in estimators] + params = list(estimators[0].parameters()) + for l in range(1, args.feature_levels): + params += list(estimators[l].parameters()) + optimizer1 = torch.optim.Adam(params, lr=args.lr, weight_decay=0.0005) + scheduler1 = torch.optim.lr_scheduler.MultiStepLR(optimizer1, milestones=[70, 90], gamma=0.1) + + best_img_auc = 0 + N_batch = 8192 + for epoch in range(args.epochs): + adapters.train() #追加1/8 + vq_ops.train() + constraintor.train() + for estimator in estimators: + estimator.train() + if epoch < FIRST_STAGE_EPOCH: + train_loader = train_loader1 + else: + train_loader = train_loader2 + train_loss_total, total_num = 0, 0 + progress_bar = tqdm(total=len(train_loader)) + progress_bar.set_description(f"Epoch[{epoch}/{args.epochs}]") + for step, batch in enumerate(train_loader): + progress_bar.update(1) + images, _, masks, class_names = batch + + images = images.to(args.device) + masks = masks.to(args.device) + + with torch.no_grad(): + features = encoder(images) + features_ad = encoder(images)#追加1/8 + + + ref_features = get_mc_reference_features(encoder, args.train_dataset_dir, class_names, images.device, args.train_ref_shot) + mfeatures = get_mc_matched_ref_features(features, class_names, ref_features) + rfeatures = get_residual_features(features, mfeatures, pos_flag=True) + + features = [adapters[i](features[i]) for i in range(len(features))] #追加1/8 + ref_features_ad = get_mc_reference_features(encoder, args.train_dataset_dir, class_names, images.device, args.train_ref_shot)#追加1/8 + mfeatures_ad = get_mc_matched_ref_features(features_ad, class_names, ref_features_ad)#追加1/8 + rfeatures_ad = get_residual_features(features_ad, mfeatures_ad, pos_flag=True)#追加1/8 + + + lvl_masks = [] + for l in range(args.feature_levels): + _, _, h, w = rfeatures[l].size() + m = F.interpolate(masks, size=(h, w), mode='nearest').squeeze(1) + lvl_masks.append(m) + rfeatures_t = [rfeature.detach().clone() for rfeature in rfeatures] + rfeatures_ad_t = [rfeature.detach().clone() for rfeature in rfeatures_ad] #追加1/8 + + + loss_vq = vq_ops(rfeatures, lvl_masks, train=True) + train_loss_total += loss_vq.item() + total_num += 1 + optimizer_vq.zero_grad() + loss_vq.backward() + optimizer_vq.step() + + rfeatures = constraintor(*rfeatures) + loss = 0 + for l in range(args.feature_levels): # backward svdd loss + e = rfeatures[l] + t = rfeatures_t[l] + bs, dim, h, w = e.size() + e = e.permute(0, 2, 3, 1).reshape(-1, dim) + t = t.permute(0, 2, 3, 1).reshape(-1, dim) + m = lvl_masks[l] + m = m.reshape(-1) + + loss_i, _, _ = calculate_log_barrier_bi_occ_loss(e, m, t) + loss += loss_i + + optimizer_ada.zero_grad() #追加1/8 + optimizer0.zero_grad() + loss.backward() + optimizer0.step() + optimizer_ada.step() #追加1/8 + + train_loss_total += loss.item() + total_num += 1 + + # detach the rfeatures for flow optimization + rfeatures = [rfeature.detach().clone() for rfeature in rfeatures] + # train flow corresponding to with neck + loss, num = train(args, rfeatures, estimators, optimizer1, masks, boundary_ops, epoch, N_batch=N_batch, FIRST_STAGE_EPOCH=FIRST_STAGE_EPOCH) + train_loss_total += loss + total_num += num + + scheduler_vq.step() + scheduler_ada.step() #追加1/8 + scheduler0.step() + scheduler1.step() + + progress_bar.close() + print(f"Epoch[{epoch}/{args.epochs}]: train_loss: {train_loss_total / total_num}") + + if (epoch + 1) % args.eval_freq == 0: + s1_res, s2_res, s_res = [], [], [] + test_ref_features = load_mc_reference_features(args.test_ref_feature_dir, CLASSES['unseen'], args.device, args.num_ref_shot) + for class_name in CLASSES['unseen']: + if args.classes == 'capsules': + test_dataset = CAPSULES(args.test_dataset_dir, class_name=class_name, train=False, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in MVTEC.CLASS_NAMES: + test_dataset = MVTEC(args.test_dataset_dir, class_name=class_name, train=False, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in VISA.CLASS_NAMES: + test_dataset = VISA(args.test_dataset_dir, class_name=class_name, train=False, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in BTAD.CLASS_NAMES: + test_dataset = BTAD(args.test_dataset_dir, class_name=class_name, train=False, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in MVTEC3D.CLASS_NAMES: + test_dataset = MVTEC3D(args.test_dataset_dir, class_name=class_name, train=False, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in MPDD.CLASS_NAMES: + test_dataset = MPDD(args.test_dataset_dir, class_name=class_name, train=False, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in MVTECLOCO.CLASS_NAMES: + test_dataset = MVTECLOCO(args.test_dataset_dir, class_name=class_name, train=False, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in BRATS.CLASS_NAMES: + test_dataset = BRATS(args.test_dataset_dir, class_name=class_name, train=False, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + else: + raise ValueError('Unrecognized class name: {}'.format(class_name)) + test_loader = DataLoader( + test_dataset, batch_size=1, shuffle=False, num_workers=8, drop_last=False + ) + metrics = validate(args, encoder, vq_ops, constraintor, estimators, test_loader, test_ref_features[class_name], args.device, class_name) + img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro = metrics['scores'] + + print("Epoch: {}, Class Name: {}, Image AUC | AP | F1_Score: {} | {} | {}, Pixel AUC | AP | F1_Score | AUPRO: {} | {} | {} | {}".format( + epoch, class_name, img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro)) + s1_res.append(metrics['scores1']) + s2_res.append(metrics['scores2']) + s_res.append(metrics['scores']) + + s1_res = np.array(s1_res) + s2_res = np.array(s2_res) + s_res = np.array(s_res) + img_auc1, img_ap1, img_f1_score1, pix_auc1, pix_ap1, pix_f1_score1, pix_aupro1 = np.mean(s1_res, axis=0) + img_auc2, img_ap2, img_f1_score2, pix_auc2, pix_ap2, pix_f1_score2, pix_aupro2 = np.mean(s2_res, axis=0) + img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro = np.mean(s_res, axis=0) + print('(Logps) Average Image AUC | AP | F1_Score: {:.3f} | {:.3f} | {:.3f}, Average Pixel AUC | AP | F1_Score | AUPRO: {:.3f} | {:.3f} | {:.3f} | {:.3f}'.format( + img_auc1, img_ap1, img_f1_score1, pix_auc1, pix_ap1, pix_f1_score1, pix_aupro1)) + print('(BScores) Average Image AUC | AP | F1_Score: {:.3f} | {:.3f} | {:.3f}, Average Pixel AUC | AP | F1_Score | AUPRO: {:.3f} | {:.3f} | {:.3f} | {:.3f}'.format( + img_auc2, img_ap2, img_f1_score2, pix_auc2, pix_ap2, pix_f1_score2, pix_aupro2)) + print('(Merged) Average Image AUC | AP | F1_Score: {:.3f} | {:.3f} | {:.3f}, Average Pixel AUC | AP | F1_Score | AUPRO: {:.3f} | {:.3f} | {:.3f} | {:.3f}'.format( + img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro)) + + if img_auc > best_img_auc: + os.makedirs(args.checkpoint_path, exist_ok=True) + best_img_auc = img_auc + state_dict = {'vq_ops': vq_ops.state_dict(), + 'adapter_state_dict': [adapter.state_dict() for adapter in adapters], #追加1/8 + 'constraintor': constraintor.state_dict(), + 'estimators': [estimator.state_dict() for estimator in estimators]} + torch.save(state_dict, os.path.join(args.checkpoint_path, f'{args.setting}_epoch_{epoch}_checkpoints.pth')) + #torch.save(state_dict, os.path.join(args.checkpoint_path, f'{args.setting}_checkpoints.pth')) + + +def load_mc_reference_features(root_dir: str, class_names, device: torch.device, num_shot=4): + refs = {} + for class_name in class_names: + layer1_refs = np.load(os.path.join(root_dir, class_name, 'layer1.npy')) + layer2_refs = np.load(os.path.join(root_dir, class_name, 'layer2.npy')) + layer3_refs = np.load(os.path.join(root_dir, class_name, 'layer3.npy')) + + layer1_refs = torch.from_numpy(layer1_refs).to(device) + layer2_refs = torch.from_numpy(layer2_refs).to(device) + layer3_refs = torch.from_numpy(layer3_refs).to(device) + + K1 = (layer1_refs.shape[0] // TOTAL_SHOT) * num_shot + layer1_refs = layer1_refs[:K1, :] + K2 = (layer2_refs.shape[0] // TOTAL_SHOT) * num_shot + layer2_refs = layer2_refs[:K2, :] + K3 = (layer3_refs.shape[0] // TOTAL_SHOT) * num_shot + layer3_refs = layer3_refs[:K3, :] + + refs[class_name] = (layer1_refs, layer2_refs, layer3_refs) + + return refs + + +if __name__ == "__main__": + parser = argparse.ArgumentParser() + parser.add_argument('--setting', type=str, default="visa_to_mvtec") + parser.add_argument('--classes', type=str, default="none") + parser.add_argument('--train_dataset_dir', type=str, default="") + parser.add_argument('--test_dataset_dir', type=str, default="") + parser.add_argument('--test_ref_feature_dir', type=str, default="./ref_features/w50/mvtec_4shot") + parser.add_argument('--bgadweight_dir', type=str, default="none")# 12/16追加 + parser.add_argument('--batch_size', type=int, default=32) + parser.add_argument('--lr', type=float, default=1e-5) + parser.add_argument('--epochs', type=int, default=100) + parser.add_argument('--device', type=str, default="cuda:0") + parser.add_argument('--checkpoint_path', type=str, default="./checkpoints/") + parser.add_argument('--bgad_weight_dir', type=str, default="none") # 1/8追加 + parser.add_argument('--resad_weight_dir', type=str, default="none") # 1/8追加 + parser.add_argument('--eval_freq', type=int, default=1) + parser.add_argument('--backbone', type=str, default="wide_resnet50_2") + + # flow parameters + parser.add_argument('--flow_arch', type=str, default='conditional_flow_model') + parser.add_argument('--feature_levels', default=3, type=int) + parser.add_argument('--coupling_layers', type=int, default=10) + parser.add_argument('--clamp_alpha', type=float, default=1.9) + parser.add_argument('--pos_embed_dim', type=int, default=256) + parser.add_argument('--pos_beta', type=float, default=0.05) + parser.add_argument('--margin_tau', type=float, default=0.1) + parser.add_argument('--bgspp_lambda', type=float, default=1) + + parser.add_argument('--fdm_alpha', type=float, default=0.4) # low value, more training distribution + parser.add_argument('--num_embeddings', type=int, default=1536) # VQ embeddings + parser.add_argument("--train_ref_shot", type=int, default=4) + parser.add_argument("--num_ref_shot", type=int, default=4) + + args = parser.parse_args() + init_seeds(42) + + main(args) diff --git a/main_attention.py b/main_attention.py new file mode 100644 index 0000000..d2119b7 --- /dev/null +++ b/main_attention.py @@ -0,0 +1,326 @@ +import os +import warnings +import argparse +from tqdm import tqdm +import numpy as np +import torch +import timm +import torch.nn.functional as F +from torch.utils.data import DataLoader + +from train import train +from validate_attention import validate +from datasets.mvtec import MVTEC, MVTECANO +from datasets.visa import VISA, VISAANO +from datasets.btad import BTAD +from datasets.mvtec_3d import MVTEC3D +from datasets.mpdd import MPDD +from datasets.mvtec_loco import MVTECLOCO +from datasets.brats import BRATS +from datasets.capsules import CAPSULES, CAPSULESANO + +from models.fc_flow import load_flow_model +from models.modules import MultiScaleConv +from models.vq import MultiScaleVQ +from utils import init_seeds, get_residual_features, get_mc_matched_ref_features, get_mc_reference_features,get_mc_soft_matched_features +from utils import BoundaryAverager +from losses.loss import calculate_log_barrier_bi_occ_loss +from classes import VISA_TO_MVTEC, MVTEC_TO_VISA, MVTEC_TO_BTAD, MVTEC_TO_MVTEC3D +from classes import MVTEC_TO_MPDD, MVTEC_TO_MVTECLOCO, MVTEC_TO_BRATS +from classes import MVTEC_TO_MVTEC, VISA_TO_VISA +from classes import CAPSULES_TO_CAPSULES + +warnings.filterwarnings('ignore') + +TOTAL_SHOT = 4 # total few-shot reference samples +FIRST_STAGE_EPOCH = 10 +SETTINGS = {'visa_to_mvtec': VISA_TO_MVTEC, 'mvtec_to_visa': MVTEC_TO_VISA, + 'mvtec_to_btad': MVTEC_TO_BTAD, 'mvtec_to_mvtec3d': MVTEC_TO_MVTEC3D, + 'mvtec_to_mpdd': MVTEC_TO_MPDD, 'mvtec_to_mvtecloco': MVTEC_TO_MVTECLOCO, + 'mvtec_to_brats': MVTEC_TO_BRATS,'mvtec_to_mvtec':MVTEC_TO_MVTEC, 'visa_to_visa':VISA_TO_VISA, 'capsules_to_capsules': CAPSULES_TO_CAPSULES} + + +def main(args): + if args.setting in SETTINGS.keys(): + CLASSES = SETTINGS[args.setting] + else: + raise ValueError(f"Dataset setting must be in {SETTINGS.keys()}, but got {args.setting}.") + + if args.classes == 'capsules': # from mvtec to other datasets # from mvtec to other datasets + train_dataset1 = CAPSULES(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, + normalize="w50", + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader1 = DataLoader( + train_dataset1, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True + ) + train_dataset2 = CAPSULESANO(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader2 = DataLoader( + train_dataset2, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True + ) + + elif CLASSES['seen'][0] in MVTEC.CLASS_NAMES: # from mvtec to other datasets + train_dataset1 = MVTEC(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, + normalize="w50", + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader1 = DataLoader( + train_dataset1, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True + ) + train_dataset2 = MVTECANO(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader2 = DataLoader( + train_dataset2, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True + ) + else: # from visa to mvtec + train_dataset1 = VISA(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, + normalize="w50", + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader1 = DataLoader( + train_dataset1, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True + ) + train_dataset2 = VISAANO(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, + normalize="w50", + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader2 = DataLoader( + train_dataset2, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True + ) + if args.backbone == 'wide_resnet50_2': + encoder = timm.create_model('wide_resnet50_2', features_only=True, + out_indices=(1, 2, 3), pretrained=True).eval() # the pretrained checkpoint will be in /home/.cache/torch/hub/checkpoints/ + encoder = encoder.to(args.device) + feat_dims = encoder.feature_info.channels() + elif args.backbone == 'tf_efficientnet_b6':#10/26追加 + encoder = timm.create_model('tf_efficientnet_b6', features_only=True, + out_indices=(1, 2, 3), pretrained=True).eval() # the pretrained checkpoint will be in /home/.cache/torch/hub/checkpoints/ + encoder = encoder.to(args.device) + feat_dims = encoder.feature_info.channels() + boundary_ops = BoundaryAverager(num_levels=args.feature_levels) + vq_ops = MultiScaleVQ(num_embeddings=args.num_embeddings, channels=feat_dims).to(args.device) + optimizer_vq = torch.optim.Adam(vq_ops.parameters(), lr=args.lr, weight_decay=0.0005) + scheduler_vq = torch.optim.lr_scheduler.MultiStepLR(optimizer_vq, milestones=[70, 90], gamma=0.1) + + constraintor = MultiScaleConv(feat_dims).to(args.device) + # weight_decay is the l2 weight penalty lambda, weight_decay = lambda / 2 + optimizer0 = torch.optim.Adam(constraintor.parameters(), lr=args.lr, weight_decay=0.0005) + scheduler0 = torch.optim.lr_scheduler.MultiStepLR(optimizer0, milestones=[70, 90], gamma=0.1) + + # Normflow decoder + estimators = [load_flow_model(args, feat_dim) for feat_dim in feat_dims] + estimators = [decoder.to(args.device) for decoder in estimators] + params = list(estimators[0].parameters()) + for l in range(1, args.feature_levels): + params += list(estimators[l].parameters()) + optimizer1 = torch.optim.Adam(params, lr=args.lr, weight_decay=0.0005) + scheduler1 = torch.optim.lr_scheduler.MultiStepLR(optimizer1, milestones=[70, 90], gamma=0.1) + + best_img_auc = 0 + N_batch = 8192 + for epoch in range(args.epochs): + vq_ops.train() + constraintor.train() + for estimator in estimators: + estimator.train() + if epoch < FIRST_STAGE_EPOCH: + train_loader = train_loader1 + else: + train_loader = train_loader2 + train_loss_total, total_num = 0, 0 + progress_bar = tqdm(total=len(train_loader)) + progress_bar.set_description(f"Epoch[{epoch}/{args.epochs}]") + for step, batch in enumerate(train_loader): + progress_bar.update(1) + images, _, masks, class_names = batch + + images = images.to(args.device) + masks = masks.to(args.device) + + with torch.no_grad(): + features = encoder(images) + + ref_features = get_mc_reference_features(encoder, args.train_dataset_dir, class_names, images.device, args.train_ref_shot) + mfeatures = get_mc_soft_matched_features(features, class_names, ref_features) + rfeatures = get_residual_features(features, mfeatures, pos_flag=True) + + lvl_masks = [] + for l in range(args.feature_levels): + _, _, h, w = rfeatures[l].size() + m = F.interpolate(masks, size=(h, w), mode='nearest').squeeze(1) + lvl_masks.append(m) + rfeatures_t = [rfeature.detach().clone() for rfeature in rfeatures] + + loss_vq = vq_ops(rfeatures, lvl_masks, train=True) + train_loss_total += loss_vq.item() + total_num += 1 + optimizer_vq.zero_grad() + loss_vq.backward() + optimizer_vq.step() + + rfeatures = constraintor(*rfeatures) + loss = 0 + for l in range(args.feature_levels): # backward svdd loss + e = rfeatures[l] + t = rfeatures_t[l] + bs, dim, h, w = e.size() + e = e.permute(0, 2, 3, 1).reshape(-1, dim) + t = t.permute(0, 2, 3, 1).reshape(-1, dim) + m = lvl_masks[l] + m = m.reshape(-1) + + loss_i, _, _ = calculate_log_barrier_bi_occ_loss(e, m, t) + loss += loss_i + optimizer0.zero_grad() + loss.backward() + optimizer0.step() + + train_loss_total += loss.item() + total_num += 1 + + # detach the rfeatures for flow optimization + rfeatures = [rfeature.detach().clone() for rfeature in rfeatures] + # train flow corresponding to with neck + loss, num = train(args, rfeatures, estimators, optimizer1, masks, boundary_ops, epoch, N_batch=N_batch, FIRST_STAGE_EPOCH=FIRST_STAGE_EPOCH) + train_loss_total += loss + total_num += num + + scheduler_vq.step() + scheduler0.step() + scheduler1.step() + + progress_bar.close() + print(f"Epoch[{epoch}/{args.epochs}]: train_loss: {train_loss_total / total_num}") + + if (epoch + 1) % args.eval_freq == 0: + s1_res, s2_res, s_res = [], [], [] + test_ref_features = load_mc_reference_features(args.test_ref_feature_dir, CLASSES['unseen'], args.device, args.num_ref_shot) + for class_name in CLASSES['unseen']: + if args.classes == 'capsules': + test_dataset = CAPSULES(args.test_dataset_dir, class_name=class_name, train=False, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in MVTEC.CLASS_NAMES: + test_dataset = MVTEC(args.test_dataset_dir, class_name=class_name, train=False, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in VISA.CLASS_NAMES: + test_dataset = VISA(args.test_dataset_dir, class_name=class_name, train=False, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in BTAD.CLASS_NAMES: + test_dataset = BTAD(args.test_dataset_dir, class_name=class_name, train=False, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in MVTEC3D.CLASS_NAMES: + test_dataset = MVTEC3D(args.test_dataset_dir, class_name=class_name, train=False, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in MPDD.CLASS_NAMES: + test_dataset = MPDD(args.test_dataset_dir, class_name=class_name, train=False, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in MVTECLOCO.CLASS_NAMES: + test_dataset = MVTECLOCO(args.test_dataset_dir, class_name=class_name, train=False, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in BRATS.CLASS_NAMES: + test_dataset = BRATS(args.test_dataset_dir, class_name=class_name, train=False, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + else: + raise ValueError('Unrecognized class name: {}'.format(class_name)) + test_loader = DataLoader( + test_dataset, batch_size=1, shuffle=False, num_workers=8, drop_last=False + ) + metrics = validate(args, encoder, vq_ops, constraintor, estimators, test_loader, test_ref_features[class_name], args.device, class_name) + img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro = metrics['scores'] + + print("Epoch: {}, Class Name: {}, Image AUC | AP | F1_Score: {} | {} | {}, Pixel AUC | AP | F1_Score | AUPRO: {} | {} | {} | {}".format( + epoch, class_name, img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro)) + s1_res.append(metrics['scores1']) + s2_res.append(metrics['scores2']) + s_res.append(metrics['scores']) + + s1_res = np.array(s1_res) + s2_res = np.array(s2_res) + s_res = np.array(s_res) + img_auc1, img_ap1, img_f1_score1, pix_auc1, pix_ap1, pix_f1_score1, pix_aupro1 = np.mean(s1_res, axis=0) + img_auc2, img_ap2, img_f1_score2, pix_auc2, pix_ap2, pix_f1_score2, pix_aupro2 = np.mean(s2_res, axis=0) + img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro = np.mean(s_res, axis=0) + print('(Logps) Average Image AUC | AP | F1_Score: {:.3f} | {:.3f} | {:.3f}, Average Pixel AUC | AP | F1_Score | AUPRO: {:.3f} | {:.3f} | {:.3f} | {:.3f}'.format( + img_auc1, img_ap1, img_f1_score1, pix_auc1, pix_ap1, pix_f1_score1, pix_aupro1)) + print('(BScores) Average Image AUC | AP | F1_Score: {:.3f} | {:.3f} | {:.3f}, Average Pixel AUC | AP | F1_Score | AUPRO: {:.3f} | {:.3f} | {:.3f} | {:.3f}'.format( + img_auc2, img_ap2, img_f1_score2, pix_auc2, pix_ap2, pix_f1_score2, pix_aupro2)) + print('(Merged) Average Image AUC | AP | F1_Score: {:.3f} | {:.3f} | {:.3f}, Average Pixel AUC | AP | F1_Score | AUPRO: {:.3f} | {:.3f} | {:.3f} | {:.3f}'.format( + img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro)) + + if img_auc > best_img_auc: + os.makedirs(args.checkpoint_path, exist_ok=True) + best_img_auc = img_auc + state_dict = {'vq_ops': vq_ops.state_dict(), + 'constraintor': constraintor.state_dict(), + 'estimators': [estimator.state_dict() for estimator in estimators]} + torch.save(state_dict, os.path.join(args.checkpoint_path, f'{args.setting}_epoch_{epoch}_checkpoints.pth')) + #torch.save(state_dict, os.path.join(args.checkpoint_path, f'{args.setting}_checkpoints.pth')) + + +def load_mc_reference_features(root_dir: str, class_names, device: torch.device, num_shot=4): + refs = {} + for class_name in class_names: + layer1_refs = np.load(os.path.join(root_dir, class_name, 'layer1.npy')) + layer2_refs = np.load(os.path.join(root_dir, class_name, 'layer2.npy')) + layer3_refs = np.load(os.path.join(root_dir, class_name, 'layer3.npy')) + + layer1_refs = torch.from_numpy(layer1_refs).to(device) + layer2_refs = torch.from_numpy(layer2_refs).to(device) + layer3_refs = torch.from_numpy(layer3_refs).to(device) + + K1 = (layer1_refs.shape[0] // TOTAL_SHOT) * num_shot + layer1_refs = layer1_refs[:K1, :] + K2 = (layer2_refs.shape[0] // TOTAL_SHOT) * num_shot + layer2_refs = layer2_refs[:K2, :] + K3 = (layer3_refs.shape[0] // TOTAL_SHOT) * num_shot + layer3_refs = layer3_refs[:K3, :] + + refs[class_name] = (layer1_refs, layer2_refs, layer3_refs) + + return refs + + +if __name__ == "__main__": + parser = argparse.ArgumentParser() + parser.add_argument('--setting', type=str, default="visa_to_mvtec") + parser.add_argument('--classes', type=str, default="none") + parser.add_argument('--train_dataset_dir', type=str, default="") + parser.add_argument('--test_dataset_dir', type=str, default="") + parser.add_argument('--test_ref_feature_dir', type=str, default="./ref_features/w50/mvtec_4shot") + parser.add_argument('--bgadweight_dir', type=str, default="none")# 12/16追加 + parser.add_argument('--batch_size', type=int, default=32) + parser.add_argument('--lr', type=float, default=1e-5) + parser.add_argument('--epochs', type=int, default=100) + parser.add_argument('--device', type=str, default="cuda:0") + parser.add_argument('--checkpoint_path', type=str, default="./checkpoints/") + parser.add_argument('--eval_freq', type=int, default=1) + parser.add_argument('--backbone', type=str, default="wide_resnet50_2") + parser.add_argument('--rank', type=int, default="0") + + # flow parameters + parser.add_argument('--flow_arch', type=str, default='conditional_flow_model') + parser.add_argument('--feature_levels', default=3, type=int) + parser.add_argument('--coupling_layers', type=int, default=10) + parser.add_argument('--clamp_alpha', type=float, default=1.9) + parser.add_argument('--pos_embed_dim', type=int, default=256) + parser.add_argument('--pos_beta', type=float, default=0.05) + parser.add_argument('--margin_tau', type=float, default=0.1) + parser.add_argument('--bgspp_lambda', type=float, default=1) + + parser.add_argument('--fdm_alpha', type=float, default=0.4) # low value, more training distribution + parser.add_argument('--num_embeddings', type=int, default=1536) # VQ embeddings + parser.add_argument("--train_ref_shot", type=int, default=4) + parser.add_argument("--num_ref_shot", type=int, default=4) + + args = parser.parse_args() + init_seeds(42) + + main(args) + diff --git a/main_freq_blend.py b/main_freq_blend.py new file mode 100644 index 0000000..97e0aed --- /dev/null +++ b/main_freq_blend.py @@ -0,0 +1,352 @@ +import os +import warnings +import argparse +from tqdm import tqdm +import numpy as np +import torch +import torch.nn as nn +import timm +import torch.nn.functional as F +from torch.utils.data import DataLoader + +from train import train +# ※ファイル名に合わせて変更してください +from validate_freq_blend import validate +from datasets.mvtec import MVTEC, MVTECANO +from datasets.visa import VISA, VISAANO +from datasets.btad import BTAD +from datasets.mvtec_3d import MVTEC3D +from datasets.mpdd import MPDD +from datasets.mvtec_loco import MVTECLOCO +from datasets.brats import BRATS +from datasets.capsules import CAPSULES, CAPSULESANO + +from models.fc_flow import load_flow_model +from models.modules import ConvBnAct, get_position_encoding +from utils import init_seeds +from utils import get_random_normal_images, load_and_transform_vision_data +from utils import BoundaryAverager +from losses.loss import calculate_log_barrier_bi_occ_loss +from classes import VISA_TO_MVTEC, MVTEC_TO_VISA, MVTEC_TO_BTAD, MVTEC_TO_MVTEC3D +from classes import MVTEC_TO_MPDD, MVTEC_TO_MVTECLOCO, MVTEC_TO_BRATS +from classes import MVTEC_TO_MVTEC, VISA_TO_VISA +from classes import CAPSULES_TO_CAPSULES + +warnings.filterwarnings('ignore') + +TOTAL_SHOT = 4 +FIRST_STAGE_EPOCH = 10 +SETTINGS = {'visa_to_mvtec': VISA_TO_MVTEC, 'mvtec_to_visa': MVTEC_TO_VISA, + 'mvtec_to_btad': MVTEC_TO_BTAD, 'mvtec_to_mvtec3d': MVTEC_TO_MVTEC3D, + 'mvtec_to_mpdd': MVTEC_TO_MPDD, 'mvtec_to_mvtecloco': MVTECLOCO, + 'mvtec_to_brats': MVTEC_TO_BRATS,'mvtec_to_mvtec':MVTEC_TO_MVTEC, 'visa_to_visa':VISA_TO_VISA, 'capsules_to_capsules': CAPSULES_TO_CAPSULES} + +# 元の Constraintor (CNN版) +class MultiScaleConv(nn.Module): + def __init__(self, channels): + super().__init__() + self.projs = nn.ModuleList([ConvBnAct(c, c) for c in channels]) + + def forward(self, *features): + return tuple([self.projs[i](features[i]) for i in range(len(features))]) + +class HaarWaveletFilter2Component(nn.Module): + def __init__(self): + super().__init__() + ll = torch.tensor([[0.5, 0.5], [0.5, 0.5]]) + hl = torch.tensor([[-0.5, -0.5], [0.5, 0.5]]) + lh = torch.tensor([[-0.5, 0.5], [-0.5, 0.5]]) + hh = torch.tensor([[0.5, -0.5], [-0.5, 0.5]]) + + self.register_buffer('k_ll', ll.view(1, 1, 2, 2)) + self.register_buffer('k_hl', hl.view(1, 1, 2, 2)) + self.register_buffer('k_lh', lh.view(1, 1, 2, 2)) + self.register_buffer('k_hh', hh.view(1, 1, 2, 2)) + + def get_LF_HF(self, x): + B, C, H, W = x.shape + ll = F.conv2d(x, self.k_ll.expand(C, 1, 2, 2), stride=2, groups=C) + hl = F.conv2d(x, self.k_hl.expand(C, 1, 2, 2), stride=2, groups=C) + lh = F.conv2d(x, self.k_lh.expand(C, 1, 2, 2), stride=2, groups=C) + hh = F.conv2d(x, self.k_hh.expand(C, 1, 2, 2), stride=2, groups=C) + + lf = F.conv_transpose2d(ll, self.k_ll.expand(C, 1, 2, 2), stride=2, groups=C) + hf = F.conv_transpose2d(hl, self.k_hl.expand(C, 1, 2, 2), stride=2, groups=C) + \ + F.conv_transpose2d(lh, self.k_lh.expand(C, 1, 2, 2), stride=2, groups=C) + \ + F.conv_transpose2d(hh, self.k_hh.expand(C, 1, 2, 2), stride=2, groups=C) + return lf, hf + +# ========================================== +# アイデアB: 分割統治カンペ抽出関数 +# ========================================== +def get_freq_reference_features(encoder, wav_filter, root, class_names, device, num_shot=4): + ref_lf_dict = {} + ref_hf_dict = {} + class_names = np.unique(class_names) + for class_name in class_names: + normal_paths = get_random_normal_images(root, class_name, num_shot) + images = load_and_transform_vision_data(normal_paths, device) + with torch.no_grad(): + features_raw = encoder(images) + lf_list, hf_list = [], [] + for l in range(len(features_raw)): + lf, hf = wav_filter.get_LF_HF(features_raw[l]) + bs, c, h, w = lf.shape + lf_list.append(lf.permute(0, 2, 3, 1).reshape(-1, c)) + hf_list.append(hf.permute(0, 2, 3, 1).reshape(-1, c)) + ref_lf_dict[class_name] = lf_list + ref_hf_dict[class_name] = hf_list + return ref_lf_dict, ref_hf_dict + +# ========================================== +# アイデアB: LFマッチングと固定比率ブレンド残差 +# ========================================== +def get_freq_matched_residuals(test_lf_list, test_hf_list, ref_lf_list, ref_hf_list, alpha=0.5, pos_flag=True): + rfeatures = [] + for l in range(len(test_lf_list)): + t_lf = test_lf_list[l] + t_hf = test_hf_list[l] + r_lf_all = ref_lf_list[l] + r_hf_all = ref_hf_list[l] + + B, C, H, W = t_lf.shape + t_lf_flat = t_lf.permute(0, 2, 3, 1).reshape(-1, C).contiguous() + + # ★ マッチングは LF (低周波=構造) のみで行う ★ + t_lf_n = F.normalize(t_lf_flat, p=2, dim=1) + r_lf_n = F.normalize(r_lf_all, p=2, dim=1) + + dist = t_lf_n @ r_lf_n.T + cidx = torch.argmax(dist, dim=1) + + # ★ LFとHFのカンペを「同じインデックス」で引き当てる ★ + m_lf = r_lf_all[cidx].reshape(B, H, W, C).permute(0, 3, 1, 2) + m_hf = r_hf_all[cidx].reshape(B, H, W, C).permute(0, 3, 1, 2) + + # 各周波数領域での残差 + res_lf = t_lf - m_lf + res_hf = t_hf - m_hf + + # 固定比率でブレンド (元の次元数Cのまま) + rfeature = alpha * res_lf + (1.0 - alpha) * res_hf + + if pos_flag: + pos_embed = get_position_encoding(C, H, W).to(t_lf.device).unsqueeze(0).repeat(B, 1, 1, 1) + rfeature = rfeature + pos_embed + + rfeatures.append(rfeature) + return rfeatures + +def main(args): + CLASSES = SETTINGS[args.setting] + + if args.classes == 'capsules': + train_dataset1 = CAPSULES(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, normalize="w50", img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader1 = DataLoader(train_dataset1, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True) + train_dataset2 = CAPSULESANO(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader2 = DataLoader(train_dataset2, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True) + elif CLASSES['seen'][0] in MVTEC.CLASS_NAMES: + train_dataset1 = MVTEC(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, normalize="w50", img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader1 = DataLoader(train_dataset1, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True) + train_dataset2 = MVTECANO(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader2 = DataLoader(train_dataset2, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True) + else: + train_dataset1 = VISA(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, normalize="w50", img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader1 = DataLoader(train_dataset1, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True) + train_dataset2 = VISAANO(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, normalize="w50", img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader2 = DataLoader(train_dataset2, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True) + + if args.backbone == 'wide_resnet50_2': + encoder = timm.create_model('wide_resnet50_2', features_only=True, out_indices=(1, 2, 3), pretrained=True).eval().to(args.device) + feat_dims = encoder.feature_info.channels() + elif args.backbone == 'tf_efficientnet_b6': + encoder = timm.create_model('tf_efficientnet_b6', features_only=True, out_indices=(1, 2, 3), pretrained=True).eval().to(args.device) + feat_dims = encoder.feature_info.channels() + + boundary_ops = BoundaryAverager(num_levels=args.feature_levels) + wav_filter = HaarWaveletFilter2Component().to(args.device) + + constraintor = MultiScaleConv(feat_dims).to(args.device) + optimizer0 = torch.optim.Adam(constraintor.parameters(), lr=args.lr, weight_decay=0.0005) + scheduler0 = torch.optim.lr_scheduler.MultiStepLR(optimizer0, milestones=[70, 90], gamma=0.1) + + estimators = [load_flow_model(args, feat_dim).to(args.device) for feat_dim in feat_dims] + params_f = list(estimators[0].parameters()) + for l in range(1, args.feature_levels): + params_f += list(estimators[l].parameters()) + optimizer1 = torch.optim.Adam(params_f, lr=args.lr, weight_decay=0.0005) + scheduler1 = torch.optim.lr_scheduler.MultiStepLR(optimizer1, milestones=[70, 90], gamma=0.1) + + best_img_auc = 0 + N_batch = 8192 + + for epoch in range(args.epochs): + constraintor.train() + for estimator in estimators: estimator.train() + + train_loader = train_loader1 if epoch < FIRST_STAGE_EPOCH else train_loader2 + train_loss_total, total_num = 0, 0 + progress_bar = tqdm(total=len(train_loader)) + progress_bar.set_description(f"Epoch[{epoch}/{args.epochs}]") + + for step, batch in enumerate(train_loader): + progress_bar.update(1) + images, _, masks, class_names = batch[0], batch[1], batch[2], batch[3] + images, masks = images.to(args.device), masks.to(args.device) + + with torch.no_grad(): + features_raw = encoder(images) + + # テスト画像の分離 + test_lf_list, test_hf_list = [], [] + for l in range(args.feature_levels): + lf, hf = wav_filter.get_LF_HF(features_raw[l]) + test_lf_list.append(lf) + test_hf_list.append(hf) + + # カンペの分離 + ref_lf_dict, ref_hf_dict = get_freq_reference_features(encoder, wav_filter, args.train_dataset_dir, class_names, images.device, args.train_ref_shot) + + # 分割統治マッチングと固定比率ブレンド + rfeatures = [] + for i, class_name in enumerate(class_names): + t_lf = [feat[i:i+1] for feat in test_lf_list] + t_hf = [feat[i:i+1] for feat in test_hf_list] + r_lf = ref_lf_dict[class_name] + r_hf = ref_hf_dict[class_name] + + res = get_freq_matched_residuals(t_lf, t_hf, r_lf, r_hf, alpha=args.blend_alpha, pos_flag=True) + + if i == 0: + rfeatures = res + else: + rfeatures = [torch.cat([rfeatures[l], res[l]], dim=0) for l in range(args.feature_levels)] + + lvl_masks = [] + for l in range(args.feature_levels): + _, _, h, w = rfeatures[l].size() + m = F.interpolate(masks, size=(h, w), mode='nearest').squeeze(1) + lvl_masks.append(m) + rfeatures_t = [rfeature.detach().clone() for rfeature in rfeatures] + + rfeatures = constraintor(*rfeatures) + loss = 0 + for l in range(args.feature_levels): + e = rfeatures[l] + t = rfeatures_t[l] + bs, dim, h, w = e.size() + e = e.permute(0, 2, 3, 1).reshape(-1, dim) + t = t.permute(0, 2, 3, 1).reshape(-1, dim) + m = lvl_masks[l] + m = m.reshape(-1) + + loss_i, _, _ = calculate_log_barrier_bi_occ_loss(e, m, t) + loss += loss_i + optimizer0.zero_grad() + loss.backward() + optimizer0.step() + + train_loss_total += loss.item() + total_num += 1 + + rfeatures = [rfeature.detach().clone() for rfeature in rfeatures] + loss, num = train(args, rfeatures, estimators, optimizer1, masks, boundary_ops, epoch, N_batch=N_batch, FIRST_STAGE_EPOCH=FIRST_STAGE_EPOCH) + train_loss_total += loss + total_num += num + + scheduler0.step() + scheduler1.step() + + progress_bar.close() + print(f"Epoch[{epoch}/{args.epochs}]: train_loss: {train_loss_total / total_num}") + + if (epoch + 1) % args.eval_freq == 0: + s1_res, s2_res, s_res = [], [], [] + test_ref_features = load_mc_reference_features(args.test_ref_feature_dir, CLASSES['unseen'], args.device, args.num_ref_shot) + + for class_name in CLASSES['unseen']: + if args.classes == 'capsules': test_dataset = CAPSULES(args.test_dataset_dir, class_name=class_name, train=False, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in MVTEC.CLASS_NAMES: test_dataset = MVTEC(args.test_dataset_dir, class_name=class_name, train=False, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in VISA.CLASS_NAMES: test_dataset = VISA(args.test_dataset_dir, class_name=class_name, train=False, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in BTAD.CLASS_NAMES: test_dataset = BTAD(args.test_dataset_dir, class_name=class_name, train=False, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in MVTEC3D.CLASS_NAMES: test_dataset = MVTEC3D(args.test_dataset_dir, class_name=class_name, train=False, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in MPDD.CLASS_NAMES: test_dataset = MPDD(args.test_dataset_dir, class_name=class_name, train=False, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in MVTECLOCO.CLASS_NAMES: test_dataset = MVTECLOCO(args.test_dataset_dir, class_name=class_name, train=False, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in BRATS.CLASS_NAMES: test_dataset = BRATS(args.test_dataset_dir, class_name=class_name, train=False, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + else: raise ValueError('Unrecognized class name: {}'.format(class_name)) + + test_loader = DataLoader(test_dataset, batch_size=1, shuffle=False, num_workers=8, drop_last=False) + + metrics = validate(args, encoder, constraintor, wav_filter, estimators, test_loader, test_ref_features[class_name], args.device, class_name) + + img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro = metrics['scores'] + print("Epoch: {}, Class Name: {}, Image AUC | AP | F1_Score: {} | {} | {}, Pixel AUC | AP | F1_Score | AUPRO: {} | {} | {} | {}".format( + epoch, class_name, img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro)) + s1_res.append(metrics['scores1']) + s2_res.append(metrics['scores2']) + s_res.append(metrics['scores']) + + s1_res = np.array(s1_res) + s2_res = np.array(s2_res) + s_res = np.array(s_res) + img_auc1, img_ap1, img_f1_score1, pix_auc1, pix_ap1, pix_f1_score1, pix_aupro1 = np.mean(s1_res, axis=0) + img_auc2, img_ap2, img_f1_score2, pix_auc2, pix_ap2, pix_f1_score2, pix_aupro2 = np.mean(s2_res, axis=0) + img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro = np.mean(s_res, axis=0) + + print('(Merged) Average Image AUC | AP | F1_Score: {:.3f} | {:.3f} | {:.3f}, Average Pixel AUC | AP | F1_Score | AUPRO: {:.3f} | {:.3f} | {:.3f} | {:.3f}'.format( + img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro)) + + if img_auc > best_img_auc: + os.makedirs(args.checkpoint_path, exist_ok=True) + best_img_auc = img_auc + state_dict = {'constraintor': constraintor.state_dict(), + 'estimators': [estimator.state_dict() for estimator in estimators]} + torch.save(state_dict, os.path.join(args.checkpoint_path, f'{args.setting}_epoch_{epoch}_checkpoints.pth')) + +def load_mc_reference_features(root_dir: str, class_names, device: torch.device, num_shot=4): + refs = {} + for class_name in class_names: + layer1_refs = torch.from_numpy(np.load(os.path.join(root_dir, class_name, 'layer1.npy'))).to(device) + layer2_refs = torch.from_numpy(np.load(os.path.join(root_dir, class_name, 'layer2.npy'))).to(device) + layer3_refs = torch.from_numpy(np.load(os.path.join(root_dir, class_name, 'layer3.npy'))).to(device) + K1 = (layer1_refs.shape[0] // TOTAL_SHOT) * num_shot + K2 = (layer2_refs.shape[0] // TOTAL_SHOT) * num_shot + K3 = (layer3_refs.shape[0] // TOTAL_SHOT) * num_shot + refs[class_name] = (layer1_refs[:K1, :], layer2_refs[:K2, :], layer3_refs[:K3, :]) + return refs + +if __name__ == "__main__": + parser = argparse.ArgumentParser() + parser.add_argument('--setting', type=str, default="visa_to_mvtec") + parser.add_argument('--classes', type=str, default="none") + parser.add_argument('--train_dataset_dir', type=str, default="") + parser.add_argument('--test_dataset_dir', type=str, default="") + parser.add_argument('--test_ref_feature_dir', type=str, default="./ref_features/w50/mvtec_4shot") + parser.add_argument('--bgadweight_dir', type=str, default="none") + parser.add_argument('--batch_size', type=int, default=32) + parser.add_argument('--lr', type=float, default=1e-5) + parser.add_argument('--epochs', type=int, default=100) + parser.add_argument('--device', type=str, default="cuda:0") + parser.add_argument('--checkpoint_path', type=str, default="./checkpoints/") + parser.add_argument('--eval_freq', type=int, default=1) + parser.add_argument('--backbone', type=str, default="wide_resnet50_2") + parser.add_argument('--rank', type=int, default="0") + + # ★ LFとHFのブレンド比率 (1.0 = LFのみ, 0.5 = 半々, 0.0 = HFのみ) + parser.add_argument('--blend_alpha', type=float, default=0.5) + + parser.add_argument('--flow_arch', type=str, default='conditional_flow_model') + parser.add_argument('--feature_levels', default=3, type=int) + parser.add_argument('--coupling_layers', type=int, default=10) + parser.add_argument('--clamp_alpha', type=float, default=1.9) + parser.add_argument('--pos_embed_dim', type=int, default=256) + parser.add_argument('--pos_beta', type=float, default=0.05) + parser.add_argument('--margin_tau', type=float, default=0.1) + parser.add_argument('--bgspp_lambda', type=float, default=1) + parser.add_argument('--fdm_alpha', type=float, default=0.4) + parser.add_argument("--train_ref_shot", type=int, default=4) + parser.add_argument("--num_ref_shot", type=int, default=4) + + args = parser.parse_args() + init_seeds(42) + main(args) diff --git a/main_global.py b/main_global.py new file mode 100644 index 0000000..5fe0e9e --- /dev/null +++ b/main_global.py @@ -0,0 +1,320 @@ +import os +import warnings +import argparse +from tqdm import tqdm +import numpy as np +import torch +import torch.nn as nn +import timm +import torch.nn.functional as F +from torch.utils.data import DataLoader + +from train import train +from validate_global import validate +from datasets.mvtec import MVTEC, MVTECANO +from datasets.visa import VISA, VISAANO +from datasets.btad import BTAD +from datasets.mvtec_3d import MVTEC3D +from datasets.mpdd import MPDD +from datasets.mvtec_loco import MVTECLOCO +from datasets.brats import BRATS +from datasets.capsules import CAPSULES, CAPSULESANO + +from models.fc_flow import load_flow_model +from utils import init_seeds, get_residual_features, get_mc_matched_ref_features, get_mc_reference_features +from utils import BoundaryAverager +from losses.loss import calculate_log_barrier_bi_occ_loss +from classes import VISA_TO_MVTEC, MVTEC_TO_VISA, MVTEC_TO_BTAD, MVTEC_TO_MVTEC3D +from classes import MVTEC_TO_MPDD, MVTEC_TO_MVTECLOCO, MVTEC_TO_BRATS +from classes import MVTEC_TO_MVTEC, VISA_TO_VISA +from classes import CAPSULES_TO_CAPSULES + +warnings.filterwarnings('ignore') + +TOTAL_SHOT = 4 +FIRST_STAGE_EPOCH = 10 +SETTINGS = {'visa_to_mvtec': VISA_TO_MVTEC, 'mvtec_to_visa': MVTEC_TO_VISA, + 'mvtec_to_btad': MVTEC_TO_BTAD, 'mvtec_to_mvtec3d': MVTEC_TO_MVTEC3D, + 'mvtec_to_mpdd': MVTEC_TO_MPDD, 'mvtec_to_mvtecloco': MVTEC_TO_MVTECLOCO, + 'mvtec_to_brats': MVTEC_TO_BRATS,'mvtec_to_mvtec':MVTEC_TO_MVTEC, 'visa_to_visa':VISA_TO_VISA, 'capsules_to_capsules': CAPSULES_TO_CAPSULES} + +class GlobalContextConstraintor(nn.Module): + def __init__(self, feat_dims, num_heads=4, num_layers=1): + super().__init__() + self.num_levels = len(feat_dims) + self.local_convs = nn.ModuleList() + self.transformers = nn.ModuleList() + + self.apply_transformer = [l > 0 for l in range(self.num_levels)] + + for i, dim in enumerate(feat_dims): + conv = nn.Sequential( + nn.Conv2d(dim, dim, kernel_size=3, padding=1, bias=False), + nn.BatchNorm2d(dim), + nn.ReLU(inplace=True), + nn.Conv2d(dim, dim, kernel_size=3, padding=1, bias=False), + nn.BatchNorm2d(dim), + nn.ReLU(inplace=True) + ) + self.local_convs.append(conv) + + if self.apply_transformer[i]: + heads = num_heads if dim % num_heads == 0 else 1 + encoder_layer = nn.TransformerEncoderLayer( + d_model=dim, + nhead=heads, + dim_feedforward=dim * 2, + activation='relu', + batch_first=True, + dropout=0.1 + ) + self.transformers.append(nn.TransformerEncoder(encoder_layer, num_layers=num_layers)) + else: + self.transformers.append(nn.Identity()) + + def forward(self, *features): + out_features = [] + for i in range(self.num_levels): + x = features[i] + B, C, H, W = x.shape + + x_local = self.local_convs[i](x) + x + + if self.apply_transformer[i]: + pos_embed = self.get_2d_sincos_pos_embed(C, H, W, x.device) + pos_embed = pos_embed.unsqueeze(0).expand(B, -1, -1) + + x_flat = x_local.view(B, C, -1).permute(0, 2, 1) + x_flat = x_flat + pos_embed + + x_global = self.transformers[i](x_flat) + + x_out = x_global.permute(0, 2, 1).view(B, C, H, W) + out_features.append(x_out + x_local) + else: + out_features.append(x_local) + + return tuple(out_features) + + def get_2d_sincos_pos_embed(self, embed_dim, grid_size_h, grid_size_w, device): + grid_h = torch.arange(grid_size_h, dtype=torch.float32, device=device) + grid_w = torch.arange(grid_size_w, dtype=torch.float32, device=device) + grid_h, grid_w = torch.meshgrid(grid_h, grid_w, indexing='ij') + + grid_h = grid_h.reshape(-1) + grid_w = grid_w.reshape(-1) + + emb_h = self.get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid_h) + emb_w = self.get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid_w) + + pos_embed = torch.cat([emb_h, emb_w], dim=1) + if embed_dim % 2 != 0: + pos_embed = F.pad(pos_embed, (0, 1)) + return pos_embed + + def get_1d_sincos_pos_embed_from_grid(self, embed_dim, pos): + omega = torch.arange(embed_dim // 2, dtype=torch.float32, device=pos.device) + omega /= (embed_dim / 2.) + omega = 1. / (10000 ** omega) + + out = torch.einsum('m,d->md', pos, omega) + emb_sin = torch.sin(out) + emb_cos = torch.cos(out) + + emb = torch.cat([emb_sin, emb_cos], dim=1) + if embed_dim % 2 != 0: + emb = F.pad(emb, (0, 1)) + return emb + +def main(args): + CLASSES = SETTINGS[args.setting] + + if args.classes == 'capsules': + train_dataset1 = CAPSULES(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, normalize="w50", img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader1 = DataLoader(train_dataset1, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True) + train_dataset2 = CAPSULESANO(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader2 = DataLoader(train_dataset2, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True) + elif CLASSES['seen'][0] in MVTEC.CLASS_NAMES: + train_dataset1 = MVTEC(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, normalize="w50", img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader1 = DataLoader(train_dataset1, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True) + train_dataset2 = MVTECANO(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader2 = DataLoader(train_dataset2, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True) + else: + train_dataset1 = VISA(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, normalize="w50", img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader1 = DataLoader(train_dataset1, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True) + train_dataset2 = VISAANO(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, normalize="w50", img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader2 = DataLoader(train_dataset2, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True) + + if args.backbone == 'wide_resnet50_2': + encoder = timm.create_model('wide_resnet50_2', features_only=True, out_indices=(1, 2, 3), pretrained=True).eval().to(args.device) + feat_dims = encoder.feature_info.channels() + elif args.backbone == 'tf_efficientnet_b6': + encoder = timm.create_model('tf_efficientnet_b6', features_only=True, out_indices=(1, 2, 3), pretrained=True).eval().to(args.device) + feat_dims = encoder.feature_info.channels() + + boundary_ops = BoundaryAverager(num_levels=args.feature_levels) + + constraintor = GlobalContextConstraintor(feat_dims).to(args.device) + + optimizer0 = torch.optim.Adam(constraintor.parameters(), lr=args.lr, weight_decay=0.0005) + scheduler0 = torch.optim.lr_scheduler.MultiStepLR(optimizer0, milestones=[70, 90], gamma=0.1) + + estimators = [load_flow_model(args, feat_dim).to(args.device) for feat_dim in feat_dims] + params_f = list(estimators[0].parameters()) + for l in range(1, args.feature_levels): + params_f += list(estimators[l].parameters()) + optimizer1 = torch.optim.Adam(params_f, lr=args.lr, weight_decay=0.0005) + scheduler1 = torch.optim.lr_scheduler.MultiStepLR(optimizer1, milestones=[70, 90], gamma=0.1) + + best_img_auc = 0 + N_batch = 8192 + + for epoch in range(args.epochs): + constraintor.train() + for estimator in estimators: estimator.train() + + train_loader = train_loader1 if epoch < FIRST_STAGE_EPOCH else train_loader2 + train_loss_total, total_num = 0, 0 + progress_bar = tqdm(total=len(train_loader)) + progress_bar.set_description(f"Epoch[{epoch}/{args.epochs}]") + + for step, batch in enumerate(train_loader): + progress_bar.update(1) + images, _, masks, class_names = batch[0], batch[1], batch[2], batch[3] + images, masks = images.to(args.device), masks.to(args.device) + + with torch.no_grad(): + features = encoder(images) + + ref_features = get_mc_reference_features(encoder, args.train_dataset_dir, class_names, images.device, args.train_ref_shot) + + mfeatures = get_mc_matched_ref_features(features, class_names, ref_features) + + rfeatures = get_residual_features(features, mfeatures, pos_flag=True) + + lvl_masks = [] + for l in range(args.feature_levels): + _, _, h, w = rfeatures[l].size() + m = F.interpolate(masks, size=(h, w), mode='nearest').squeeze(1) + lvl_masks.append(m) + rfeatures_t = [rfeature.detach().clone() for rfeature in rfeatures] + + rfeatures = constraintor(*rfeatures) + loss = 0 + for l in range(args.feature_levels): + e = rfeatures[l] + t = rfeatures_t[l] + bs, dim, h, w = e.size() + e = e.permute(0, 2, 3, 1).reshape(-1, dim) + t = t.permute(0, 2, 3, 1).reshape(-1, dim) + m = lvl_masks[l] + m = m.reshape(-1) + + loss_i, _, _ = calculate_log_barrier_bi_occ_loss(e, m, t) + loss += loss_i + optimizer0.zero_grad() + loss.backward() + optimizer0.step() + + train_loss_total += loss.item() + total_num += 1 + + rfeatures = [rfeature.detach().clone() for rfeature in rfeatures] + loss, num = train(args, rfeatures, estimators, optimizer1, masks, boundary_ops, epoch, N_batch=N_batch, FIRST_STAGE_EPOCH=FIRST_STAGE_EPOCH) + train_loss_total += loss + total_num += num + + scheduler0.step() + scheduler1.step() + + progress_bar.close() + print(f"Epoch[{epoch}/{args.epochs}]: train_loss: {train_loss_total / total_num}") + + if (epoch + 1) % args.eval_freq == 0: + s1_res, s2_res, s_res = [], [], [] + test_ref_features = load_mc_reference_features(args.test_ref_feature_dir, CLASSES['unseen'], args.device, args.num_ref_shot) + + for class_name in CLASSES['unseen']: + if args.classes == 'capsules': test_dataset = CAPSULES(args.test_dataset_dir, class_name=class_name, train=False, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in MVTEC.CLASS_NAMES: test_dataset = MVTEC(args.test_dataset_dir, class_name=class_name, train=False, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in VISA.CLASS_NAMES: test_dataset = VISA(args.test_dataset_dir, class_name=class_name, train=False, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in BTAD.CLASS_NAMES: test_dataset = BTAD(args.test_dataset_dir, class_name=class_name, train=False, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in MVTEC3D.CLASS_NAMES: test_dataset = MVTEC3D(args.test_dataset_dir, class_name=class_name, train=False, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in MPDD.CLASS_NAMES: test_dataset = MPDD(args.test_dataset_dir, class_name=class_name, train=False, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in MVTECLOCO.CLASS_NAMES: test_dataset = MVTECLOCO(args.test_dataset_dir, class_name=class_name, train=False, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in BRATS.CLASS_NAMES: test_dataset = BRATS(args.test_dataset_dir, class_name=class_name, train=False, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + else: raise ValueError('Unrecognized class name: {}'.format(class_name)) + + test_loader = DataLoader(test_dataset, batch_size=1, shuffle=False, num_workers=8, drop_last=False) + + metrics = validate(args, encoder, constraintor, estimators, test_loader, test_ref_features[class_name], args.device, class_name) + + img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro = metrics['scores'] + print("Epoch: {}, Class Name: {}, Image AUC | AP | F1_Score: {} | {} | {}, Pixel AUC | AP | F1_Score | AUPRO: {} | {} | {} | {}".format( + epoch, class_name, img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro)) + s1_res.append(metrics['scores1']) + s2_res.append(metrics['scores2']) + s_res.append(metrics['scores']) + + s1_res = np.array(s1_res) + s2_res = np.array(s2_res) + s_res = np.array(s_res) + img_auc1, img_ap1, img_f1_score1, pix_auc1, pix_ap1, pix_f1_score1, pix_aupro1 = np.mean(s1_res, axis=0) + img_auc2, img_ap2, img_f1_score2, pix_auc2, pix_ap2, pix_f1_score2, pix_aupro2 = np.mean(s2_res, axis=0) + img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro = np.mean(s_res, axis=0) + + print('(Merged) Average Image AUC | AP | F1_Score: {:.3f} | {:.3f} | {:.3f}, Average Pixel AUC | AP | F1_Score | AUPRO: {:.3f} | {:.3f} | {:.3f} | {:.3f}'.format( + img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro)) + + if img_auc > best_img_auc: + os.makedirs(args.checkpoint_path, exist_ok=True) + best_img_auc = img_auc + state_dict = {'constraintor': constraintor.state_dict(), + 'estimators': [estimator.state_dict() for estimator in estimators]} + torch.save(state_dict, os.path.join(args.checkpoint_path, f'{args.setting}_epoch_{epoch}_checkpoints.pth')) + +def load_mc_reference_features(root_dir: str, class_names, device: torch.device, num_shot=4): + refs = {} + for class_name in class_names: + layer1_refs = torch.from_numpy(np.load(os.path.join(root_dir, class_name, 'layer1.npy'))).to(device) + layer2_refs = torch.from_numpy(np.load(os.path.join(root_dir, class_name, 'layer2.npy'))).to(device) + layer3_refs = torch.from_numpy(np.load(os.path.join(root_dir, class_name, 'layer3.npy'))).to(device) + K1 = (layer1_refs.shape[0] // TOTAL_SHOT) * num_shot + K2 = (layer2_refs.shape[0] // TOTAL_SHOT) * num_shot + K3 = (layer3_refs.shape[0] // TOTAL_SHOT) * num_shot + refs[class_name] = (layer1_refs[:K1, :], layer2_refs[:K2, :], layer3_refs[:K3, :]) + return refs + +if __name__ == "__main__": + parser = argparse.ArgumentParser() + parser.add_argument('--setting', type=str, default="visa_to_mvtec") + parser.add_argument('--classes', type=str, default="none") + parser.add_argument('--train_dataset_dir', type=str, default="") + parser.add_argument('--test_dataset_dir', type=str, default="") + parser.add_argument('--test_ref_feature_dir', type=str, default="./ref_features/w50/mvtec_4shot") + parser.add_argument('--bgadweight_dir', type=str, default="none") + parser.add_argument('--batch_size', type=int, default=32) + parser.add_argument('--lr', type=float, default=1e-5) + parser.add_argument('--epochs', type=int, default=100) + parser.add_argument('--device', type=str, default="cuda:0") + parser.add_argument('--checkpoint_path', type=str, default="./checkpoints/") + parser.add_argument('--eval_freq', type=int, default=1) + parser.add_argument('--backbone', type=str, default="wide_resnet50_2") + parser.add_argument('--rank', type=int, default="0") + + parser.add_argument('--flow_arch', type=str, default='conditional_flow_model') + parser.add_argument('--feature_levels', default=3, type=int) + parser.add_argument('--coupling_layers', type=int, default=10) + parser.add_argument('--clamp_alpha', type=float, default=1.9) + parser.add_argument('--pos_embed_dim', type=int, default=256) + parser.add_argument('--pos_beta', type=float, default=0.05) + parser.add_argument('--margin_tau', type=float, default=0.1) + parser.add_argument('--bgspp_lambda', type=float, default=1) + parser.add_argument('--fdm_alpha', type=float, default=0.4) + parser.add_argument("--train_ref_shot", type=int, default=4) + parser.add_argument("--num_ref_shot", type=int, default=4) + + args = parser.parse_args() + init_seeds(42) + main(args) + diff --git a/main_ib.py b/main_ib.py index 3a07eed..c808333 100644 --- a/main_ib.py +++ b/main_ib.py @@ -25,10 +25,13 @@ from models.modules import MultiScaleOrthogonalProjector from models.vq import MultiScaleVQ4 from utils import init_seeds, get_residual_features, get_mc_matched_ref_features, get_random_normal_images +from utils import load_weights_ada from utils import BoundaryAverager from losses.loss import calculate_log_barrier_bi_occ_loss, calculate_orthogonal_regularizer from classes import VISA_TO_MVTEC, MVTEC_TO_VISA, MVTEC_TO_BTAD, MVTEC_TO_MVTEC3D -from classes import MVTEC_TO_MPDD, MVTEC_TO_MVTECLOCO, MVTEC_TO_BRATS +from classes import MVTEC_TO_MPDD, MVTEC_TO_MVTECLOCO, MVTEC_TO_BRATS, MVTEC_TO_MVTEC,MVTECFEW_TO_MVTEC +# visualizerのインポート +from visualizer import Visualizer, denormalization warnings.filterwarnings('ignore') @@ -37,7 +40,7 @@ SETTINGS = {'visa_to_mvtec': VISA_TO_MVTEC, 'mvtec_to_visa': MVTEC_TO_VISA, 'mvtec_to_btad': MVTEC_TO_BTAD, 'mvtec_to_mvtec3d': MVTEC_TO_MVTEC3D, 'mvtec_to_mpdd': MVTEC_TO_MPDD, 'mvtec_to_mvtecloco': MVTEC_TO_MVTECLOCO, - 'mvtec_to_brats': MVTEC_TO_BRATS} + 'mvtec_to_brats': MVTEC_TO_BRATS, 'mvtec_to_mvtec': MVTEC_TO_MVTEC,'mvtecfew_to_mvtec': MVTECFEW_TO_MVTEC} def main(args): @@ -45,28 +48,41 @@ def main(args): CLASSES = SETTINGS[args.setting] else: raise ValueError(f"Dataset setting must be in {SETTINGS.keys()}, but got {args.setting}.") - if CLASSES['seen'][0] in MVTEC.CLASS_NAMES: # from mvtec to other datasets + if args.train_dataset == 'mvtec_few': + train_dataset1 = MVTECFEW(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, + normalize="w50", + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader1 = DataLoader( + train_dataset1, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True + ) + train_dataset2 = MVTECFEWANO(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader2 = DataLoader( + train_dataset2, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True + ) + elif CLASSES['seen'][0] in MVTEC.CLASS_NAMES: # from mvtec to other datasets train_dataset1 = MVTEC(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, - normalize="imagebind", - img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + normalize="w50", + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) train_loader1 = DataLoader( train_dataset1, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True ) train_dataset2 = MVTECANO(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, - normalize='imagebind', - img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) train_loader2 = DataLoader( train_dataset2, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True ) else: # from visa to mvtec train_dataset1 = VISA(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, - normalize="imagebind", + normalize="w50", img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) train_loader1 = DataLoader( train_dataset1, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True ) train_dataset2 = VISAANO(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, - normalize="imagebind", + normalize="w50", img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) train_loader2 = DataLoader( train_dataset2, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True @@ -96,7 +112,17 @@ def main(args): scheduler1 = torch.optim.lr_scheduler.MultiStepLR(optimizer1, milestones=[70, 90], gamma=0.1) best_pro = 0 + best_img_auc = 0 N_batch = 16 * 16 * 32 + # 可視化オブジェクトの初期化 + # 可視化結果を保存するディレクトリを指定 + visualization_output_dir = os.path.join(args.checkpoint_path, 'visualizations') + os.makedirs(visualization_output_dir, exist_ok=True) + my_visualizer = Visualizer(root=visualization_output_dir) # + # 最良モデルのエポックで保存するためのデータ保持用 + best_epoch_class_data = {} + + for epoch in range(args.epochs): vq_ops.train() constraintor.train() @@ -116,7 +142,7 @@ def main(args): progress_bar.set_description(f"Epoch[{epoch}/{args.epochs}]") for step, batch in enumerate(train_loader): progress_bar.update(1) - images, _, masks, class_names = batch + images, _, masks, class_names,anomaly_types = batch images = images.to(args.device) masks = masks.to(args.device) @@ -130,6 +156,10 @@ def main(args): ref_features = get_mc_reference_features(encoder, args.train_dataset_dir, class_names, images.device, args.train_ref_shot) mfeatures = get_mc_matched_ref_features(features, class_names, ref_features) rfeatures = get_residual_features(features, mfeatures) + if args.residual=='False': + rfeatures = features + else: + rfeatures = rfeatures lvl_masks = [] for l in range(args.feature_levels): @@ -187,56 +217,94 @@ def main(args): scheduler1.step() progress_bar.close() - print(f"Epoch[{epoch}/{args.epochs}]: VQ loss: {train_loss_total_vq / total_num_vq}, OCC loss: {train_loss_total_occ / total_num_occ} (n: {train_loss_total_occn / total_num_occn}, a: {train_loss_total_occa / total_num_occa}), " \ - f"Ort loss: {train_loss_total_ort / total_num_ort}, " \ - f"Flow loss: {train_loss_total_flow / total_num_flow}") + vq_loss_avg = train_loss_total_vq / total_num_vq if total_num_vq > 0 else 0 + occ_loss_avg = train_loss_total_occ / total_num_occ if total_num_occ > 0 else 0 + occn_loss_avg = train_loss_total_occn / total_num_occn if total_num_occn > 0 else 0 + occa_loss_avg = train_loss_total_occa / total_num_occa if total_num_occa > 0 else 0 + ort_loss_avg = train_loss_total_ort / total_num_ort if total_num_ort > 0 else 0 + flow_loss_avg = train_loss_total_flow / total_num_flow if total_num_flow > 0 else 0 + + print(f"Epoch[{epoch}/{args.epochs}]: VQ loss: {vq_loss_avg}, OCC loss: {occ_loss_avg} (n: {occn_loss_avg}, a: {occa_loss_avg}), " \ + f"Ort loss: {ort_loss_avg}, " \ + f"Flow loss: {flow_loss_avg}") + #print(f"Epoch[{epoch}/{args.epochs}]: VQ loss: {train_loss_total_vq / total_num_vq}, OCC loss: {train_loss_total_occ / total_num_occ} (n: {train_loss_total_occn / total_num_occn}, a: {train_loss_total_occa / total_num_occa}), " \ + #f"Ort loss: {train_loss_total_ort / total_num_ort}, " \ + #f"Flow loss: {train_loss_total_flow / total_num_flow}") if (epoch + 1) % args.eval_freq == 0: s1_res, s2_res, s_res = [], [], [] test_proto_features = load_mc_reference_features(args.test_ref_feature_dir, CLASSES['unseen'], args.device, args.num_ref_shot) - for class_name in CLASSES['unseen']: - if class_name in MVTEC.CLASS_NAMES: - test_dataset = MVTEC(args.test_dataset_dir, class_name=class_name, train=False, + # 各クラスの評価結果とデータを一時的に保持する辞書 + current_epoch_class_data_for_saving = {} + + for class_name_eval in CLASSES['unseen']: + if class_name_eval in MVTEC.CLASS_NAMES: + test_dataset = MVTEC(args.test_dataset_dir, class_name=class_name_eval, train=False, normalize='imagebind', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) - elif class_name in VISA.CLASS_NAMES: - test_dataset = VISA(args.test_dataset_dir, class_name=class_name, train=False, + elif class_name_eval in VISA.CLASS_NAMES: + test_dataset = VISA(args.test_dataset_dir, class_name=class_name_eval, train=False, normalize='imagebind', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) - elif class_name in BTAD.CLASS_NAMES: - test_dataset = BTAD(args.test_dataset_dir, class_name=class_name, train=False, + elif class_name_eval in BTAD.CLASS_NAMES: + test_dataset = BTAD(args.test_dataset_dir, class_name=class_name_eval, train=False, normalize='imagebind', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) - elif class_name in MVTEC3D.CLASS_NAMES: - test_dataset = MVTEC3D(args.test_dataset_dir, class_name=class_name, train=False, + elif class_name_eval in MVTEC3D.CLASS_NAMES: + test_dataset = MVTEC3D(args.test_dataset_dir, class_name=class_name_eval, train=False, normalize='imagebind', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) - elif class_name in MPDD.CLASS_NAMES: - test_dataset = MPDD(args.test_dataset_dir, class_name=class_name, train=False, + elif class_name_eval in MPDD.CLASS_NAMES: + test_dataset = MPDD(args.test_dataset_dir, class_name=class_name_eval, train=False, normalize='imagebind', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) - elif class_name in MVTECLOCO.CLASS_NAMES: - test_dataset = MVTECLOCO(args.test_dataset_dir, class_name=class_name, train=False, + elif class_name_eval in MVTECLOCO.CLASS_NAMES: + test_dataset = MVTECLOCO(args.test_dataset_dir, class_name=class_name_eval, train=False, normalize='imagebind', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) - elif class_name in BRATS.CLASS_NAMES: - test_dataset = BRATS(args.test_dataset_dir, class_name=class_name, train=False, + elif class_name_eval in BRATS.CLASS_NAMES: + test_dataset = BRATS(args.test_dataset_dir, class_name=class_name_eval, train=False, normalize='imagebind', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) else: - raise ValueError('Unrecognized class name: {}'.format(class_name)) + raise ValueError('Unrecognized class name: {}'.format(class_name_eval)) test_loader = DataLoader( test_dataset, batch_size=1, shuffle=False, num_workers=8, drop_last=False ) - metrics = validate(args, encoder, vq_ops, constraintor, estimators, test_loader, test_proto_features[class_name], args.device, class_name) + metrics = validate(args, encoder, vq_ops, constraintor, estimators, test_loader, + test_proto_features[class_name_eval], args.device, class_name_eval) + #metrics = validate(args, encoder, vq_ops, constraintor, estimators, test_loader, test_proto_features[class_name], args.device, class_name) img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro = metrics['scores'] print("Epoch: {}, Class Name: {}, Image AUC | AP | F1_Score: {} | {} | {}, Pixel AUC | AP | F1_Score | AUPRO: {} | {} | {} | {}".format( - epoch, class_name, img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro)) + epoch, class_name_eval, img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro)) s1_res.append(metrics['scores1']) s2_res.append(metrics['scores2']) s_res.append(metrics['scores']) - + # 可視化結果を保存するのは最終エポックのみ + if epoch == args.epochs - 1: # 最終エポックの場合のみ可視化を保存 + # `validate` 関数が返す `metrics` から必要なデータを取得 + scores = metrics['scores_map'] + gts_masks = metrics['gt_masks_raw'] + images_raw = metrics['images_raw'] + + # Visualizerを使ってプロット + # クラスごとにサブディレクトリを作成 + output_class_dir = os.path.join(visualization_output_dir, class_name_eval, f'final_epoch') # ディレクトリ名を'final_epoch'に固定 + os.makedirs(output_class_dir, exist_ok=True) + my_visualizer.set_prefix(f'{class_name_eval}_final_epoch') # プレフィックスをクラス名と最終エポックに設定 + my_visualizer.root = output_class_dir # 保存先ディレクトリを更新 + + my_visualizer.plot(images_raw, scores, gts_masks) + print(f" - クラス '{class_name_eval}': 最終エポックの可視化結果を {output_class_dir} に保存しました。") + # --- 変更点ここまで --- + + # 各クラスの評価結果から特徴量とラベルデータを一時的に保存 + current_epoch_class_data_for_saving[class_name_eval] = { + 'features': metrics['features'], + 'anomaly_types': metrics['anomaly_types'], + 'gts_labels': metrics['gts_labels'] + } s1_res = np.array(s1_res) s2_res = np.array(s2_res) s_res = np.array(s_res) @@ -250,13 +318,43 @@ def main(args): print('(Merged) Average Image AUC | AP | F1_Score: {:.3f} | {:.3f} | {:.3f}, Average Pixel AUC | AP | F1_Score | AUPRO: {:.3f} | {:.3f} | {:.3f} | {:.3f}'.format( img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro)) - if pix_aupro > best_pro: + if img_auc > best_img_auc: #pix_aupro > best_pro: os.makedirs(args.checkpoint_path, exist_ok=True) - best_pro = pix_aupro + best_img_auc = img_auc #best_pro = pix_aupro state_dict = {'vq_ops': vq_ops.state_dict(), 'constraintor': constraintor.state_dict(), 'estimators': [estimator.state_dict() for estimator in estimators]} torch.save(state_dict, os.path.join(args.checkpoint_path, f'{args.setting}_checkpoints.pth')) + # 新しい特徴量保存ディレクトリを作成 + features_save_dir = os.path.join(args.checkpoint_path, 'features_for_analysis') + os.makedirs(features_save_dir, exist_ok=True) + # -- ここから変更 -- + # 各クラスについて、最良スコア時のデータを保存 + for class_name_to_save, data in current_epoch_class_data_for_saving.items(): + # クラス名の下にファイルを保存するパスを構築 + class_specific_save_dir = os.path.join(features_save_dir, class_name_to_save) + os.makedirs(class_specific_save_dir, exist_ok=True) + + # ファイル名から epoch 番号を削除 + # これにより、常に同じファイル名で上書き保存され、 + # 最終的にベストスコア時のデータだけが残る + features_filename = os.path.join(class_specific_save_dir, 'best_features.npy') + anomaly_types_filename = os.path.join(class_specific_save_dir, 'best_anomaly_types.npy') + gts_labels_filename = os.path.join(class_specific_save_dir, 'best_gts_labels.npy') + + np.save(features_filename, data['features']) + + # anomaly_types をテキストファイルとして保存 + with open(anomaly_types_filename, 'w') as f: + for item in data['anomaly_types']: + f.write(str(item) + '\n') # 各要素を1行ずつ書き込む + + np.save(gts_labels_filename, data['gts_labels']) + print(f" - クラス '{class_name_to_save}': 最良スコア時のデータを {class_specific_save_dir} に上書き保存しました。") + # -- 変更ここまで -- + + # 最良エポックのデータなので、今後の可視化のためにこれを覚えておく + best_epoch_class_data = current_epoch_class_data_for_saving.copy() def load_mc_reference_features(root_dir: str, class_names, device: torch.device, num_shot=4): @@ -322,6 +420,7 @@ def load_and_transform_vision_data(image_paths, device): if __name__ == "__main__": parser = argparse.ArgumentParser() + parser.add_argument('--train_dataset', type=str, default='mvtec') parser.add_argument('--setting', type=str, default="visa_to_mvtec") parser.add_argument('--train_dataset_dir', type=str, default="") parser.add_argument('--test_dataset_dir', type=str, default="") @@ -334,6 +433,7 @@ def load_and_transform_vision_data(image_paths, device): parser.add_argument('--checkpoint_path', type=str, default="./checkpoints/") parser.add_argument('--eval_freq', type=int, default=1) parser.add_argument('--backbone', type=str, default="imagebind") + parser.add_argument('--residual', type = True, default=True) # flow parameters parser.add_argument('--flow_arch', type=str, default='flow_model') @@ -362,4 +462,4 @@ def load_and_transform_vision_data(image_paths, device): - \ No newline at end of file + diff --git a/main_osp.py b/main_osp.py new file mode 100644 index 0000000..bc38aa7 --- /dev/null +++ b/main_osp.py @@ -0,0 +1,339 @@ +import os +import warnings +import argparse +from tqdm import tqdm +import numpy as np +import torch +import timm +import torch.nn.functional as F +from torch.utils.data import DataLoader + +from validate_osp import validate # validate_osp をインポート +from datasets.mvtec import MVTEC, MVTECANO +from datasets.visa import VISA, VISAANO +from datasets.btad import BTAD +from datasets.mvtec_3d import MVTEC3D +from datasets.mpdd import MPDD +from datasets.mvtec_loco import MVTECLOCO +from datasets.brats import BRATS +from datasets.capsules import CAPSULES, CAPSULESANO + +from models.fc_flow import load_flow_model +from models.modules import MultiScaleConv +from utils import init_seeds, get_residual_features, get_mc_matched_ref_features, get_mc_reference_features +from utils import BoundaryAverager +from losses.loss import calculate_log_barrier_bi_occ_loss +from classes import VISA_TO_MVTEC, MVTEC_TO_VISA, MVTEC_TO_BTAD, MVTEC_TO_MVTEC3D +from classes import MVTEC_TO_MPDD, MVTEC_TO_MVTECLOCO, MVTEC_TO_BRATS +from classes import MVTEC_TO_MVTEC, VISA_TO_VISA +from classes import CAPSULES_TO_CAPSULES + +warnings.filterwarnings('ignore') + +TOTAL_SHOT = 4 +FIRST_STAGE_EPOCH = 10 +SETTINGS = {'visa_to_mvtec': VISA_TO_MVTEC, 'mvtec_to_visa': MVTEC_TO_VISA, + 'mvtec_to_btad': MVTEC_TO_BTAD, 'mvtec_to_mvtec3d': MVTEC_TO_MVTEC3D, + 'mvtec_to_mpdd': MVTEC_TO_MPDD, 'mvtec_to_mvtecloco': MVTEC_TO_MVTECLOCO, + 'mvtec_to_brats': MVTEC_TO_BRATS,'mvtec_to_mvtec':MVTEC_TO_MVTEC, 'visa_to_visa':VISA_TO_VISA, 'capsules_to_capsules': CAPSULES_TO_CAPSULES} + +# ========================================== +# OSP (Orthogonal Subspace Projection) 用関数 +# ========================================== +def compute_osp_matrices(ref_features_tuple, keep_variance=0.95): + proj_matrices, means = [], [] + for ref_feat in ref_features_tuple: + if ref_feat.dim() == 4: + B, C, H, W = ref_feat.shape + flat_ref = ref_feat.permute(0, 2, 3, 1).reshape(-1, C) + else: + flat_ref = ref_feat + C = flat_ref.shape[-1] + mean = flat_ref.mean(dim=0, keepdim=True) + centered = flat_ref - mean + U, S, V = torch.linalg.svd(centered.cpu(), full_matrices=False) + var = (S ** 2) / (centered.size(0) - 1) + cum_var = torch.cumsum(var, dim=0) / var.sum() + k = torch.searchsorted(cum_var, keep_variance).item() + 1 + basis = V[:k, :].T.to(ref_feat.device) + proj_matrix = torch.mm(basis, basis.T) + proj_matrices.append(proj_matrix) + means.append(mean.to(ref_feat.device)) + return proj_matrices, means + +def apply_osp(residuals_list, proj_matrices, means): + osp_results = [] + for i, res in enumerate(residuals_list): + B, C, H, W = res.shape + res_flat = res.permute(0, 2, 3, 1).reshape(-1, C) + res_centered = res_flat - means[i] + res_parallel = torch.mm(res_centered, proj_matrices[i]) + res_ortho = res_centered - res_parallel + osp_results.append(res_ortho.reshape(B, H, W, C).permute(0, 3, 1, 2)) + return osp_results +# ========================================== + + +def main(args): + if args.setting in SETTINGS.keys(): + CLASSES = SETTINGS[args.setting] + else: + raise ValueError(f"Dataset setting must be in {SETTINGS.keys()}, but got {args.setting}.") + + if args.classes == 'capsules': + train_dataset1 = CAPSULES(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, + normalize="w50", img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader1 = DataLoader(train_dataset1, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True) + train_dataset2 = CAPSULESANO(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, + normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader2 = DataLoader(train_dataset2, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True) + + elif CLASSES['seen'][0] in MVTEC.CLASS_NAMES: + train_dataset1 = MVTEC(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, + normalize="w50", img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader1 = DataLoader(train_dataset1, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True) + train_dataset2 = MVTECANO(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, + normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader2 = DataLoader(train_dataset2, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True) + else: + train_dataset1 = VISA(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, + normalize="w50", img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader1 = DataLoader(train_dataset1, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True) + train_dataset2 = VISAANO(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, + normalize="w50", img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader2 = DataLoader(train_dataset2, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True) + + if args.backbone == 'wide_resnet50_2': + encoder = timm.create_model('wide_resnet50_2', features_only=True, out_indices=(1, 2, 3), pretrained=True).eval() + encoder = encoder.to(args.device) + feat_dims = encoder.feature_info.channels() + elif args.backbone == 'tf_efficientnet_b6': + encoder = timm.create_model('tf_efficientnet_b6', features_only=True, out_indices=(1, 2, 3), pretrained=True).eval() + encoder = encoder.to(args.device) + feat_dims = encoder.feature_info.channels() + + boundary_ops = BoundaryAverager(num_levels=args.feature_levels) + + # constraintor のみ保持 (vq_ops は削除) + constraintor = MultiScaleConv(feat_dims).to(args.device) + optimizer0 = torch.optim.Adam(constraintor.parameters(), lr=args.lr, weight_decay=0.0005) + scheduler0 = torch.optim.lr_scheduler.MultiStepLR(optimizer0, milestones=[70, 90], gamma=0.1) + + # Normflow decoder + estimators = [load_flow_model(args, feat_dim) for feat_dim in feat_dims] + estimators = [decoder.to(args.device) for decoder in estimators] + params = list(estimators[0].parameters()) + for l in range(1, args.feature_levels): + params += list(estimators[l].parameters()) + optimizer1 = torch.optim.Adam(params, lr=args.lr, weight_decay=0.0005) + scheduler1 = torch.optim.lr_scheduler.MultiStepLR(optimizer1, milestones=[70, 90], gamma=0.1) + + # SVD行列をキャッシュする辞書 + osp_cache = {} + + from train import train + + best_img_auc = 0 + N_batch = 8192 + for epoch in range(args.epochs): + constraintor.train() + for estimator in estimators: + estimator.train() + + if epoch < FIRST_STAGE_EPOCH: + train_loader = train_loader1 + else: + train_loader = train_loader2 + + train_loss_total, total_num = 0, 0 + progress_bar = tqdm(total=len(train_loader)) + progress_bar.set_description(f"Epoch[{epoch}/{args.epochs}]") + + for step, batch in enumerate(train_loader): + progress_bar.update(1) + images, _, masks, class_names = batch + + images = images.to(args.device) + masks = masks.to(args.device) + + with torch.no_grad(): + features = encoder(images) + + ref_features = get_mc_reference_features(encoder, args.train_dataset_dir, class_names, images.device, args.train_ref_shot) + mfeatures = get_mc_matched_ref_features(features, class_names, ref_features) + rfeatures = get_residual_features(features, mfeatures, pos_flag=True) + + # --- OSPの適用 (クラスごとに計算・キャッシュしてバッチ内の該当スライスに適用) --- + rfeatures_osp = [torch.zeros_like(rf) for rf in rfeatures] + for i, c_name in enumerate(class_names): + if c_name not in osp_cache: + c_refs = ref_features[c_name] + osp_cache[c_name] = compute_osp_matrices(c_refs, keep_variance=0.95) # 0.95で95%のズレを吸収 + + proj_matrices, means = osp_cache[c_name] + single_rfeat = [rf[i:i+1] for rf in rfeatures] + single_rfeat_osp = apply_osp(single_rfeat, proj_matrices, means) + + for l in range(args.feature_levels): + rfeatures_osp[l][i] = single_rfeat_osp[l][0] + + rfeatures = rfeatures_osp + # -------------------------------------------------------------------------- + + lvl_masks = [] + for l in range(args.feature_levels): + _, _, h, w = rfeatures[l].size() + m = F.interpolate(masks, size=(h, w), mode='nearest').squeeze(1) + lvl_masks.append(m) + rfeatures_t = [rfeature.detach().clone() for rfeature in rfeatures] + + # constraintor適用 + rfeatures = constraintor(*rfeatures) + # 学習時のみ、特徴量に微小なノイズを付加して過学習を防ぐ + noise_std = 0.01 # ノイズの強さ(0.005 〜 0.05あたりで調整) + rfeatures_noisy = [rf + torch.randn_like(rf) * noise_std for rf in rfeatures] + loss = 0 + for l in range(args.feature_levels): + e = rfeatures[l] + t = rfeatures_t[l] + bs, dim, h, w = e.size() + e = e.permute(0, 2, 3, 1).reshape(-1, dim) + t = t.permute(0, 2, 3, 1).reshape(-1, dim) + m = lvl_masks[l] + m = m.reshape(-1) + + loss_i, _, _ = calculate_log_barrier_bi_occ_loss(e, m, t) + loss += loss_i + + optimizer0.zero_grad() + loss.backward() + optimizer0.step() + + train_loss_total += loss.item() + total_num += 1 + + rfeatures = [rfeature.detach().clone() for rfeature in rfeatures] + loss, num = train(args, rfeatures, estimators, optimizer1, masks, boundary_ops, epoch, N_batch=N_batch, FIRST_STAGE_EPOCH=FIRST_STAGE_EPOCH) + train_loss_total += loss + total_num += num + + scheduler0.step() + scheduler1.step() + + progress_bar.close() + print(f"Epoch[{epoch}/{args.epochs}]: train_loss: {train_loss_total / total_num}") + + if (epoch + 1) % args.eval_freq == 0: + s1_res, s2_res, s_res = [], [], [] + test_ref_features = load_mc_reference_features(args.test_ref_feature_dir, CLASSES['unseen'], args.device, args.num_ref_shot) + + for class_name in CLASSES['unseen']: + if class_name not in osp_cache: # 評価時のOSP行列をキャッシュ + osp_cache[class_name] = compute_osp_matrices(test_ref_features[class_name], keep_variance=0.95) + proj_matrices, means = osp_cache[class_name] + + if args.classes == 'capsules': + test_dataset = CAPSULES(args.test_dataset_dir, class_name=class_name, train=False, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in MVTEC.CLASS_NAMES: + test_dataset = MVTEC(args.test_dataset_dir, class_name=class_name, train=False, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in VISA.CLASS_NAMES: + test_dataset = VISA(args.test_dataset_dir, class_name=class_name, train=False, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in BTAD.CLASS_NAMES: + test_dataset = BTAD(args.test_dataset_dir, class_name=class_name, train=False, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in MVTEC3D.CLASS_NAMES: + test_dataset = MVTEC3D(args.test_dataset_dir, class_name=class_name, train=False, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in MPDD.CLASS_NAMES: + test_dataset = MPDD(args.test_dataset_dir, class_name=class_name, train=False, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in MVTECLOCO.CLASS_NAMES: + test_dataset = MVTECLOCO(args.test_dataset_dir, class_name=class_name, train=False, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in BRATS.CLASS_NAMES: + test_dataset = BRATS(args.test_dataset_dir, class_name=class_name, train=False, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + else: + raise ValueError('Unrecognized class name: {}'.format(class_name)) + + test_loader = DataLoader(test_dataset, batch_size=1, shuffle=False, num_workers=8, drop_last=False) + + # validate_osp の validate 関数を呼び出し (proj_matrices, means を渡す) + metrics = validate(args, encoder, constraintor, estimators, test_loader, test_ref_features[class_name], args.device, class_name, proj_matrices, means) + + img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro = metrics['scores'] + + print("Epoch: {}, Class Name: {}, Image AUC | AP | F1_Score: {} | {} | {}, Pixel AUC | AP | F1_Score | AUPRO: {} | {} | {} | {}".format( + epoch, class_name, img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro)) + s1_res.append(metrics['scores1']) + s2_res.append(metrics['scores2']) + s_res.append(metrics['scores']) + + s1_res = np.array(s1_res) + s2_res = np.array(s2_res) + s_res = np.array(s_res) + img_auc1, img_ap1, img_f1_score1, pix_auc1, pix_ap1, pix_f1_score1, pix_aupro1 = np.mean(s1_res, axis=0) + img_auc2, img_ap2, img_f1_score2, pix_auc2, pix_ap2, pix_f1_score2, pix_aupro2 = np.mean(s2_res, axis=0) + img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro = np.mean(s_res, axis=0) + print('(Logps) Average Image AUC | AP | F1_Score: {:.3f} | {:.3f} | {:.3f}, Average Pixel AUC | AP | F1_Score | AUPRO: {:.3f} | {:.3f} | {:.3f} | {:.3f}'.format( + img_auc1, img_ap1, img_f1_score1, pix_auc1, pix_ap1, pix_f1_score1, pix_aupro1)) + print('(BScores) Average Image AUC | AP | F1_Score: {:.3f} | {:.3f} | {:.3f}, Average Pixel AUC | AP | F1_Score | AUPRO: {:.3f} | {:.3f} | {:.3f} | {:.3f}'.format( + img_auc2, img_ap2, img_f1_score2, pix_auc2, pix_ap2, pix_f1_score2, pix_aupro2)) + print('(Merged) Average Image AUC | AP | F1_Score: {:.3f} | {:.3f} | {:.3f}, Average Pixel AUC | AP | F1_Score | AUPRO: {:.3f} | {:.3f} | {:.3f} | {:.3f}'.format( + img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro)) + if img_auc > best_img_auc: + os.makedirs(args.checkpoint_path, exist_ok=True) + best_img_auc = img_auc + state_dict = {'constraintor': constraintor.state_dict(), + 'estimators': [estimator.state_dict() for estimator in estimators]} + torch.save(state_dict, os.path.join(args.checkpoint_path, f'{args.setting}_epoch_{epoch}_checkpoints.pth')) + +def load_mc_reference_features(root_dir: str, class_names, device: torch.device, num_shot=4): + refs = {} + for class_name in class_names: + layer1_refs = np.load(os.path.join(root_dir, class_name, 'layer1.npy')) + layer2_refs = np.load(os.path.join(root_dir, class_name, 'layer2.npy')) + layer3_refs = np.load(os.path.join(root_dir, class_name, 'layer3.npy')) + + layer1_refs = torch.from_numpy(layer1_refs).to(device) + layer2_refs = torch.from_numpy(layer2_refs).to(device) + layer3_refs = torch.from_numpy(layer3_refs).to(device) + + K1 = (layer1_refs.shape[0] // TOTAL_SHOT) * num_shot + layer1_refs = layer1_refs[:K1, :] + K2 = (layer2_refs.shape[0] // TOTAL_SHOT) * num_shot + layer2_refs = layer2_refs[:K2, :] + K3 = (layer3_refs.shape[0] // TOTAL_SHOT) * num_shot + layer3_refs = layer3_refs[:K3, :] + + refs[class_name] = (layer1_refs, layer2_refs, layer3_refs) + return refs + +if __name__ == "__main__": + parser = argparse.ArgumentParser() + parser.add_argument('--setting', type=str, default="visa_to_mvtec") + parser.add_argument('--classes', type=str, default="none") + parser.add_argument('--train_dataset_dir', type=str, default="") + parser.add_argument('--test_dataset_dir', type=str, default="") + parser.add_argument('--test_ref_feature_dir', type=str, default="./ref_features/w50/mvtec_4shot") + parser.add_argument('--bgadweight_dir', type=str, default="none") + parser.add_argument('--batch_size', type=int, default=32) + parser.add_argument('--lr', type=float, default=1e-5) + parser.add_argument('--epochs', type=int, default=100) + parser.add_argument('--device', type=str, default="cuda:0") + parser.add_argument('--checkpoint_path', type=str, default="./checkpoints/") + parser.add_argument('--eval_freq', type=int, default=1) + parser.add_argument('--backbone', type=str, default="wide_resnet50_2") + parser.add_argument('--rank', type=int, default="0") + + # flow parameters + parser.add_argument('--flow_arch', type=str, default='conditional_flow_model') + parser.add_argument('--feature_levels', default=3, type=int) + parser.add_argument('--coupling_layers', type=int, default=10) + parser.add_argument('--clamp_alpha', type=float, default=1.9) + parser.add_argument('--pos_embed_dim', type=int, default=256) + parser.add_argument('--pos_beta', type=float, default=0.05) + parser.add_argument('--margin_tau', type=float, default=0.1) + parser.add_argument('--bgspp_lambda', type=float, default=1) + + parser.add_argument("--train_ref_shot", type=int, default=4) + parser.add_argument("--num_ref_shot", type=int, default=4) + + args = parser.parse_args() + init_seeds(42) + main(args) diff --git a/main_vit.py b/main_vit.py new file mode 100644 index 0000000..73c6c04 --- /dev/null +++ b/main_vit.py @@ -0,0 +1,361 @@ +import os +import warnings +import argparse +from tqdm import tqdm +import numpy as np +import torch +import timm +import torch.nn.functional as F +from torch.utils.data import DataLoader + +from train import train +from validate import validate +from datasets.mvtec import MVTEC, MVTECANO +from datasets.visa import VISA, VISAANO +from datasets.btad import BTAD +from datasets.mvtec_3d import MVTEC3D +from datasets.mpdd import MPDD +from datasets.mvtec_loco import MVTECLOCO +from datasets.brats import BRATS +from datasets.capsules import CAPSULES, CAPSULESANO + +from models.fc_flow import load_flow_model +from models.modules import MultiScaleConv +from models.vq import MultiScaleVQ +from utils import init_seeds, get_residual_features, get_mc_matched_ref_features, get_mc_reference_features +from utils import BoundaryAverager +from losses.loss import calculate_log_barrier_bi_occ_loss +from classes import VISA_TO_MVTEC, MVTEC_TO_VISA, MVTEC_TO_BTAD, MVTEC_TO_MVTEC3D +from classes import MVTEC_TO_MPDD, MVTEC_TO_MVTECLOCO, MVTEC_TO_BRATS +from classes import MVTEC_TO_MVTEC, VISA_TO_VISA +from classes import CAPSULES_TO_CAPSULES +import torch.nn as nn + +warnings.filterwarnings('ignore') +# --- import文の下、定数定義(TOTAL_SHOTなど)の上あたりに追加 --- +class ViTFeatureExtractor(nn.Module): + def __init__(self, model_name='deit_base_patch16_224', out_indices=(3, 7, 11)): + super().__init__() + self.vit = timm.create_model(model_name, pretrained=True) + self.out_indices = out_indices + self.patch_size = self.vit.patch_embed.patch_size[0] + self.embed_dim = self.vit.embed_dim + + def forward(self, x): + B, C, H, W = x.shape + h_out, w_out = H // self.patch_size, W // self.patch_size + + x = self.vit.patch_embed(x) + x = self.vit._pos_embed(x) + x = self.vit.norm_pre(x) + + features = [] + for i, blk in enumerate(self.vit.blocks): + x = blk(x) + if i in self.out_indices: + num_prefix_tokens = self.vit.num_prefix_tokens + tokens = x[:, num_prefix_tokens:] + feat = tokens.transpose(1, 2).reshape(B, self.embed_dim, h_out, w_out) + features.append(feat) + + return features +# ------------------------------------------------------------- +TOTAL_SHOT = 4 # total few-shot reference samples +FIRST_STAGE_EPOCH = 10 +SETTINGS = {'visa_to_mvtec': VISA_TO_MVTEC, 'mvtec_to_visa': MVTEC_TO_VISA, + 'mvtec_to_btad': MVTEC_TO_BTAD, 'mvtec_to_mvtec3d': MVTEC_TO_MVTEC3D, + 'mvtec_to_mpdd': MVTEC_TO_MPDD, 'mvtec_to_mvtecloco': MVTEC_TO_MVTECLOCO, + 'mvtec_to_brats': MVTEC_TO_BRATS,'mvtec_to_mvtec':MVTEC_TO_MVTEC, 'visa_to_visa':VISA_TO_VISA, 'capsules_to_capsules': CAPSULES_TO_CAPSULES} + + +def main(args): + if args.setting in SETTINGS.keys(): + CLASSES = SETTINGS[args.setting] + else: + raise ValueError(f"Dataset setting must be in {SETTINGS.keys()}, but got {args.setting}.") + + if args.classes == 'capsules': # from mvtec to other datasets # from mvtec to other datasets + train_dataset1 = CAPSULES(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, + normalize="w50", + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader1 = DataLoader( + train_dataset1, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True + ) + train_dataset2 = CAPSULESANO(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader2 = DataLoader( + train_dataset2, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True + ) + + elif CLASSES['seen'][0] in MVTEC.CLASS_NAMES: # from mvtec to other datasets + train_dataset1 = MVTEC(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, + normalize="w50", + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader1 = DataLoader( + train_dataset1, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True + ) + train_dataset2 = MVTECANO(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader2 = DataLoader( + train_dataset2, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True + ) + else: # from visa to mvtec + train_dataset1 = VISA(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, + normalize="w50", + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader1 = DataLoader( + train_dataset1, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True + ) + train_dataset2 = VISAANO(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, + normalize="w50", + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader2 = DataLoader( + train_dataset2, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True + ) + if args.backbone == 'wide_resnet50_2': + encoder = timm.create_model('wide_resnet50_2', features_only=True, + out_indices=(1, 2, 3), pretrained=True).eval() # the pretrained checkpoint will be in /home/.cache/torch/hub/checkpoints/ + encoder = encoder.to(args.device) + feat_dims = encoder.feature_info.channels() + elif args.backbone == 'tf_efficientnet_b6':#10/26追加 + encoder = timm.create_model('tf_efficientnet_b6', features_only=True, + out_indices=(1, 2, 3), pretrained=True).eval() # the pretrained checkpoint will be in /home/.cache/torch/hub/checkpoints/ + encoder = encoder.to(args.device) + feat_dims = encoder.feature_info.channels() + elif args.backbone == 'vit_base_patch14':#4/15追加 + #encoder = ViTFeatureExtractor(model_name='vit_base_patch14_reg4_dinov2.lvd142m', out_indices=(3, 7, 11)).eval() + encoder = ViTFeatureExtractor(model_name='vit_base_patch16_224_dino', out_indices=(3, 7, 11)).eval() + encoder = encoder.to(args.device) + feat_dims = [encoder.embed_dim] * len(encoder.out_indices) + + boundary_ops = BoundaryAverager(num_levels=args.feature_levels) + vq_ops = MultiScaleVQ(num_embeddings=args.num_embeddings, channels=feat_dims).to(args.device) + optimizer_vq = torch.optim.Adam(vq_ops.parameters(), lr=args.lr, weight_decay=0.0005) + scheduler_vq = torch.optim.lr_scheduler.MultiStepLR(optimizer_vq, milestones=[70, 90], gamma=0.1) + + constraintor = MultiScaleConv(feat_dims).to(args.device) + # weight_decay is the l2 weight penalty lambda, weight_decay = lambda / 2 + optimizer0 = torch.optim.Adam(constraintor.parameters(), lr=args.lr, weight_decay=0.0005) + scheduler0 = torch.optim.lr_scheduler.MultiStepLR(optimizer0, milestones=[70, 90], gamma=0.1) + + # Normflow decoder + estimators = [load_flow_model(args, feat_dim) for feat_dim in feat_dims] + estimators = [decoder.to(args.device) for decoder in estimators] + params = list(estimators[0].parameters()) + for l in range(1, args.feature_levels): + params += list(estimators[l].parameters()) + optimizer1 = torch.optim.Adam(params, lr=args.lr, weight_decay=0.0005) + scheduler1 = torch.optim.lr_scheduler.MultiStepLR(optimizer1, milestones=[70, 90], gamma=0.1) + + best_img_auc = 0 + N_batch = 8192 + for epoch in range(args.epochs): + vq_ops.train() + constraintor.train() + for estimator in estimators: + estimator.train() + if epoch < FIRST_STAGE_EPOCH: + train_loader = train_loader1 + else: + train_loader = train_loader2 + train_loss_total, total_num = 0, 0 + progress_bar = tqdm(total=len(train_loader)) + progress_bar.set_description(f"Epoch[{epoch}/{args.epochs}]") + for step, batch in enumerate(train_loader): + progress_bar.update(1) + #images, _, masks, class_names,_ = batch + images, _, masks, class_names = batch[0], batch[1], batch[2], batch[3] + + images = images.to(args.device) + masks = masks.to(args.device) + + with torch.no_grad(): + features = encoder(images) + + ref_features = get_mc_reference_features(encoder, args.train_dataset_dir, class_names, images.device, args.train_ref_shot) + mfeatures = get_mc_matched_ref_features(features, class_names, ref_features) + rfeatures = get_residual_features(features, mfeatures, pos_flag=True) + + lvl_masks = [] + for l in range(args.feature_levels): + _, _, h, w = rfeatures[l].size() + m = F.interpolate(masks, size=(h, w), mode='nearest').squeeze(1) + lvl_masks.append(m) + rfeatures_t = [rfeature.detach().clone() for rfeature in rfeatures] + + loss_vq = vq_ops(rfeatures, lvl_masks, train=True) + train_loss_total += loss_vq.item() + total_num += 1 + optimizer_vq.zero_grad() + loss_vq.backward() + optimizer_vq.step() + + rfeatures = constraintor(*rfeatures) + loss = 0 + for l in range(args.feature_levels): # backward svdd loss + e = rfeatures[l] + t = rfeatures_t[l] + bs, dim, h, w = e.size() + e = e.permute(0, 2, 3, 1).reshape(-1, dim) + t = t.permute(0, 2, 3, 1).reshape(-1, dim) + m = lvl_masks[l] + m = m.reshape(-1) + + loss_i, _, _ = calculate_log_barrier_bi_occ_loss(e, m, t) + loss += loss_i + optimizer0.zero_grad() + loss.backward() + optimizer0.step() + + train_loss_total += loss.item() + total_num += 1 + + # detach the rfeatures for flow optimization + rfeatures = [rfeature.detach().clone() for rfeature in rfeatures] + # train flow corresponding to with neck + loss, num = train(args, rfeatures, estimators, optimizer1, masks, boundary_ops, epoch, N_batch=N_batch, FIRST_STAGE_EPOCH=FIRST_STAGE_EPOCH) + train_loss_total += loss + total_num += num + + scheduler_vq.step() + scheduler0.step() + scheduler1.step() + + progress_bar.close() + print(f"Epoch[{epoch}/{args.epochs}]: train_loss: {train_loss_total / total_num}") + + if (epoch + 1) % args.eval_freq == 0: + s1_res, s2_res, s_res = [], [], [] + test_ref_features = load_mc_reference_features(args.test_ref_feature_dir, CLASSES['unseen'], args.device, args.num_ref_shot) + for class_name in CLASSES['unseen']: + if args.classes == 'capsules': + test_dataset = CAPSULES(args.test_dataset_dir, class_name=class_name, train=False, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in MVTEC.CLASS_NAMES: + test_dataset = MVTEC(args.test_dataset_dir, class_name=class_name, train=False, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in VISA.CLASS_NAMES: + test_dataset = VISA(args.test_dataset_dir, class_name=class_name, train=False, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in BTAD.CLASS_NAMES: + test_dataset = BTAD(args.test_dataset_dir, class_name=class_name, train=False, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in MVTEC3D.CLASS_NAMES: + test_dataset = MVTEC3D(args.test_dataset_dir, class_name=class_name, train=False, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in MPDD.CLASS_NAMES: + test_dataset = MPDD(args.test_dataset_dir, class_name=class_name, train=False, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in MVTECLOCO.CLASS_NAMES: + test_dataset = MVTECLOCO(args.test_dataset_dir, class_name=class_name, train=False, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in BRATS.CLASS_NAMES: + test_dataset = BRATS(args.test_dataset_dir, class_name=class_name, train=False, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + else: + raise ValueError('Unrecognized class name: {}'.format(class_name)) + test_loader = DataLoader( + test_dataset, batch_size=1, shuffle=False, num_workers=8, drop_last=False + ) + metrics = validate(args, encoder, vq_ops, constraintor, estimators, test_loader, test_ref_features[class_name], args.device, class_name) + img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro = metrics['scores'] + + print("Epoch: {}, Class Name: {}, Image AUC | AP | F1_Score: {} | {} | {}, Pixel AUC | AP | F1_Score | AUPRO: {} | {} | {} | {}".format( + epoch, class_name, img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro)) + s1_res.append(metrics['scores1']) + s2_res.append(metrics['scores2']) + s_res.append(metrics['scores']) + + s1_res = np.array(s1_res) + s2_res = np.array(s2_res) + s_res = np.array(s_res) + img_auc1, img_ap1, img_f1_score1, pix_auc1, pix_ap1, pix_f1_score1, pix_aupro1 = np.mean(s1_res, axis=0) + img_auc2, img_ap2, img_f1_score2, pix_auc2, pix_ap2, pix_f1_score2, pix_aupro2 = np.mean(s2_res, axis=0) + img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro = np.mean(s_res, axis=0) + print('(Logps) Average Image AUC | AP | F1_Score: {:.3f} | {:.3f} | {:.3f}, Average Pixel AUC | AP | F1_Score | AUPRO: {:.3f} | {:.3f} | {:.3f} | {:.3f}'.format( + img_auc1, img_ap1, img_f1_score1, pix_auc1, pix_ap1, pix_f1_score1, pix_aupro1)) + print('(BScores) Average Image AUC | AP | F1_Score: {:.3f} | {:.3f} | {:.3f}, Average Pixel AUC | AP | F1_Score | AUPRO: {:.3f} | {:.3f} | {:.3f} | {:.3f}'.format( + img_auc2, img_ap2, img_f1_score2, pix_auc2, pix_ap2, pix_f1_score2, pix_aupro2)) + print('(Merged) Average Image AUC | AP | F1_Score: {:.3f} | {:.3f} | {:.3f}, Average Pixel AUC | AP | F1_Score | AUPRO: {:.3f} | {:.3f} | {:.3f} | {:.3f}'.format( + img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro)) + + if img_auc > best_img_auc: + os.makedirs(args.checkpoint_path, exist_ok=True) + best_img_auc = img_auc + state_dict = {'vq_ops': vq_ops.state_dict(), + 'constraintor': constraintor.state_dict(), + 'estimators': [estimator.state_dict() for estimator in estimators]} + torch.save(state_dict, os.path.join(args.checkpoint_path, f'{args.setting}_epoch_{epoch}_checkpoints.pth')) + #torch.save(state_dict, os.path.join(args.checkpoint_path, f'{args.setting}_checkpoints.pth')) + + +def load_mc_reference_features(root_dir: str, class_names, device: torch.device, num_shot=4): + refs = {} + for class_name in class_names: + layer1_refs = np.load(os.path.join(root_dir, class_name, 'layer1.npy')) + layer2_refs = np.load(os.path.join(root_dir, class_name, 'layer2.npy')) + layer3_refs = np.load(os.path.join(root_dir, class_name, 'layer3.npy')) + + layer1_refs = torch.from_numpy(layer1_refs).to(device) + layer2_refs = torch.from_numpy(layer2_refs).to(device) + layer3_refs = torch.from_numpy(layer3_refs).to(device) + + K1 = (layer1_refs.shape[0] // TOTAL_SHOT) * num_shot + layer1_refs = layer1_refs[:K1, :] + K2 = (layer2_refs.shape[0] // TOTAL_SHOT) * num_shot + layer2_refs = layer2_refs[:K2, :] + K3 = (layer3_refs.shape[0] // TOTAL_SHOT) * num_shot + layer3_refs = layer3_refs[:K3, :] + + refs[class_name] = (layer1_refs, layer2_refs, layer3_refs) + + return refs + + +if __name__ == "__main__": + parser = argparse.ArgumentParser() + parser.add_argument('--setting', type=str, default="visa_to_mvtec") + parser.add_argument('--classes', type=str, default="none") + parser.add_argument('--train_dataset_dir', type=str, default="") + parser.add_argument('--test_dataset_dir', type=str, default="") + parser.add_argument('--test_ref_feature_dir', type=str, default="./ref_features/w50/mvtec_4shot") + parser.add_argument('--bgadweight_dir', type=str, default="none")# 12/16追加 + parser.add_argument('--batch_size', type=int, default=32) + parser.add_argument('--lr', type=float, default=1e-5) + parser.add_argument('--epochs', type=int, default=100) + parser.add_argument('--device', type=str, default="cuda:0") + parser.add_argument('--checkpoint_path', type=str, default="./checkpoints/") + parser.add_argument('--eval_freq', type=int, default=1) + parser.add_argument('--backbone', type=str, default="wide_resnet50_2") + parser.add_argument('--rank', type=int, default="0") + + # flow parameters + parser.add_argument('--flow_arch', type=str, default='conditional_flow_model') + parser.add_argument('--feature_levels', default=3, type=int) + parser.add_argument('--coupling_layers', type=int, default=10) + parser.add_argument('--clamp_alpha', type=float, default=1.9) + parser.add_argument('--pos_embed_dim', type=int, default=256) + parser.add_argument('--pos_beta', type=float, default=0.05) + parser.add_argument('--margin_tau', type=float, default=0.1) + parser.add_argument('--bgspp_lambda', type=float, default=1) + + parser.add_argument('--fdm_alpha', type=float, default=0.4) # low value, more training distribution + parser.add_argument('--num_embeddings', type=int, default=1536) # VQ embeddings + parser.add_argument("--train_ref_shot", type=int, default=4) + parser.add_argument("--num_ref_shot", type=int, default=4) + + args = parser.parse_args() + init_seeds(42) + + main(args) + diff --git a/main_wav.py b/main_wav.py new file mode 100644 index 0000000..6836f1d --- /dev/null +++ b/main_wav.py @@ -0,0 +1,344 @@ +import os +import warnings +import argparse +from tqdm import tqdm +import numpy as np +import torch +import torch.nn as nn +import timm +import torch.nn.functional as F +from torch.utils.data import DataLoader + +from train import train +from validate_wav import validate +from datasets.mvtec import MVTEC, MVTECANO +from datasets.visa import VISA, VISAANO +from datasets.btad import BTAD +from datasets.mvtec_3d import MVTEC3D +from datasets.mpdd import MPDD +from datasets.mvtec_loco import MVTECLOCO +from datasets.brats import BRATS +from datasets.capsules import CAPSULES, CAPSULESANO + +from models.fc_flow import load_flow_model +from models.modules import MultiScaleConv +from utils import init_seeds, get_residual_features, get_mc_matched_ref_features, get_mc_reference_features_wav +from utils import BoundaryAverager +from losses.loss import calculate_log_barrier_bi_occ_loss +from classes import VISA_TO_MVTEC, MVTEC_TO_VISA, MVTEC_TO_BTAD, MVTEC_TO_MVTEC3D +from classes import MVTEC_TO_MPDD, MVTEC_TO_MVTECLOCO, MVTEC_TO_BRATS +from classes import MVTEC_TO_MVTEC, VISA_TO_VISA +from classes import CAPSULES_TO_CAPSULES + +warnings.filterwarnings('ignore') + +TOTAL_SHOT = 4 +FIRST_STAGE_EPOCH = 10 +SETTINGS = {'visa_to_mvtec': VISA_TO_MVTEC, 'mvtec_to_visa': MVTEC_TO_VISA, + 'mvtec_to_btad': MVTEC_TO_BTAD, 'mvtec_to_mvtec3d': MVTEC_TO_MVTEC3D, + 'mvtec_to_mpdd': MVTEC_TO_MPDD, 'mvtec_to_mvtecloco': MVTEC_TO_MVTECLOCO, + 'mvtec_to_brats': MVTEC_TO_BRATS,'mvtec_to_mvtec':MVTEC_TO_MVTEC, 'visa_to_visa':VISA_TO_VISA, 'capsules_to_capsules': CAPSULES_TO_CAPSULES} + +# ========================================== +# Haar Wavelet Filter +# ========================================== +class HaarWaveletFilter(nn.Module): + def __init__(self, low_freq_weight=0.1, high_freq_weight=1.2): + super().__init__() + self.lf_w = low_freq_weight + self.hf_w = high_freq_weight + + ll = torch.tensor([[0.5, 0.5], [0.5, 0.5]]) + hl = torch.tensor([[-0.5, -0.5], [0.5, 0.5]]) + lh = torch.tensor([[-0.5, 0.5], [-0.5, 0.5]]) + hh = torch.tensor([[0.5, -0.5], [-0.5, 0.5]]) + + self.register_buffer('k_ll', ll.view(1, 1, 2, 2)) + self.register_buffer('k_hl', hl.view(1, 1, 2, 2)) + self.register_buffer('k_lh', lh.view(1, 1, 2, 2)) + self.register_buffer('k_hh', hh.view(1, 1, 2, 2)) + + def forward(self, x): + B, C, H, W = x.shape + ll = F.conv2d(x, self.k_ll.expand(C, 1, 2, 2), stride=2, groups=C) + hl = F.conv2d(x, self.k_hl.expand(C, 1, 2, 2), stride=2, groups=C) + lh = F.conv2d(x, self.k_lh.expand(C, 1, 2, 2), stride=2, groups=C) + hh = F.conv2d(x, self.k_hh.expand(C, 1, 2, 2), stride=2, groups=C) + + ll = ll * self.lf_w + hl = hl * self.hf_w + lh = lh * self.hf_w + hh = hh * self.hf_w + + out = F.conv_transpose2d(ll, self.k_ll.expand(C, 1, 2, 2), stride=2, groups=C) + \ + F.conv_transpose2d(hl, self.k_hl.expand(C, 1, 2, 2), stride=2, groups=C) + \ + F.conv_transpose2d(lh, self.k_lh.expand(C, 1, 2, 2), stride=2, groups=C) + \ + F.conv_transpose2d(hh, self.k_hh.expand(C, 1, 2, 2), stride=2, groups=C) + return out + + +def main(args): + if args.setting in SETTINGS.keys(): + CLASSES = SETTINGS[args.setting] + else: + raise ValueError(f"Dataset setting must be in {SETTINGS.keys()}, but got {args.setting}.") + + if args.classes == 'capsules': + train_dataset1 = CAPSULES(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, + normalize="w50", + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader1 = DataLoader( + train_dataset1, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True + ) + train_dataset2 = CAPSULESANO(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader2 = DataLoader( + train_dataset2, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True + ) + elif CLASSES['seen'][0] in MVTEC.CLASS_NAMES: + train_dataset1 = MVTEC(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, + normalize="w50", + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader1 = DataLoader( + train_dataset1, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True + ) + train_dataset2 = MVTECANO(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, + normalize='w50', + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader2 = DataLoader( + train_dataset2, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True + ) + else: + train_dataset1 = VISA(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, + normalize="w50", + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader1 = DataLoader( + train_dataset1, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True + ) + train_dataset2 = VISAANO(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, + normalize="w50", + img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader2 = DataLoader( + train_dataset2, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True + ) + + if args.backbone == 'wide_resnet50_2': + encoder = timm.create_model('wide_resnet50_2', features_only=True, + out_indices=(1, 2, 3), pretrained=True).eval() + encoder = encoder.to(args.device) + feat_dims = encoder.feature_info.channels() + elif args.backbone == 'tf_efficientnet_b6': + encoder = timm.create_model('tf_efficientnet_b6', features_only=True, + out_indices=(1, 2, 3), pretrained=True).eval() + encoder = encoder.to(args.device) + feat_dims = encoder.feature_info.channels() + + boundary_ops = BoundaryAverager(num_levels=args.feature_levels) + + # ウェーブレットフィルタの初期化 + wav_filter = HaarWaveletFilter(low_freq_weight=args.lf_weight, high_freq_weight=args.hf_weight).to(args.device) + wav_filter.eval() + + # constraintorの初期化 (元のmain.pyに準拠) + constraintor = MultiScaleConv(feat_dims).to(args.device) + optimizer0 = torch.optim.Adam(constraintor.parameters(), lr=args.lr, weight_decay=0.0005) + scheduler0 = torch.optim.lr_scheduler.MultiStepLR(optimizer0, milestones=[70, 90], gamma=0.1) + + # NFの初期化 + estimators = [load_flow_model(args, feat_dim) for feat_dim in feat_dims] + estimators = [decoder.to(args.device) for decoder in estimators] + params = list(estimators[0].parameters()) + for l in range(1, args.feature_levels): + params += list(estimators[l].parameters()) + optimizer1 = torch.optim.Adam(params, lr=args.lr, weight_decay=0.0005) + scheduler1 = torch.optim.lr_scheduler.MultiStepLR(optimizer1, milestones=[70, 90], gamma=0.1) + + best_img_auc = 0 + N_batch = 8192 + + for epoch in range(args.epochs): + constraintor.train() + for estimator in estimators: + estimator.train() + + if epoch < FIRST_STAGE_EPOCH: + train_loader = train_loader1 + else: + train_loader = train_loader2 + + train_loss_total, total_num = 0, 0 + progress_bar = tqdm(total=len(train_loader)) + progress_bar.set_description(f"Epoch[{epoch}/{args.epochs}]") + + for step, batch in enumerate(train_loader): + progress_bar.update(1) + images, _, masks, class_names = batch + + images = images.to(args.device) + masks = masks.to(args.device) + + with torch.no_grad(): + features = encoder(images) + # --- ウェーブレット変換 (Pre-filter) --- + features = [wav_filter(f) for f in features] + + # --- utilsの関数内で自動変換させる --- + ref_features = get_mc_reference_features_wav(encoder, args.train_dataset_dir, class_names, images.device, args.train_ref_shot, wav_filter=wav_filter) + mfeatures = get_mc_matched_ref_features(features, class_names, ref_features) + rfeatures = get_residual_features(features, mfeatures, pos_flag=True) + + lvl_masks = [] + for l in range(args.feature_levels): + _, _, h, w = rfeatures[l].size() + m = F.interpolate(masks, size=(h, w), mode='nearest').squeeze(1) + lvl_masks.append(m) + rfeatures_t = [rfeature.detach().clone() for rfeature in rfeatures] + + # constraintor 適用 + rfeatures = constraintor(*rfeatures) + loss = 0 + for l in range(args.feature_levels): + e = rfeatures[l] + t = rfeatures_t[l] + bs, dim, h, w = e.size() + e = e.permute(0, 2, 3, 1).reshape(-1, dim) + t = t.permute(0, 2, 3, 1).reshape(-1, dim) + m = lvl_masks[l] + m = m.reshape(-1) + + loss_i, _, _ = calculate_log_barrier_bi_occ_loss(e, m, t) + loss += loss_i + optimizer0.zero_grad() + loss.backward() + optimizer0.step() + + train_loss_total += loss.item() + total_num += 1 + + rfeatures = [rfeature.detach().clone() for rfeature in rfeatures] + loss, num = train(args, rfeatures, estimators, optimizer1, masks, boundary_ops, epoch, N_batch=N_batch, FIRST_STAGE_EPOCH=FIRST_STAGE_EPOCH) + train_loss_total += loss + total_num += num + + scheduler0.step() + scheduler1.step() + + progress_bar.close() + print(f"Epoch[{epoch}/{args.epochs}]: train_loss: {train_loss_total / total_num}") + + # --- 評価フェーズ --- + if (epoch + 1) % args.eval_freq == 0: + s1_res, s2_res, s_res = [], [], [] + test_ref_features = load_mc_reference_features(args.test_ref_feature_dir, CLASSES['unseen'], args.device, args.num_ref_shot) + + for class_name in CLASSES['unseen']: + if args.classes == 'capsules': + test_dataset = CAPSULES(args.test_dataset_dir, class_name=class_name, train=False, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in MVTEC.CLASS_NAMES: + test_dataset = MVTEC(args.test_dataset_dir, class_name=class_name, train=False, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in VISA.CLASS_NAMES: + test_dataset = VISA(args.test_dataset_dir, class_name=class_name, train=False, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in BTAD.CLASS_NAMES: + test_dataset = BTAD(args.test_dataset_dir, class_name=class_name, train=False, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in MVTEC3D.CLASS_NAMES: + test_dataset = MVTEC3D(args.test_dataset_dir, class_name=class_name, train=False, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in MPDD.CLASS_NAMES: + test_dataset = MPDD(args.test_dataset_dir, class_name=class_name, train=False, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in MVTECLOCO.CLASS_NAMES: + test_dataset = MVTECLOCO(args.test_dataset_dir, class_name=class_name, train=False, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in BRATS.CLASS_NAMES: + test_dataset = BRATS(args.test_dataset_dir, class_name=class_name, train=False, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + else: + raise ValueError('Unrecognized class name: {}'.format(class_name)) + + test_loader = DataLoader(test_dataset, batch_size=1, shuffle=False, num_workers=8, drop_last=False) + + # validate関数に wav_filter を渡す + metrics = validate(args, encoder, constraintor, wav_filter, estimators, test_loader, test_ref_features[class_name], args.device, class_name) + + img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro = metrics['scores'] + + # 元の main.py と同じ詳細な print 文 + print("Epoch: {}, Class Name: {}, Image AUC | AP | F1_Score: {} | {} | {}, Pixel AUC | AP | F1_Score | AUPRO: {} | {} | {} | {}".format( + epoch, class_name, img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro)) + s1_res.append(metrics['scores1']) + s2_res.append(metrics['scores2']) + s_res.append(metrics['scores']) + + s1_res = np.array(s1_res) + s2_res = np.array(s2_res) + s_res = np.array(s_res) + img_auc1, img_ap1, img_f1_score1, pix_auc1, pix_ap1, pix_f1_score1, pix_aupro1 = np.mean(s1_res, axis=0) + img_auc2, img_ap2, img_f1_score2, pix_auc2, pix_ap2, pix_f1_score2, pix_aupro2 = np.mean(s2_res, axis=0) + img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro = np.mean(s_res, axis=0) + + # 元の main.py と同じ平均値の print 文 + print('(Logps) Average Image AUC | AP | F1_Score: {:.3f} | {:.3f} | {:.3f}, Average Pixel AUC | AP | F1_Score | AUPRO: {:.3f} | {:.3f} | {:.3f} | {:.3f}'.format( + img_auc1, img_ap1, img_f1_score1, pix_auc1, pix_ap1, pix_f1_score1, pix_aupro1)) + print('(BScores) Average Image AUC | AP | F1_Score: {:.3f} | {:.3f} | {:.3f}, Average Pixel AUC | AP | F1_Score | AUPRO: {:.3f} | {:.3f} | {:.3f} | {:.3f}'.format( + img_auc2, img_ap2, img_f1_score2, pix_auc2, pix_ap2, pix_f1_score2, pix_aupro2)) + print('(Merged) Average Image AUC | AP | F1_Score: {:.3f} | {:.3f} | {:.3f}, Average Pixel AUC | AP | F1_Score | AUPRO: {:.3f} | {:.3f} | {:.3f} | {:.3f}'.format( + img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro)) + + if img_auc > best_img_auc: + os.makedirs(args.checkpoint_path, exist_ok=True) + best_img_auc = img_auc + state_dict = {'constraintor': constraintor.state_dict(), + 'estimators': [estimator.state_dict() for estimator in estimators]} + torch.save(state_dict, os.path.join(args.checkpoint_path, f'{args.setting}_epoch_{epoch}_checkpoints.pth')) + + +def load_mc_reference_features(root_dir: str, class_names, device: torch.device, num_shot=4): + refs = {} + for class_name in class_names: + layer1_refs = torch.from_numpy(np.load(os.path.join(root_dir, class_name, 'layer1.npy'))).to(device) + layer2_refs = torch.from_numpy(np.load(os.path.join(root_dir, class_name, 'layer2.npy'))).to(device) + layer3_refs = torch.from_numpy(np.load(os.path.join(root_dir, class_name, 'layer3.npy'))).to(device) + K1 = (layer1_refs.shape[0] // TOTAL_SHOT) * num_shot + K2 = (layer2_refs.shape[0] // TOTAL_SHOT) * num_shot + K3 = (layer3_refs.shape[0] // TOTAL_SHOT) * num_shot + refs[class_name] = (layer1_refs[:K1, :], layer2_refs[:K2, :], layer3_refs[:K3, :]) + return refs + + +if __name__ == "__main__": + parser = argparse.ArgumentParser() + parser.add_argument('--setting', type=str, default="visa_to_mvtec") + parser.add_argument('--classes', type=str, default="none") + parser.add_argument('--train_dataset_dir', type=str, default="") + parser.add_argument('--test_dataset_dir', type=str, default="") + parser.add_argument('--test_ref_feature_dir', type=str, default="./ref_features/w50/mvtec_4shot_wav") + parser.add_argument('--bgadweight_dir', type=str, default="none") + parser.add_argument('--batch_size', type=int, default=32) + parser.add_argument('--lr', type=float, default=1e-5) + parser.add_argument('--epochs', type=int, default=100) + parser.add_argument('--device', type=str, default="cuda:0") + parser.add_argument('--checkpoint_path', type=str, default="./checkpoints/") + parser.add_argument('--eval_freq', type=int, default=1) + parser.add_argument('--backbone', type=str, default="wide_resnet50_2") + parser.add_argument('--rank', type=int, default="0") + + # flow parameters + parser.add_argument('--flow_arch', type=str, default='conditional_flow_model') + parser.add_argument('--feature_levels', default=3, type=int) + parser.add_argument('--coupling_layers', type=int, default=10) + parser.add_argument('--clamp_alpha', type=float, default=1.9) + parser.add_argument('--pos_embed_dim', type=int, default=256) + parser.add_argument('--pos_beta', type=float, default=0.05) + parser.add_argument('--margin_tau', type=float, default=0.1) + parser.add_argument('--bgspp_lambda', type=float, default=1) + + parser.add_argument('--fdm_alpha', type=float, default=0.4) + parser.add_argument('--num_embeddings', type=int, default=1536) + parser.add_argument("--train_ref_shot", type=int, default=4) + parser.add_argument("--num_ref_shot", type=int, default=4) + + # --- 追加: ウェーブレット用パラメータ --- + parser.add_argument("--lf_weight", type=float, default=0.1) + parser.add_argument("--hf_weight", type=float, default=1.2) + + args = parser.parse_args() + init_seeds(42) + + main(args) diff --git a/main_wav1.py b/main_wav1.py new file mode 100644 index 0000000..44d7ccb --- /dev/null +++ b/main_wav1.py @@ -0,0 +1,336 @@ +import os +import warnings +import argparse +from tqdm import tqdm +import numpy as np +import torch +import torch.nn as nn +import timm +import torch.nn.functional as F +from torch.utils.data import DataLoader + +from train import train +from validate_wav1 import validate +from datasets.mvtec import MVTEC, MVTECANO +from datasets.visa import VISA, VISAANO +from datasets.btad import BTAD +from datasets.mvtec_3d import MVTEC3D +from datasets.mpdd import MPDD +from datasets.mvtec_loco import MVTECLOCO +from datasets.brats import BRATS +from datasets.capsules import CAPSULES, CAPSULESANO + +from models.fc_flow import load_flow_model +from models.modules import MultiScaleConv +# 注意: get_mc_reference_features_wav ではなく、元の生特徴量を扱う get_mc_reference_features をインポート +from utils import init_seeds, get_residual_features, get_mc_matched_ref_features, get_mc_reference_features +from utils import BoundaryAverager +from losses.loss import calculate_log_barrier_bi_occ_loss +from classes import VISA_TO_MVTEC, MVTEC_TO_VISA, MVTEC_TO_BTAD, MVTEC_TO_MVTEC3D +from classes import MVTEC_TO_MPDD, MVTEC_TO_MVTECLOCO, MVTEC_TO_BRATS +from classes import MVTEC_TO_MVTEC, VISA_TO_VISA +from classes import CAPSULES_TO_CAPSULES + +warnings.filterwarnings('ignore') + +TOTAL_SHOT = 4 +FIRST_STAGE_EPOCH = 10 +SETTINGS = {'visa_to_mvtec': VISA_TO_MVTEC, 'mvtec_to_visa': MVTEC_TO_VISA, + 'mvtec_to_btad': MVTEC_TO_BTAD, 'mvtec_to_mvtec3d': MVTEC_TO_MVTEC3D, + 'mvtec_to_mpdd': MVTEC_TO_MPDD, 'mvtec_to_mvtecloco': MVTEC_TO_MVTECLOCO, + 'mvtec_to_brats': MVTEC_TO_BRATS,'mvtec_to_mvtec':MVTEC_TO_MVTEC, 'visa_to_visa':VISA_TO_VISA, 'capsules_to_capsules': CAPSULES_TO_CAPSULES} + +# ========================================== +# 1. Frequency Gating Network (層独立) +# ========================================== +class FrequencyGatingNetwork(nn.Module): + def __init__(self, in_channels): + super().__init__() + self.net = nn.Sequential( + nn.AdaptiveAvgPool2d(1), + nn.Flatten(), + nn.Linear(in_channels, 128), + nn.ReLU(), + nn.Linear(128, 6) # Layer1, Layer2, Layer3 のそれぞれに対する LF/HF で計6出力 + ) + + def forward(self, x): + out = self.net(x) + # LF(低周波)は 0.0 ~ 0.5、HF(高周波)は 1.0 ~ 2.0 に制限 + lf_w = torch.sigmoid(out[:, 0::2]) * 0.5 + hf_w = 1.0 + torch.sigmoid(out[:, 1::2]) * 1.0 + + # (Layer1_lf, Layer1_hf), (Layer2_lf, Layer2_hf), (Layer3_lf, Layer3_hf) の順で返す + return (lf_w[:, 0].view(-1,1,1,1), hf_w[:, 0].view(-1,1,1,1)), \ + (lf_w[:, 1].view(-1,1,1,1), hf_w[:, 1].view(-1,1,1,1)), \ + (lf_w[:, 2].view(-1,1,1,1), hf_w[:, 2].view(-1,1,1,1)) + +# ========================================== +# 2. Dynamic Haar Wavelet Filter +# ========================================== +class HaarWaveletFilterDynamic(nn.Module): + def __init__(self): + super().__init__() + ll = torch.tensor([[0.5, 0.5], [0.5, 0.5]]) + hl = torch.tensor([[-0.5, -0.5], [0.5, 0.5]]) + lh = torch.tensor([[-0.5, 0.5], [-0.5, 0.5]]) + hh = torch.tensor([[0.5, -0.5], [-0.5, 0.5]]) + + self.register_buffer('k_ll', ll.view(1, 1, 2, 2)) + self.register_buffer('k_hl', hl.view(1, 1, 2, 2)) + self.register_buffer('k_lh', lh.view(1, 1, 2, 2)) + self.register_buffer('k_hh', hh.view(1, 1, 2, 2)) + + def forward(self, x, lf_w, hf_w): + B, C, H, W = x.shape + ll = F.conv2d(x, self.k_ll.expand(C, 1, 2, 2), stride=2, groups=C) + hl = F.conv2d(x, self.k_hl.expand(C, 1, 2, 2), stride=2, groups=C) + lh = F.conv2d(x, self.k_lh.expand(C, 1, 2, 2), stride=2, groups=C) + hh = F.conv2d(x, self.k_hh.expand(C, 1, 2, 2), stride=2, groups=C) + + ll = ll * lf_w + hl = hl * hf_w + lh = lh * hf_w + hh = hh * hf_w + + out = F.conv_transpose2d(ll, self.k_ll.expand(C, 1, 2, 2), stride=2, groups=C) + \ + F.conv_transpose2d(hl, self.k_hl.expand(C, 1, 2, 2), stride=2, groups=C) + \ + F.conv_transpose2d(lh, self.k_lh.expand(C, 1, 2, 2), stride=2, groups=C) + \ + F.conv_transpose2d(hh, self.k_hh.expand(C, 1, 2, 2), stride=2, groups=C) + return out + + +def main(args): + if args.setting in SETTINGS.keys(): + CLASSES = SETTINGS[args.setting] + else: + raise ValueError(f"Dataset setting must be in {SETTINGS.keys()}, but got {args.setting}.") + + if args.classes == 'capsules': + train_dataset1 = CAPSULES(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, normalize="w50", img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader1 = DataLoader(train_dataset1, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True) + train_dataset2 = CAPSULESANO(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader2 = DataLoader(train_dataset2, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True) + elif CLASSES['seen'][0] in MVTEC.CLASS_NAMES: + train_dataset1 = MVTEC(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, normalize="w50", img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader1 = DataLoader(train_dataset1, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True) + train_dataset2 = MVTECANO(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader2 = DataLoader(train_dataset2, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True) + else: + train_dataset1 = VISA(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, normalize="w50", img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader1 = DataLoader(train_dataset1, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True) + train_dataset2 = VISAANO(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, normalize="w50", img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader2 = DataLoader(train_dataset2, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True) + + if args.backbone == 'wide_resnet50_2': + encoder = timm.create_model('wide_resnet50_2', features_only=True, out_indices=(1, 2, 3), pretrained=True).eval() + encoder = encoder.to(args.device) + feat_dims = encoder.feature_info.channels() + elif args.backbone == 'tf_efficientnet_b6': + encoder = timm.create_model('tf_efficientnet_b6', features_only=True, out_indices=(1, 2, 3), pretrained=True).eval() + encoder = encoder.to(args.device) + feat_dims = encoder.feature_info.channels() + + boundary_ops = BoundaryAverager(num_levels=args.feature_levels) + + # --- モジュールの初期化 --- + gating_in_channels = feat_dims[-1] + gating_net = FrequencyGatingNetwork(in_channels=gating_in_channels).to(args.device) + wav_filter = HaarWaveletFilterDynamic().to(args.device) + + constraintor = MultiScaleConv(feat_dims).to(args.device) + optimizer0 = torch.optim.Adam(constraintor.parameters(), lr=args.lr, weight_decay=0.0005) + scheduler0 = torch.optim.lr_scheduler.MultiStepLR(optimizer0, milestones=[70, 90], gamma=0.1) + + estimators = [load_flow_model(args, feat_dim) for feat_dim in feat_dims] + estimators = [decoder.to(args.device) for decoder in estimators] + params = list(estimators[0].parameters()) + for l in range(1, args.feature_levels): + params += list(estimators[l].parameters()) + + # ★ ゲーティングネットワークのパラメータをオプティマイザに追加 + params += list(gating_net.parameters()) + + optimizer1 = torch.optim.Adam(params, lr=args.lr, weight_decay=0.0005) + scheduler1 = torch.optim.lr_scheduler.MultiStepLR(optimizer1, milestones=[70, 90], gamma=0.1) + + best_img_auc = 0 + N_batch = 8192 + + for epoch in range(args.epochs): + gating_net.train() # 学習モード + constraintor.train() + for estimator in estimators: + estimator.train() + + if epoch < FIRST_STAGE_EPOCH: + train_loader = train_loader1 + else: + train_loader = train_loader2 + + train_loss_total, total_num = 0, 0 + progress_bar = tqdm(total=len(train_loader)) + progress_bar.set_description(f"Epoch[{epoch}/{args.epochs}]") + + for step, batch in enumerate(train_loader): + progress_bar.update(1) + images, _, masks, class_names = batch + images, masks = images.to(args.device), masks.to(args.device) + + with torch.no_grad(): + # 1. 生の特徴量を抽出 (まだウェーブレットはかけない) + features_raw = encoder(images) + + # 2. 画像の深い特徴量から、層ごとの重みを予測 [w1, w2, w3] + w1, w2, w3 = gating_net(features_raw[-1].detach()) + weights = [w1, w2, w3] + + # 3. 生のカンペを取得し、生のテスト画像とマッチング + ref_features_raw = get_mc_reference_features(encoder, args.train_dataset_dir, class_names, images.device, args.train_ref_shot) + mfeatures_raw = get_mc_matched_ref_features(features_raw, class_names, ref_features_raw) + + # 4. テスト画像とカンペの両方に、予測した「層ごとの同じ重み」でフィルタをかける + features_wav = [wav_filter(features_raw[i], weights[i][0], weights[i][1]) for i in range(args.feature_levels)] + mfeatures_wav = [wav_filter(mfeatures_raw[i], weights[i][0], weights[i][1]) for i in range(args.feature_levels)] + + # 5. フィルタリング後の特徴量同士で残差を計算 + rfeatures = get_residual_features(features_wav, mfeatures_wav, pos_flag=True) + + lvl_masks = [] + for l in range(args.feature_levels): + _, _, h, w = rfeatures[l].size() + m = F.interpolate(masks, size=(h, w), mode='nearest').squeeze(1) + lvl_masks.append(m) + rfeatures_t = [rfeature.detach().clone() for rfeature in rfeatures] + + rfeatures = constraintor(*rfeatures) + loss = 0 + for l in range(args.feature_levels): + e = rfeatures[l] + t = rfeatures_t[l] + bs, dim, h, w = e.size() + e = e.permute(0, 2, 3, 1).reshape(-1, dim) + t = t.permute(0, 2, 3, 1).reshape(-1, dim) + m = lvl_masks[l] + m = m.reshape(-1) + + loss_i, _, _ = calculate_log_barrier_bi_occ_loss(e, m, t) + loss += loss_i + optimizer0.zero_grad() + loss.backward() + optimizer0.step() + + train_loss_total += loss.item() + total_num += 1 + + rfeatures = [rfeature.detach().clone() for rfeature in rfeatures] + loss, num = train(args, rfeatures, estimators, optimizer1, masks, boundary_ops, epoch, N_batch=N_batch, FIRST_STAGE_EPOCH=FIRST_STAGE_EPOCH) + train_loss_total += loss + total_num += num + + scheduler0.step() + scheduler1.step() + + progress_bar.close() + print(f"Epoch[{epoch}/{args.epochs}]: train_loss: {train_loss_total / total_num}") + + # --- 評価フェーズ --- + if (epoch + 1) % args.eval_freq == 0: + s1_res, s2_res, s_res = [], [], [] + # 注意: テスト用のカンペも、ウェーブレット変換をしていない「元のResADの生特徴量」をロードします + test_ref_features = load_mc_reference_features(args.test_ref_feature_dir, CLASSES['unseen'], args.device, args.num_ref_shot) + + for class_name in CLASSES['unseen']: + if args.classes == 'capsules': test_dataset = CAPSULES(args.test_dataset_dir, class_name=class_name, train=False, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in MVTEC.CLASS_NAMES: test_dataset = MVTEC(args.test_dataset_dir, class_name=class_name, train=False, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in VISA.CLASS_NAMES: test_dataset = VISA(args.test_dataset_dir, class_name=class_name, train=False, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in BTAD.CLASS_NAMES: test_dataset = BTAD(args.test_dataset_dir, class_name=class_name, train=False, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in MVTEC3D.CLASS_NAMES: test_dataset = MVTEC3D(args.test_dataset_dir, class_name=class_name, train=False, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in MPDD.CLASS_NAMES: test_dataset = MPDD(args.test_dataset_dir, class_name=class_name, train=False, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in MVTECLOCO.CLASS_NAMES: test_dataset = MVTECLOCO(args.test_dataset_dir, class_name=class_name, train=False, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in BRATS.CLASS_NAMES: test_dataset = BRATS(args.test_dataset_dir, class_name=class_name, train=False, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + else: raise ValueError('Unrecognized class name: {}'.format(class_name)) + + test_loader = DataLoader(test_dataset, batch_size=1, shuffle=False, num_workers=8, drop_last=False) + + # validate関数に gating_net を追加して渡す + metrics = validate(args, encoder, constraintor, gating_net, wav_filter, estimators, test_loader, test_ref_features[class_name], args.device, class_name) + + img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro = metrics['scores'] + + print("Epoch: {}, Class Name: {}, Image AUC | AP | F1_Score: {} | {} | {}, Pixel AUC | AP | F1_Score | AUPRO: {} | {} | {} | {}".format( + epoch, class_name, img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro)) + s1_res.append(metrics['scores1']) + s2_res.append(metrics['scores2']) + s_res.append(metrics['scores']) + + s1_res = np.array(s1_res) + s2_res = np.array(s2_res) + s_res = np.array(s_res) + img_auc1, img_ap1, img_f1_score1, pix_auc1, pix_ap1, pix_f1_score1, pix_aupro1 = np.mean(s1_res, axis=0) + img_auc2, img_ap2, img_f1_score2, pix_auc2, pix_ap2, pix_f1_score2, pix_aupro2 = np.mean(s2_res, axis=0) + img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro = np.mean(s_res, axis=0) + + print('(Logps) Average Image AUC | AP | F1_Score: {:.3f} | {:.3f} | {:.3f}, Average Pixel AUC | AP | F1_Score | AUPRO: {:.3f} | {:.3f} | {:.3f} | {:.3f}'.format( + img_auc1, img_ap1, img_f1_score1, pix_auc1, pix_ap1, pix_f1_score1, pix_aupro1)) + print('(BScores) Average Image AUC | AP | F1_Score: {:.3f} | {:.3f} | {:.3f}, Average Pixel AUC | AP | F1_Score | AUPRO: {:.3f} | {:.3f} | {:.3f} | {:.3f}'.format( + img_auc2, img_ap2, img_f1_score2, pix_auc2, pix_ap2, pix_f1_score2, pix_aupro2)) + print('(Merged) Average Image AUC | AP | F1_Score: {:.3f} | {:.3f} | {:.3f}, Average Pixel AUC | AP | F1_Score | AUPRO: {:.3f} | {:.3f} | {:.3f} | {:.3f}'.format( + img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro)) + + if img_auc > best_img_auc: + os.makedirs(args.checkpoint_path, exist_ok=True) + best_img_auc = img_auc + # ★ ここに gating_net の値の保存処理を追加しています + state_dict = {'gating_net': gating_net.state_dict(), + 'constraintor': constraintor.state_dict(), + 'estimators': [estimator.state_dict() for estimator in estimators]} + torch.save(state_dict, os.path.join(args.checkpoint_path, f'{args.setting}_epoch_{epoch}_checkpoints.pth')) + +def load_mc_reference_features(root_dir: str, class_names, device: torch.device, num_shot=4): + refs = {} + for class_name in class_names: + layer1_refs = torch.from_numpy(np.load(os.path.join(root_dir, class_name, 'layer1.npy'))).to(device) + layer2_refs = torch.from_numpy(np.load(os.path.join(root_dir, class_name, 'layer2.npy'))).to(device) + layer3_refs = torch.from_numpy(np.load(os.path.join(root_dir, class_name, 'layer3.npy'))).to(device) + K1 = (layer1_refs.shape[0] // TOTAL_SHOT) * num_shot + K2 = (layer2_refs.shape[0] // TOTAL_SHOT) * num_shot + K3 = (layer3_refs.shape[0] // TOTAL_SHOT) * num_shot + refs[class_name] = (layer1_refs[:K1, :], layer2_refs[:K2, :], layer3_refs[:K3, :]) + return refs + +if __name__ == "__main__": + parser = argparse.ArgumentParser() + parser.add_argument('--setting', type=str, default="visa_to_mvtec") + parser.add_argument('--classes', type=str, default="none") + parser.add_argument('--train_dataset_dir', type=str, default="") + parser.add_argument('--test_dataset_dir', type=str, default="") + # ★ 注意: ここで指定するディレクトリは「ウェーブレット処理をしていない生特徴量」のカンペフォルダを指定してください + parser.add_argument('--test_ref_feature_dir', type=str, default="./ref_features/w50/mvtec_4shot") + parser.add_argument('--bgadweight_dir', type=str, default="none") + parser.add_argument('--batch_size', type=int, default=32) + parser.add_argument('--lr', type=float, default=1e-5) + parser.add_argument('--epochs', type=int, default=100) + parser.add_argument('--device', type=str, default="cuda:0") + parser.add_argument('--checkpoint_path', type=str, default="./checkpoints/") + parser.add_argument('--eval_freq', type=int, default=1) + parser.add_argument('--backbone', type=str, default="wide_resnet50_2") + parser.add_argument('--rank', type=int, default="0") + + # flow parameters + parser.add_argument('--flow_arch', type=str, default='conditional_flow_model') + parser.add_argument('--feature_levels', default=3, type=int) + parser.add_argument('--coupling_layers', type=int, default=10) + parser.add_argument('--clamp_alpha', type=float, default=1.9) + parser.add_argument('--pos_embed_dim', type=int, default=256) + parser.add_argument('--pos_beta', type=float, default=0.05) + parser.add_argument('--margin_tau', type=float, default=0.1) + parser.add_argument('--bgspp_lambda', type=float, default=1) + parser.add_argument('--fdm_alpha', type=float, default=0.4) + parser.add_argument("--train_ref_shot", type=int, default=4) + parser.add_argument("--num_ref_shot", type=int, default=4) + + args = parser.parse_args() + init_seeds(42) + + main(args) diff --git a/main_wav_cf.py b/main_wav_cf.py new file mode 100644 index 0000000..36596b7 --- /dev/null +++ b/main_wav_cf.py @@ -0,0 +1,318 @@ +import os +import warnings +import argparse +from tqdm import tqdm +import numpy as np +import torch +import torch.nn as nn +import timm +import torch.nn.functional as F +from torch.utils.data import DataLoader + +from train import train +from validate_wav_cf import validate +from datasets.mvtec import MVTEC, MVTECANO +from datasets.visa import VISA, VISAANO +from datasets.btad import BTAD +from datasets.mvtec_3d import MVTEC3D +from datasets.mpdd import MPDD +from datasets.mvtec_loco import MVTECLOCO +from datasets.brats import BRATS +from datasets.capsules import CAPSULES, CAPSULESANO + +from models.fc_flow import load_flow_model +from models.modules import MultiScaleConv +# ★ get_random_normal_images, load_and_transform_vision_data を追加インポート +from utils import init_seeds, get_residual_features, get_mc_matched_ref_features +from utils import get_random_normal_images, load_and_transform_vision_data +from utils import BoundaryAverager +from losses.loss import calculate_log_barrier_bi_occ_loss +from classes import VISA_TO_MVTEC, MVTEC_TO_VISA, MVTEC_TO_BTAD, MVTEC_TO_MVTEC3D +from classes import MVTEC_TO_MPDD, MVTEC_TO_MVTECLOCO, MVTEC_TO_BRATS +from classes import MVTEC_TO_MVTEC, VISA_TO_VISA +from classes import CAPSULES_TO_CAPSULES + +warnings.filterwarnings('ignore') + +TOTAL_SHOT = 4 +FIRST_STAGE_EPOCH = 10 +SETTINGS = {'visa_to_mvtec': VISA_TO_MVTEC, 'mvtec_to_visa': MVTEC_TO_VISA, + 'mvtec_to_btad': MVTEC_TO_BTAD, 'mvtec_to_mvtec3d': MVTEC_TO_MVTEC3D, + 'mvtec_to_mpdd': MVTEC_TO_MPDD, 'mvtec_to_mvtecloco': MVTEC_TO_MVTECLOCO, + 'mvtec_to_brats': MVTEC_TO_BRATS,'mvtec_to_mvtec':MVTEC_TO_MVTEC, 'visa_to_visa':VISA_TO_VISA, 'capsules_to_capsules': CAPSULES_TO_CAPSULES} + +# ========================================== +# ★ CFモジュール用の専用カンペ抽出関数(平坦化前にCFを通す) +# ========================================== +def get_cf_reference_features(encoder, wav_filter, cf_modules, root, class_names, device, num_shot=4): + reference_features = {} + class_names = np.unique(class_names) + for class_name in class_names: + normal_paths = get_random_normal_images(root, class_name, num_shot) + images = load_and_transform_vision_data(normal_paths, device) + with torch.no_grad(): + features_raw = encoder(images) + features_cat = [] + for l in range(len(features_raw)): + # 空間次元(H, W)を保ったままフィルタとCFを通す + r_lf, r_hf = wav_filter.get_LF_HF(features_raw[l]) + r_lf, r_hf = cf_modules[l](r_lf, r_hf) + cat_f = torch.cat([r_lf, r_hf], dim=1) + + # CFを通した後に、マッチング用に平坦化(Flatten)する + bs, c, h, w = cat_f.shape + cat_f = cat_f.permute(0, 2, 3, 1).reshape(-1, c) + features_cat.append(cat_f) + reference_features[class_name] = features_cat + return reference_features + +class HaarWaveletFilter2Component(nn.Module): + def __init__(self): + super().__init__() + ll = torch.tensor([[0.5, 0.5], [0.5, 0.5]]) + hl = torch.tensor([[-0.5, -0.5], [0.5, 0.5]]) + lh = torch.tensor([[-0.5, 0.5], [-0.5, 0.5]]) + hh = torch.tensor([[0.5, -0.5], [-0.5, 0.5]]) + + self.register_buffer('k_ll', ll.view(1, 1, 2, 2)) + self.register_buffer('k_hl', hl.view(1, 1, 2, 2)) + self.register_buffer('k_lh', lh.view(1, 1, 2, 2)) + self.register_buffer('k_hh', hh.view(1, 1, 2, 2)) + + def get_LF_HF(self, x): + B, C, H, W = x.shape + ll = F.conv2d(x, self.k_ll.expand(C, 1, 2, 2), stride=2, groups=C) + hl = F.conv2d(x, self.k_hl.expand(C, 1, 2, 2), stride=2, groups=C) + lh = F.conv2d(x, self.k_lh.expand(C, 1, 2, 2), stride=2, groups=C) + hh = F.conv2d(x, self.k_hh.expand(C, 1, 2, 2), stride=2, groups=C) + + lf = F.conv_transpose2d(ll, self.k_ll.expand(C, 1, 2, 2), stride=2, groups=C) + hf_hl = F.conv_transpose2d(hl, self.k_hl.expand(C, 1, 2, 2), stride=2, groups=C) + hf_lh = F.conv_transpose2d(lh, self.k_lh.expand(C, 1, 2, 2), stride=2, groups=C) + hf_hh = F.conv_transpose2d(hh, self.k_hh.expand(C, 1, 2, 2), stride=2, groups=C) + hf = hf_hl + hf_lh + hf_hh + return lf, hf + +class CrossFrequencyModule(nn.Module): + def __init__(self, in_channels): + super().__init__() + self.conv_block = nn.Sequential( + nn.Conv2d(in_channels * 2, in_channels * 2, kernel_size=3, padding=1), + nn.Conv2d(in_channels * 2, in_channels * 2, kernel_size=3, padding=1), + nn.LeakyReLU(0.2, inplace=True) + ) + + def forward(self, lf, hf): + x = torch.cat([lf, hf], dim=1) + x_out = self.conv_block(x) + lf_out, hf_out = torch.chunk(x_out, 2, dim=1) + lf_refined = lf + lf_out + hf_refined = hf + hf_out + return lf_refined, hf_refined + +def main(args): + CLASSES = SETTINGS[args.setting] + + if args.classes == 'capsules': + train_dataset1 = CAPSULES(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, normalize="w50", img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader1 = DataLoader(train_dataset1, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True) + train_dataset2 = CAPSULESANO(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader2 = DataLoader(train_dataset2, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True) + elif CLASSES['seen'][0] in MVTEC.CLASS_NAMES: + train_dataset1 = MVTEC(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, normalize="w50", img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader1 = DataLoader(train_dataset1, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True) + train_dataset2 = MVTECANO(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader2 = DataLoader(train_dataset2, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True) + else: + train_dataset1 = VISA(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, normalize="w50", img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader1 = DataLoader(train_dataset1, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True) + train_dataset2 = VISAANO(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, normalize="w50", img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + train_loader2 = DataLoader(train_dataset2, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True) + + if args.backbone == 'wide_resnet50_2': + encoder = timm.create_model('wide_resnet50_2', features_only=True, out_indices=(1, 2, 3), pretrained=True).eval().to(args.device) + feat_dims = encoder.feature_info.channels() + elif args.backbone == 'tf_efficientnet_b6': + encoder = timm.create_model('tf_efficientnet_b6', features_only=True, out_indices=(1, 2, 3), pretrained=True).eval().to(args.device) + feat_dims = encoder.feature_info.channels() + + feat_dims_cat = [dim * 2 for dim in feat_dims] + + boundary_ops = BoundaryAverager(num_levels=args.feature_levels) + wav_filter = HaarWaveletFilter2Component().to(args.device) + cf_modules = nn.ModuleList([CrossFrequencyModule(dim) for dim in feat_dims]).to(args.device) + + constraintor = MultiScaleConv(feat_dims_cat).to(args.device) + params_c = list(constraintor.parameters()) + list(cf_modules.parameters()) + optimizer0 = torch.optim.Adam(params_c, lr=args.lr, weight_decay=0.0005) + scheduler0 = torch.optim.lr_scheduler.MultiStepLR(optimizer0, milestones=[70, 90], gamma=0.1) + + estimators = [load_flow_model(args, feat_dim).to(args.device) for feat_dim in feat_dims_cat] + params_f = list(estimators[0].parameters()) + for l in range(1, args.feature_levels): + params_f += list(estimators[l].parameters()) + optimizer1 = torch.optim.Adam(params_f, lr=args.lr, weight_decay=0.0005) + scheduler1 = torch.optim.lr_scheduler.MultiStepLR(optimizer1, milestones=[70, 90], gamma=0.1) + + best_img_auc = 0 + N_batch = 8192 + + for epoch in range(args.epochs): + constraintor.train() + cf_modules.train() + for estimator in estimators: estimator.train() + + train_loader = train_loader1 if epoch < FIRST_STAGE_EPOCH else train_loader2 + train_loss_total, total_num = 0, 0 + progress_bar = tqdm(total=len(train_loader)) + progress_bar.set_description(f"Epoch[{epoch}/{args.epochs}]") + + for step, batch in enumerate(train_loader): + progress_bar.update(1) + images, _, masks, class_names = batch[0], batch[1], batch[2], batch[3] + images, masks = images.to(args.device), masks.to(args.device) + + with torch.no_grad(): + features_raw = encoder(images) + + # --- 修正箇所:専用関数で安全にカンペを抽出 --- + ref_features_cat_dict = get_cf_reference_features( + encoder, wav_filter, cf_modules, args.train_dataset_dir, class_names, images.device, args.train_ref_shot + ) + + features_cat = [] + for l in range(args.feature_levels): + test_lf, test_hf = wav_filter.get_LF_HF(features_raw[l]) + test_lf, test_hf = cf_modules[l](test_lf, test_hf) + features_cat.append(torch.cat([test_lf, test_hf], dim=1)) + + mfeatures = get_mc_matched_ref_features(features_cat, class_names, ref_features_cat_dict) + rfeatures = get_residual_features(features_cat, mfeatures, pos_flag=True) + # ---------------------------------------------- + + lvl_masks = [] + for l in range(args.feature_levels): + _, _, h, w = rfeatures[l].size() + m = F.interpolate(masks, size=(h, w), mode='nearest').squeeze(1) + lvl_masks.append(m) + rfeatures_t = [rfeature.detach().clone() for rfeature in rfeatures] + + rfeatures = constraintor(*rfeatures) + loss = 0 + for l in range(args.feature_levels): + e = rfeatures[l] + t = rfeatures_t[l] + bs, dim, h, w = e.size() + e = e.permute(0, 2, 3, 1).reshape(-1, dim) + t = t.permute(0, 2, 3, 1).reshape(-1, dim) + m = lvl_masks[l] + m = m.reshape(-1) + + loss_i, _, _ = calculate_log_barrier_bi_occ_loss(e, m, t) + loss += loss_i + optimizer0.zero_grad() + loss.backward() + optimizer0.step() + + train_loss_total += loss.item() + total_num += 1 + + rfeatures = [rfeature.detach().clone() for rfeature in rfeatures] + loss, num = train(args, rfeatures, estimators, optimizer1, masks, boundary_ops, epoch, N_batch=N_batch, FIRST_STAGE_EPOCH=FIRST_STAGE_EPOCH) + train_loss_total += loss + total_num += num + + scheduler0.step() + scheduler1.step() + + progress_bar.close() + print(f"Epoch[{epoch}/{args.epochs}]: train_loss: {train_loss_total / total_num}") + + if (epoch + 1) % args.eval_freq == 0: + s1_res, s2_res, s_res = [], [], [] + test_ref_features = load_mc_reference_features(args.test_ref_feature_dir, CLASSES['unseen'], args.device, args.num_ref_shot) + + for class_name in CLASSES['unseen']: + if args.classes == 'capsules': test_dataset = CAPSULES(args.test_dataset_dir, class_name=class_name, train=False, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in MVTEC.CLASS_NAMES: test_dataset = MVTEC(args.test_dataset_dir, class_name=class_name, train=False, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in VISA.CLASS_NAMES: test_dataset = VISA(args.test_dataset_dir, class_name=class_name, train=False, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in BTAD.CLASS_NAMES: test_dataset = BTAD(args.test_dataset_dir, class_name=class_name, train=False, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in MVTEC3D.CLASS_NAMES: test_dataset = MVTEC3D(args.test_dataset_dir, class_name=class_name, train=False, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in MPDD.CLASS_NAMES: test_dataset = MPDD(args.test_dataset_dir, class_name=class_name, train=False, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in MVTECLOCO.CLASS_NAMES: test_dataset = MVTECLOCO(args.test_dataset_dir, class_name=class_name, train=False, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + elif class_name in BRATS.CLASS_NAMES: test_dataset = BRATS(args.test_dataset_dir, class_name=class_name, train=False, normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) + else: raise ValueError('Unrecognized class name: {}'.format(class_name)) + + test_loader = DataLoader(test_dataset, batch_size=1, shuffle=False, num_workers=8, drop_last=False) + + metrics = validate(args, encoder, constraintor, wav_filter, cf_modules, estimators, test_loader, test_ref_features[class_name], args.device, class_name) + + img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro = metrics['scores'] + print("Epoch: {}, Class Name: {}, Image AUC | AP | F1_Score: {} | {} | {}, Pixel AUC | AP | F1_Score | AUPRO: {} | {} | {} | {}".format( + epoch, class_name, img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro)) + s1_res.append(metrics['scores1']) + s2_res.append(metrics['scores2']) + s_res.append(metrics['scores']) + + s1_res = np.array(s1_res) + s2_res = np.array(s2_res) + s_res = np.array(s_res) + img_auc1, img_ap1, img_f1_score1, pix_auc1, pix_ap1, pix_f1_score1, pix_aupro1 = np.mean(s1_res, axis=0) + img_auc2, img_ap2, img_f1_score2, pix_auc2, pix_ap2, pix_f1_score2, pix_aupro2 = np.mean(s2_res, axis=0) + img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro = np.mean(s_res, axis=0) + + print('(Merged) Average Image AUC | AP | F1_Score: {:.3f} | {:.3f} | {:.3f}, Average Pixel AUC | AP | F1_Score | AUPRO: {:.3f} | {:.3f} | {:.3f} | {:.3f}'.format( + img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro)) + + if img_auc > best_img_auc: + os.makedirs(args.checkpoint_path, exist_ok=True) + best_img_auc = img_auc + state_dict = {'cf_modules': cf_modules.state_dict(), + 'constraintor': constraintor.state_dict(), + 'estimators': [estimator.state_dict() for estimator in estimators]} + torch.save(state_dict, os.path.join(args.checkpoint_path, f'{args.setting}_epoch_{epoch}_checkpoints.pth')) + +def load_mc_reference_features(root_dir: str, class_names, device: torch.device, num_shot=4): + refs = {} + for class_name in class_names: + layer1_refs = torch.from_numpy(np.load(os.path.join(root_dir, class_name, 'layer1.npy'))).to(device) + layer2_refs = torch.from_numpy(np.load(os.path.join(root_dir, class_name, 'layer2.npy'))).to(device) + layer3_refs = torch.from_numpy(np.load(os.path.join(root_dir, class_name, 'layer3.npy'))).to(device) + K1 = (layer1_refs.shape[0] // TOTAL_SHOT) * num_shot + K2 = (layer2_refs.shape[0] // TOTAL_SHOT) * num_shot + K3 = (layer3_refs.shape[0] // TOTAL_SHOT) * num_shot + refs[class_name] = (layer1_refs[:K1, :], layer2_refs[:K2, :], layer3_refs[:K3, :]) + return refs + +if __name__ == "__main__": + parser = argparse.ArgumentParser() + parser.add_argument('--setting', type=str, default="visa_to_mvtec") + parser.add_argument('--classes', type=str, default="none") + parser.add_argument('--train_dataset_dir', type=str, default="") + parser.add_argument('--test_dataset_dir', type=str, default="") + parser.add_argument('--test_ref_feature_dir', type=str, default="./ref_features/w50/mvtec_4shot") + parser.add_argument('--bgadweight_dir', type=str, default="none") + parser.add_argument('--batch_size', type=int, default=32) + parser.add_argument('--lr', type=float, default=1e-5) + parser.add_argument('--epochs', type=int, default=100) + parser.add_argument('--device', type=str, default="cuda:0") + parser.add_argument('--checkpoint_path', type=str, default="./checkpoints/") + parser.add_argument('--eval_freq', type=int, default=1) + parser.add_argument('--backbone', type=str, default="wide_resnet50_2") + parser.add_argument('--rank', type=int, default="0") + + parser.add_argument('--flow_arch', type=str, default='conditional_flow_model') + parser.add_argument('--feature_levels', default=3, type=int) + parser.add_argument('--coupling_layers', type=int, default=10) + parser.add_argument('--clamp_alpha', type=float, default=1.9) + parser.add_argument('--pos_embed_dim', type=int, default=256) + parser.add_argument('--pos_beta', type=float, default=0.05) + parser.add_argument('--margin_tau', type=float, default=0.1) + parser.add_argument('--bgspp_lambda', type=float, default=1) + parser.add_argument('--fdm_alpha', type=float, default=0.4) + parser.add_argument("--train_ref_shot", type=int, default=4) + parser.add_argument("--num_ref_shot", type=int, default=4) + + args = parser.parse_args() + init_seeds(42) + main(args) diff --git a/models/imagebind.py b/models/imagebind.py index cd9608e..6a4acf3 100644 --- a/models/imagebind.py +++ b/models/imagebind.py @@ -10,8 +10,8 @@ class ImageBindModel(nn.Module): def __init__(self, device='cuda:0'): super(ImageBindModel, self).__init__() - - imagebind_ckpt_path = './pretrained_weights/imagebind/imagebind_huge.pth' +#DLboxで使うため + imagebind_ckpt_path = '/home/ueno/pretrained_weights/imagebind/imagebind_huge.pth' print (f'Initializing visual encoder from {imagebind_ckpt_path} ...') self.visual_encoder, self.visual_hidden_size = imagebind_model.imagebind_huge({}) @@ -139,4 +139,4 @@ def generate(self, inputs, web_demo=False): """ anomaly_map = self.prepare_generation_embedding(inputs, web_demo) - return anomaly_map \ No newline at end of file + return anomaly_map diff --git a/models/modules.py b/models/modules.py index 6e3353b..47cda41 100644 --- a/models/modules.py +++ b/models/modules.py @@ -389,4 +389,4 @@ def forward(self, layer1_x, layer2_x, layer3_x, layer4_x): return out1, out2, out3, out4 - \ No newline at end of file + diff --git a/scripts/eval_text_ref_map.py b/scripts/eval_text_ref_map.py new file mode 100644 index 0000000..aad711e --- /dev/null +++ b/scripts/eval_text_ref_map.py @@ -0,0 +1,846 @@ +import argparse +import csv +import importlib +import json +import math +import os +import random +import subprocess +import sys +from contextlib import contextmanager +from pathlib import Path + +import numpy as np +import torch +import torch.nn.functional as F +from PIL import Image +from torch.utils.data import DataLoader, Dataset +from torchvision import transforms + + +PROJECT_ROOT = Path(__file__).resolve().parents[1] +if str(PROJECT_ROOT) not in sys.path: + sys.path.insert(0, str(PROJECT_ROOT)) + +IMAGE_EXTS = {".png", ".jpg", ".jpeg", ".bmp", ".tif", ".tiff"} +IMAGENET_MEAN = (0.485, 0.456, 0.406) +IMAGENET_STD = (0.229, 0.224, 0.225) +CLIP_MEAN = (0.48145466, 0.4578275, 0.40821073) +CLIP_STD = (0.26862954, 0.26130258, 0.27577711) +OPENAI_CLIP_BPE_URL = "https://openaipublic.azureedge.net/clip/bpe_simple_vocab_16e6.txt.gz" + + +MVTec_CLASSES = [ + "bottle", + "cable", + "capsule", + "carpet", + "grid", + "hazelnut", + "leather", + "metal_nut", + "pill", + "screw", + "tile", + "toothbrush", + "transistor", + "wood", + "zipper", +] + +VISA_CLASSES = [ + "candle", + "capsules", + "cashew", + "chewinggum", + "fryum", + "macaroni1", + "macaroni2", + "pcb1", + "pcb2", + "pcb3", + "pcb4", + "pipe_fryum", +] + + +def parse_args(): + parser = argparse.ArgumentParser( + description="Evaluate AdaCLIP text-reference anomaly maps without ResAD/Flow fusion." + ) + parser.add_argument("--dataset", type=str, default="mvtec", choices=["mvtec", "visa"]) + parser.add_argument("--data_root", type=str, required=True) + parser.add_argument("--class_name", type=str, default="all") + parser.add_argument("--num_ref_shot", type=int, default=4) + parser.add_argument("--ref_root", type=str, default="") + parser.add_argument("--ref_selection", type=str, default="first", choices=["first", "random"]) + parser.add_argument("--seed", type=int, default=42) + + parser.add_argument("--score_mode", type=str, default="text_ref", + choices=["text_ref", "cos_only", "residual_norm", "adaclip_text"]) + parser.add_argument("--image_score", type=str, default="topk", choices=["max", "topk", "mean"]) + parser.add_argument("--topk_ratio", type=float, default=0.01) + parser.add_argument("--score_norm", type=str, default="raw", choices=["raw", "image_minmax", "both"]) + parser.add_argument("--gaussian_sigma", type=float, default=0.0) + parser.add_argument("--chunk_size", type=int, default=8192) + + parser.add_argument("--clip_layer", type=int, default=24) + parser.add_argument("--clip_image_size", type=int, default=336) + parser.add_argument("--adaclip_repo_url", type=str, default="https://github.com/tomo082/AdaCLIP_res") + parser.add_argument("--adaclip_repo_path", type=str, default="") + parser.add_argument("--adaclip_checkpoint", type=str, default="") + parser.add_argument("--adaclip_checkpoint_url", type=str, default="") + parser.add_argument("--adaclip_cache_dir", type=str, default="~/.cache/adaclip_res") + parser.add_argument("--adaclip_model", type=str, default="ViT-L-14-336") + parser.add_argument("--adaclip_prompt_mode", type=str, default="hybrid", + choices=["hybrid", "static_only", "dynamic_only"]) + + parser.add_argument("--batch_size", type=int, default=1) + parser.add_argument("--num_workers", type=int, default=0) + parser.add_argument("--device", type=str, default="cuda:0") + parser.add_argument("--save_dir", type=str, required=True) + parser.add_argument("--save_visuals", action="store_true") + parser.add_argument("--max_visuals_per_class", type=int, default=25) + return parser.parse_args() + + +class ReferenceImageDataset(Dataset): + def __init__(self, image_paths, image_size): + self.image_paths = [Path(p) for p in image_paths] + self.transform = transforms.Compose([ + transforms.Resize((image_size, image_size), transforms.InterpolationMode.BICUBIC), + transforms.ToTensor(), + transforms.Normalize(IMAGENET_MEAN, IMAGENET_STD), + ]) + + def __len__(self): + return len(self.image_paths) + + def __getitem__(self, idx): + image = Image.open(self.image_paths[idx]).convert("RGB") + return self.transform(image), str(self.image_paths[idx]) + + +class AdaCLIPTextRefExtractor(torch.nn.Module): + """Small AdaCLIP runtime for projected patch tokens and text features. + + The AdaCLIP checkpoint and repository are only used for feature extraction. + ResAD constraintor, VQ, and flow modules are intentionally not constructed. + """ + + def __init__( + self, + repo_url, + repo_path, + checkpoint, + checkpoint_url, + cache_dir, + model_name, + layer, + image_size, + prompt_mode, + device, + ): + super().__init__() + self.repo_url = repo_url + self.repo_path = repo_path + self.checkpoint = checkpoint + self.checkpoint_url = checkpoint_url + self.cache_dir = Path(cache_dir).expanduser() + self.model_name = model_name + self.layer = layer + self.image_size = image_size + self.prompt_mode = prompt_mode + self.device = torch.device(device) + + self._clip_mean = torch.tensor(CLIP_MEAN).view(1, 3, 1, 1) + self._clip_std = torch.tensor(CLIP_STD).view(1, 3, 1, 1) + self._imagenet_mean = torch.tensor(IMAGENET_MEAN).view(1, 3, 1, 1) + self._imagenet_std = torch.tensor(IMAGENET_STD).view(1, 3, 1, 1) + + repo = self._resolve_repo_path() + ckpt = self._resolve_checkpoint_path() + self.trainer = self._build_trainer(repo, ckpt) + self.model = self.trainer.clip_model + self.model.eval() + for param in self.model.parameters(): + param.requires_grad = False + + def _resolve_repo_path(self): + if self.repo_path: + repo = Path(self.repo_path).expanduser().resolve() + if not repo.exists(): + raise FileNotFoundError(f"AdaCLIP repo path not found: {repo}") + return repo + + repo = self.cache_dir / "repos" / "AdaCLIP_res" + if repo.exists(): + return repo.resolve() + + repo.parent.mkdir(parents=True, exist_ok=True) + print(f"[AdaCLIP] cloning repo to {repo}") + subprocess.run(["git", "clone", self.repo_url, str(repo)], check=True) + return repo.resolve() + + def _resolve_checkpoint_path(self): + if self.checkpoint: + path = Path(self.checkpoint).expanduser().resolve() + if not path.exists(): + raise FileNotFoundError(f"AdaCLIP checkpoint not found: {path}") + return path + + path = self.cache_dir / "adaclip_checkpoint.pth" + if path.exists(): + return path.resolve() + if not self.checkpoint_url: + raise ValueError( + "Provide --adaclip_checkpoint or --adaclip_checkpoint_url for AdaCLIP weights." + ) + + path.parent.mkdir(parents=True, exist_ok=True) + print(f"[AdaCLIP] downloading checkpoint to {path}") + torch.hub.download_url_to_file(self.checkpoint_url, str(path), progress=True) + return path.resolve() + + def _build_trainer(self, repo_path, checkpoint_path): + self._ensure_tokenizer_assets(repo_path) + sys.path.insert(0, str(repo_path)) + try: + method_module = importlib.import_module("method") + trainer_cls = getattr(method_module, "AdaCLIP_Trainer") + except Exception as exc: + raise ImportError(f"Failed to import AdaCLIP_Trainer from {repo_path}") from exc + + config = self._load_model_config(repo_path) + trainer = trainer_cls( + backbone=self.model_name, + feat_list=[self.layer], + input_dim=config["vision_cfg"]["width"], + output_dim=config["embed_dim"], + learning_rate=0.0, + device=str(self.device), + image_size=self.image_size, + prompting_depth=4, + prompting_length=5, + prompting_branch="VL", + prompting_type="SD", + use_hsf=True, + k_clusters=20, + ) + trainer.load(str(checkpoint_path)) + trainer.clip_model.to(self.device) + trainer.clip_model.eval() + return trainer + + def _ensure_tokenizer_assets(self, repo_path): + bpe_path = Path(repo_path) / "method" / "bpe_simple_vocab_16e6.txt.gz" + if self._is_valid_gzip(bpe_path): + return + + bpe_path.parent.mkdir(parents=True, exist_ok=True) + tmp_path = bpe_path.with_suffix(bpe_path.suffix + ".download") + print(f"[AdaCLIP] tokenizer BPE missing or invalid; downloading to {bpe_path}") + try: + torch.hub.download_url_to_file(OPENAI_CLIP_BPE_URL, str(tmp_path), progress=True) + if not self._is_valid_gzip(tmp_path): + raise RuntimeError("downloaded BPE file is not a valid gzip archive") + os.replace(tmp_path, bpe_path) + except Exception as exc: + if tmp_path.exists(): + tmp_path.unlink() + raise RuntimeError( + "Failed to prepare AdaCLIP tokenizer BPE. If the cached repo contains a " + "Git LFS pointer, install git-lfs and run `git lfs pull`, or manually " + f"download {OPENAI_CLIP_BPE_URL} to {bpe_path}." + ) from exc + + @staticmethod + def _is_valid_gzip(path): + path = Path(path) + if not path.is_file(): + return False + try: + with path.open("rb") as handle: + return handle.read(2) == b"\x1f\x8b" + except OSError: + return False + + def _load_model_config(self, repo_path): + config_path = Path(repo_path) / "model_configs" / f"{self.model_name}.json" + if not config_path.is_file(): + raise FileNotFoundError(f"AdaCLIP model config does not exist: {config_path}") + with config_path.open("r", encoding="utf-8") as handle: + return json.load(handle) + + def _normalize_for_adaclip(self, images): + mean = self._imagenet_mean.to(images.device, images.dtype) + std = self._imagenet_std.to(images.device, images.dtype) + clip_mean = self._clip_mean.to(images.device, images.dtype) + clip_std = self._clip_std.to(images.device, images.dtype) + images = images * std + mean + return (images - clip_mean) / clip_std + + @staticmethod + def _prompt_type_for_mode(original_prompt_type, prompt_mode): + if prompt_mode == "hybrid": + return original_prompt_type + if prompt_mode == "static_only": + if "S" not in original_prompt_type: + raise ValueError("prompt_mode='static_only' requires static prompts in prompting_type.") + return "S" + if prompt_mode == "dynamic_only": + if "D" not in original_prompt_type: + raise ValueError("prompt_mode='dynamic_only' requires dynamic prompts in prompting_type.") + return "D" + raise ValueError(f"Unsupported prompt mode: {prompt_mode}") + + def _capture_prompt_state(self): + state = {} + for name in ("prompting_type",): + if hasattr(self.model, name): + state[name] = getattr(self.model, name) + for name in ("text_prompter", "visual_prompter"): + module = getattr(self.model, name, None) + if module is not None and hasattr(module, "prompting_type"): + state[f"{name}.prompting_type"] = module.prompting_type + return state + + def _set_prompt_type(self, prompt_type): + if hasattr(self.model, "prompting_type"): + self.model.prompting_type = prompt_type + for name in ("text_prompter", "visual_prompter"): + module = getattr(self.model, name, None) + if module is not None and hasattr(module, "prompting_type"): + module.prompting_type = prompt_type + + def _restore_prompt_state(self, state): + if "prompting_type" in state: + self.model.prompting_type = state["prompting_type"] + for name in ("text_prompter", "visual_prompter"): + key = f"{name}.prompting_type" + module = getattr(self.model, name, None) + if key in state and module is not None and hasattr(module, "prompting_type"): + module.prompting_type = state[key] + + @contextmanager + def _prompt_mode_context(self): + state = self._capture_prompt_state() + original = state.get("prompting_type", getattr(self.model, "prompting_type", "SD")) + prompt_type = self._prompt_type_for_mode(original, self.prompt_mode) + self._set_prompt_type(prompt_type) + try: + yield + finally: + self._restore_prompt_state(state) + + @torch.no_grad() + def extract(self, images, class_names): + images = self._normalize_for_adaclip(images.to(self.device)) + if isinstance(class_names, str): + class_names = [class_names] * images.shape[0] + else: + class_names = list(class_names) + + with self._prompt_mode_context(): + with torch.cuda.amp.autocast(enabled=images.is_cuda): + _, proj_patch_tokens, text_features = self.model.extract_feat(images, class_names) + + tokens = self._select_layer_tokens(proj_patch_tokens) + tokens = tokens.float() + text_features = text_features.float() + return tokens, text_features + + def _select_layer_tokens(self, proj_patch_tokens): + if isinstance(proj_patch_tokens, (list, tuple)): + if len(proj_patch_tokens) != 1: + return proj_patch_tokens[-1] + return proj_patch_tokens[0] + return proj_patch_tokens + + +def list_image_files(folder): + folder = Path(folder) + if not folder.is_dir(): + raise FileNotFoundError(f"Directory not found: {folder}") + return sorted(p for p in folder.iterdir() if p.suffix.lower() in IMAGE_EXTS) + + +def get_ref_image_paths(args, class_name): + root = Path(args.ref_root or args.data_root) + if args.dataset == "mvtec": + candidates = list_image_files(root / class_name / "train" / "good") + elif args.dataset == "visa": + candidates = get_visa_train_normal_paths(root, class_name) + else: + raise ValueError(f"Unsupported dataset: {args.dataset}") + + if len(candidates) < args.num_ref_shot: + raise ValueError( + f"{class_name}: requested {args.num_ref_shot} reference shots, found {len(candidates)}" + ) + if args.ref_selection == "random": + rng = random.Random(args.seed) + candidates = rng.sample(candidates, args.num_ref_shot) + else: + candidates = candidates[:args.num_ref_shot] + return candidates + + +def get_visa_train_normal_paths(root, class_name): + csv_path = Path(root) / "split_csv" / "1cls.csv" + if not csv_path.exists(): + # Few-shot folders are often exported in an MVTec-like layout. + fallback = Path(root) / class_name / "train" / "good" + return list_image_files(fallback) + + paths = [] + with csv_path.open("r", encoding="utf-8") as f: + reader = csv.DictReader(f) + for row in reader: + if row.get("object") != class_name: + continue + if row.get("split") != "train" or row.get("label") != "normal": + continue + image = row.get("image", "") + image = image[1:] if image.startswith("/") else image + paths.append(Path(root) / image) + return sorted(paths) + + +def get_classes(dataset, class_name): + classes = MVTec_CLASSES if dataset == "mvtec" else VISA_CLASSES + if class_name == "all": + return classes + requested = [c.strip() for c in class_name.split(",") if c.strip()] + unknown = [c for c in requested if c not in classes] + if unknown: + raise ValueError(f"Unknown class(es) for {dataset}: {unknown}") + return requested + + +def build_test_dataset(args, class_name): + kwargs = dict( + root=args.data_root, + class_name=class_name, + train=False, + normalize="w50", + img_size=args.clip_image_size, + crp_size=args.clip_image_size, + msk_size=args.clip_image_size, + msk_crp_size=args.clip_image_size, + ) + if args.dataset == "mvtec": + from datasets.mvtec import MVTEC + return MVTEC(**kwargs) + if args.dataset == "visa": + from datasets.visa import VISA + return VISA(**kwargs) + raise ValueError(f"Unsupported dataset: {args.dataset}") + + +def infer_patch_grid(tokens): + n = tokens.shape[1] + side = int(math.sqrt(n)) + if side * side == n: + return tokens, side, side + + side = int(math.sqrt(n - 1)) + if side * side == n - 1: + return tokens[:, 1:, :], side, side + + raise ValueError(f"Cannot infer square patch grid from token count {n}") + + +def extract_text_pair(text_features, batch_size): + if text_features.dim() == 2: + text_features = text_features.unsqueeze(0).expand(batch_size, -1, -1) + + if text_features.dim() != 3: + raise ValueError(f"Unsupported text_features shape: {tuple(text_features.shape)}") + + if text_features.shape[-1] == 2: + t_n = text_features[:, :, 0] + t_a = text_features[:, :, 1] + elif text_features.shape[1] == 2: + t_n = text_features[:, 0, :] + t_a = text_features[:, 1, :] + else: + raise ValueError(f"Cannot identify normal/abnormal text features: {tuple(text_features.shape)}") + + if t_n.shape[0] == 1 and batch_size > 1: + t_n = t_n.expand(batch_size, -1) + t_a = t_a.expand(batch_size, -1) + return t_n, t_a + + +def nearest_reference(query_tokens, memory_tokens, chunk_size): + query_n = F.normalize(query_tokens, p=2, dim=1) + memory_n = F.normalize(memory_tokens, p=2, dim=1) + matched = [] + for start in range(0, query_tokens.shape[0], chunk_size): + end = min(start + chunk_size, query_tokens.shape[0]) + sim = query_n[start:end] @ memory_n.T + idx = torch.argmax(sim, dim=1) + matched.append(memory_tokens[idx]) + return torch.cat(matched, dim=0) + + +def compute_patch_scores(tokens, text_features, memory_tokens, score_mode, chunk_size): + tokens, grid_h, grid_w = infer_patch_grid(tokens) + b, n, c = tokens.shape + flat_q = tokens.reshape(-1, c) + + t_n, t_a = extract_text_pair(text_features, b) + t_n = t_n.to(tokens.device, tokens.dtype) + t_a = t_a.to(tokens.device, tokens.dtype) + + if score_mode == "adaclip_text": + q_n = F.normalize(tokens, p=2, dim=-1) + tn = F.normalize(t_n, p=2, dim=1).unsqueeze(1) + ta = F.normalize(t_a, p=2, dim=1).unsqueeze(1) + logits = torch.stack([ + (q_n * tn).sum(dim=-1), + (q_n * ta).sum(dim=-1), + ], dim=-1) + scores = torch.softmax(logits, dim=-1)[..., 1] + return scores.reshape(b, grid_h, grid_w) + + matched = nearest_reference(flat_q, memory_tokens.to(tokens.device, tokens.dtype), chunk_size) + residual = flat_q - matched + residual_norm = torch.linalg.norm(residual, dim=1) + + rt = (t_a - t_n).repeat_interleave(n, dim=0) + cos = (F.normalize(residual, p=2, dim=1) * F.normalize(rt, p=2, dim=1)).sum(dim=1) + cos = torch.relu(cos) + + if score_mode == "text_ref": + scores = residual_norm * cos + elif score_mode == "cos_only": + scores = cos + elif score_mode == "residual_norm": + scores = residual_norm + else: + raise ValueError(f"Unsupported score_mode: {score_mode}") + + return scores.reshape(b, grid_h, grid_w) + + +def upsample_maps(patch_maps, target_hw): + maps = patch_maps.unsqueeze(1) + maps = F.interpolate(maps, size=target_hw, mode="bilinear", align_corners=False) + return maps[:, 0] + + +def normalize_imagewise(maps): + out = maps.copy() + for i in range(out.shape[0]): + mn = float(out[i].min()) + mx = float(out[i].max()) + if mx > mn: + out[i] = (out[i] - mn) / (mx - mn) + else: + out[i] = 0.0 + return out + + +def maybe_smooth(maps, sigma): + if sigma <= 0: + return maps + try: + from scipy.ndimage import gaussian_filter + except ImportError as exc: + raise ImportError("--gaussian_sigma requires scipy") from exc + return np.stack([gaussian_filter(m, sigma=sigma) for m in maps], axis=0) + + +def aggregate_image_scores(maps, mode, topk_ratio): + flat = maps.reshape(maps.shape[0], -1) + if mode == "max": + return flat.max(axis=1) + if mode == "mean": + return flat.mean(axis=1) + if mode == "topk": + k = max(1, int(flat.shape[1] * topk_ratio)) + idx = np.argpartition(flat, -k, axis=1)[:, -k:] + return np.take_along_axis(flat, idx, axis=1).mean(axis=1) + raise ValueError(f"Unsupported image_score: {mode}") + + +def safe_roc_auc(labels, scores): + from sklearn.metrics import roc_auc_score + if len(np.unique(labels)) < 2: + return float("nan") + return float(roc_auc_score(labels, scores)) + + +def safe_average_precision(labels, scores): + from sklearn.metrics import average_precision_score + if len(np.unique(labels)) < 2: + return float("nan") + return float(average_precision_score(labels, scores)) + + +def compute_aupro(gt_masks, scores): + if gt_masks.sum() <= 0 or float(scores.max()) <= float(scores.min()): + return float("nan") + try: + from utils import calculate_aupro + return float(calculate_aupro(gt_masks, scores)) + except Exception: + return float("nan") + + +def compute_metrics(labels, gt_masks, maps, image_score, topk_ratio): + labels = np.asarray(labels).astype(np.int64) + gt_masks = np.asarray(gt_masks).astype(np.uint8) + maps = np.asarray(maps).astype(np.float32) + image_scores = aggregate_image_scores(maps, image_score, topk_ratio) + + return { + "image_auc": safe_roc_auc(labels, image_scores), + "image_ap": safe_average_precision(labels, image_scores), + "pixel_auc": safe_roc_auc(gt_masks.reshape(-1), maps.reshape(-1)), + "pixel_ap": safe_average_precision(gt_masks.reshape(-1), maps.reshape(-1)), + "aupro": compute_aupro(gt_masks, maps), + } + + +def denorm_imagenet(image_tensor): + image = image_tensor.detach().cpu().float() + mean = torch.tensor(IMAGENET_MEAN).view(3, 1, 1) + std = torch.tensor(IMAGENET_STD).view(3, 1, 1) + image = image * std + mean + image = image.clamp(0, 1) + return (image.permute(1, 2, 0).numpy() * 255).astype(np.uint8) + + +def save_visual(path, image_tensor, mask, score_map, title): + import matplotlib.pyplot as plt + + image = denorm_imagenet(image_tensor) + fig, axes = plt.subplots(1, 3, figsize=(10, 3.4)) + axes[0].imshow(image) + axes[0].set_title("input") + axes[1].imshow(mask, cmap="gray") + axes[1].set_title("gt") + im = axes[2].imshow(score_map, cmap="jet") + axes[2].set_title(title) + for ax in axes: + ax.axis("off") + fig.colorbar(im, ax=axes[2], fraction=0.046, pad=0.04) + fig.tight_layout() + path.parent.mkdir(parents=True, exist_ok=True) + fig.savefig(path, dpi=160) + plt.close(fig) + + +def unpack_batch(batch): + images = batch[0] + labels = batch[1] + masks = batch[2] + return images, labels, masks + + +@torch.no_grad() +def build_memory_bank(args, extractor, class_name): + paths = get_ref_image_paths(args, class_name) + dataset = ReferenceImageDataset(paths, args.clip_image_size) + loader = DataLoader(dataset, batch_size=args.batch_size, shuffle=False, num_workers=args.num_workers) + chunks = [] + for images, _ in loader: + images = images.to(args.device) + class_names = [class_name] * images.shape[0] + tokens, _ = extractor.extract(images, class_names) + tokens, _, _ = infer_patch_grid(tokens) + chunks.append(tokens.reshape(-1, tokens.shape[-1]).cpu()) + memory = torch.cat(chunks, dim=0).float() + print(f"[Reference] {class_name}: {len(paths)} images, memory={tuple(memory.shape)}") + return memory + + +@torch.no_grad() +def evaluate_class(args, extractor, class_name): + save_dir = Path(args.save_dir) + class_dir = save_dir / class_name + maps_dir = class_dir / "maps" + maps_dir.mkdir(parents=True, exist_ok=True) + + memory = None + if args.score_mode != "adaclip_text": + memory = build_memory_bank(args, extractor, class_name) + dataset = build_test_dataset(args, class_name) + loader = DataLoader(dataset, batch_size=args.batch_size, shuffle=False, num_workers=args.num_workers) + + all_raw_maps = [] + all_labels = [] + all_masks = [] + all_paths = [] + visual_count = 0 + sample_idx = 0 + + for batch in loader: + images, labels, masks = unpack_batch(batch) + images = images.to(args.device) + masks = masks.float() + if masks.dim() == 4 and masks.shape[1] == 1: + masks = masks[:, 0] + class_names = [class_name] * images.shape[0] + + tokens, text_features = extractor.extract(images, class_names) + patch_maps = compute_patch_scores( + tokens=tokens, + text_features=text_features, + memory_tokens=memory, + score_mode=args.score_mode, + chunk_size=args.chunk_size, + ) + score_maps = upsample_maps(patch_maps, target_hw=masks.shape[-2:]).cpu().numpy() + + labels_np = labels.detach().cpu().numpy() + masks_np = masks.detach().cpu().numpy() + score_maps = maybe_smooth(score_maps.astype(np.float32), args.gaussian_sigma) + + for b in range(score_maps.shape[0]): + image_path = getattr(dataset, "image_paths", [None] * len(dataset))[sample_idx] + image_path = str(image_path) if image_path is not None else f"{class_name}_{sample_idx:06d}" + stem = Path(image_path).stem + np.save(maps_dir / f"{sample_idx:06d}_{stem}_raw.npy", score_maps[b]) + all_paths.append(image_path) + if args.save_visuals and visual_count < args.max_visuals_per_class: + save_visual( + class_dir / "visuals" / f"{sample_idx:06d}_{stem}_raw.png", + images[b], + masks_np[b], + score_maps[b], + args.score_mode, + ) + visual_count += 1 + sample_idx += 1 + + all_raw_maps.append(score_maps) + all_labels.append(labels_np) + all_masks.append(masks_np) + + raw_maps = np.concatenate(all_raw_maps, axis=0) + labels = np.concatenate(all_labels, axis=0) + masks = np.concatenate(all_masks, axis=0) + + norm_variants = ["raw", "image_minmax"] if args.score_norm == "both" else [args.score_norm] + rows = [] + for norm_name in norm_variants: + maps = raw_maps if norm_name == "raw" else normalize_imagewise(raw_maps) + if norm_name == "image_minmax": + for idx, path in enumerate(all_paths): + stem = Path(path).stem + np.save(maps_dir / f"{idx:06d}_{stem}_image_minmax.npy", maps[idx]) + metrics = compute_metrics(labels, masks, maps, args.image_score, args.topk_ratio) + row = { + "class_name": class_name, + "score_mode": args.score_mode, + "score_norm": norm_name, + "image_score": args.image_score, + "topk_ratio": args.topk_ratio, + **metrics, + } + rows.append(row) + print( + f"[Class {class_name} | {args.num_ref_shot}-shot | {args.score_mode} | {norm_name}] " + f"Image AUROC: {metrics['image_auc']:.4f} | AP: {metrics['image_ap']:.4f} | " + f"Pixel AUROC: {metrics['pixel_auc']:.4f} | AP: {metrics['pixel_ap']:.4f} | " + f"AUPRO: {metrics['aupro']:.4f}" + ) + + with (class_dir / "metrics.json").open("w", encoding="utf-8") as f: + json.dump(rows, f, indent=2, ensure_ascii=False) + return rows + + +def save_metrics_csv(path, rows): + path = Path(path) + path.parent.mkdir(parents=True, exist_ok=True) + fieldnames = [ + "class_name", + "score_mode", + "score_norm", + "image_score", + "topk_ratio", + "image_auc", + "image_ap", + "pixel_auc", + "pixel_ap", + "aupro", + ] + with path.open("w", newline="", encoding="utf-8") as f: + writer = csv.DictWriter(f, fieldnames=fieldnames) + writer.writeheader() + for row in rows: + writer.writerow(row) + + +def add_average_rows(rows): + average_rows = [] + keys = sorted(set((r["score_norm"], r["image_score"]) for r in rows)) + for score_norm, image_score in keys: + subset = [r for r in rows if r["score_norm"] == score_norm and r["image_score"] == image_score] + avg = { + "class_name": "average", + "score_mode": subset[0]["score_mode"], + "score_norm": score_norm, + "image_score": image_score, + "topk_ratio": subset[0]["topk_ratio"], + } + for metric in ("image_auc", "image_ap", "pixel_auc", "pixel_ap", "aupro"): + avg[metric] = float(np.nanmean([r[metric] for r in subset])) + average_rows.append(avg) + print( + f"[Average | {avg['score_mode']} | {score_norm}] " + f"Image AUROC: {avg['image_auc']:.4f} | AP: {avg['image_ap']:.4f} | " + f"Pixel AUROC: {avg['pixel_auc']:.4f} | AP: {avg['pixel_ap']:.4f} | " + f"AUPRO: {avg['aupro']:.4f}" + ) + return average_rows + + +def main(): + args = parse_args() + random.seed(args.seed) + np.random.seed(args.seed) + torch.manual_seed(args.seed) + + Path(args.save_dir).mkdir(parents=True, exist_ok=True) + classes = get_classes(args.dataset, args.class_name) + + print("[TextRefMap] dataset:", args.dataset) + print("[TextRefMap] classes:", classes) + print("[TextRefMap] score_mode:", args.score_mode) + print("[TextRefMap] image_score:", args.image_score) + print("[TextRefMap] score_norm:", args.score_norm) + print("[AdaCLIP] layer:", args.clip_layer) + print("[AdaCLIP] prompt_mode:", args.adaclip_prompt_mode) + + extractor = AdaCLIPTextRefExtractor( + repo_url=args.adaclip_repo_url, + repo_path=args.adaclip_repo_path, + checkpoint=args.adaclip_checkpoint, + checkpoint_url=args.adaclip_checkpoint_url, + cache_dir=args.adaclip_cache_dir, + model_name=args.adaclip_model, + layer=args.clip_layer, + image_size=args.clip_image_size, + prompt_mode=args.adaclip_prompt_mode, + device=args.device, + ) + + all_rows = [] + for class_name in classes: + all_rows.extend(evaluate_class(args, extractor, class_name)) + + all_rows_with_avg = all_rows + add_average_rows(all_rows) + save_metrics_csv(Path(args.save_dir) / "metrics.csv", all_rows_with_avg) + with (Path(args.save_dir) / "metrics.json").open("w", encoding="utf-8") as f: + json.dump(all_rows_with_avg, f, indent=2, ensure_ascii=False) + print("[TextRefMap] saved:", Path(args.save_dir).resolve()) + + +if __name__ == "__main__": + main() diff --git a/utils.py b/utils.py index e465752..d272a4f 100644 --- a/utils.py +++ b/utils.py @@ -54,7 +54,37 @@ def get_matched_ref_features(features: List[Tensor], ref_features: List[Tensor]) matched_ref_features.append(index_feats) return matched_ref_features +def get_matched_ref_features_top(features: List[Tensor], ref_features: List[Tensor], rank: int = 0) -> List[Tensor]: + """ + Get matched reference features for one class. + Args: + rank (int): 取得する類似度の順位。0で最も似ているもの、1で2番目...を指定。 + """ + matched_ref_features = [] + for layer_id in range(len(features)): + feature = features[layer_id] + B, C, H, W = feature.shape + feature = feature.permute(0, 2, 3, 1).reshape(-1, C).contiguous() # (N1, C) + feature_n = F.normalize(feature, p=2, dim=1) + coreset = ref_features[layer_id] # (N2, C) + coreset_n = F.normalize(coreset, p=2, dim=1) + dist = feature_n @ coreset_n.T + # --- 変更箇所: rankに応じてインデックスを取得 --- + if rank == 0: + cidx = torch.argmax(dist, dim=1) + else: + # 上位 (rank + 1) 個を取得し、その最後の要素(指定された順位のもの)を取得 + # kがメモリバンクサイズを超えないように注意が必要ですが、通常は十分大きいためこのまま実装します + _, topk_indices = torch.topk(dist, k=rank + 1, dim=1) + cidx = topk_indices[:, -1] + # ---------------------------------------------- + + index_feats = coreset[cidx] + index_feats = index_feats.reshape(B, H, W, C).permute(0, 3, 1, 2) + matched_ref_features.append(index_feats) + + return matched_ref_features def get_residual_features(features: List[Tensor], ref_features: List[Tensor], pos_flag: bool = False) -> List[Tensor]: residual_features = [] @@ -69,7 +99,91 @@ def get_residual_features(features: List[Tensor], ref_features: List[Tensor], po residual_features.append(ri) return residual_features +import torch + +def get_image_level_matched_features(features, ref_features): + matched_ref_features = [] + for layer_id in range(len(features)): + feature = features[layer_id] + B, C, H, W = feature.shape + + coreset = ref_features[layer_id] + K = coreset.shape[0] // (H * W) + coreset_spatial = coreset.view(K, H, W, C).permute(0, 3, 1, 2).contiguous() + + feat_flat = feature.reshape(B, 1, -1) + core_flat = coreset_spatial.reshape(1, K, -1) + + feat_norm = F.normalize(feat_flat, p=2, dim=2) + core_norm = F.normalize(core_flat, p=2, dim=2) + sim = torch.sum(feat_norm * core_norm, dim=2) + + best_idx = torch.argmax(sim, dim=1) + matched = coreset_spatial[best_idx] + matched_ref_features.append(matched) + + return matched_ref_features + +def get_mc_image_level_matched_features(features, class_names, ref_features): + matched_ref_features = [[] for _ in range(len(features))] + for idx, c in enumerate(class_names): + ref_features_c = ref_features[c] + + for layer_id in range(len(features)): + feature = features[layer_id][idx:idx+1] + _, C, H, W = feature.shape + + coreset = ref_features_c[layer_id] + K = coreset.shape[0] // (H * W) + coreset_spatial = coreset.view(K, H, W, C).permute(0, 3, 1, 2).contiguous() + + feat_flat = feature.reshape(1, 1, -1) + core_flat = coreset_spatial.reshape(1, K, -1) + + feat_norm = F.normalize(feat_flat, p=2, dim=2) + core_norm = F.normalize(core_flat, p=2, dim=2) + sim = torch.sum(feat_norm * core_norm, dim=2) + + best_idx = torch.argmax(sim, dim=1) + matched = coreset_spatial[best_idx].squeeze(0) + matched_ref_features[layer_id].append(matched) + + matched_ref_features = [torch.stack(item, dim=0) for item in matched_ref_features] + return matched_ref_features +def get_fourier_residual_features(features, mfeatures, pos_flag=True): + """ + 周波数領域で残差を計算する関数 + features: テスト画像の特徴量リスト [B, C, H, W] + mfeatures: マッチングされた参照画像の特徴量リスト [B, C, H, W] + """ + rfeatures = [] + for i in range(len(features)): + f, mf = features[i], mfeatures[i] + + # 1. 2D FFTで空間から周波数領域へ変換 + f_fft = torch.fft.fft2(f, norm="ortho") + mf_fft = torch.fft.fft2(mf, norm="ortho") + # 2. 振幅 (Amplitude) と 位相 (Phase) に分離 + f_amp, f_pha = torch.abs(f_fft), torch.angle(f_fft) + mf_amp, mf_pha = torch.abs(mf_fft), torch.angle(mf_fft) + + # 3. 振幅の残差を計算 (ここがキズに反応する) + # 位相の残差は空間構造の歪みを表すが、位置ズレに寛容にするためテスト画像の位相を保持する + res_amp = torch.abs(f_amp - mf_amp) + + # 4. 残差振幅とテスト画像の位相を結合して複素数に戻す + res_fft_new = res_amp * torch.exp(1j * f_pha) + + # 5. 逆フーリエ変換 (IFFT) で空間領域の残差マップに戻す + res_f = torch.fft.ifft2(res_fft_new, norm="ortho").real + #元のMSEに合わせるために2乗する + if pos_flag: + res_f = torch.pow(res_f, 2) # または torch.abs(res_f) + + rfeatures.append(res_f) + + return rfeatures def load_reference_features(root_dir: str, class_name: str, device: torch.device) -> List[Tensor]: """ @@ -119,7 +233,29 @@ def get_mc_reference_features(encoder, root, class_names, device, num_shot=4): features[l] = features[l].permute(0, 2, 3, 1).reshape(-1, c) reference_features[class_name] = features return reference_features - +def get_mc_reference_features_wav(encoder, root, class_names, device, num_shot=4, wav_filter=None): + """ + Get reference features for multiple classes. + """ + reference_features = {} + class_names = np.unique(class_names) + for class_name in class_names: + normal_paths = get_random_normal_images(root, class_name, num_shot) + images = load_and_transform_vision_data(normal_paths, device) + with torch.no_grad(): + features = encoder(images) + + # ======== 【追加】 ======== + # 平坦化(reshape)される前に空間情報(H, W)を保ったままウェーブレット変換を適用 + if wav_filter is not None: + features = [wav_filter(f) for f in features] + # ======================== + + for l in range(len(features)): + bs, c, h, w = features[l].shape + features[l] = features[l].permute(0, 2, 3, 1).reshape(-1, c) + reference_features[class_name] = features + return reference_features def load_and_transform_vision_data(image_paths, device): if image_paths is None: @@ -175,6 +311,22 @@ def calculate_metrics(scores, labels, gt_masks, pro=True, only_max_value=True): labels (np.ndarray): shape (N, ), 0 for normal, 1 for abnormal. gt_masks (np.ndarray): shape (N, H, W). """ + #pixel AUC error analysis + from scipy.stats import rankdata + if False: + scores_flat = np.maximum(scores.reshape(-1,1),0) + ranks_flat = rankdata(scores_flat) + scores = scores_flat.reshape(N,H,W) + ranks = ranks_flat.reshape(N,H,W) + mask_rank = (gt_masks*ranks).sum(dim=(1,2)) + unmask_rank = ((1-gt_masks)*ranks).sum(dim=(1,2)) + fn_mask=(rankdata(rankdata(mask_rank))[labels==1])[-10:] + fp_unmask = (rankdata(rankdata(unmask_rank))[labels==0])[:10] + + score_max = scores.max(dim=(1,2)) + false_negatives=(rankdata(rankdata(score_max))[labels==1])[-10:] + false_positives=(rankdata(rankdata(score_max))[labels==0])[:10] + # average precision pix_ap = round(average_precision_score(gt_masks.flatten(), scores.flatten()), 5) # f1 score, f1 score is to balance the precision and recall @@ -203,12 +355,8 @@ def calculate_metrics(scores, labels, gt_masks, pro=True, only_max_value=True): img_aps.append(img_ap) img_aucs.append(img_auc) img_f1_scores.append(img_f1_score) - img_ap, img_auc, img_f1_score = np.max(img_aps), np.max(img_aucs), np.max(img_f1_scores) - - if pro: - pix_aupro = calculate_aupro(gt_masks, scores) - else: - pix_aupro = -1 + img_ap, img_auc, img_f1_score = np.max(img_aps), np.max(img_aucs), np.max(img_f1_scores) + pix_aupro = calculate_aupro(gt_masks, scores) return img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro @@ -275,4 +423,213 @@ def applying_EFDM(input_features_list, ref_features_list, alpha=0.5): aligned_features = aligned_features.view(B, C, W, H) aligned_features_list.append(aligned_features) - return aligned_features_list \ No newline at end of file + return aligned_features_list + ''' +def load_weights(encoder, decoders, filename):#12/16追加 + #path = os.path.join(WEIGHT_DIR, filename) + state = torch.load(filename) + encoder.load_state_dict(state['encoder_state_dict'], strict=False) + decoders = [decoder.load_state_dict(state, strict=False) for decoder, state in zip(decoders, state['decoder_state_dict'])] #12/16変更 + print('Loading weights from {}'.format(filename)) + ''' +def load_weights(encoder, decoders, filename): + state = torch.load(filename, map_location='cpu') + print(f'Loading weights from {filename}') + + # --- Encoder のロード --- + if 'encoder_state_dict' in state: + enc_msg = encoder.load_state_dict(state['encoder_state_dict'], strict=False) + all_keys = len(encoder.state_dict().keys()) + loaded_keys = all_keys - len(enc_msg.missing_keys) + print(f"--- Encoder Load Report ---") + print(f" Successfully loaded: {loaded_keys}/{all_keys}") + else: + print(" [ERROR] 'encoder_state_dict' not found.") + + # --- Decoders のロード --- + if 'decoder_state_dict' in state: + print(f"--- Decoders Detailed Report ---") + for i, (decoder, d_state) in enumerate(zip(decoders, state['decoder_state_dict'])): + dec_msg = decoder.load_state_dict(d_state, strict=False) + + all_keys = set(decoder.state_dict().keys()) + missing_keys = set(dec_msg.missing_keys) + loaded_keys_count = len(all_keys) - len(missing_keys) + + print(f" Decoder {i}: Loaded {loaded_keys_count}/{len(all_keys)} parameters.") + # 修正箇所: インデントを削除しました + for key in sorted(list(all_keys)): + print(f" - {key}") + + if len(missing_keys) > 0: + print(f" [Missing parameters in Decoder {i}]") + # 読み込めなかった変数名をすべて書き出す + for key in sorted(list(missing_keys)): + print(f" - {key}") + else: + print(f" All parameters loaded successfully for Decoder {i}.") + else: + print(" [ERROR] 'decoder_state_dict' not found.") + +def load_weights_ada(adapter, filename): + #path = os.path.join(WEIGHT_DIR, filename) + state = torch.load(filename) + #encoder.load_state_dict(state['encoder_state_dict'], strict=False) + #decoders = [decoder.load_state_dict(state, strict=False) for decoder, state in zip(decoders, state['decoder_state_dict'])] + adapters = [adapters.load_state_dict(state, strict=False) for adapter, state in zip(adapters, state['adapter_state_dict'])]#modified 1/8 + print('Loading weights from {}'.format(filename)) +def get_soft_matched_features(features: List[Tensor], ref_features: List[Tensor], tau=0.05) -> List[Tensor]: + """ + 評価用: Soft Attentionを用いた滑らかな特徴マッチング + tau (Temperature): 値が小さいほどArgmaxに近づき(鮮明)、大きいほど平均に近づく(滑らか) + """ + matched_ref_features = [] + for layer_id in range(len(features)): + feature = features[layer_id] + B, C, H, W = feature.shape + + # [B*H*W, C] に変形 + feature_flat = feature.permute(0, 2, 3, 1).reshape(-1, C).contiguous() + feature_n = F.normalize(feature_flat, p=2, dim=1) + + coreset = ref_features[layer_id] # [K*H*W, C] (Kはショット数) + coreset_n = F.normalize(coreset, p=2, dim=1) + + # 1. すべての参照ピクセルとの類似度を計算 + sim = feature_n @ coreset_n.T # [B*H*W, K*H*W] + + # 2. Temperature付きSoftmaxで重み(Attention)を計算 + attn = F.softmax(sim / tau, dim=1) # [B*H*W, K*H*W] + + # 3. 参照ピクセルを重み付きでブレンド + index_feats = attn @ coreset # [B*H*W, C] + + # 4. 元の画像形状に戻す + index_feats = index_feats.reshape(B, H, W, C).permute(0, 3, 1, 2).contiguous() + matched_ref_features.append(index_feats) + + return matched_ref_features + +def get_mc_soft_matched_features(features: List[Tensor], class_names: List[str], ref_features: Dict[str, List[Tensor]], tau=0.05) -> List[Tensor]: + """ + 学習用: マルチクラス対応のSoft Attentionマッチング + """ + matched_ref_features = [[] for _ in range(len(features))] + for idx, c in enumerate(class_names): + ref_features_c = ref_features[c] + + for layer_id in range(len(features)): + feature = features[layer_id][idx:idx+1] + _, C, H, W = feature.shape + + feature_flat = feature.permute(0, 2, 3, 1).reshape(-1, C).contiguous() + feature_n = F.normalize(feature_flat, p=2, dim=1) + + coreset = ref_features_c[layer_id] + coreset_n = F.normalize(coreset, p=2, dim=1) + + sim = feature_n @ coreset_n.T + attn = F.softmax(sim / tau, dim=1) + + index_feats = attn @ coreset + index_feats = index_feats.reshape(1, H, W, C).permute(0, 3, 1, 2).contiguous() + matched_ref_features[layer_id].append(index_feats) + + matched_ref_features = [torch.cat(item, dim=0) for item in matched_ref_features] + return matched_ref_features +def compute_osp_matrices_from_refs(ref_features_tuple, keep_variance=0.95): + """ + 参照特徴量 (正常データ) から、レイヤーごとの射影行列と平均ベクトルを計算する。 + ref_features_tuple: (layer1_refs, layer2_refs, layer3_refs) などのタプル + 各 tensor は [N, C] または [N, C, H, W] などの形状 + """ + proj_matrices = [] + means = [] + + for ref_feat in ref_features_tuple: + # 形状を [Batch*H*W, Channels] の2次元に平坦化する + if ref_feat.dim() == 4: + B, C, H, W = ref_feat.shape + flat_ref = ref_feat.permute(0, 2, 3, 1).reshape(-1, C) + elif ref_feat.dim() == 3: + B, L, C = ref_feat.shape + flat_ref = ref_feat.reshape(-1, C) + else: + flat_ref = ref_feat + C = flat_ref.shape[-1] + + mean = flat_ref.mean(dim=0, keepdim=True) + centered = flat_ref - mean + + # 特異値分解 (SVD) を計算 + U, S, V = torch.linalg.svd(centered.cpu(), full_matrices=False) + + # 寄与率から上位 k 次元を決定 + var = (S ** 2) / (centered.size(0) - 1) + cum_var = torch.cumsum(var, dim=0) / var.sum() + k = torch.searchsorted(cum_var, keep_variance).item() + 1 + + basis = V[:k, :].T.to(ref_feat.device) # [C, k] + proj_matrix = torch.mm(basis, basis.T) # [C, C] + + proj_matrices.append(proj_matrix) + means.append(mean.to(ref_feat.device)) + + return proj_matrices, means + +def apply_osp(rfeatures_list, proj_matrices, means): + """ + 抽出された残差リストに対して、正常空間成分を削り落とす。 + """ + osp_residuals = [] + for i, rfeat in enumerate(rfeatures_list): + B, C, H, W = rfeat.shape + r_flat = rfeat.permute(0, 2, 3, 1).reshape(-1, C) + + # 平行成分(正常なズレ)を計算 + r_parallel = torch.mm(r_flat - means[i], proj_matrices[i]) + + # 直交成分(純粋な異常)を残す + r_orthogonal = (r_flat - means[i]) - r_parallel + + # 元の形状に戻す + r_orthogonal = r_orthogonal.reshape(B, H, W, C).permute(0, 3, 1, 2) + osp_residuals.append(r_orthogonal) + + return osp_residuals +import torch +import torch.nn.functional as F +import torch.nn as nn +class HaarWaveletFilter(nn.Module): + def __init__(self, low_freq_weight=0.1, high_freq_weight=1.2): + super().__init__() + self.lf_w = low_freq_weight + self.hf_w = high_freq_weight + + ll = torch.tensor([[0.5, 0.5], [0.5, 0.5]]) + hl = torch.tensor([[-0.5, -0.5], [0.5, 0.5]]) + lh = torch.tensor([[-0.5, 0.5], [-0.5, 0.5]]) + hh = torch.tensor([[0.5, -0.5], [-0.5, 0.5]]) + + self.register_buffer('k_ll', ll.view(1, 1, 2, 2)) + self.register_buffer('k_hl', hl.view(1, 1, 2, 2)) + self.register_buffer('k_lh', lh.view(1, 1, 2, 2)) + self.register_buffer('k_hh', hh.view(1, 1, 2, 2)) + + def forward(self, x): + B, C, H, W = x.shape + ll = F.conv2d(x, self.k_ll.expand(C, 1, 2, 2), stride=2, groups=C) + hl = F.conv2d(x, self.k_hl.expand(C, 1, 2, 2), stride=2, groups=C) + lh = F.conv2d(x, self.k_lh.expand(C, 1, 2, 2), stride=2, groups=C) + hh = F.conv2d(x, self.k_hh.expand(C, 1, 2, 2), stride=2, groups=C) + + ll = ll * self.lf_w + hl = hl * self.hf_w + lh = lh * self.hf_w + hh = hh * self.hf_w + + out = F.conv_transpose2d(ll, self.k_ll.expand(C, 1, 2, 2), stride=2, groups=C) + \ + F.conv_transpose2d(hl, self.k_hl.expand(C, 1, 2, 2), stride=2, groups=C) + \ + F.conv_transpose2d(lh, self.k_lh.expand(C, 1, 2, 2), stride=2, groups=C) + \ + F.conv_transpose2d(hh, self.k_hh.expand(C, 1, 2, 2), stride=2, groups=C) + return out diff --git a/validate.py b/validate.py index 667fd6f..ed3600a 100644 --- a/validate.py +++ b/validate.py @@ -7,7 +7,7 @@ from models.modules import get_position_encoding from models.utils import get_logp -from utils import get_residual_features, get_matched_ref_features +from utils import get_residual_features, get_matched_ref_features,get_fourier_residual_features from utils import calculate_metrics, applying_EFDM from losses.utils import get_logp_a @@ -28,7 +28,7 @@ def validate(args, encoder, vq_ops, constraintor, estimators, test_loader, ref_f for idx, batch in enumerate(test_loader): progress_bar.update(1) - image, label, mask, _ = batch + image, label, mask = batch[0], batch[1], batch[2] gt_mask_list.append(mask.squeeze(1).cpu().numpy().astype(bool)) label_list.append(label.cpu().numpy().astype(bool).ravel()) @@ -40,6 +40,14 @@ def validate(args, encoder, vq_ops, constraintor, estimators, test_loader, ref_f features = encoder(image) mfeatures = get_matched_ref_features(features, ref_features) rfeatures = get_residual_features(features, mfeatures, pos_flag=True) + elif args.backbone == 'tf_efficientnet_b6':#10/26追加 + features = encoder(image) + mfeatures = get_matched_ref_features(features, ref_features) + rfeatures = get_residual_features(features, mfeatures, pos_flag=True) + elif args.backbone == 'vit_base_patch14': + features = encoder(image) + mfeatures = get_matched_ref_features(features, ref_features) + rfeatures = get_residual_features(features, mfeatures, pos_flag=True) else: features = encoder.encode_image_from_tensors(image) for i in range(len(features)): @@ -139,4 +147,4 @@ def aggregate_anomaly_scores(logps_list, feature_levels=3, class_name=None, size for i in range(scores.shape[0]): scores[i] = gaussian_filter(scores[i], sigma=4) - return scores \ No newline at end of file + return scores diff --git a/validate1.py b/validate1.py new file mode 100644 index 0000000..270653e --- /dev/null +++ b/validate1.py @@ -0,0 +1,148 @@ +#最も近い参照特徴量以外にも対応するため + +import warnings +from tqdm import tqdm +from scipy.ndimage import gaussian_filter +import numpy as np +import torch +import torch.nn.functional as F + +from models.modules import get_position_encoding +from models.utils import get_logp +from utils import get_residual_features, get_matched_ref_features,get_matched_ref_features_top +from utils import calculate_metrics, applying_EFDM +from losses.utils import get_logp_a + +warnings.filterwarnings('ignore') + + +def validate1(args, encoder, vq_ops, constraintor, estimators, test_loader, ref_features, device, class_name): + vq_ops.eval() + constraintor.eval() + for estimator in estimators: + estimator.eval() + + label_list, gt_mask_list = [], [] + logps1_list = [list() for _ in range(args.feature_levels)] + logps2_list = [list() for _ in range(args.feature_levels)] + progress_bar = tqdm(total=len(test_loader)) + progress_bar.set_description(f"Evaluating") + for idx, batch in enumerate(test_loader): + progress_bar.update(1) + + image, label, mask= batch[:3] + gt_mask_list.append(mask.squeeze(1).cpu().numpy().astype(bool)) + label_list.append(label.cpu().numpy().astype(bool).ravel()) + + image = image.to(device) + size = image.shape[-1] + + with torch.no_grad(): + if args.backbone == 'wide_resnet50_2': + features = encoder(image) + mfeatures = get_matched_ref_features(features, ref_features) + rfeatures = get_residual_features(features, mfeatures, pos_flag=True) + elif args.backbone == 'tf_efficientnet_b6':#10/26追加 + features = encoder(image) + mfeatures = get_matched_ref_features_top(features, ref_features,args.rank) + rfeatures = get_residual_features(features, mfeatures, pos_flag=True) + else: + features = encoder.encode_image_from_tensors(image) + for i in range(len(features)): + b, l, c = features[i].shape + features[i] = features[i].permute(0, 2, 1).reshape(b, c, 16, 16) + mfeatures = get_matched_ref_features(features, ref_features) + rfeatures = get_residual_features(features, mfeatures) + + fdm_features = vq_ops(rfeatures, train=False) + rfeatures = applying_EFDM(rfeatures, fdm_features, alpha=args.fdm_alpha) + rfeatures = constraintor(*rfeatures) + + for l in range(args.feature_levels): + e = rfeatures[l] # BxCxHxW + bs, dim, h, w = e.size() + e = e.permute(0, 2, 3, 1).reshape(-1, dim) + + # (bs, 128, h, w) + pos_embed = get_position_encoding(args.pos_embed_dim, h, w).to(args.device).unsqueeze(0).repeat(bs, 1, 1, 1) + pos_embed = pos_embed.permute(0, 2, 3, 1).reshape(-1, args.pos_embed_dim) + estimator = estimators[l] + + if args.flow_arch == 'flow_model': + z, log_jac_det = estimator(e) + else: + z, log_jac_det = estimator(e, [pos_embed, ]) + + logps = get_logp(dim, z, log_jac_det) + logps = logps / dim + logps1_list[l].append(logps.reshape(bs, h, w)) + + logps_a = get_logp_a(dim, z, log_jac_det) # logps corresponding to abnormal distribution + logits = torch.stack([logps, logps_a], dim=-1) # (N, 2) + sa = torch.softmax(logits, dim=-1)[:, 1] + logps2_list[l].append(sa.reshape(bs, h, w)) + + progress_bar.close() + + labels = np.concatenate(label_list) + gt_masks = np.concatenate(gt_mask_list, axis=0) + scores1 = convert_to_anomaly_scores(logps1_list, feature_levels=args.feature_levels, class_name=class_name, size=size) + scores2 = aggregate_anomaly_scores(logps2_list, feature_levels=args.feature_levels, class_name=class_name, size=size) + + img_auc1, img_ap1, img_f1_score1, pix_auc1, pix_ap1, pix_f1_score1, pix_aupro1 = calculate_metrics(scores1, labels, gt_masks, pro=False, only_max_value=True) + img_auc2, img_ap2, img_f1_score2, pix_auc2, pix_ap2, pix_f1_score2, pix_aupro2 = calculate_metrics(scores2, labels, gt_masks, pro=False, only_max_value=True) + + scores = (scores1 + scores2) / 2 + img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro = calculate_metrics(scores, labels, gt_masks, pro=False, only_max_value=True) + + metrics = {} + metrics['scores1'] = [img_auc1, img_ap1, img_f1_score1, pix_auc1, pix_ap1, pix_f1_score1, pix_aupro1] + metrics['scores2'] = [img_auc2, img_ap2, img_f1_score2, pix_auc2, pix_ap2, pix_f1_score2, pix_aupro2] + metrics['scores'] = [img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro] + + return metrics + + +def convert_to_anomaly_scores(logps_list, feature_levels=3, class_name=None, size=224): + normal_map = [list() for _ in range(feature_levels)] + for l in range(feature_levels): + logps = torch.cat(logps_list[l], dim=0) + logps-= torch.max(logps) # normalize log-likelihoods to (-Inf:0] by subtracting a constant + probs = torch.exp(logps) # convert to probs in range [0:1] + # upsample + normal_map[l] = F.interpolate(probs.unsqueeze(1), + size=size, mode='bilinear', align_corners=True).squeeze().cpu().numpy() + + # score aggregation + scores = np.zeros_like(normal_map[0]) + for l in range(feature_levels): + scores += normal_map[l] + + # normality score to anomaly score + scores = scores.max() - scores + + #if class_name in ['pill', 'cable', 'capsule', 'screw']: + for i in range(scores.shape[0]): + scores[i] = gaussian_filter(scores[i], sigma=4) + + return scores + + +def aggregate_anomaly_scores(logps_list, feature_levels=3, class_name=None, size=224): + abnormal_map = [list() for _ in range(feature_levels)] + for l in range(feature_levels): + probs = torch.cat(logps_list[l], dim=0) + # upsample + abnormal_map[l] = F.interpolate(probs.unsqueeze(1), + size=size, mode='bilinear', align_corners=True).squeeze().cpu().numpy() + + # score aggregation + scores = np.zeros_like(abnormal_map[0]) + for l in range(feature_levels): + scores += abnormal_map[l] + scores /= feature_levels + + for i in range(scores.shape[0]): + scores[i] = gaussian_filter(scores[i], sigma=4) + + return scores diff --git a/validate_Fourier.py b/validate_Fourier.py new file mode 100644 index 0000000..de2ebaa --- /dev/null +++ b/validate_Fourier.py @@ -0,0 +1,154 @@ +import warnings +from tqdm import tqdm +from scipy.ndimage import gaussian_filter +import numpy as np +import torch +import torch.nn.functional as F + +from models.modules import get_position_encoding +from models.utils import get_logp +from utils import get_residual_features, get_matched_ref_features,get_fourier_residual_features +from utils import get_residual_features, get_image_level_matched_features, get_fourier_residual_features +from utils import calculate_metrics, applying_EFDM +from losses.utils import get_logp_a + +warnings.filterwarnings('ignore') + + +def validate(args, encoder, vq_ops, constraintor, estimators, test_loader, ref_features, device, class_name): + vq_ops.eval() + constraintor.eval() + for estimator in estimators: + estimator.eval() + + label_list, gt_mask_list = [], [] + logps1_list = [list() for _ in range(args.feature_levels)] + logps2_list = [list() for _ in range(args.feature_levels)] + progress_bar = tqdm(total=len(test_loader)) + progress_bar.set_description(f"Evaluating") + for idx, batch in enumerate(test_loader): + progress_bar.update(1) + + image, label, mask = batch[0], batch[1], batch[2] + gt_mask_list.append(mask.squeeze(1).cpu().numpy().astype(bool)) + label_list.append(label.cpu().numpy().astype(bool).ravel()) + + image = image.to(device) + size = image.shape[-1] + + with torch.no_grad(): + if args.backbone == 'wide_resnet50_2': + features = encoder(image) + #mfeatures = get_matched_ref_features(features, ref_features) + mfeatures = get_image_level_matched_features(features, ref_features) + rfeatures = get_fourier_residual_features(features, mfeatures, pos_flag=True) + elif args.backbone == 'tf_efficientnet_b6':#10/26追加 + features = encoder(image) + #mfeatures = get_matched_ref_features(features, ref_features) + mfeatures = get_image_level_matched_features(features, ref_features) + rfeatures = get_fourier_residual_features(features, mfeatures, pos_flag=True) + elif args.backbone == 'vit_base_patch14': + features = encoder(image) + #mfeatures = get_matched_ref_features(features, ref_features) + mfeatures = get_image_level_matched_features(features, ref_features) + rfeatures = get_fourier_residual_features(features, mfeatures, pos_flag=True) + else: + features = encoder.encode_image_from_tensors(image) + for i in range(len(features)): + b, l, c = features[i].shape + features[i] = features[i].permute(0, 2, 1).reshape(b, c, 16, 16) + mfeatures = get_matched_ref_features(features, ref_features) + rfeatures = get_residual_features(features, mfeatures) + + fdm_features = vq_ops(rfeatures, train=False) + rfeatures = applying_EFDM(rfeatures, fdm_features, alpha=args.fdm_alpha) + rfeatures = constraintor(*rfeatures) + + for l in range(args.feature_levels): + e = rfeatures[l] # BxCxHxW + bs, dim, h, w = e.size() + e = e.permute(0, 2, 3, 1).reshape(-1, dim) + + # (bs, 128, h, w) + pos_embed = get_position_encoding(args.pos_embed_dim, h, w).to(args.device).unsqueeze(0).repeat(bs, 1, 1, 1) + pos_embed = pos_embed.permute(0, 2, 3, 1).reshape(-1, args.pos_embed_dim) + estimator = estimators[l] + + if args.flow_arch == 'flow_model': + z, log_jac_det = estimator(e) + else: + z, log_jac_det = estimator(e, [pos_embed, ]) + + logps = get_logp(dim, z, log_jac_det) + logps = logps / dim + logps1_list[l].append(logps.reshape(bs, h, w)) + + logps_a = get_logp_a(dim, z, log_jac_det) # logps corresponding to abnormal distribution + logits = torch.stack([logps, logps_a], dim=-1) # (N, 2) + sa = torch.softmax(logits, dim=-1)[:, 1] + logps2_list[l].append(sa.reshape(bs, h, w)) + + progress_bar.close() + + labels = np.concatenate(label_list) + gt_masks = np.concatenate(gt_mask_list, axis=0) + scores1 = convert_to_anomaly_scores(logps1_list, feature_levels=args.feature_levels, class_name=class_name, size=size) + scores2 = aggregate_anomaly_scores(logps2_list, feature_levels=args.feature_levels, class_name=class_name, size=size) + + img_auc1, img_ap1, img_f1_score1, pix_auc1, pix_ap1, pix_f1_score1, pix_aupro1 = calculate_metrics(scores1, labels, gt_masks, pro=False, only_max_value=True) + img_auc2, img_ap2, img_f1_score2, pix_auc2, pix_ap2, pix_f1_score2, pix_aupro2 = calculate_metrics(scores2, labels, gt_masks, pro=False, only_max_value=True) + + scores = (scores1 + scores2) / 2 + img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro = calculate_metrics(scores, labels, gt_masks, pro=False, only_max_value=True) + + metrics = {} + metrics['scores1'] = [img_auc1, img_ap1, img_f1_score1, pix_auc1, pix_ap1, pix_f1_score1, pix_aupro1] + metrics['scores2'] = [img_auc2, img_ap2, img_f1_score2, pix_auc2, pix_ap2, pix_f1_score2, pix_aupro2] + metrics['scores'] = [img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro] + + return metrics + + +def convert_to_anomaly_scores(logps_list, feature_levels=3, class_name=None, size=224): + normal_map = [list() for _ in range(feature_levels)] + for l in range(feature_levels): + logps = torch.cat(logps_list[l], dim=0) + logps-= torch.max(logps) # normalize log-likelihoods to (-Inf:0] by subtracting a constant + probs = torch.exp(logps) # convert to probs in range [0:1] + # upsample + normal_map[l] = F.interpolate(probs.unsqueeze(1), + size=size, mode='bilinear', align_corners=True).squeeze().cpu().numpy() + + # score aggregation + scores = np.zeros_like(normal_map[0]) + for l in range(feature_levels): + scores += normal_map[l] + + # normality score to anomaly score + scores = scores.max() - scores + + #if class_name in ['pill', 'cable', 'capsule', 'screw']: + for i in range(scores.shape[0]): + scores[i] = gaussian_filter(scores[i], sigma=4) + + return scores + + +def aggregate_anomaly_scores(logps_list, feature_levels=3, class_name=None, size=224): + abnormal_map = [list() for _ in range(feature_levels)] + for l in range(feature_levels): + probs = torch.cat(logps_list[l], dim=0) + # upsample + abnormal_map[l] = F.interpolate(probs.unsqueeze(1), + size=size, mode='bilinear', align_corners=True).squeeze().cpu().numpy() + + # score aggregation + scores = np.zeros_like(abnormal_map[0]) + for l in range(feature_levels): + scores += abnormal_map[l] + scores /= feature_levels + + for i in range(scores.shape[0]): + scores[i] = gaussian_filter(scores[i], sigma=4) + + return scores diff --git a/validate_attention.py b/validate_attention.py new file mode 100644 index 0000000..8b43a6c --- /dev/null +++ b/validate_attention.py @@ -0,0 +1,151 @@ +import warnings +from tqdm import tqdm +from scipy.ndimage import gaussian_filter +import numpy as np +import torch +import torch.nn.functional as F + +from models.modules import get_position_encoding +from models.utils import get_logp +from utils import get_residual_features, get_matched_ref_features,get_fourier_residual_features +from utils import get_residual_features, get_soft_matched_features +from utils import calculate_metrics, applying_EFDM +from losses.utils import get_logp_a + +warnings.filterwarnings('ignore') + + +def validate(args, encoder, vq_ops, constraintor, estimators, test_loader, ref_features, device, class_name): + vq_ops.eval() + constraintor.eval() + for estimator in estimators: + estimator.eval() + + label_list, gt_mask_list = [], [] + logps1_list = [list() for _ in range(args.feature_levels)] + logps2_list = [list() for _ in range(args.feature_levels)] + progress_bar = tqdm(total=len(test_loader)) + progress_bar.set_description(f"Evaluating") + for idx, batch in enumerate(test_loader): + progress_bar.update(1) + + image, label, mask = batch[0], batch[1], batch[2] + gt_mask_list.append(mask.squeeze(1).cpu().numpy().astype(bool)) + label_list.append(label.cpu().numpy().astype(bool).ravel()) + + image = image.to(device) + size = image.shape[-1] + + with torch.no_grad(): + if args.backbone == 'wide_resnet50_2': + features = encoder(image) + mfeatures = get_soft_matched_features(features, ref_features) + rfeatures = get_residual_features(features, mfeatures, pos_flag=True) + elif args.backbone == 'tf_efficientnet_b6':#10/26追加 + features = encoder(image) + mfeatures = get_soft_matched_features(features, ref_features) + rfeatures = get_residual_features(features, mfeatures, pos_flag=True) + elif args.backbone == 'vit_base_patch14': + features = encoder(image) + mfeatures = get_soft_matched_features(features, ref_features) + rfeatures = get_residual_features(features, mfeatures, pos_flag=True) + else: + features = encoder.encode_image_from_tensors(image) + for i in range(len(features)): + b, l, c = features[i].shape + features[i] = features[i].permute(0, 2, 1).reshape(b, c, 16, 16) + mfeatures = get_matched_ref_features(features, ref_features) + rfeatures = get_residual_features(features, mfeatures) + + fdm_features = vq_ops(rfeatures, train=False) + rfeatures = applying_EFDM(rfeatures, fdm_features, alpha=args.fdm_alpha) + rfeatures = constraintor(*rfeatures) + + for l in range(args.feature_levels): + e = rfeatures[l] # BxCxHxW + bs, dim, h, w = e.size() + e = e.permute(0, 2, 3, 1).reshape(-1, dim) + + # (bs, 128, h, w) + pos_embed = get_position_encoding(args.pos_embed_dim, h, w).to(args.device).unsqueeze(0).repeat(bs, 1, 1, 1) + pos_embed = pos_embed.permute(0, 2, 3, 1).reshape(-1, args.pos_embed_dim) + estimator = estimators[l] + + if args.flow_arch == 'flow_model': + z, log_jac_det = estimator(e) + else: + z, log_jac_det = estimator(e, [pos_embed, ]) + + logps = get_logp(dim, z, log_jac_det) + logps = logps / dim + logps1_list[l].append(logps.reshape(bs, h, w)) + + logps_a = get_logp_a(dim, z, log_jac_det) # logps corresponding to abnormal distribution + logits = torch.stack([logps, logps_a], dim=-1) # (N, 2) + sa = torch.softmax(logits, dim=-1)[:, 1] + logps2_list[l].append(sa.reshape(bs, h, w)) + + progress_bar.close() + + labels = np.concatenate(label_list) + gt_masks = np.concatenate(gt_mask_list, axis=0) + scores1 = convert_to_anomaly_scores(logps1_list, feature_levels=args.feature_levels, class_name=class_name, size=size) + scores2 = aggregate_anomaly_scores(logps2_list, feature_levels=args.feature_levels, class_name=class_name, size=size) + + img_auc1, img_ap1, img_f1_score1, pix_auc1, pix_ap1, pix_f1_score1, pix_aupro1 = calculate_metrics(scores1, labels, gt_masks, pro=False, only_max_value=True) + img_auc2, img_ap2, img_f1_score2, pix_auc2, pix_ap2, pix_f1_score2, pix_aupro2 = calculate_metrics(scores2, labels, gt_masks, pro=False, only_max_value=True) + + scores = (scores1 + scores2) / 2 + img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro = calculate_metrics(scores, labels, gt_masks, pro=False, only_max_value=True) + + metrics = {} + metrics['scores1'] = [img_auc1, img_ap1, img_f1_score1, pix_auc1, pix_ap1, pix_f1_score1, pix_aupro1] + metrics['scores2'] = [img_auc2, img_ap2, img_f1_score2, pix_auc2, pix_ap2, pix_f1_score2, pix_aupro2] + metrics['scores'] = [img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro] + + return metrics + + +def convert_to_anomaly_scores(logps_list, feature_levels=3, class_name=None, size=224): + normal_map = [list() for _ in range(feature_levels)] + for l in range(feature_levels): + logps = torch.cat(logps_list[l], dim=0) + logps-= torch.max(logps) # normalize log-likelihoods to (-Inf:0] by subtracting a constant + probs = torch.exp(logps) # convert to probs in range [0:1] + # upsample + normal_map[l] = F.interpolate(probs.unsqueeze(1), + size=size, mode='bilinear', align_corners=True).squeeze().cpu().numpy() + + # score aggregation + scores = np.zeros_like(normal_map[0]) + for l in range(feature_levels): + scores += normal_map[l] + + # normality score to anomaly score + scores = scores.max() - scores + + #if class_name in ['pill', 'cable', 'capsule', 'screw']: + for i in range(scores.shape[0]): + scores[i] = gaussian_filter(scores[i], sigma=4) + + return scores + + +def aggregate_anomaly_scores(logps_list, feature_levels=3, class_name=None, size=224): + abnormal_map = [list() for _ in range(feature_levels)] + for l in range(feature_levels): + probs = torch.cat(logps_list[l], dim=0) + # upsample + abnormal_map[l] = F.interpolate(probs.unsqueeze(1), + size=size, mode='bilinear', align_corners=True).squeeze().cpu().numpy() + + # score aggregation + scores = np.zeros_like(abnormal_map[0]) + for l in range(feature_levels): + scores += abnormal_map[l] + scores /= feature_levels + + for i in range(scores.shape[0]): + scores[i] = gaussian_filter(scores[i], sigma=4) + + return scores diff --git a/validate_freq_blend.py b/validate_freq_blend.py new file mode 100644 index 0000000..80948e4 --- /dev/null +++ b/validate_freq_blend.py @@ -0,0 +1,180 @@ +import warnings +from tqdm import tqdm +from scipy.ndimage import gaussian_filter +import numpy as np +import torch +import torch.nn.functional as F + +from models.modules import get_position_encoding +from models.utils import get_logp +from utils import calculate_metrics +from losses.utils import get_logp_a + +warnings.filterwarnings('ignore') + +# 推論時専用の分割統治マッチング関数 +def get_freq_matched_residuals_infer(test_lf_list, test_hf_list, ref_lf_list, ref_hf_list, alpha=0.5, pos_flag=True, device='cuda:0'): + rfeatures = [] + for l in range(len(test_lf_list)): + t_lf = test_lf_list[l] + t_hf = test_hf_list[l] + r_lf_all = ref_lf_list[l] + r_hf_all = ref_hf_list[l] + + B, C, H, W = t_lf.shape + t_lf_flat = t_lf.permute(0, 2, 3, 1).reshape(-1, C).contiguous() + + t_lf_n = F.normalize(t_lf_flat, p=2, dim=1) + r_lf_n = F.normalize(r_lf_all, p=2, dim=1) + + dist = t_lf_n @ r_lf_n.T + cidx = torch.argmax(dist, dim=1) + + m_lf = r_lf_all[cidx].reshape(B, H, W, C).permute(0, 3, 1, 2) + m_hf = r_hf_all[cidx].reshape(B, H, W, C).permute(0, 3, 1, 2) + + res_lf = t_lf - m_lf + res_hf = t_hf - m_hf + + rfeature = alpha * res_lf + (1.0 - alpha) * res_hf + + if pos_flag: + pos_embed = get_position_encoding(C, H, W).to(device).unsqueeze(0).repeat(B, 1, 1, 1) + rfeature = rfeature + pos_embed + + rfeatures.append(rfeature) + return rfeatures + + +def validate(args, encoder, constraintor, wav_filter, estimators, test_loader, ref_features, device, class_name): + constraintor.eval() + for estimator in estimators: + estimator.eval() + + ref_lf_cat = [] + ref_hf_cat = [] + + with torch.no_grad(): + dummy_img = torch.zeros(1, 3, 224, 224).to(device) + dummy_feats = encoder(dummy_img) + + for l in range(args.feature_levels): + _, C, H, W = dummy_feats[l].shape + + K = ref_features[l].shape[0] // (H * W) + ref_4d = ref_features[l].view(K, H, W, C).permute(0, 3, 1, 2) + + ref_lf, ref_hf = wav_filter.get_LF_HF(ref_4d) + + ref_lf_cat.append(ref_lf.permute(0, 2, 3, 1).reshape(-1, C)) + ref_hf_cat.append(ref_hf.permute(0, 2, 3, 1).reshape(-1, C)) + + label_list, gt_mask_list = [], [] + logps1_list = [list() for _ in range(args.feature_levels)] + logps2_list = [list() for _ in range(args.feature_levels)] + + progress_bar = tqdm(total=len(test_loader)) + progress_bar.set_description(f"Evaluating {class_name}") + + for idx, batch in enumerate(test_loader): + progress_bar.update(1) + + image, label, mask = batch[0], batch[1], batch[2] + gt_mask_list.append(mask.squeeze(1).cpu().numpy().astype(bool)) + label_list.append(label.cpu().numpy().astype(bool).ravel()) + + image = image.to(device) + size = image.shape[-1] + + with torch.no_grad(): + features_raw = encoder(image) + test_lf_list, test_hf_list = [], [] + + for l in range(args.feature_levels): + lf, hf = wav_filter.get_LF_HF(features_raw[l]) + test_lf_list.append(lf) + test_hf_list.append(hf) + + rfeatures = get_freq_matched_residuals_infer( + test_lf_list, test_hf_list, ref_lf_cat, ref_hf_cat, + alpha=args.blend_alpha, pos_flag=True, device=device + ) + + rfeatures = constraintor(*rfeatures) + + for l in range(args.feature_levels): + e = rfeatures[l] + bs, dim, h, w = e.size() + e = e.permute(0, 2, 3, 1).reshape(-1, dim) + + # ★ 修正ポイント: estimatorの条件付けには一律で args.pos_embed_dim を使用する + pos_embed = get_position_encoding(args.pos_embed_dim, h, w).to(device).unsqueeze(0).repeat(bs, 1, 1, 1) + pos_embed = pos_embed.permute(0, 2, 3, 1).reshape(-1, args.pos_embed_dim) + estimator = estimators[l] + + if args.flow_arch == 'flow_model': + z, log_jac_det = estimator(e) + else: + z, log_jac_det = estimator(e, [pos_embed, ]) + + logps = get_logp(dim, z, log_jac_det) + logps = logps / dim + logps1_list[l].append(logps.reshape(bs, h, w)) + + logps_a = get_logp_a(dim, z, log_jac_det) + logits = torch.stack([logps, logps_a], dim=-1) + sa = torch.softmax(logits, dim=-1)[:, 1] + logps2_list[l].append(sa.reshape(bs, h, w)) + + progress_bar.close() + + labels = np.concatenate(label_list) + gt_masks = np.concatenate(gt_mask_list, axis=0) + + scores1 = convert_to_anomaly_scores(logps1_list, feature_levels=args.feature_levels, size=size) + scores2 = aggregate_anomaly_scores(logps2_list, feature_levels=args.feature_levels, size=size) + + img_auc1, img_ap1, img_f1_score1, pix_auc1, pix_ap1, pix_f1_score1, pix_aupro1 = calculate_metrics(scores1, labels, gt_masks, pro=False, only_max_value=True) + img_auc2, img_ap2, img_f1_score2, pix_auc2, pix_ap2, pix_f1_score2, pix_aupro2 = calculate_metrics(scores2, labels, gt_masks, pro=False, only_max_value=True) + + scores = (scores1 + scores2) / 2 + img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro = calculate_metrics(scores, labels, gt_masks, pro=False, only_max_value=True) + + metrics = {} + metrics['scores1'] = [img_auc1, img_ap1, img_f1_score1, pix_auc1, pix_ap1, pix_f1_score1, pix_aupro1] + metrics['scores2'] = [img_auc2, img_ap2, img_f1_score2, pix_auc2, pix_ap2, pix_f1_score2, pix_aupro2] + metrics['scores'] = [img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro] + + return metrics + +def convert_to_anomaly_scores(logps_list, feature_levels=3, size=224): + normal_map = [list() for _ in range(feature_levels)] + for l in range(feature_levels): + logps = torch.cat(logps_list[l], dim=0) + logps -= torch.max(logps) + probs = torch.exp(logps) + normal_map[l] = F.interpolate(probs.unsqueeze(1), size=size, mode='bilinear', align_corners=True).squeeze().cpu().numpy() + + scores = np.zeros_like(normal_map[0]) + for l in range(feature_levels): + scores += normal_map[l] + scores = scores.max() - scores + + for i in range(scores.shape[0]): + scores[i] = gaussian_filter(scores[i], sigma=4) + return scores + +def aggregate_anomaly_scores(logps_list, feature_levels=3, size=224): + abnormal_map = [list() for _ in range(feature_levels)] + for l in range(feature_levels): + probs = torch.cat(logps_list[l], dim=0) + abnormal_map[l] = F.interpolate(probs.unsqueeze(1), size=size, mode='bilinear', align_corners=True).squeeze().cpu().numpy() + + scores = np.zeros_like(abnormal_map[0]) + for l in range(feature_levels): + scores += abnormal_map[l] + scores /= feature_levels + + for i in range(scores.shape[0]): + scores[i] = gaussian_filter(scores[i], sigma=4) + return scores diff --git a/validate_global.py b/validate_global.py new file mode 100644 index 0000000..916c15c --- /dev/null +++ b/validate_global.py @@ -0,0 +1,137 @@ +import warnings +from tqdm import tqdm +from scipy.ndimage import gaussian_filter +import numpy as np +import torch +import torch.nn.functional as F + +from models.modules import get_position_encoding +from models.utils import get_logp +from utils import get_residual_features +from utils import calculate_metrics +from losses.utils import get_logp_a + +warnings.filterwarnings('ignore') + +def get_matched_ref_features(features, ref_features): + matched_ref_features = [] + for layer_id in range(len(features)): + feature = features[layer_id] + B, C, H, W = feature.shape + feature = feature.permute(0, 2, 3, 1).reshape(-1, C).contiguous() + feature_n = F.normalize(feature, p=2, dim=1) + coreset = ref_features[layer_id] + coreset_n = F.normalize(coreset, p=2, dim=1) + dist = feature_n @ coreset_n.T + cidx = torch.argmax(dist, dim=1) + index_feats = coreset[cidx] + index_feats = index_feats.reshape(B, H, W, C).permute(0, 3, 1, 2) + matched_ref_features.append(index_feats) + return matched_ref_features + +def validate(args, encoder, constraintor, estimators, test_loader, ref_features, device, class_name): + constraintor.eval() + for estimator in estimators: + estimator.eval() + + label_list, gt_mask_list = [], [] + logps1_list = [list() for _ in range(args.feature_levels)] + logps2_list = [list() for _ in range(args.feature_levels)] + + progress_bar = tqdm(total=len(test_loader)) + progress_bar.set_description(f"Evaluating {class_name}") + + for idx, batch in enumerate(test_loader): + progress_bar.update(1) + + image, label, mask = batch[0], batch[1], batch[2] + gt_mask_list.append(mask.squeeze(1).cpu().numpy().astype(bool)) + label_list.append(label.cpu().numpy().astype(bool).ravel()) + + image = image.to(device) + size = image.shape[-1] + + with torch.no_grad(): + features = encoder(image) + + mfeatures = get_matched_ref_features(features, ref_features) + + rfeatures = get_residual_features(features, mfeatures, pos_flag=True) + + rfeatures = constraintor(*rfeatures) + + for l in range(args.feature_levels): + e = rfeatures[l] + bs, dim, h, w = e.size() + e = e.permute(0, 2, 3, 1).reshape(-1, dim) + + pos_embed = get_position_encoding(args.pos_embed_dim, h, w).to(args.device).unsqueeze(0).repeat(bs, 1, 1, 1) + pos_embed = pos_embed.permute(0, 2, 3, 1).reshape(-1, args.pos_embed_dim) + estimator = estimators[l] + + if args.flow_arch == 'flow_model': + z, log_jac_det = estimator(e) + else: + z, log_jac_det = estimator(e, [pos_embed, ]) + + logps = get_logp(dim, z, log_jac_det) + logps = logps / dim + logps1_list[l].append(logps.reshape(bs, h, w)) + + logps_a = get_logp_a(dim, z, log_jac_det) + logits = torch.stack([logps, logps_a], dim=-1) + sa = torch.softmax(logits, dim=-1)[:, 1] + logps2_list[l].append(sa.reshape(bs, h, w)) + + progress_bar.close() + + labels = np.concatenate(label_list) + gt_masks = np.concatenate(gt_mask_list, axis=0) + + scores1 = convert_to_anomaly_scores(logps1_list, feature_levels=args.feature_levels, size=size) + scores2 = aggregate_anomaly_scores(logps2_list, feature_levels=args.feature_levels, size=size) + + img_auc1, img_ap1, img_f1_score1, pix_auc1, pix_ap1, pix_f1_score1, pix_aupro1 = calculate_metrics(scores1, labels, gt_masks, pro=False, only_max_value=True) + img_auc2, img_ap2, img_f1_score2, pix_auc2, pix_ap2, pix_f1_score2, pix_aupro2 = calculate_metrics(scores2, labels, gt_masks, pro=False, only_max_value=True) + + scores = (scores1 + scores2) / 2 + img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro = calculate_metrics(scores, labels, gt_masks, pro=False, only_max_value=True) + + metrics = {} + metrics['scores1'] = [img_auc1, img_ap1, img_f1_score1, pix_auc1, pix_ap1, pix_f1_score1, pix_aupro1] + metrics['scores2'] = [img_auc2, img_ap2, img_f1_score2, pix_auc2, pix_ap2, pix_f1_score2, pix_aupro2] + metrics['scores'] = [img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro] + + return metrics + +def convert_to_anomaly_scores(logps_list, feature_levels=3, size=224): + normal_map = [list() for _ in range(feature_levels)] + for l in range(feature_levels): + logps = torch.cat(logps_list[l], dim=0) + logps -= torch.max(logps) + probs = torch.exp(logps) + normal_map[l] = F.interpolate(probs.unsqueeze(1), size=size, mode='bilinear', align_corners=True).squeeze().cpu().numpy() + + scores = np.zeros_like(normal_map[0]) + for l in range(feature_levels): + scores += normal_map[l] + scores = scores.max() - scores + + for i in range(scores.shape[0]): + scores[i] = gaussian_filter(scores[i], sigma=4) + return scores + +def aggregate_anomaly_scores(logps_list, feature_levels=3, size=224): + abnormal_map = [list() for _ in range(feature_levels)] + for l in range(feature_levels): + probs = torch.cat(logps_list[l], dim=0) + abnormal_map[l] = F.interpolate(probs.unsqueeze(1), size=size, mode='bilinear', align_corners=True).squeeze().cpu().numpy() + + scores = np.zeros_like(abnormal_map[0]) + for l in range(feature_levels): + scores += abnormal_map[l] + scores /= feature_levels + + for i in range(scores.shape[0]): + scores[i] = gaussian_filter(scores[i], sigma=4) + return scores diff --git a/validate_osp.py b/validate_osp.py new file mode 100644 index 0000000..f110192 --- /dev/null +++ b/validate_osp.py @@ -0,0 +1,149 @@ +import warnings +from tqdm import tqdm +from scipy.ndimage import gaussian_filter +import numpy as np +import torch +import torch.nn.functional as F + +from models.modules import get_position_encoding +from models.utils import get_logp +from utils import get_residual_features, get_matched_ref_features, get_fourier_residual_features +from utils import calculate_metrics +from losses.utils import get_logp_a + +warnings.filterwarnings('ignore') + +# ========================================== +# OSP 適用関数 (main_osp.pyと同じもの) +# ========================================== +def apply_osp(residuals_list, proj_matrices, means): + osp_results = [] + for i, res in enumerate(residuals_list): + B, C, H, W = res.shape + res_flat = res.permute(0, 2, 3, 1).reshape(-1, C) + res_centered = res_flat - means[i] + res_parallel = torch.mm(res_centered, proj_matrices[i]) + res_ortho = res_centered - res_parallel + osp_results.append(res_ortho.reshape(B, H, W, C).permute(0, 3, 1, 2)) + return osp_results +# ========================================== + +def validate(args, encoder, constraintor, estimators, test_loader, ref_features, device, class_name, osp_proj_matrices, osp_means): + constraintor.eval() + for estimator in estimators: + estimator.eval() + + label_list, gt_mask_list = [], [] + logps1_list = [list() for _ in range(args.feature_levels)] + logps2_list = [list() for _ in range(args.feature_levels)] + progress_bar = tqdm(total=len(test_loader)) + progress_bar.set_description(f"Evaluating") + + for idx, batch in enumerate(test_loader): + progress_bar.update(1) + + image, label, mask = batch[0], batch[1], batch[2] + gt_mask_list.append(mask.squeeze(1).cpu().numpy().astype(bool)) + label_list.append(label.cpu().numpy().astype(bool).ravel()) + + image = image.to(device) + size = image.shape[-1] + + with torch.no_grad(): + if args.backbone in ['wide_resnet50_2', 'tf_efficientnet_b6', 'vit_base_patch14']: + features = encoder(image) + mfeatures = get_matched_ref_features(features, ref_features) + rfeatures = get_residual_features(features, mfeatures, pos_flag=True) + else: + features = encoder.encode_image_from_tensors(image) + for i in range(len(features)): + b, l, c = features[i].shape + features[i] = features[i].permute(0, 2, 1).reshape(b, c, 16, 16) + mfeatures = get_matched_ref_features(features, ref_features) + rfeatures = get_residual_features(features, mfeatures) + + # --- VQとEFDMの代わりにOSPを適用 --- + rfeatures = apply_osp(rfeatures, osp_proj_matrices, osp_means) + + # --- constraintorで滑らかに補正 --- + rfeatures = constraintor(*rfeatures) + + for l in range(args.feature_levels): + e = rfeatures[l] + bs, dim, h, w = e.size() + e = e.permute(0, 2, 3, 1).reshape(-1, dim) + + pos_embed = get_position_encoding(args.pos_embed_dim, h, w).to(args.device).unsqueeze(0).repeat(bs, 1, 1, 1) + pos_embed = pos_embed.permute(0, 2, 3, 1).reshape(-1, args.pos_embed_dim) + estimator = estimators[l] + + if args.flow_arch == 'flow_model': + z, log_jac_det = estimator(e) + else: + z, log_jac_det = estimator(e, [pos_embed, ]) + + logps = get_logp(dim, z, log_jac_det) + logps = logps / dim + logps1_list[l].append(logps.reshape(bs, h, w)) + + logps_a = get_logp_a(dim, z, log_jac_det) + logits = torch.stack([logps, logps_a], dim=-1) + sa = torch.softmax(logits, dim=-1)[:, 1] + logps2_list[l].append(sa.reshape(bs, h, w)) + + progress_bar.close() + + labels = np.concatenate(label_list) + gt_masks = np.concatenate(gt_mask_list, axis=0) + scores1 = convert_to_anomaly_scores(logps1_list, feature_levels=args.feature_levels, class_name=class_name, size=size) + scores2 = aggregate_anomaly_scores(logps2_list, feature_levels=args.feature_levels, class_name=class_name, size=size) + + img_auc1, img_ap1, img_f1_score1, pix_auc1, pix_ap1, pix_f1_score1, pix_aupro1 = calculate_metrics(scores1, labels, gt_masks, pro=False, only_max_value=True) + img_auc2, img_ap2, img_f1_score2, pix_auc2, pix_ap2, pix_f1_score2, pix_aupro2 = calculate_metrics(scores2, labels, gt_masks, pro=False, only_max_value=True) + + scores = (scores1 + scores2) / 2 + img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro = calculate_metrics(scores, labels, gt_masks, pro=False, only_max_value=True) + + metrics = {} + metrics['scores1'] = [img_auc1, img_ap1, img_f1_score1, pix_auc1, pix_ap1, pix_f1_score1, pix_aupro1] + metrics['scores2'] = [img_auc2, img_ap2, img_f1_score2, pix_auc2, pix_ap2, pix_f1_score2, pix_aupro2] + metrics['scores'] = [img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro] + + return metrics + + +def convert_to_anomaly_scores(logps_list, feature_levels=3, class_name=None, size=224): + normal_map = [list() for _ in range(feature_levels)] + for l in range(feature_levels): + logps = torch.cat(logps_list[l], dim=0) + logps-= torch.max(logps) + probs = torch.exp(logps) + normal_map[l] = F.interpolate(probs.unsqueeze(1), size=size, mode='bilinear', align_corners=True).squeeze().cpu().numpy() + + scores = np.zeros_like(normal_map[0]) + for l in range(feature_levels): + scores += normal_map[l] + + scores = scores.max() - scores + + for i in range(scores.shape[0]): + scores[i] = gaussian_filter(scores[i], sigma=4) + + return scores + + +def aggregate_anomaly_scores(logps_list, feature_levels=3, class_name=None, size=224): + abnormal_map = [list() for _ in range(feature_levels)] + for l in range(feature_levels): + probs = torch.cat(logps_list[l], dim=0) + abnormal_map[l] = F.interpolate(probs.unsqueeze(1), size=size, mode='bilinear', align_corners=True).squeeze().cpu().numpy() + + scores = np.zeros_like(abnormal_map[0]) + for l in range(feature_levels): + scores += abnormal_map[l] + scores /= feature_levels + + for i in range(scores.shape[0]): + scores[i] = gaussian_filter(scores[i], sigma=4) + + return scores diff --git a/validate_wav.py b/validate_wav.py new file mode 100644 index 0000000..2b732e0 --- /dev/null +++ b/validate_wav.py @@ -0,0 +1,126 @@ +import warnings +from tqdm import tqdm +from scipy.ndimage import gaussian_filter +import numpy as np +import torch +import torch.nn.functional as F + +from models.modules import get_position_encoding +from models.utils import get_logp +from utils import get_residual_features, get_matched_ref_features +from utils import calculate_metrics +from losses.utils import get_logp_a + +warnings.filterwarnings('ignore') + +def validate(args, encoder, constraintor, wav_filter, estimators, test_loader, ref_features, device, class_name): + constraintor.eval() + wav_filter.eval() + for estimator in estimators: + estimator.eval() + + label_list, gt_mask_list = [], [] + logps1_list = [list() for _ in range(args.feature_levels)] + logps2_list = [list() for _ in range(args.feature_levels)] + + progress_bar = tqdm(total=len(test_loader)) + progress_bar.set_description(f"Evaluating {class_name}") + + for idx, batch in enumerate(test_loader): + progress_bar.update(1) + + image, label, mask = batch[0], batch[1], batch[2] + gt_mask_list.append(mask.squeeze(1).cpu().numpy().astype(bool)) + label_list.append(label.cpu().numpy().astype(bool).ravel()) + + image = image.to(device) + size = image.shape[-1] + + with torch.no_grad(): + features = encoder(image) + + # --- 追加: テスト画像をウェーブレット変換 (Pre-filter) --- + features = [wav_filter(f) for f in features] + + mfeatures = get_matched_ref_features(features, ref_features) + rfeatures = get_residual_features(features, mfeatures, pos_flag=True) + + rfeatures = constraintor(*rfeatures) + + for l in range(args.feature_levels): + e = rfeatures[l] + bs, dim, h, w = e.size() + e = e.permute(0, 2, 3, 1).reshape(-1, dim) + + pos_embed = get_position_encoding(args.pos_embed_dim, h, w).to(args.device).unsqueeze(0).repeat(bs, 1, 1, 1) + pos_embed = pos_embed.permute(0, 2, 3, 1).reshape(-1, args.pos_embed_dim) + estimator = estimators[l] + + if args.flow_arch == 'flow_model': + z, log_jac_det = estimator(e) + else: + z, log_jac_det = estimator(e, [pos_embed, ]) + + logps = get_logp(dim, z, log_jac_det) + logps = logps / dim + logps1_list[l].append(logps.reshape(bs, h, w)) + + logps_a = get_logp_a(dim, z, log_jac_det) + logits = torch.stack([logps, logps_a], dim=-1) + sa = torch.softmax(logits, dim=-1)[:, 1] + logps2_list[l].append(sa.reshape(bs, h, w)) + + progress_bar.close() + + labels = np.concatenate(label_list) + gt_masks = np.concatenate(gt_mask_list, axis=0) + + scores1 = convert_to_anomaly_scores(logps1_list, feature_levels=args.feature_levels, size=size) + scores2 = aggregate_anomaly_scores(logps2_list, feature_levels=args.feature_levels, size=size) + + img_auc1, img_ap1, img_f1_score1, pix_auc1, pix_ap1, pix_f1_score1, pix_aupro1 = calculate_metrics(scores1, labels, gt_masks, pro=False, only_max_value=True) + img_auc2, img_ap2, img_f1_score2, pix_auc2, pix_ap2, pix_f1_score2, pix_aupro2 = calculate_metrics(scores2, labels, gt_masks, pro=False, only_max_value=True) + + scores = (scores1 + scores2) / 2 + img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro = calculate_metrics(scores, labels, gt_masks, pro=False, only_max_value=True) + + metrics = {} + metrics['scores1'] = [img_auc1, img_ap1, img_f1_score1, pix_auc1, pix_ap1, pix_f1_score1, pix_aupro1] + metrics['scores2'] = [img_auc2, img_ap2, img_f1_score2, pix_auc2, pix_ap2, pix_f1_score2, pix_aupro2] + metrics['scores'] = [img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro] + + return metrics + + +def convert_to_anomaly_scores(logps_list, feature_levels=3, size=224): + normal_map = [list() for _ in range(feature_levels)] + for l in range(feature_levels): + logps = torch.cat(logps_list[l], dim=0) + logps -= torch.max(logps) + probs = torch.exp(logps) + normal_map[l] = F.interpolate(probs.unsqueeze(1), size=size, mode='bilinear', align_corners=True).squeeze().cpu().numpy() + + scores = np.zeros_like(normal_map[0]) + for l in range(feature_levels): + scores += normal_map[l] + scores = scores.max() - scores + + for i in range(scores.shape[0]): + scores[i] = gaussian_filter(scores[i], sigma=4) + return scores + + +def aggregate_anomaly_scores(logps_list, feature_levels=3, size=224): + abnormal_map = [list() for _ in range(feature_levels)] + for l in range(feature_levels): + probs = torch.cat(logps_list[l], dim=0) + abnormal_map[l] = F.interpolate(probs.unsqueeze(1), size=size, mode='bilinear', align_corners=True).squeeze().cpu().numpy() + + scores = np.zeros_like(abnormal_map[0]) + for l in range(feature_levels): + scores += abnormal_map[l] + scores /= feature_levels + + for i in range(scores.shape[0]): + scores[i] = gaussian_filter(scores[i], sigma=4) + return scores diff --git a/validate_wav1.py b/validate_wav1.py new file mode 100644 index 0000000..9c62ba9 --- /dev/null +++ b/validate_wav1.py @@ -0,0 +1,134 @@ +import warnings +from tqdm import tqdm +from scipy.ndimage import gaussian_filter +import numpy as np +import torch +import torch.nn.functional as F + +from models.modules import get_position_encoding +from models.utils import get_logp +# utilsから get_matched_ref_features をインポート +from utils import get_residual_features, get_matched_ref_features +from utils import calculate_metrics +from losses.utils import get_logp_a + +warnings.filterwarnings('ignore') + +# 引数に gating_net を追加 +def validate(args, encoder, constraintor, gating_net, wav_filter, estimators, test_loader, ref_features, device, class_name): + gating_net.eval() # ゲーティングネットワークも評価モードに + constraintor.eval() + for estimator in estimators: + estimator.eval() + + label_list, gt_mask_list = [], [] + logps1_list = [list() for _ in range(args.feature_levels)] + logps2_list = [list() for _ in range(args.feature_levels)] + + progress_bar = tqdm(total=len(test_loader)) + progress_bar.set_description(f"Evaluating {class_name}") + + for idx, batch in enumerate(test_loader): + progress_bar.update(1) + + image, label, mask = batch[0], batch[1], batch[2] + gt_mask_list.append(mask.squeeze(1).cpu().numpy().astype(bool)) + label_list.append(label.cpu().numpy().astype(bool).ravel()) + + image = image.to(device) + size = image.shape[-1] + + with torch.no_grad(): + # 1. 生の特徴量を抽出 + features_raw = encoder(image) + + # 2. 画像の深い特徴量から層ごとの重みを予測 [w1, w2, w3] + w1, w2, w3 = gating_net(features_raw[-1].detach()) + weights = [w1, w2, w3] + + # 3. 生のカンペと生の特徴量をマッチング + mfeatures_raw = get_matched_ref_features(features_raw, ref_features) + + # 4. テスト画像とカンペの両方に、層ごとの重みでウェーブレットを適用 + features_wav = [wav_filter(features_raw[i], weights[i][0], weights[i][1]) for i in range(args.feature_levels)] + mfeatures_wav = [wav_filter(mfeatures_raw[i], weights[i][0], weights[i][1]) for i in range(args.feature_levels)] + + # 5. フィルタリング後の特徴量で残差を計算 + rfeatures = get_residual_features(features_wav, mfeatures_wav, pos_flag=True) + rfeatures = constraintor(*rfeatures) + + for l in range(args.feature_levels): + e = rfeatures[l] + bs, dim, h, w = e.size() + e = e.permute(0, 2, 3, 1).reshape(-1, dim) + + pos_embed = get_position_encoding(args.pos_embed_dim, h, w).to(args.device).unsqueeze(0).repeat(bs, 1, 1, 1) + pos_embed = pos_embed.permute(0, 2, 3, 1).reshape(-1, args.pos_embed_dim) + estimator = estimators[l] + + if args.flow_arch == 'flow_model': + z, log_jac_det = estimator(e) + else: + z, log_jac_det = estimator(e, [pos_embed, ]) + + logps = get_logp(dim, z, log_jac_det) + logps = logps / dim + logps1_list[l].append(logps.reshape(bs, h, w)) + + logps_a = get_logp_a(dim, z, log_jac_det) + logits = torch.stack([logps, logps_a], dim=-1) + sa = torch.softmax(logits, dim=-1)[:, 1] + logps2_list[l].append(sa.reshape(bs, h, w)) + + progress_bar.close() + + labels = np.concatenate(label_list) + gt_masks = np.concatenate(gt_mask_list, axis=0) + + scores1 = convert_to_anomaly_scores(logps1_list, feature_levels=args.feature_levels, size=size) + scores2 = aggregate_anomaly_scores(logps2_list, feature_levels=args.feature_levels, size=size) + + img_auc1, img_ap1, img_f1_score1, pix_auc1, pix_ap1, pix_f1_score1, pix_aupro1 = calculate_metrics(scores1, labels, gt_masks, pro=False, only_max_value=True) + img_auc2, img_ap2, img_f1_score2, pix_auc2, pix_ap2, pix_f1_score2, pix_aupro2 = calculate_metrics(scores2, labels, gt_masks, pro=False, only_max_value=True) + + scores = (scores1 + scores2) / 2 + img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro = calculate_metrics(scores, labels, gt_masks, pro=False, only_max_value=True) + + metrics = {} + metrics['scores1'] = [img_auc1, img_ap1, img_f1_score1, pix_auc1, pix_ap1, pix_f1_score1, pix_aupro1] + metrics['scores2'] = [img_auc2, img_ap2, img_f1_score2, pix_auc2, pix_ap2, pix_f1_score2, pix_aupro2] + metrics['scores'] = [img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro] + + return metrics + +def convert_to_anomaly_scores(logps_list, feature_levels=3, size=224): + normal_map = [list() for _ in range(feature_levels)] + for l in range(feature_levels): + logps = torch.cat(logps_list[l], dim=0) + logps -= torch.max(logps) + probs = torch.exp(logps) + normal_map[l] = F.interpolate(probs.unsqueeze(1), size=size, mode='bilinear', align_corners=True).squeeze().cpu().numpy() + + scores = np.zeros_like(normal_map[0]) + for l in range(feature_levels): + scores += normal_map[l] + scores = scores.max() - scores + + for i in range(scores.shape[0]): + scores[i] = gaussian_filter(scores[i], sigma=4) + return scores + +def aggregate_anomaly_scores(logps_list, feature_levels=3, size=224): + abnormal_map = [list() for _ in range(feature_levels)] + for l in range(feature_levels): + probs = torch.cat(logps_list[l], dim=0) + abnormal_map[l] = F.interpolate(probs.unsqueeze(1), size=size, mode='bilinear', align_corners=True).squeeze().cpu().numpy() + + scores = np.zeros_like(abnormal_map[0]) + for l in range(feature_levels): + scores += abnormal_map[l] + scores /= feature_levels + + for i in range(scores.shape[0]): + scores[i] = gaussian_filter(scores[i], sigma=4) + return scores diff --git a/validate_wav_cf.py b/validate_wav_cf.py new file mode 100644 index 0000000..087d0d2 --- /dev/null +++ b/validate_wav_cf.py @@ -0,0 +1,156 @@ +import warnings +from tqdm import tqdm +from scipy.ndimage import gaussian_filter +import numpy as np +import torch +import torch.nn.functional as F + +from models.modules import get_position_encoding +from models.utils import get_logp +from utils import get_residual_features, get_matched_ref_features +from utils import calculate_metrics +from losses.utils import get_logp_a + +warnings.filterwarnings('ignore') + +def validate(args, encoder, constraintor, wav_filter, cf_modules, estimators, test_loader, ref_features, device, class_name): + constraintor.eval() + cf_modules.eval() + for estimator in estimators: + estimator.eval() + + ref_features_cat = [] + + # --- 修正箇所: ダミー画像を使って(H, W)を逆算し、2次元カンペを4次元に復元 --- + with torch.no_grad(): + dummy_img = torch.zeros(1, 3, 224, 224).to(device) + dummy_feats = encoder(dummy_img) + + for l in range(args.feature_levels): + # エンコーダの出力から正しい解像度(H, W)を取得 + _, C, H, W = dummy_feats[l].shape + + # .npyから読み込んだ2次元テンソル(K*H*W, C)を、元の4次元(K, C, H, W)に再構築 + K = ref_features[l].shape[0] // (H * W) + ref_4d = ref_features[l].view(K, H, W, C).permute(0, 3, 1, 2) + + # 4次元テンソルに対してフィルタとCFを適用 + ref_lf, ref_hf = wav_filter.get_LF_HF(ref_4d) + ref_lf, ref_hf = cf_modules[l](ref_lf, ref_hf) + + # 連結して再びマッチング用の平坦な2次元に戻す + cat_4d = torch.cat([ref_lf, ref_hf], dim=1) + bs, c, h, w = cat_4d.shape + cat_2d = cat_4d.permute(0, 2, 3, 1).reshape(-1, c) + + ref_features_cat.append(cat_2d) + # -------------------------------------------------------------------------- + + label_list, gt_mask_list = [], [] + logps1_list = [list() for _ in range(args.feature_levels)] + logps2_list = [list() for _ in range(args.feature_levels)] + + progress_bar = tqdm(total=len(test_loader)) + progress_bar.set_description(f"Evaluating {class_name}") + + for idx, batch in enumerate(test_loader): + progress_bar.update(1) + + image, label, mask = batch[0], batch[1], batch[2] + gt_mask_list.append(mask.squeeze(1).cpu().numpy().astype(bool)) + label_list.append(label.cpu().numpy().astype(bool).ravel()) + + image = image.to(device) + size = image.shape[-1] + + with torch.no_grad(): + features_raw = encoder(image) + features_cat = [] + + for l in range(args.feature_levels): + test_lf, test_hf = wav_filter.get_LF_HF(features_raw[l]) + test_lf, test_hf = cf_modules[l](test_lf, test_hf) + + # テスト画像側は空間を維持したまま連結してマッチングに送る + features_cat.append(torch.cat([test_lf, test_hf], dim=1)) + + mfeatures_cat = get_matched_ref_features(features_cat, ref_features_cat) + rfeatures = get_residual_features(features_cat, mfeatures_cat, pos_flag=True) + + rfeatures = constraintor(*rfeatures) + + for l in range(args.feature_levels): + e = rfeatures[l] + bs, dim, h, w = e.size() + e = e.permute(0, 2, 3, 1).reshape(-1, dim) + + pos_embed = get_position_encoding(args.pos_embed_dim, h, w).to(args.device).unsqueeze(0).repeat(bs, 1, 1, 1) + pos_embed = pos_embed.permute(0, 2, 3, 1).reshape(-1, args.pos_embed_dim) + estimator = estimators[l] + + if args.flow_arch == 'flow_model': + z, log_jac_det = estimator(e) + else: + z, log_jac_det = estimator(e, [pos_embed, ]) + + logps = get_logp(dim, z, log_jac_det) + logps = logps / dim + logps1_list[l].append(logps.reshape(bs, h, w)) + + logps_a = get_logp_a(dim, z, log_jac_det) + logits = torch.stack([logps, logps_a], dim=-1) + sa = torch.softmax(logits, dim=-1)[:, 1] + logps2_list[l].append(sa.reshape(bs, h, w)) + + progress_bar.close() + + labels = np.concatenate(label_list) + gt_masks = np.concatenate(gt_mask_list, axis=0) + + scores1 = convert_to_anomaly_scores(logps1_list, feature_levels=args.feature_levels, size=size) + scores2 = aggregate_anomaly_scores(logps2_list, feature_levels=args.feature_levels, size=size) + + img_auc1, img_ap1, img_f1_score1, pix_auc1, pix_ap1, pix_f1_score1, pix_aupro1 = calculate_metrics(scores1, labels, gt_masks, pro=False, only_max_value=True) + img_auc2, img_ap2, img_f1_score2, pix_auc2, pix_ap2, pix_f1_score2, pix_aupro2 = calculate_metrics(scores2, labels, gt_masks, pro=False, only_max_value=True) + + scores = (scores1 + scores2) / 2 + img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro = calculate_metrics(scores, labels, gt_masks, pro=False, only_max_value=True) + + metrics = {} + metrics['scores1'] = [img_auc1, img_ap1, img_f1_score1, pix_auc1, pix_ap1, pix_f1_score1, pix_aupro1] + metrics['scores2'] = [img_auc2, img_ap2, img_f1_score2, pix_auc2, pix_ap2, pix_f1_score2, pix_aupro2] + metrics['scores'] = [img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro] + + return metrics + +def convert_to_anomaly_scores(logps_list, feature_levels=3, size=224): + normal_map = [list() for _ in range(feature_levels)] + for l in range(feature_levels): + logps = torch.cat(logps_list[l], dim=0) + logps -= torch.max(logps) + probs = torch.exp(logps) + normal_map[l] = F.interpolate(probs.unsqueeze(1), size=size, mode='bilinear', align_corners=True).squeeze().cpu().numpy() + + scores = np.zeros_like(normal_map[0]) + for l in range(feature_levels): + scores += normal_map[l] + scores = scores.max() - scores + + for i in range(scores.shape[0]): + scores[i] = gaussian_filter(scores[i], sigma=4) + return scores + +def aggregate_anomaly_scores(logps_list, feature_levels=3, size=224): + abnormal_map = [list() for _ in range(feature_levels)] + for l in range(feature_levels): + probs = torch.cat(logps_list[l], dim=0) + abnormal_map[l] = F.interpolate(probs.unsqueeze(1), size=size, mode='bilinear', align_corners=True).squeeze().cpu().numpy() + + scores = np.zeros_like(abnormal_map[0]) + for l in range(feature_levels): + scores += abnormal_map[l] + scores /= feature_levels + + for i in range(scores.shape[0]): + scores[i] = gaussian_filter(scores[i], sigma=4) + return scores diff --git a/visualizer.py b/visualizer.py index 9d1d470..a5a9a4e 100644 --- a/visualizer.py +++ b/visualizer.py @@ -4,7 +4,7 @@ from scipy.ndimage import gaussian_filter import matplotlib import matplotlib.pyplot as plt - +from utils import get_image_scores def denormalization(x, mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]): mean = np.array(mean) @@ -33,6 +33,9 @@ def plot(self, test_imgs, scores, gt_masks): vmin = scores.min() * 255. + 80 vmax = vmax - 20 norm = matplotlib.colors.Normalize(vmin=vmin, vmax=vmax) + img_scores = get_image_scores(scores, topk=10)#6.27 + rank = np.argsort(img_scores) + rank = np.argsort(rank) for i in range(len(scores)): img = test_imgs[i] img = denormalization(img) @@ -54,7 +57,7 @@ def plot(self, test_imgs, scores, gt_masks): ax_img[1].title.set_text('GroundTruth') ax_img[2].imshow(heat_map, cmap='jet', norm=norm, interpolation='none') ax_img[2].imshow(img, cmap='gray', alpha=0.7, interpolation='none') - ax_img[2].title.set_text('Segmentation') + ax_img[2].title.set_text('Segmentation' + str(rank[i]) + '/' + str(len(rank)) + '/' + str(img_scores[i])) fig_img.savefig(os.path.join(self.root, str(i) + '.png'), dpi=300) # if img_types[i] == 'good': @@ -68,4 +71,4 @@ def plot(self, test_imgs, scores, gt_masks): # else: # fig_img.savefig(os.path.join(self.root, 'anomaly_nok', img_types[i] + '_' + file_names[i]), dpi=300) - plt.close() \ No newline at end of file + plt.close()