From ea6f51611e176f335de02f17863e6288ef5554bc Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Thu, 29 May 2025 02:14:02 +0900 Subject: [PATCH 001/258] Create change_plan --- change_plan | 3 +++ 1 file changed, 3 insertions(+) create mode 100644 change_plan diff --git a/change_plan b/change_plan new file mode 100644 index 0000000..1f3c977 --- /dev/null +++ b/change_plan @@ -0,0 +1,3 @@ +ResADの場合、例えばmvtecでtestする場合、visaの全データを用いてtrainしている。このときtrainをnormal+dataaug anomalt or normal + anomaly + dataaug anomalyにする? +fewshotなので他のデータセットのクラスの異常にも対応できるようにしたいcutpasteよりはdream?dreamを物体があるに場所しか適応できないように変える。maskのデータはあるからperlinnoiseを生成してそこから絞り込み。 + From b91612565c7f7904e0058d81422c3c3a27b9f156 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Thu, 29 May 2025 11:35:54 +0900 Subject: [PATCH 002/258] Update change_plan --- change_plan | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/change_plan b/change_plan index 1f3c977..9bef973 100644 --- a/change_plan +++ b/change_plan @@ -1,3 +1,4 @@ ResADの場合、例えばmvtecでtestする場合、visaの全データを用いてtrainしている。このときtrainをnormal+dataaug anomalt or normal + anomaly + dataaug anomalyにする? fewshotなので他のデータセットのクラスの異常にも対応できるようにしたいcutpasteよりはdream?dreamを物体があるに場所しか適応できないように変える。maskのデータはあるからperlinnoiseを生成してそこから絞り込み。 - +NSA: cutpasteの応用、コピー元が正常画像、切り取った画像をノイズやブレンドなどを行い異常風に生成、それをほかの正常画像にはりつける。 + From f8e297e2ff84929d270188131499f675f7d389ac Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Tue, 10 Jun 2025 20:11:28 +0900 Subject: [PATCH 003/258] Update classes.py add MVTEC_TO_MVTEC --- classes.py | 9 ++++++++- 1 file changed, 8 insertions(+), 1 deletion(-) diff --git a/classes.py b/classes.py index 6ab46be..3112bba 100644 --- a/classes.py +++ b/classes.py @@ -36,4 +36,11 @@ 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']} From 179f686157df2553e875add5e0f38bdd7da3b4d6 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Tue, 10 Jun 2025 20:13:28 +0900 Subject: [PATCH 004/258] Update main.py add mvtec_to_mvtec --- main.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/main.py b/main.py index 3468d6f..8f8fe49 100644 --- a/main.py +++ b/main.py @@ -34,7 +34,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} def main(args): @@ -299,4 +299,4 @@ def load_mc_reference_features(root_dir: str, class_names, device: torch.device, - \ No newline at end of file + From 5d14063542fd1aa0aa71ebe21436c3370b130beb Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Tue, 10 Jun 2025 22:11:57 +0900 Subject: [PATCH 005/258] Update main.py --- main.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/main.py b/main.py index 8f8fe49..e9579f3 100644 --- a/main.py +++ b/main.py @@ -25,7 +25,7 @@ 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_MPDD, MVTEC_TO_MVTECLOCO, MVTEC_TO_BRATS, MVTEC_TO_MVTEC warnings.filterwarnings('ignore') From 119b3b710f1bc89d89d9d0465d6aeb09a831447c Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 11 Jun 2025 10:38:45 +0900 Subject: [PATCH 006/258] Update main.py --- main.py | 1 + 1 file changed, 1 insertion(+) diff --git a/main.py b/main.py index e9579f3..6e66ca2 100644 --- a/main.py +++ b/main.py @@ -108,6 +108,7 @@ def main(args): train_loss_total, total_num = 0, 0 progress_bar = tqdm(total=len(train_loader)) progress_bar.set_description(f"Epoch[{epoch}/{args.epochs}]") +#data aug を適応するならここ? supervisedだからあまり意味ない? for step, batch in enumerate(train_loader): progress_bar.update(1) images, _, masks, class_names = batch From 08d9c3a274211d4874d9cf042e8e80d4758f434e Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 11 Jun 2025 10:44:09 +0900 Subject: [PATCH 007/258] Update change_plan --- change_plan | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/change_plan b/change_plan index 9bef973..55eecf3 100644 --- a/change_plan +++ b/change_plan @@ -1,4 +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の設定を適応するにはどうするか? +新しいテスト用のデータセットを作る From 89d61b7b2eb067fbec716aee2a9fbad0ec1951fd Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Thu, 12 Jun 2025 18:27:59 +0900 Subject: [PATCH 008/258] Update extract_ref_features.py MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit imagebind使用 --- extract_ref_features.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/extract_ref_features.py b/extract_ref_features.py index 384d475..4a71b23 100644 --- a/extract_ref_features.py +++ b/extract_ref_features.py @@ -226,4 +226,7 @@ def main2(args): parser.add_argument('--save_dir', type=str, default="./ref_features/w50/mvtec_4shot") args = parser.parse_args() - main(args) \ No newline at end of file + if args.mode == 'main': + main(args) + elif args.mode == 'main2': + main2(args) From f486bac8aecc6af76369ed8c2e14344b17dcb67f Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Thu, 12 Jun 2025 18:34:16 +0900 Subject: [PATCH 009/258] Update imagebind.py --- models/imagebind.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) 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 From eb9453d00719dc083be52b5d0da928881c150133 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Thu, 12 Jun 2025 18:46:04 +0900 Subject: [PATCH 010/258] Update extract_ref_features.py --- extract_ref_features.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/extract_ref_features.py b/extract_ref_features.py index 4a71b23..2adad1b 100644 --- a/extract_ref_features.py +++ b/extract_ref_features.py @@ -219,11 +219,13 @@ def main2(args): 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('--save_dir', type=str, default="./ref_features/w50/mvtec_4shot") + parser.add_argument('--mode', type=str, default='main') args = parser.parse_args() if args.mode == 'main': From 7df7162fd6f032c63e5bbbfb766e7946e792ddb1 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sun, 15 Jun 2025 21:38:32 +0900 Subject: [PATCH 011/258] Update main_ib.py --- main_ib.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/main_ib.py b/main_ib.py index 3a07eed..d9d4619 100644 --- a/main_ib.py +++ b/main_ib.py @@ -37,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} def main(args): @@ -362,4 +362,4 @@ def load_and_transform_vision_data(image_paths, device): - \ No newline at end of file + From dfef06753f00552227e7307f7e62e8b5b29a9d6c Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sun, 15 Jun 2025 21:43:04 +0900 Subject: [PATCH 012/258] Update main_ib.py --- main_ib.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/main_ib.py b/main_ib.py index d9d4619..3830bb6 100644 --- a/main_ib.py +++ b/main_ib.py @@ -28,7 +28,7 @@ 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 warnings.filterwarnings('ignore') From 3fc5b09846c9f9e4bc399066ad2166b595b7d055 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Thu, 19 Jun 2025 18:45:45 +0900 Subject: [PATCH 013/258] Update main.py MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit pix_auproが-1になってしまい、checkpointに保存されないので、img_aucを基準にするように変更 --- main.py | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) diff --git a/main.py b/main.py index 6e66ca2..96f99fe 100644 --- a/main.py +++ b/main.py @@ -94,7 +94,8 @@ 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_pro = 0 #-1になってしまうので変更 + best_img_auc = 0 N_batch = 8192 for epoch in range(args.epochs): vq_ops.train() @@ -230,9 +231,9 @@ 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]} From 6d86c70a4f837c0c0ca5abe4f399c9163e802455 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Thu, 19 Jun 2025 18:46:58 +0900 Subject: [PATCH 014/258] Update main.py --- main.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/main.py b/main.py index 96f99fe..f008be6 100644 --- a/main.py +++ b/main.py @@ -231,7 +231,7 @@ 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 img_auc > best_img_auc #pix_aupro > best_pro: + if img_auc > best_img_auc: #pix_aupro > best_pro: os.makedirs(args.checkpoint_path, exist_ok=True) best_img_auc = img_auc #best_pro = pix_aupro state_dict = {'vq_ops': vq_ops.state_dict(), From 42848f0dc570e169c4fb048cf293d8171e40e64d Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Thu, 19 Jun 2025 20:29:12 +0900 Subject: [PATCH 015/258] Update mvtec.py --- datasets/mvtec.py | 44 ++++++++++++++++++++++++++------------------ 1 file changed, 26 insertions(+), 18 deletions(-) diff --git a/datasets/mvtec.py b/datasets/mvtec.py index c52ef6b..b9f19cc 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_type[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,7 @@ 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 = [], [], [], [] for phase in ['train', 'test']: image_dir = os.path.join(self.root, class_name, phase) @@ -157,6 +157,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,15 +165,17 @@ 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) @@ -180,6 +183,7 @@ def _load_all_data(self, class_names=None): all_labels.extend(labels) all_mask_paths.extend(mask_paths) all_class_names.extend(class_names) + all_anomaly_types.extend(anomaly_types) # anomaly_types を追加 return all_image_paths, all_labels, all_mask_paths, all_class_names @@ -201,11 +205,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 +243,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 +267,7 @@ 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 = [], [], [], [] phase = 'train' if self.train else 'test' image_dir = os.path.join(self.root, class_name, phase) @@ -284,6 +288,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,15 +296,17 @@ 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, 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) @@ -307,7 +314,8 @@ def _load_all_data(self, class_names=None): 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) # 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 +334,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 From 0e738b04bd8ac75a51ffeeebba867d5d3a35ec37 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Thu, 19 Jun 2025 20:40:59 +0900 Subject: [PATCH 016/258] Update validate.py MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit UMAP可視化のために戻り値を追加 --- validate.py | 34 +++++++++++++++++++++++++++++----- 1 file changed, 29 insertions(+), 5 deletions(-) diff --git a/validate.py b/validate.py index 667fd6f..cc6bf28 100644 --- a/validate.py +++ b/validate.py @@ -19,7 +19,11 @@ def validate(args, encoder, vq_ops, constraintor, estimators, test_loader, ref_f constraintor.eval() for estimator in estimators: estimator.eval() - + + # UMAP可視化のために追加するリスト + all_features_to_return = [] + all_anomaly_types_to_return = [] + all_gts_to_return = [] # 0/1の画像レベルのラベル label_list, gt_mask_list = [], [] logps1_list = [list() for _ in range(args.feature_levels)] logps2_list = [list() for _ in range(args.feature_levels)] @@ -27,8 +31,15 @@ def validate(args, encoder, vq_ops, constraintor, estimators, test_loader, ref_f progress_bar.set_description(f"Evaluating") for idx, batch in enumerate(test_loader): progress_bar.update(1) + # データセットから返される値を修正したMVTEC/MVTECANOクラスの__getitem__を想定 + # image: 画像テンソル + # label: 画像レベルのGTラベル (0:正常, 1:異常) + # mask: ピクセルレベルのGTマスク + # class_name_batch: (元のコードの_に対応) その画像のクラス名 (str) - バッチ内の全画像で同じはず + # anomaly_type_batch: (追加) その画像の異常タイプ名 (str, 例: 'scratch', 'hole', 'good') + image, label, mask, class_name_batch, anomaly_type_batch = batch # ここを変更 + #image, label, mask, _ = batch - image, label, mask, _ = batch gt_mask_list.append(mask.squeeze(1).cpu().numpy().astype(bool)) label_list.append(label.cpu().numpy().astype(bool).ravel()) @@ -48,10 +59,19 @@ def validate(args, encoder, vq_ops, constraintor, estimators, test_loader, ref_f mfeatures = get_matched_ref_features(features, ref_features) rfeatures = get_residual_features(features, mfeatures) + + # --- UMAP可視化のために特徴量と異常タイプ名を収集 --- + # ここで、UMAPに渡す特徴量を決定します。 + # 通常、最も深い層(最後の要素)の特徴量をフラットにして使います。 + current_features_flat = rfeatures[-1].cpu().numpy().reshape(image.shape[0], -1) + all_features_to_return.append(current_features_flat) + all_anomaly_types_to_return.extend(anomaly_type_batch) # リストのままextend + all_gts_to_return.extend(label.cpu().numpy()) # labelは0/1のGTラベル + fdm_features = vq_ops(rfeatures, train=False) rfeatures = applying_EFDM(rfeatures, fdm_features, alpha=args.fdm_alpha) - rfeatures = constraintor(*rfeatures) - + rfeatures = constraintor(*rfeatures) + for l in range(args.feature_levels): e = rfeatures[l] # BxCxHxW bs, dim, h, w = e.size() @@ -93,6 +113,10 @@ def validate(args, encoder, vq_ops, constraintor, estimators, test_loader, ref_f 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] + # UMAP可視化のために追加した戻り値 + metrics['features'] = np.concatenate(all_features_to_return, axis=0) + metrics['anomaly_types'] = np.array(all_anomaly_types_to_return, dtype=object) # 文字列を含むのでobject型 + metrics['gts_labels'] = np.array(all_gts_to_return) # 0/1のGTラベル return metrics @@ -139,4 +163,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 From 95858c0b05e6f7e11a4997aac915db65f442bec8 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Thu, 19 Jun 2025 20:51:22 +0900 Subject: [PATCH 017/258] Update main.py --- main.py | 38 ++++++++++++++++++++++++++++++++++++-- 1 file changed, 36 insertions(+), 2 deletions(-) diff --git a/main.py b/main.py index f008be6..b71d799 100644 --- a/main.py +++ b/main.py @@ -97,6 +97,10 @@ def main(args): #best_pro = 0 #-1になってしまうので変更 best_img_auc = 0 N_batch = 8192 + + # 最良モデルのエポックで保存するためのデータ保持用 + best_epoch_class_data = {} + for epoch in range(args.epochs): vq_ops.train() constraintor.train() @@ -175,6 +179,9 @@ def main(args): 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) + # 各クラスの評価結果とデータを一時的に保持する辞書 + current_epoch_class_data_for_saving = {} + 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, @@ -209,7 +216,10 @@ def main(args): 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) + #metrics = validate(args, encoder, vq_ops, constraintor, estimators, test_loader, test_ref_features[class_name], args.device, class_name) +   metrics = validate(args, encoder, vq_ops, constraintor, estimators, test_loader, + test_ref_features[class_name_eval], args.device, class_name_eval) + 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( @@ -217,7 +227,12 @@ def main(args): s1_res.append(metrics['scores1']) s2_res.append(metrics['scores2']) s_res.append(metrics['scores']) - + # 各クラスの評価結果から特徴量とラベルデータを一時的に保存 + 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) @@ -238,6 +253,25 @@ def main(args): '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) + + features_filename = os.path.join(class_specific_save_dir, f'epoch{epoch}_features.npy') + anomaly_types_filename = os.path.join(class_specific_save_dir, f'epoch{epoch}_anomaly_types.npy') + gts_labels_filename = os.path.join(class_specific_save_dir, f'epoch{epoch}_gts_labels.npy') + + np.save(features_filename, data['features']) + np.save(anomaly_types_filename, data['anomaly_types']) + np.save(gts_labels_filename, data['gts_labels']) + print(f" - クラス '{class_name_to_save}': Epoch {epoch} のデータを {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): From 7eca7ce8e605cb78db9e906da4fcae0098c4d140 Mon Sep 17 00:00:00 2001 From: Tomoya Ueno Date: Fri, 20 Jun 2025 00:19:48 +0900 Subject: [PATCH 018/258] =?UTF-8?q?=E7=B4=9B=E3=82=89=E3=82=8F=E3=81=97?= =?UTF-8?q?=E3=81=84=E3=81=AE=E3=81=A7class=5Fname=E3=82=92class=5Fname=5F?= =?UTF-8?q?eval=E3=81=AB=E5=A4=89=E6=9B=B4?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- main.py | 36 ++++++++++++++++++------------------ 1 file changed, 18 insertions(+), 18 deletions(-) diff --git a/main.py b/main.py index b71d799..2741e40 100644 --- a/main.py +++ b/main.py @@ -182,48 +182,48 @@ def main(args): # 各クラスの評価結果とデータを一時的に保持する辞書 current_epoch_class_data_for_saving = {} - 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, + 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='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, + elif class_name_eval in VISA.CLASS_NAMES: + test_dataset = VISA(args.test_dataset_dir, class_name=class_name_eval, 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, + elif class_name_eval in BTAD.CLASS_NAMES: + test_dataset = BTAD(args.test_dataset_dir, class_name=class_name_eval, 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, + elif class_name_eval in MVTEC3D.CLASS_NAMES: + test_dataset = MVTEC3D(args.test_dataset_dir, class_name=class_name_eval, 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, + elif class_name_eval in MPDD.CLASS_NAMES: + test_dataset = MPDD(args.test_dataset_dir, class_name=class_name_eval, 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, + elif class_name_eval in MVTECLOCO.CLASS_NAMES: + test_dataset = MVTECLOCO(args.test_dataset_dir, class_name=class_name_eval, 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, + elif class_name_eval in BRATS.CLASS_NAMES: + test_dataset = BRATS(args.test_dataset_dir, class_name=class_name_eval, 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)) + 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_ref_features[class_name], args.device, class_name) -   metrics = validate(args, encoder, vq_ops, constraintor, estimators, test_loader, + metrics = validate(args, encoder, vq_ops, constraintor, estimators, test_loader, test_ref_features[class_name_eval], args.device, class_name_eval) 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']) From 32162cf1b66354507583096142216980ab6e7804 Mon Sep 17 00:00:00 2001 From: Tomoya Ueno Date: Fri, 20 Jun 2025 00:46:26 +0900 Subject: [PATCH 019/258] =?UTF-8?q?class=5Fnames=E3=81=AE=E3=82=A8?= =?UTF-8?q?=E3=83=A9=E3=83=BC=E3=81=AB=E3=81=A4=E3=81=84=E3=81=A6?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- datasets/mvtec.py | 14 +++++++++----- 1 file changed, 9 insertions(+), 5 deletions(-) diff --git a/datasets/mvtec.py b/datasets/mvtec.py index b9f19cc..065561e 100644 --- a/datasets/mvtec.py +++ b/datasets/mvtec.py @@ -137,6 +137,8 @@ def __len__(self): 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) @@ -178,12 +180,12 @@ def _load_all_data(self, class_names=None): 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) - all_anomaly_types.extend(anomaly_types) # anomaly_types を追加 + all_anomaly_types.extend(anomaly_types_from_load_data) # anomaly_types を追加 return all_image_paths, all_labels, all_mask_paths, all_class_names @@ -268,6 +270,8 @@ def _load_image_and_mask(self, image_path, label, mask_path): 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) @@ -299,7 +303,7 @@ def _load_data(self, class_name): 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, anomaly_types # anomaly_types も返す + return image_paths, labels, mask_paths, class_names_list, anomaly_types # anomaly_types も返す def _load_all_data(self, class_names=None): all_image_paths = [] @@ -309,12 +313,12 @@ def _load_all_data(self, class_names=None): 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) - all_anomaly_types.extend(anomaly_types) # anomaly_types を追加 + 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 も返す From c96043b3c8ced6cbb33aea16292ee47d7c11c996 Mon Sep 17 00:00:00 2001 From: Tomoya Ueno Date: Fri, 20 Jun 2025 00:50:11 +0900 Subject: [PATCH 020/258] =?UTF-8?q?unpack=E3=81=AE=E3=82=A8=E3=83=A9?= =?UTF-8?q?=E3=83=BC?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- datasets/mvtec.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/datasets/mvtec.py b/datasets/mvtec.py index 065561e..a085214 100644 --- a/datasets/mvtec.py +++ b/datasets/mvtec.py @@ -186,7 +186,7 @@ def _load_all_data(self, class_names=None): 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 + return all_image_paths, all_labels, all_mask_paths, all_class_names, all_anomaly_types # anomaly_types も返す class MVTEC(Dataset): From 54cd05a0e9a11cc0b4f548832137769283488e42 Mon Sep 17 00:00:00 2001 From: Tomoya Ueno Date: Fri, 20 Jun 2025 00:53:15 +0900 Subject: [PATCH 021/258] =?UTF-8?q?anomaly=5Ftypes=E3=82=92=E8=BF=BD?= =?UTF-8?q?=E5=8A=A0?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- main.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/main.py b/main.py index 2741e40..971731e 100644 --- a/main.py +++ b/main.py @@ -116,7 +116,7 @@ def main(args): #data aug を適応するならここ? supervisedだからあまり意味ない? 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) From 7fdc0374fd95df722e94d9c6dcc0279aa2745b99 Mon Sep 17 00:00:00 2001 From: Tomoya Ueno Date: Fri, 20 Jun 2025 01:10:05 +0900 Subject: [PATCH 022/258] =?UTF-8?q?anomaly=5Ftypes=E3=82=92=E3=83=86?= =?UTF-8?q?=E3=82=AD=E3=82=B9=E3=83=88=E3=81=A7=E4=BF=9D=E5=AD=98?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- main.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/main.py b/main.py index 971731e..b8fc9dc 100644 --- a/main.py +++ b/main.py @@ -266,7 +266,10 @@ def main(args): gts_labels_filename = os.path.join(class_specific_save_dir, f'epoch{epoch}_gts_labels.npy') np.save(features_filename, data['features']) - np.save(anomaly_types_filename, data['anomaly_types']) + # 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}': Epoch {epoch} のデータを {class_specific_save_dir} に保存しました。") From dd4da74cdd1e8d9a444cd172d37ac108737f25cf Mon Sep 17 00:00:00 2001 From: Tomoya Ueno Date: Fri, 20 Jun 2025 01:43:46 +0900 Subject: [PATCH 023/258] =?UTF-8?q?=E7=89=B9=E5=BE=B4=E9=87=8F=E3=82=92?= =?UTF-8?q?=E4=B8=8A=E6=9B=B8=E3=81=8D=E4=BF=9D=E5=AD=98=E3=81=99=E3=82=8B?= =?UTF-8?q?=E3=82=88=E3=81=86=E3=81=AB=E5=A4=89=E6=9B=B4?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- main.py | 16 ++++++++++++---- 1 file changed, 12 insertions(+), 4 deletions(-) diff --git a/main.py b/main.py index b8fc9dc..4efd636 100644 --- a/main.py +++ b/main.py @@ -256,22 +256,30 @@ def main(args): # 新しい特徴量保存ディレクトリを作成 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) - features_filename = os.path.join(class_specific_save_dir, f'epoch{epoch}_features.npy') - anomaly_types_filename = os.path.join(class_specific_save_dir, f'epoch{epoch}_anomaly_types.npy') - gts_labels_filename = os.path.join(class_specific_save_dir, f'epoch{epoch}_gts_labels.npy') + # ファイル名から 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}': Epoch {epoch} のデータを {class_specific_save_dir} に保存しました。") + print(f" - クラス '{class_name_to_save}': 最良スコア時のデータを {class_specific_save_dir} に上書き保存しました。") + # -- 変更ここまで -- # 最良エポックのデータなので、今後の可視化のためにこれを覚えておく best_epoch_class_data = current_epoch_class_data_for_saving.copy() From 1979bd12b28d96c6e61df9979e4bd2fda58417ad Mon Sep 17 00:00:00 2001 From: Tomoya Ueno Date: Fri, 20 Jun 2025 01:53:59 +0900 Subject: [PATCH 024/258] =?UTF-8?q?=E5=8F=96=E3=81=A3=E3=81=A6=E3=81=8F?= =?UTF-8?q?=E3=82=8B=E7=89=B9=E5=BE=B4=E9=87=8F=E3=82=92=E5=A4=89=E6=9B=B4?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- validate.py | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) diff --git a/validate.py b/validate.py index cc6bf28..94254a2 100644 --- a/validate.py +++ b/validate.py @@ -63,14 +63,15 @@ def validate(args, encoder, vq_ops, constraintor, estimators, test_loader, ref_f # --- UMAP可視化のために特徴量と異常タイプ名を収集 --- # ここで、UMAPに渡す特徴量を決定します。 # 通常、最も深い層(最後の要素)の特徴量をフラットにして使います。 + fdm_features = vq_ops(rfeatures, train=False) + rfeatures = applying_EFDM(rfeatures, fdm_features, alpha=args.fdm_alpha) + rfeatures = constraintor(*rfeatures) + current_features_flat = rfeatures[-1].cpu().numpy().reshape(image.shape[0], -1) all_features_to_return.append(current_features_flat) all_anomaly_types_to_return.extend(anomaly_type_batch) # リストのままextend all_gts_to_return.extend(label.cpu().numpy()) # labelは0/1のGTラベル - 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 From 1d049a87b3d03ef3f54e391d680a781ff970bb7e Mon Sep 17 00:00:00 2001 From: Tomoya Ueno Date: Fri, 20 Jun 2025 12:56:23 +0900 Subject: [PATCH 025/258] =?UTF-8?q?visualizer=E9=81=A9=E5=BF=9C?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- main.py | 32 ++++++++++++++++++++++++++++---- validate.py | 13 ++++++++++--- 2 files changed, 38 insertions(+), 7 deletions(-) diff --git a/main.py b/main.py index 4efd636..9faa99e 100644 --- a/main.py +++ b/main.py @@ -26,7 +26,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, MVTEC_TO_MVTEC - +# visualizerのインポート +from visualizer import Visualizer, denormalization warnings.filterwarnings('ignore') TOTAL_SHOT = 4 # total few-shot reference samples @@ -97,7 +98,12 @@ def main(args): #best_pro = 0 #-1になってしまうので変更 best_img_auc = 0 N_batch = 8192 - + + # 可視化オブジェクトの初期化 + # 可視化結果を保存するディレクトリを指定 + 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 = {} @@ -216,10 +222,10 @@ def main(args): 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) metrics = validate(args, encoder, vq_ops, constraintor, estimators, test_loader, test_ref_features[class_name_eval], args.device, class_name_eval) - + #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( @@ -227,6 +233,24 @@ def main(args): s1_res.append(metrics['scores1']) s2_res.append(metrics['scores2']) s_res.append(metrics['scores']) + # ここで可視化処理を追加 + # `validate` 関数が返す `metrics` から必要なデータを取得 + scores = metrics['scores_map'] # validate.pyでscores_mapとして返すように修正が必要 + gts_masks = metrics['gt_masks_raw'] # validate.pyでgt_masks_rawとして返すように修正が必要 + images_raw = metrics['images_raw'] # validate.pyでimages_rawとして返すように修正が必要 + # Visualizerを使ってプロット + # クラスごとにサブディレクトリを作成 + output_class_dir = os.path.join(visualization_output_dir, class_name_eval, f'epoch_{epoch}') + os.makedirs(output_class_dir, exist_ok=True) + my_visualizer.set_prefix(f'{class_name_eval}_epoch{epoch}') # プレフィックスをクラス名とエポックに設定 + my_visualizer.root = output_class_dir # 保存先ディレクトリを更新 + # ここで images_raw はまだ正規化されている可能性があるので、denormalizationを適用 + # validate.py の中で生画像を保存するか、ここでロードし直す方が良いかもしれません + # 今回はvalidate.pyがテスト画像をそのまま返すように仮定 + # scores は NumPy 配列、gt_masks も NumPy 配列であることを確認 + 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'], diff --git a/validate.py b/validate.py index 94254a2..cf0b96c 100644 --- a/validate.py +++ b/validate.py @@ -24,6 +24,10 @@ def validate(args, encoder, vq_ops, constraintor, estimators, test_loader, ref_f all_features_to_return = [] all_anomaly_types_to_return = [] all_gts_to_return = [] # 0/1の画像レベルのラベル + # 可視化のために追加 + all_images_raw = [] # 生の画像データ + all_scores_map = [] # スコアマップ + label_list, gt_mask_list = [], [] logps1_list = [list() for _ in range(args.feature_levels)] logps2_list = [list() for _ in range(args.feature_levels)] @@ -39,7 +43,7 @@ def validate(args, encoder, vq_ops, constraintor, estimators, test_loader, ref_f # anomaly_type_batch: (追加) その画像の異常タイプ名 (str, 例: 'scratch', 'hole', 'good') image, label, mask, class_name_batch, anomaly_type_batch = batch # ここを変更 #image, label, mask, _ = batch - + all_images_raw.append(image.cpu().numpy()) gt_mask_list.append(mask.squeeze(1).cpu().numpy().astype(bool)) label_list.append(label.cpu().numpy().astype(bool).ravel()) @@ -109,7 +113,7 @@ def validate(args, encoder, vq_ops, constraintor, estimators, test_loader, ref_f 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) - + #visualizerを使えるようにするためにtest_imgsを返す 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] @@ -118,7 +122,10 @@ def validate(args, encoder, vq_ops, constraintor, estimators, test_loader, ref_f metrics['features'] = np.concatenate(all_features_to_return, axis=0) metrics['anomaly_types'] = np.array(all_anomaly_types_to_return, dtype=object) # 文字列を含むのでobject型 metrics['gts_labels'] = np.array(all_gts_to_return) # 0/1のGTラベル - + # 可視化のために追加 + metrics['images_raw'] = np.concatenate(all_images_raw, axis=0) + metrics['scores_map'] = scores + metrics['gt_masks_raw'] = gt_masks return metrics From 447af863075e0808d54c559dccf21c312626a128 Mon Sep 17 00:00:00 2001 From: Tomoya Ueno Date: Fri, 20 Jun 2025 13:14:27 +0900 Subject: [PATCH 026/258] =?UTF-8?q?=E5=8F=AF=E8=A6=96=E5=8C=96=E3=81=AF?= =?UTF-8?q?=E6=9C=80=E5=BE=8C=E3=81=AE=E3=82=A8=E3=83=9D=E3=83=83=E3=82=AF?= =?UTF-8?q?=E3=81=AE=E3=81=BF=E9=81=A9=E5=BF=9C?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- main.py | 35 ++++++++++++++++++----------------- 1 file changed, 18 insertions(+), 17 deletions(-) diff --git a/main.py b/main.py index 9faa99e..383c3ef 100644 --- a/main.py +++ b/main.py @@ -233,23 +233,24 @@ def main(args): s1_res.append(metrics['scores1']) s2_res.append(metrics['scores2']) s_res.append(metrics['scores']) - # ここで可視化処理を追加 - # `validate` 関数が返す `metrics` から必要なデータを取得 - scores = metrics['scores_map'] # validate.pyでscores_mapとして返すように修正が必要 - gts_masks = metrics['gt_masks_raw'] # validate.pyでgt_masks_rawとして返すように修正が必要 - images_raw = metrics['images_raw'] # validate.pyでimages_rawとして返すように修正が必要 - # Visualizerを使ってプロット - # クラスごとにサブディレクトリを作成 - output_class_dir = os.path.join(visualization_output_dir, class_name_eval, f'epoch_{epoch}') - os.makedirs(output_class_dir, exist_ok=True) - my_visualizer.set_prefix(f'{class_name_eval}_epoch{epoch}') # プレフィックスをクラス名とエポックに設定 - my_visualizer.root = output_class_dir # 保存先ディレクトリを更新 - # ここで images_raw はまだ正規化されている可能性があるので、denormalizationを適用 - # validate.py の中で生画像を保存するか、ここでロードし直す方が良いかもしれません - # 今回はvalidate.pyがテスト画像をそのまま返すように仮定 - # scores は NumPy 配列、gt_masks も NumPy 配列であることを確認 - my_visualizer.plot(images_raw, scores, gts_masks) # - print(f" - クラス '{class_name_eval}': 可視化結果を {output_class_dir} に保存しました。") + + # 可視化結果を保存するのは最終エポックのみ + 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] = { From 27005030879ba54568eb62b187cf9f4a4248c151 Mon Sep 17 00:00:00 2001 From: Tomoya Ueno Date: Fri, 20 Jun 2025 16:29:09 +0900 Subject: [PATCH 027/258] =?UTF-8?q?=E3=82=BF=E3=82=A4=E3=83=97=E3=83=9F?= =?UTF-8?q?=E3=82=B9=E4=BF=AE=E6=AD=A3?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- datasets/mvtec.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/datasets/mvtec.py b/datasets/mvtec.py index a085214..495620b 100644 --- a/datasets/mvtec.py +++ b/datasets/mvtec.py @@ -109,7 +109,7 @@ def __init__( 12: 'transistor', 13: 'wood', 14: 'zipper'} 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_type[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 From 586bf9d51d3bb6320d3c2965b63cf3cb0dedccd0 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sat, 21 Jun 2025 18:19:04 +0900 Subject: [PATCH 028/258] Create mvtec_fewclass.py --- datasets/mvtec_fewclass.py | 341 +++++++++++++++++++++++++++++++++++++ 1 file changed, 341 insertions(+) create mode 100644 datasets/mvtec_fewclass.py diff --git a/datasets/mvtec_fewclass.py b/datasets/mvtec_fewclass.py new file mode 100644 index 0000000..495620b --- /dev/null +++ b/datasets/mvtec_fewclass.py @@ -0,0 +1,341 @@ +""" +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 MVTECANO(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 = ['bottle', 'cable', 'capsule', 'carpet', 'grid', + 'hazelnut', 'leather', 'metal_nut', 'pill', 'screw', + 'tile', 'toothbrush', 'transistor', 'wood', 'zipper'] + + 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 = {'bottle': 0, 'cable': 1, 'capsule': 2, 'carpet': 3, + 'grid': 4, 'hazelnut': 5, 'leather': 6, 'metal_nut': 7, + 'pill': 8, 'screw': 9, 'tile': 10, 'toothbrush': 11, + 'transistor': 12, 'wood': 13, 'zipper': 14} + self.idx_to_class = {0: 'bottle', 1: 'cable', 2: 'capsule', 3: 'carpet', + 4: 'grid', 5: 'hazelnut', 6: 'leather', 7: 'metal_nut', + 8: 'pill', 9: 'screw', 10: 'tile', 11: 'toothbrush', + 12: 'transistor', 13: 'wood', 14: 'zipper'} + + 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 MVTEC(Dataset): + + CLASS_NAMES = ['bottle', 'cable', 'capsule', 'carpet', 'grid', + 'hazelnut', 'leather', 'metal_nut', 'pill', 'screw', + 'tile', 'toothbrush', 'transistor', 'wood', 'zipper'] + def __init__(self, + root: str, + class_name: str = 'bottle', + 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 From 7ac66ffcac3dccc69551ab49c0e1115e26f9e53c Mon Sep 17 00:00:00 2001 From: Tomoya Ueno Date: Sat, 21 Jun 2025 18:42:51 +0900 Subject: [PATCH 029/258] MVTECFEW --- datasets/mvtec_fewclass.py | 24 +++++++----------------- 1 file changed, 7 insertions(+), 17 deletions(-) diff --git a/datasets/mvtec_fewclass.py b/datasets/mvtec_fewclass.py index 495620b..3abc510 100644 --- a/datasets/mvtec_fewclass.py +++ b/datasets/mvtec_fewclass.py @@ -25,7 +25,7 @@ IMAGENET_STD = [0.229, 0.224, 0.225] -class MVTECANO(Dataset): +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. @@ -44,9 +44,7 @@ class MVTECANO(Dataset): MVTEC_URL = 'ftp://guest:GU.205dldo@ftp.softronics.ch/mvtec_anomaly_detection/mvtec_anomaly_detection.tar.xz' - CLASS_NAMES = ['bottle', 'cable', 'capsule', 'carpet', 'grid', - 'hazelnut', 'leather', 'metal_nut', 'pill', 'screw', - 'tile', 'toothbrush', 'transistor', 'wood', 'zipper'] + CLASS_NAMES = ['capsule','screw','transistor'] def __init__( self, @@ -99,14 +97,8 @@ def __init__( T.CenterCrop(kwargs.get('msk_crp_size')), T.ToTensor()]) - self.class_to_idx = {'bottle': 0, 'cable': 1, 'capsule': 2, 'carpet': 3, - 'grid': 4, 'hazelnut': 5, 'leather': 6, 'metal_nut': 7, - 'pill': 8, 'screw': 9, 'tile': 10, 'toothbrush': 11, - 'transistor': 12, 'wood': 13, 'zipper': 14} - self.idx_to_class = {0: 'bottle', 1: 'cable', 2: 'capsule', 3: 'carpet', - 4: 'grid', 5: 'hazelnut', 6: 'leather', 7: 'metal_nut', - 8: 'pill', 9: 'screw', 10: 'tile', 11: 'toothbrush', - 12: 'transistor', 13: 'wood', 14: 'zipper'} + 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] @@ -189,14 +181,12 @@ def _load_all_data(self, class_names=None): return all_image_paths, all_labels, all_mask_paths, all_class_names, all_anomaly_types # anomaly_types も返す -class MVTEC(Dataset): +class MVTECFEW(Dataset): - CLASS_NAMES = ['bottle', 'cable', 'capsule', 'carpet', 'grid', - 'hazelnut', 'leather', 'metal_nut', 'pill', 'screw', - 'tile', 'toothbrush', 'transistor', 'wood', 'zipper'] + CLASS_NAMES = ['capsule','screw','transistor'] def __init__(self, root: str, - class_name: str = 'bottle', + class_name: str = 'capsule', train: bool = True, normalize: str = 'imagebind', **kwargs) -> None: From be8ead5b2e0b06e4560aa8f4b9fa07f42d8a3621 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sat, 21 Jun 2025 18:45:20 +0900 Subject: [PATCH 030/258] Update mvtec_fewclass.py --- datasets/mvtec_fewclass.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/datasets/mvtec_fewclass.py b/datasets/mvtec_fewclass.py index 3abc510..709f3d7 100644 --- a/datasets/mvtec_fewclass.py +++ b/datasets/mvtec_fewclass.py @@ -98,7 +98,7 @@ def __init__( T.ToTensor()]) self.class_to_idx = {'capsule': 0,'screw': 1, 'toothbrush': 2} - self.idx_to_class = { 0: 'capsule', 1: 'screw', 2: 'transistor', } + 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] From 382d9e88c28665cbd55bda2a07c62fa211bdba57 Mon Sep 17 00:00:00 2001 From: Tomoya Ueno Date: Sat, 21 Jun 2025 18:47:00 +0900 Subject: [PATCH 031/258] =?UTF-8?q?=E8=A8=AD=E5=AE=9A=E8=BF=BD=E5=8A=A0?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- classes.py | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/classes.py b/classes.py index 3112bba..225f339 100644 --- a/classes.py +++ b/classes.py @@ -44,3 +44,8 @@ 'unseen': ['bottle', 'cable', 'capsule', 'carpet', 'grid', 'hazelnut', 'leather', 'metal_nut', 'pill', 'screw', 'tile', 'toothbrush', 'transistor', 'wood', 'zipper']} + +MVTECFEW_TO_MVTEC = {'seen': ['capsule','screw','transistor'], + 'unseen': ['bottle', 'cable', 'capsule', 'carpet', 'grid', + 'hazelnut', 'leather', 'metal_nut', 'pill', 'screw', + 'tile', 'toothbrush', 'transistor', 'wood', 'zipper']} From 1e7f616b37a5f587110f1c544719411358adec32 Mon Sep 17 00:00:00 2001 From: Tomoya Ueno Date: Sat, 21 Jun 2025 18:47:50 +0900 Subject: [PATCH 032/258] =?UTF-8?q?=E8=A8=AD=E5=AE=9A=E8=BF=BD=E5=8A=A0?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- classes.py | 1 + 1 file changed, 1 insertion(+) diff --git a/classes.py b/classes.py index 225f339..bb6b7c2 100644 --- a/classes.py +++ b/classes.py @@ -49,3 +49,4 @@ 'unseen': ['bottle', 'cable', 'capsule', 'carpet', 'grid', 'hazelnut', 'leather', 'metal_nut', 'pill', 'screw', 'tile', 'toothbrush', 'transistor', 'wood', 'zipper']} +# \ No newline at end of file From 01a8beba0dfb181eb4432ae7ec56a2efe864007d Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sat, 21 Jun 2025 18:50:00 +0900 Subject: [PATCH 033/258] Update classes.py --- classes.py | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/classes.py b/classes.py index 3112bba..bb00514 100644 --- a/classes.py +++ b/classes.py @@ -44,3 +44,7 @@ 'unseen': ['bottle', 'cable', 'capsule', 'carpet', 'grid', 'hazelnut', 'leather', 'metal_nut', 'pill', 'screw', 'tile', 'toothbrush', 'transistor', 'wood', 'zipper']} +MVTECFEW_TO_MVTEC = {'seen': ['capsule','screw','transistor'], + 'unseen': ['bottle', 'cable', 'capsule', 'carpet', 'grid', + 'hazelnut', 'leather', 'metal_nut', 'pill', 'screw', + 'tile', 'toothbrush', 'transistor', 'wood', 'zipper']} From 9c226d131d5d7e1cdd03ff5fa0b21bfe1ac0530b Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sat, 21 Jun 2025 18:58:20 +0900 Subject: [PATCH 034/258] Update classes.py --- classes.py | 8 +------- 1 file changed, 1 insertion(+), 7 deletions(-) diff --git a/classes.py b/classes.py index f745a52..86f058c 100644 --- a/classes.py +++ b/classes.py @@ -44,15 +44,9 @@ 'unseen': ['bottle', 'cable', 'capsule', 'carpet', 'grid', 'hazelnut', 'leather', 'metal_nut', 'pill', 'screw', 'tile', 'toothbrush', 'transistor', 'wood', 'zipper']} -<<<<<<< HEAD -======= ->>>>>>> 01a8beba0dfb181eb4432ae7ec56a2efe864007d MVTECFEW_TO_MVTEC = {'seen': ['capsule','screw','transistor'], 'unseen': ['bottle', 'cable', 'capsule', 'carpet', 'grid', 'hazelnut', 'leather', 'metal_nut', 'pill', 'screw', 'tile', 'toothbrush', 'transistor', 'wood', 'zipper']} -<<<<<<< HEAD -# -======= ->>>>>>> 01a8beba0dfb181eb4432ae7ec56a2efe864007d + From 605f47091cb81474c325ebba04ef5fe2ed391fc3 Mon Sep 17 00:00:00 2001 From: Tomoya Ueno Date: Sun, 22 Jun 2025 02:42:04 +0900 Subject: [PATCH 035/258] a --- classes.py | 7 ------- 1 file changed, 7 deletions(-) diff --git a/classes.py b/classes.py index f745a52..225f339 100644 --- a/classes.py +++ b/classes.py @@ -44,15 +44,8 @@ 'unseen': ['bottle', 'cable', 'capsule', 'carpet', 'grid', 'hazelnut', 'leather', 'metal_nut', 'pill', 'screw', 'tile', 'toothbrush', 'transistor', 'wood', 'zipper']} -<<<<<<< HEAD -======= ->>>>>>> 01a8beba0dfb181eb4432ae7ec56a2efe864007d MVTECFEW_TO_MVTEC = {'seen': ['capsule','screw','transistor'], 'unseen': ['bottle', 'cable', 'capsule', 'carpet', 'grid', 'hazelnut', 'leather', 'metal_nut', 'pill', 'screw', 'tile', 'toothbrush', 'transistor', 'wood', 'zipper']} -<<<<<<< HEAD -# -======= ->>>>>>> 01a8beba0dfb181eb4432ae7ec56a2efe864007d From f1746b351a34787bc1fc6ca3a36339dbcde435c6 Mon Sep 17 00:00:00 2001 From: Tomoya Ueno Date: Sun, 22 Jun 2025 02:51:22 +0900 Subject: [PATCH 036/258] fewclass setting --- main.py | 23 +++++++++++++++++++---- 1 file changed, 19 insertions(+), 4 deletions(-) diff --git a/main.py b/main.py index 383c3ef..cb97522 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.mvtec_fewclass import MVTECFEWANO, MVTECFEW from models.fc_flow import load_flow_model from models.modules import MultiScaleConv @@ -25,7 +26,7 @@ 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, MVTEC_TO_MVTEC +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') @@ -35,7 +36,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_mvtec': MVTEC_TO_MVTEC} + 'mvtec_to_brats': MVTEC_TO_BRATS,'mvtec_to_mvtec': MVTEC_TO_MVTEC, 'mvtecfew_to_mvtec': MVTECFEW_TO_MVTEC} def main(args): @@ -43,8 +44,21 @@ 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="w50", img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) @@ -335,6 +349,7 @@ def load_mc_reference_features(root_dir: str, class_names, device: torch.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="") From 2ca6c7ab5ef55e5fad66700cc8ad6b85ecf61bfa Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sun, 22 Jun 2025 02:59:39 +0900 Subject: [PATCH 037/258] Update classes.py --- classes.py | 3 --- 1 file changed, 3 deletions(-) diff --git a/classes.py b/classes.py index 5b38e97..86f058c 100644 --- a/classes.py +++ b/classes.py @@ -49,7 +49,4 @@ 'unseen': ['bottle', 'cable', 'capsule', 'carpet', 'grid', 'hazelnut', 'leather', 'metal_nut', 'pill', 'screw', 'tile', 'toothbrush', 'transistor', 'wood', 'zipper']} -<<<<<<< HEAD -======= ->>>>>>> 9c226d131d5d7e1cdd03ff5fa0b21bfe1ac0530b From 41c6cbc507baef814b1cf5272650baa0549ebdeb Mon Sep 17 00:00:00 2001 From: Tomoya Ueno Date: Sat, 28 Jun 2025 01:55:09 +0900 Subject: [PATCH 038/258] =?UTF-8?q?6/27=E3=82=BC=E3=83=9F?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- classes.py | 1 - visualizer.py | 7 +++++-- 2 files changed, 5 insertions(+), 3 deletions(-) diff --git a/classes.py b/classes.py index 86f058c..225f339 100644 --- a/classes.py +++ b/classes.py @@ -49,4 +49,3 @@ 'unseen': ['bottle', 'cable', 'capsule', 'carpet', 'grid', 'hazelnut', 'leather', 'metal_nut', 'pill', 'screw', 'tile', 'toothbrush', 'transistor', 'wood', 'zipper']} - diff --git a/visualizer.py b/visualizer.py index 9d1d470..9917f50 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': From 4e7eed0a407fde10e469105ba4302886d63742f9 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Tue, 1 Jul 2025 01:02:50 +0900 Subject: [PATCH 039/258] Update main_ib.p --- main_ib.py | 101 ++++++++++++++++++++++++++++++++++++++++++++++------- 1 file changed, 89 insertions(+), 12 deletions(-) diff --git a/main_ib.py b/main_ib.py index 3830bb6..aaf5c26 100644 --- a/main_ib.py +++ b/main_ib.py @@ -28,7 +28,9 @@ 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, MVTEC_TO_MVTEC +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 +39,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_mvtec': MVTEC_TO_MVTEC} + 'mvtec_to_brats': MVTEC_TO_BRATS, 'mvtec_to_mvtec': MVTEC_TO_MVTEC,'mvtecfew_to_mvtec': MVTECFEW_TO_MVTEC} def main(args): @@ -45,28 +47,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 @@ -97,6 +112,15 @@ def main(args): best_pro = 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() @@ -236,7 +260,30 @@ def main(args): 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 +297,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): From 29e9796ba79513335f7270fd43e342715a92aac5 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Tue, 1 Jul 2025 01:08:18 +0900 Subject: [PATCH 040/258] =?UTF-8?q?main=5Fib=E3=82=92main=E3=81=A8?= =?UTF-8?q?=E6=8F=83=E3=81=88=E3=82=8B?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- main_ib.py | 1 + 1 file changed, 1 insertion(+) diff --git a/main_ib.py b/main_ib.py index aaf5c26..ef5b9a8 100644 --- a/main_ib.py +++ b/main_ib.py @@ -399,6 +399,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="") From 3b06c32b9b02227b7c6180860d11815175e0361b Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Tue, 1 Jul 2025 01:12:32 +0900 Subject: [PATCH 041/258] Update main_ib.py --- main_ib.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/main_ib.py b/main_ib.py index ef5b9a8..3cfd207 100644 --- a/main_ib.py +++ b/main_ib.py @@ -140,7 +140,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) From 54d8c9326619e25514468418b11a8e3e329651d7 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Tue, 1 Jul 2025 01:33:01 +0900 Subject: [PATCH 042/258] Update main_ib.py --- main_ib.py | 10 ++++++++++ 1 file changed, 10 insertions(+) diff --git a/main_ib.py b/main_ib.py index 3cfd207..1efe6b2 100644 --- a/main_ib.py +++ b/main_ib.py @@ -211,6 +211,16 @@ def main(args): scheduler1.step() progress_bar.close() + 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}") From cff925f37f9654900f56982c5fc39f402523312e Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Tue, 1 Jul 2025 01:49:55 +0900 Subject: [PATCH 043/258] Update main_ib.py --- main_ib.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/main_ib.py b/main_ib.py index 1efe6b2..b9bfbcf 100644 --- a/main_ib.py +++ b/main_ib.py @@ -221,9 +221,9 @@ def main(args): 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}") + #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 = [], [], [] From 1230492ce9f14d9e2428c4f7500602170f9f68e0 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Tue, 1 Jul 2025 02:11:26 +0900 Subject: [PATCH 044/258] Update main_ib.py --- main_ib.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/main_ib.py b/main_ib.py index b9bfbcf..7ac118c 100644 --- a/main_ib.py +++ b/main_ib.py @@ -228,6 +228,9 @@ def main(args): 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) + # 各クラスの評価結果とデータを一時的に保持する辞書 + current_epoch_class_data_for_saving = {} + 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, From 6142cd9bb3a7a5269904f7da65b84942df662b45 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Tue, 1 Jul 2025 02:32:13 +0900 Subject: [PATCH 045/258] Update main_ib.py --- main_ib.py | 36 +++++++++++++++++++----------------- 1 file changed, 19 insertions(+), 17 deletions(-) diff --git a/main_ib.py b/main_ib.py index 7ac118c..102894a 100644 --- a/main_ib.py +++ b/main_ib.py @@ -231,41 +231,43 @@ def main(args): # 各クラスの評価結果とデータを一時的に保持する辞書 current_epoch_class_data_for_saving = {} - 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, + 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_ref_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( From ba46e84df64a9f01fc7539271b24e128339d9f42 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Tue, 1 Jul 2025 02:49:55 +0900 Subject: [PATCH 046/258] Update main_ib.py --- main_ib.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/main_ib.py b/main_ib.py index 102894a..e195c01 100644 --- a/main_ib.py +++ b/main_ib.py @@ -266,7 +266,7 @@ def main(args): 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_eval], args.device, class_name_eval) + 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'] From cf98326ed78abc4907ddddcb4ac3ecded1eaddf1 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Tue, 1 Jul 2025 03:06:06 +0900 Subject: [PATCH 047/258] Update main_ib.py --- main_ib.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/main_ib.py b/main_ib.py index e195c01..3b16cf4 100644 --- a/main_ib.py +++ b/main_ib.py @@ -271,7 +271,7 @@ def main(args): 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']) From 53100e7c1108c7b6b54c0b780dd378f20b82aaf0 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Tue, 1 Jul 2025 03:28:30 +0900 Subject: [PATCH 048/258] Update main_ib.py --- main_ib.py | 1 + 1 file changed, 1 insertion(+) diff --git a/main_ib.py b/main_ib.py index 3b16cf4..4fc059a 100644 --- a/main_ib.py +++ b/main_ib.py @@ -111,6 +111,7 @@ 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 # 可視化オブジェクトの初期化 # 可視化結果を保存するディレクトリを指定 From 791f82dbece6330bbbc506a9f8d653831e5dfb7c Mon Sep 17 00:00:00 2001 From: Tomoya Ueno Date: Sat, 5 Jul 2025 00:21:07 +0900 Subject: [PATCH 049/258] =?UTF-8?q?=E6=AE=8B=E5=B7=AE=E3=82=92=E4=BD=BF?= =?UTF-8?q?=E3=82=8F=E3=81=AA=E3=81=84=E8=A8=AD=E5=AE=9A=E3=82=92=E8=BF=BD?= =?UTF-8?q?=E5=8A=A0?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- main.py | 8 +++++++- main_ib.py | 5 +++++ validate.py | 5 +++++ 3 files changed, 17 insertions(+), 1 deletion(-) diff --git a/main.py b/main.py index cb97522..33d1310 100644 --- a/main.py +++ b/main.py @@ -147,7 +147,12 @@ 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, pos_flag=True) - + #残差をつかうかどうか7/5 + if args.residual=='False': + rfeatures = ref_features + else: + rfeatures = rfeatures + lvl_masks = [] for l in range(args.feature_levels): _, _, h, w = rfeatures[l].size() @@ -362,6 +367,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('--residual', type=str, default=True) # flow parameters parser.add_argument('--flow_arch', type=str, default='conditional_flow_model') diff --git a/main_ib.py b/main_ib.py index 4fc059a..e6fac52 100644 --- a/main_ib.py +++ b/main_ib.py @@ -155,6 +155,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 = ref_features + else: + rfeatures = rfeatures lvl_masks = [] for l in range(args.feature_levels): @@ -428,6 +432,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') diff --git a/validate.py b/validate.py index cf0b96c..7ea527a 100644 --- a/validate.py +++ b/validate.py @@ -62,6 +62,11 @@ def validate(args, encoder, vq_ops, constraintor, estimators, test_loader, ref_f 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) + + if args.residual=='False': + rfeatures = ref_features + else: + rfeatures = rfeatures # --- UMAP可視化のために特徴量と異常タイプ名を収集 --- From c8ef4edd1c4576e2adbbb95942d41401fe880210 Mon Sep 17 00:00:00 2001 From: Tomoya Ueno Date: Sat, 5 Jul 2025 00:28:25 +0900 Subject: [PATCH 050/258] =?UTF-8?q?=E6=AE=8B=E5=B7=AE=E3=82=92=E4=BD=BF?= =?UTF-8?q?=E3=82=8F=E3=81=AA=E3=81=84=E3=82=88=E3=81=86=E3=81=AB=E5=A4=89?= =?UTF-8?q?=E6=9B=B4?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- main.py | 2 +- main_ib.py | 2 +- validate.py | 4 ++-- 3 files changed, 4 insertions(+), 4 deletions(-) diff --git a/main.py b/main.py index 33d1310..7a7e45c 100644 --- a/main.py +++ b/main.py @@ -149,7 +149,7 @@ def main(args): rfeatures = get_residual_features(features, mfeatures, pos_flag=True) #残差をつかうかどうか7/5 if args.residual=='False': - rfeatures = ref_features + rfeatures = features else: rfeatures = rfeatures diff --git a/main_ib.py b/main_ib.py index e6fac52..25bbf8b 100644 --- a/main_ib.py +++ b/main_ib.py @@ -156,7 +156,7 @@ def main(args): mfeatures = get_mc_matched_ref_features(features, class_names, ref_features) rfeatures = get_residual_features(features, mfeatures) if args.residual=='False': - rfeatures = ref_features + rfeatures = features else: rfeatures = rfeatures diff --git a/validate.py b/validate.py index 7ea527a..bd92eb0 100644 --- a/validate.py +++ b/validate.py @@ -62,9 +62,9 @@ def validate(args, encoder, vq_ops, constraintor, estimators, test_loader, ref_f 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) - + if args.residual=='False': - rfeatures = ref_features + rfeatures = features else: rfeatures = rfeatures From 203ace5f2848e15e6d9f57f73984c0b29d614605 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sat, 18 Oct 2025 11:23:34 +0900 Subject: [PATCH 051/258] Update main.py --- main.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/main.py b/main.py index 3468d6f..beaccdc 100644 --- a/main.py +++ b/main.py @@ -94,7 +94,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() @@ -229,9 +229,9 @@ 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]} @@ -299,4 +299,4 @@ def load_mc_reference_features(root_dir: str, class_names, device: torch.device, - \ No newline at end of file + From 550c029b4e4aad6dd772734bb248c1abea55cf65 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sat, 18 Oct 2025 11:27:39 +0900 Subject: [PATCH 052/258] Update classes.py --- classes.py | 14 +++++++++++++- 1 file changed, 13 insertions(+), 1 deletion(-) diff --git a/classes.py b/classes.py index 6ab46be..a1fd9cc 100644 --- a/classes.py +++ b/classes.py @@ -36,4 +36,16 @@ 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'] From 7cd89ba5b1f7de5c2357d9d28d3be13a52460f15 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sat, 18 Oct 2025 11:34:41 +0900 Subject: [PATCH 053/258] Update main.py --- main.py | 1 + 1 file changed, 1 insertion(+) diff --git a/main.py b/main.py index beaccdc..38ff98c 100644 --- a/main.py +++ b/main.py @@ -26,6 +26,7 @@ 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 warnings.filterwarnings('ignore') From e5db4c9da147caafa49eced1a9b7efcd878a7b41 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sat, 18 Oct 2025 11:35:46 +0900 Subject: [PATCH 054/258] Update main.py --- main.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/main.py b/main.py index 38ff98c..050ad57 100644 --- a/main.py +++ b/main.py @@ -35,7 +35,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} def main(args): From 340352b7ef6736c86a41defd329fb0adcc68db16 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sat, 18 Oct 2025 11:59:36 +0900 Subject: [PATCH 055/258] Update main.py --- main.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/main.py b/main.py index 050ad57..65eebbb 100644 --- a/main.py +++ b/main.py @@ -236,7 +236,8 @@ def main(args): 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): From d12ed77a76f07a3940d5b4863b8ddd165b861f33 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sat, 18 Oct 2025 12:00:54 +0900 Subject: [PATCH 056/258] Update classes.py --- classes.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/classes.py b/classes.py index a1fd9cc..1de7162 100644 --- a/classes.py +++ b/classes.py @@ -48,4 +48,4 @@ 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'] +'macaroni1', 'macaroni2', 'pcb1', 'pcb2', 'pcb3', 'pcb4', 'pipe_fryum']} From 772f01792e9e62745b8939efcc13e4d2e4526a8a Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sun, 26 Oct 2025 16:31:12 +0900 Subject: [PATCH 057/258] Update validate.py --- validate.py | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/validate.py b/validate.py index 667fd6f..dedf2e8 100644 --- a/validate.py +++ b/validate.py @@ -40,6 +40,10 @@ 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) else: features = encoder.encode_image_from_tensors(image) for i in range(len(features)): @@ -139,4 +143,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 From c21f2c3acdf75254b0c434c2cf49dfa1339ccf33 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sun, 26 Oct 2025 17:01:59 +0900 Subject: [PATCH 058/258] Update main.py --- main.py | 15 +++++++++------ 1 file changed, 9 insertions(+), 6 deletions(-) diff --git a/main.py b/main.py index 65eebbb..6c95fc5 100644 --- a/main.py +++ b/main.py @@ -70,12 +70,15 @@ 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追加 +      features = encoder(image) + mfeatures = get_matched_ref_features(features, ref_features) + rfeatures = get_residual_features(features, mfeatures, pos_flag=True) 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) From daded8892e61cab740caf3b09963b167cccf7121 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sun, 26 Oct 2025 17:03:21 +0900 Subject: [PATCH 059/258] Update main.py --- main.py | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) diff --git a/main.py b/main.py index 6c95fc5..8ba9b3c 100644 --- a/main.py +++ b/main.py @@ -76,9 +76,10 @@ def main(args):   encoder = encoder.to(args.device)   feat_dims = encoder.feature_info.channels() 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) +   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) From cf5b41a57eca0fb1264bdefc7b75eca5d5497d97 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sun, 26 Oct 2025 17:08:24 +0900 Subject: [PATCH 060/258] Update main.py --- main.py | 16 ++++++++-------- 1 file changed, 8 insertions(+), 8 deletions(-) diff --git a/main.py b/main.py index 8ba9b3c..220da6b 100644 --- a/main.py +++ b/main.py @@ -71,15 +71,15 @@ def main(args): 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() + 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() + 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) From 75359f74cbb74ba684f2195ec1a971bee6db2fb0 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sun, 26 Oct 2025 17:21:35 +0900 Subject: [PATCH 061/258] Update extract_ref_features.py --- extract_ref_features.py | 13 +++++++++---- 1 file changed, 9 insertions(+), 4 deletions(-) diff --git a/extract_ref_features.py b/extract_ref_features.py index 384d475..54bbc8a 100644 --- a/extract_ref_features.py +++ b/extract_ref_features.py @@ -113,9 +113,14 @@ 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(args.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(args.device) if args.dataset in SETTINGS.keys(): CLASS_NAMES = SETTINGS[args.dataset] @@ -226,4 +231,4 @@ def main2(args): parser.add_argument('--save_dir', type=str, default="./ref_features/w50/mvtec_4shot") args = parser.parse_args() - main(args) \ No newline at end of file + main(args) From b143c32f756e7364067c369f35093efc0db2ab77 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sun, 26 Oct 2025 17:47:21 +0900 Subject: [PATCH 062/258] Update extract_ref_features.py --- extract_ref_features.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/extract_ref_features.py b/extract_ref_features.py index 54bbc8a..a8cadf3 100644 --- a/extract_ref_features.py +++ b/extract_ref_features.py @@ -229,6 +229,6 @@ def main2(args): parser.add_argument('--dataset', type=str, default="mvtec") parser.add_argument('--few_shot_dir', type=str, default="./4shot/mvtec") 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追加 args = parser.parse_args() main(args) From 5b24649f73cbfa73adc211b9d23d240054f282cb Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sun, 26 Oct 2025 17:49:23 +0900 Subject: [PATCH 063/258] Update extract_ref_features.py --- extract_ref_features.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/extract_ref_features.py b/extract_ref_features.py index a8cadf3..802a273 100644 --- a/extract_ref_features.py +++ b/extract_ref_features.py @@ -116,11 +116,11 @@ def main(args): 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) + 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(args.device) + encoder = encoder.to(device) if args.dataset in SETTINGS.keys(): CLASS_NAMES = SETTINGS[args.dataset] From 317689c1597a9eeb61432ff17b1ce15aec25222a Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sun, 26 Oct 2025 17:56:42 +0900 Subject: [PATCH 064/258] Update extract_ref_features.py --- extract_ref_features.py | 26 +++++++++++++++----------- 1 file changed, 15 insertions(+), 11 deletions(-) diff --git a/extract_ref_features.py b/extract_ref_features.py index 802a273..9afa77e 100644 --- a/extract_ref_features.py +++ b/extract_ref_features.py @@ -142,17 +142,21 @@ def main(args): 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_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) - + #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) np.save(os.path.join(args.save_dir, class_name, 'layer1.npy'), layer1_features.cpu().numpy()) From 26e598ee0246604e1447f63a8a3ff6c20965a9b6 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sun, 26 Oct 2025 17:58:51 +0900 Subject: [PATCH 065/258] Update extract_ref_features.py --- extract_ref_features.py | 12 ++++++------ 1 file changed, 6 insertions(+), 6 deletions(-) diff --git a/extract_ref_features.py b/extract_ref_features.py index 9afa77e..8db1b98 100644 --- a/extract_ref_features.py +++ b/extract_ref_features.py @@ -142,12 +142,12 @@ def main(args): 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_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] From afbff246b97a1fdbc90bb4fc566e42d506d73224 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sun, 26 Oct 2025 18:11:56 +0900 Subject: [PATCH 066/258] Update extract_ref_features.py --- extract_ref_features.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/extract_ref_features.py b/extract_ref_features.py index 8db1b98..2602d3c 100644 --- a/extract_ref_features.py +++ b/extract_ref_features.py @@ -159,7 +159,10 @@ def main(args): #修正終わり 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()) From 8787658d5c5706be5a8ab05109bb2ceb0281089d Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sat, 13 Dec 2025 23:46:39 +0900 Subject: [PATCH 067/258] Add VISA and VISAANO dataset classes This file defines the VISA and VISAANO classes for loading and processing image datasets, including methods for data loading, transformation, and mask handling. --- datasets/capsule_visa.py | 319 +++++++++++++++++++++++++++++++++++++++ 1 file changed, 319 insertions(+) create mode 100644 datasets/capsule_visa.py diff --git a/datasets/capsule_visa.py b/datasets/capsule_visa.py new file mode 100644 index 0000000..11117f6 --- /dev/null +++ b/datasets/capsule_visa.py @@ -0,0 +1,319 @@ +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 VISA(Dataset): + + CLASS_NAMES = ['candle', 'capsules', 'cashew', 'chewinggum', 'fryum', + 'macaroni1', 'macaroni2', 'pcb1', 'pcb2', 'pcb3', 'pcb4', 'pipe_fryum'] + + 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 = {'candle': 0, 'capsules': 1, 'cashew': 2, 'chewinggum': 3, + 'fryum': 4, 'macaroni1': 5, 'macaroni2': 6, 'pcb1': 7, + 'pcb2': 8, 'pcb3': 9, 'pcb4': 10, 'pipe_fryum': 11} + self.idx_to_class = {0: 'candle', 1: 'capsules', 2: 'cashew', 3: 'chewinggum', + 4: 'fryum', 5: 'macaroni1', 6: 'macaroni2', 7: 'pcb1', + 8: 'pcb2', 9: 'pcb3', 10: 'pcb4', 11: 'pipe_fryum'} + + 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 VISAANO(Dataset): + + CLASS_NAMES = ['candle', 'capsules', 'cashew', 'chewinggum', 'fryum', + 'macaroni1', 'macaroni2', 'pcb1', 'pcb2', 'pcb3', 'pcb4', 'pipe_fryum'] + + 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 = {'candle': 0, 'capsules': 1, 'cashew': 2, 'chewinggum': 3, + 'fryum': 4, 'macaroni1': 5, 'macaroni2': 6, 'pcb1': 7, + 'pcb2': 8, 'pcb3': 9, 'pcb4': 10, 'pipe_fryum': 11} + self.idx_to_class = {0: 'candle', 1: 'capsules', 2: 'cashew', 3: 'chewinggum', + 4: 'fryum', 5: 'macaroni1', 6: 'macaroni2', 7: 'pcb1', + 8: 'pcb2', 9: 'pcb3', 10: 'pcb4', 11: 'pipe_fryum'} + + 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 From 1a8a0a9255d3610c35443e2af00836d205df05e4 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sat, 13 Dec 2025 23:49:15 +0900 Subject: [PATCH 068/258] Update capsule_visa.py --- datasets/capsule_visa.py | 26 ++++++++------------------ 1 file changed, 8 insertions(+), 18 deletions(-) diff --git a/datasets/capsule_visa.py b/datasets/capsule_visa.py index 11117f6..aa583f6 100644 --- a/datasets/capsule_visa.py +++ b/datasets/capsule_visa.py @@ -13,10 +13,9 @@ IMAGENET_STD = [0.229, 0.224, 0.225] -class VISA(Dataset): +class VISACAPSULES(Dataset): - CLASS_NAMES = ['candle', 'capsules', 'cashew', 'chewinggum', 'fryum', - 'macaroni1', 'macaroni2', 'pcb1', 'pcb2', 'pcb3', 'pcb4', 'pipe_fryum'] + CLASS_NAMES = ['capsules'] def __init__(self, root: str, @@ -68,12 +67,8 @@ def __init__(self, T.CenterCrop(kwargs.get('crp_size')), T.ToTensor()]) - self.class_to_idx = {'candle': 0, 'capsules': 1, 'cashew': 2, 'chewinggum': 3, - 'fryum': 4, 'macaroni1': 5, 'macaroni2': 6, 'pcb1': 7, - 'pcb2': 8, 'pcb3': 9, 'pcb4': 10, 'pipe_fryum': 11} - self.idx_to_class = {0: 'candle', 1: 'capsules', 2: 'cashew', 3: 'chewinggum', - 4: 'fryum', 5: 'macaroni1', 6: 'macaroni2', 7: 'pcb1', - 8: 'pcb2', 9: 'pcb3', 10: 'pcb4', 11: 'pipe_fryum'} + self.class_to_idx = {'capsules': 0 } + self.idx_to_class = {'capsules': 0 } 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] @@ -156,10 +151,9 @@ def update_class_to_idx(self, class_to_idx): self.idx_to_class = dict(zip(idxs, class_names)) -class VISAANO(Dataset): +class VISACAPSULESANO(Dataset): - CLASS_NAMES = ['candle', 'capsules', 'cashew', 'chewinggum', 'fryum', - 'macaroni1', 'macaroni2', 'pcb1', 'pcb2', 'pcb3', 'pcb4', 'pipe_fryum'] + CLASS_NAMES = ['capsules'] def __init__(self, root: str, @@ -210,12 +204,8 @@ def __init__(self, T.CenterCrop(kwargs.get('crp_size')), T.ToTensor()]) - self.class_to_idx = {'candle': 0, 'capsules': 1, 'cashew': 2, 'chewinggum': 3, - 'fryum': 4, 'macaroni1': 5, 'macaroni2': 6, 'pcb1': 7, - 'pcb2': 8, 'pcb3': 9, 'pcb4': 10, 'pipe_fryum': 11} - self.idx_to_class = {0: 'candle', 1: 'capsules', 2: 'cashew', 3: 'chewinggum', - 4: 'fryum', 5: 'macaroni1', 6: 'macaroni2', 7: 'pcb1', - 8: 'pcb2', 9: 'pcb3', 10: 'pcb4', 11: 'pipe_fryum'} + self.class_to_idx = {0:'capsules'} + 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] From 3b7d1edad09290004863f7ffa70afa8b148154a7 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sat, 13 Dec 2025 23:52:05 +0900 Subject: [PATCH 069/258] Add VISACAPSULES_TO_VISACAPSULES mapping --- classes.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/classes.py b/classes.py index 1de7162..f55bdbe 100644 --- a/classes.py +++ b/classes.py @@ -49,3 +49,6 @@ 'macaroni1', 'macaroni2', 'pcb1', 'pcb2', 'pcb3', 'pcb4', 'pipe_fryum'], 'unseen': ['candle', 'capsules', 'cashew', 'chewinggum', 'fryum', 'macaroni1', 'macaroni2', 'pcb1', 'pcb2', 'pcb3', 'pcb4', 'pipe_fryum']} + +VISACAPSULES_TO_VISACAPSULES = {'seen': ['capsuels'], + 'unseen': ['capsuels']} From 5920e0af4ab6ccd33b7c4c5306df4cc40adb197c Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sat, 13 Dec 2025 23:53:09 +0900 Subject: [PATCH 070/258] Rename class VISACAPSULESANO to CAPSULESANO --- datasets/{capsule_visa.py => capsules.py} | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) rename datasets/{capsule_visa.py => capsules.py} (99%) diff --git a/datasets/capsule_visa.py b/datasets/capsules.py similarity index 99% rename from datasets/capsule_visa.py rename to datasets/capsules.py index aa583f6..2a9b506 100644 --- a/datasets/capsule_visa.py +++ b/datasets/capsules.py @@ -13,7 +13,7 @@ IMAGENET_STD = [0.229, 0.224, 0.225] -class VISACAPSULES(Dataset): +class CAPSULES(Dataset): CLASS_NAMES = ['capsules'] @@ -151,7 +151,7 @@ def update_class_to_idx(self, class_to_idx): self.idx_to_class = dict(zip(idxs, class_names)) -class VISACAPSULESANO(Dataset): +class CAPSULESANO(Dataset): CLASS_NAMES = ['capsules'] From 49ae25fb98fee9bfa4706ca7ce20944497aa71ec Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sat, 13 Dec 2025 23:56:30 +0900 Subject: [PATCH 071/258] Update classes.py --- classes.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/classes.py b/classes.py index f55bdbe..8d06a5e 100644 --- a/classes.py +++ b/classes.py @@ -50,5 +50,5 @@ 'unseen': ['candle', 'capsules', 'cashew', 'chewinggum', 'fryum', 'macaroni1', 'macaroni2', 'pcb1', 'pcb2', 'pcb3', 'pcb4', 'pipe_fryum']} -VISACAPSULES_TO_VISACAPSULES = {'seen': ['capsuels'], +CAPSULES_TO_CAPSULES = {'seen': ['capsuels'], 'unseen': ['capsuels']} From 7d4131624666b3e5c5aae1eb74cecdfd3a5efd3d Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sun, 14 Dec 2025 00:01:42 +0900 Subject: [PATCH 072/258] Update main.py --- main.py | 21 +++++++++++++++++++-- 1 file changed, 19 insertions(+), 2 deletions(-) diff --git a/main.py b/main.py index 220da6b..b87fe9b 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 @@ -27,6 +28,7 @@ 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') @@ -35,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_mvtec':MVTEC_TO_MVTEC, 'visa_to_visa':VISA_TO_VISA} + '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): @@ -43,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 + 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) @@ -270,6 +286,7 @@ 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") From 69a75ab3393065673d9e73c7be91c1aac2b7853f Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sun, 14 Dec 2025 00:03:16 +0900 Subject: [PATCH 073/258] Update main.py --- main.py | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/main.py b/main.py index b87fe9b..1e35948 100644 --- a/main.py +++ b/main.py @@ -195,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) From deabd8058d35619282924b550128cc36764bd725 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sun, 14 Dec 2025 00:13:02 +0900 Subject: [PATCH 074/258] Update main.py --- main.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/main.py b/main.py index 1e35948..59a9350 100644 --- a/main.py +++ b/main.py @@ -46,7 +46,7 @@ def main(args): else: raise ValueError(f"Dataset setting must be in {SETTINGS.keys()}, but got {args.setting}.") - if args.classes == 'capsules' # from mvtec to other datasets + if args.classes == 'capsules': # 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) From 101a7b1ece8e370d2b6a026bd5cbe6e83db30584 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sun, 14 Dec 2025 00:15:26 +0900 Subject: [PATCH 075/258] Update main.py --- main.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/main.py b/main.py index 59a9350..c1a6aca 100644 --- a/main.py +++ b/main.py @@ -46,7 +46,7 @@ def main(args): else: raise ValueError(f"Dataset setting must be in {SETTINGS.keys()}, but got {args.setting}.") - if args.classes == 'capsules': # from mvtec to other datasets + 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) @@ -60,7 +60,7 @@ def main(args): 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 + 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) From 398aa061c93b34ee7c1457dd25ed9cf98be365eb Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sun, 14 Dec 2025 00:19:56 +0900 Subject: [PATCH 076/258] Update classes.py --- classes.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/classes.py b/classes.py index 8d06a5e..529730f 100644 --- a/classes.py +++ b/classes.py @@ -50,5 +50,5 @@ 'unseen': ['candle', 'capsules', 'cashew', 'chewinggum', 'fryum', 'macaroni1', 'macaroni2', 'pcb1', 'pcb2', 'pcb3', 'pcb4', 'pipe_fryum']} -CAPSULES_TO_CAPSULES = {'seen': ['capsuels'], - 'unseen': ['capsuels']} +CAPSULES_TO_CAPSULES = {'seen': ['capsules'], + 'unseen': ['capsules']} From e170345c5e75e7ed8ac34d67ef818c19e6d41fed Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sun, 14 Dec 2025 00:34:52 +0900 Subject: [PATCH 077/258] Update capsules.py --- datasets/capsules.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/datasets/capsules.py b/datasets/capsules.py index 2a9b506..c7111e3 100644 --- a/datasets/capsules.py +++ b/datasets/capsules.py @@ -68,7 +68,7 @@ def __init__(self, T.ToTensor()]) self.class_to_idx = {'capsules': 0 } - self.idx_to_class = {'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] @@ -204,7 +204,7 @@ def __init__(self, T.CenterCrop(kwargs.get('crp_size')), T.ToTensor()]) - self.class_to_idx = {0:'capsules'} + self.class_to_idx = {'capsules':0} self.idx_to_class = {0:'capsules'} def __getitem__(self, idx): From feef1c2355c4eafac426263c40753d9abd5be953 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Tue, 16 Dec 2025 22:10:35 +0900 Subject: [PATCH 078/258] Update extract_ref_features.py --- extract_ref_features.py | 1 + 1 file changed, 1 insertion(+) diff --git a/extract_ref_features.py b/extract_ref_features.py index 2602d3c..28e20bf 100644 --- a/extract_ref_features.py +++ b/extract_ref_features.py @@ -235,6 +235,7 @@ 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('--bgadweight_dir', type=str, default="none")# 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追加 args = parser.parse_args() From b334fe9a9f5ab3de30d57e3a018e1a6f6ac688da Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Tue, 16 Dec 2025 22:15:06 +0900 Subject: [PATCH 079/258] Add load_weights function to utils.py Added load_weights function to load model weights. --- utils.py | 9 ++++++++- 1 file changed, 8 insertions(+), 1 deletion(-) diff --git a/utils.py b/utils.py index e465752..06189ff 100644 --- a/utils.py +++ b/utils.py @@ -275,4 +275,11 @@ 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'])] + print('Loading weights from {}'.format(filename)) From 905b1c567cf65afd0c7301b2ae86ed195b8ac6af Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Tue, 16 Dec 2025 22:21:35 +0900 Subject: [PATCH 080/258] Update extract_ref_features.py --- extract_ref_features.py | 8 +++++++- 1 file changed, 7 insertions(+), 1 deletion(-) diff --git a/extract_ref_features.py b/extract_ref_features.py index 28e20bf..99c44fd 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): @@ -121,7 +122,12 @@ def main(args): 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) + + 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: From 1023ba3deaee58777050b436320cc71bcb884dd2 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 17 Dec 2025 01:20:01 +0900 Subject: [PATCH 081/258] Update extract_ref_features.py --- extract_ref_features.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/extract_ref_features.py b/extract_ref_features.py index 99c44fd..214bff8 100644 --- a/extract_ref_features.py +++ b/extract_ref_features.py @@ -122,7 +122,7 @@ def main(args): 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] From 5b8a30e2a94c42806d8cf1c8f29847ccb5a7fe2b Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 17 Dec 2025 01:23:21 +0900 Subject: [PATCH 082/258] Update extract_ref_features.py --- extract_ref_features.py | 1 + 1 file changed, 1 insertion(+) diff --git a/extract_ref_features.py b/extract_ref_features.py index 214bff8..1db4fbb 100644 --- a/extract_ref_features.py +++ b/extract_ref_features.py @@ -241,6 +241,7 @@ 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="none")# 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追加 From 02f3ad004db4bd0661fe7010955e1a20aaf0f1c2 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 17 Dec 2025 01:29:05 +0900 Subject: [PATCH 083/258] Update extract_ref_features.py --- extract_ref_features.py | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/extract_ref_features.py b/extract_ref_features.py index 1db4fbb..ee614ea 100644 --- a/extract_ref_features.py +++ b/extract_ref_features.py @@ -245,5 +245,9 @@ def main2(args): parser.add_argument('--bgadweight_dir', type=str, default="none")# 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) From 41fb21c24fbb67c280e0029670ad1303032fad62 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 17 Dec 2025 01:33:09 +0900 Subject: [PATCH 084/258] Update extract_ref_features.py --- extract_ref_features.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/extract_ref_features.py b/extract_ref_features.py index ee614ea..76fce5a 100644 --- a/extract_ref_features.py +++ b/extract_ref_features.py @@ -247,7 +247,7 @@ def main2(args): 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('--pos_embed_dim', type=int, default=128) parser.add_argument('--device', type=str, default="cuda:0") args = parser.parse_args() main(args) From c57e4094e99373121c081ab08d212d211d2542f6 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 17 Dec 2025 01:40:03 +0900 Subject: [PATCH 085/258] Update main.py --- main.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/main.py b/main.py index c1a6aca..178cf8c 100644 --- a/main.py +++ b/main.py @@ -294,7 +294,7 @@ def load_mc_reference_features(root_dir: str, class_names, device: torch.device, 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) From 77e2b17bb761a01baee65b14ae59fc60640d71c2 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 17 Dec 2025 01:53:56 +0900 Subject: [PATCH 086/258] Update extract_ref_features.py --- extract_ref_features.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/extract_ref_features.py b/extract_ref_features.py index 76fce5a..ee614ea 100644 --- a/extract_ref_features.py +++ b/extract_ref_features.py @@ -247,7 +247,7 @@ def main2(args): 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=128) + 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) From 97b1d0dec4e1ea07431d07fea699ed96cd2506e2 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 17 Dec 2025 02:02:59 +0900 Subject: [PATCH 087/258] Update utils.py --- utils.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/utils.py b/utils.py index 06189ff..3af5211 100644 --- a/utils.py +++ b/utils.py @@ -281,5 +281,5 @@ 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'])] + #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)) From c2b1012346267e753b63242d44e418d6cd29015b Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Thu, 18 Dec 2025 11:37:02 +0900 Subject: [PATCH 088/258] Update utils.py --- utils.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/utils.py b/utils.py index 3af5211..6be8f02 100644 --- a/utils.py +++ b/utils.py @@ -281,5 +281,5 @@ 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変更 + 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)) From 3ac708e7699a4c3cdee938d4fdcc95007b2c6512 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Thu, 18 Dec 2025 12:30:01 +0900 Subject: [PATCH 089/258] Update utils.py --- utils.py | 31 ++++++++++++++++++++++++++++++- 1 file changed, 30 insertions(+), 1 deletion(-) diff --git a/utils.py b/utils.py index 6be8f02..268a74f 100644 --- a/utils.py +++ b/utils.py @@ -276,10 +276,39 @@ def applying_EFDM(input_features_list, ref_features_list, alpha=0.5): aligned_features_list.append(aligned_features) 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): + # map_locationを追加してロード時のデバイス不整合を防ぐ + 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) + + # どの程度の重みがロードされたかを確認するためのログ + if len(enc_msg.missing_keys) > 0: + print(f" Encoder - Missing keys: {len(enc_msg.missing_keys)} (ignored due to strict=False)") + if len(enc_msg.unexpected_keys) > 0: + print(f" Encoder - Unexpected keys: {len(enc_msg.unexpected_keys)} (ignored due to strict=False)") + else: + print(" Warning: 'encoder_state_dict' not found in checkpoint.") + + # --- Decodersのロード --- + # 元のコードのようにリストを再代入せず、ループで各モデルを更新する + if 'decoder_state_dict' in state: + for i, (decoder, d_state) in enumerate(zip(decoders, state['decoder_state_dict'])): + dec_msg = decoder.load_state_dict(d_state, strict=False) + # 各デコーダーのロード状況も必要に応じて表示 + # print(f" Decoder {i} - Missing: {len(dec_msg.missing_keys)}, Unexpected: {len(dec_msg.unexpected_keys)}") + else: + print(" Warning: 'decoder_state_dict' not found in checkpoint.") From 1bea5ca7534c31823ca6a16f615ae1d21fb3fd7b Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Thu, 18 Dec 2025 12:50:24 +0900 Subject: [PATCH 090/258] Update utils.py --- utils.py | 41 ++++++++++++++++++++++++++--------------- 1 file changed, 26 insertions(+), 15 deletions(-) diff --git a/utils.py b/utils.py index 268a74f..b551b9d 100644 --- a/utils.py +++ b/utils.py @@ -284,31 +284,42 @@ def load_weights(encoder, decoders, filename):#12/16追加 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): - # map_locationを追加してロード時のデバイス不整合を防ぐ + # デバイス不整合を防ぐため cpu でロードしてから配分 state = torch.load(filename, map_location='cpu') - print(f'Loading weights from {filename}') - # --- Encoderのロード --- + # --- Encoder のロード --- if 'encoder_state_dict' in state: enc_msg = encoder.load_state_dict(state['encoder_state_dict'], strict=False) - # どの程度の重みがロードされたかを確認するためのログ - if len(enc_msg.missing_keys) > 0: - print(f" Encoder - Missing keys: {len(enc_msg.missing_keys)} (ignored due to strict=False)") - if len(enc_msg.unexpected_keys) > 0: - print(f" Encoder - Unexpected keys: {len(enc_msg.unexpected_keys)} (ignored due to strict=False)") + # 読み込まれたキーの数を計算 + all_keys = set(encoder.state_dict().keys()) + loaded_keys = all_keys - set(enc_msg.missing_keys) + + print(f"--- Encoder Load Report ---") + print(f" Total parameters in model: {len(all_keys)}") + print(f" Successfully loaded: {len(loaded_keys)}") + print(f" Missing (Not loaded): {len(enc_msg.missing_keys)}") + + # もし一つもロードされていなければ警告 + if len(loaded_keys) == 0: + print(" [WARNING] No weights were loaded into the Encoder! Check key names.") + elif len(enc_msg.missing_keys) > 0: + # 最初の5つだけ具体例を表示(ログが埋まるのを防ぐため) + print(f" Example of missing keys: {enc_msg.missing_keys[:5]}") else: - print(" Warning: 'encoder_state_dict' not found in checkpoint.") + print(" [ERROR] 'encoder_state_dict' not found in the file.") - # --- Decodersのロード --- - # 元のコードのようにリストを再代入せず、ループで各モデルを更新する + # --- Decoders のロード --- if 'decoder_state_dict' in state: + print(f"--- Decoders Load Report ---") + # 元のコードのバグ(decodersリストの上書き)を修正 for i, (decoder, d_state) in enumerate(zip(decoders, state['decoder_state_dict'])): dec_msg = decoder.load_state_dict(d_state, strict=False) - # 各デコーダーのロード状況も必要に応じて表示 - # print(f" Decoder {i} - Missing: {len(dec_msg.missing_keys)}, Unexpected: {len(dec_msg.unexpected_keys)}") + dec_all = len(decoder.state_dict()) + dec_loaded = dec_all - len(dec_msg.missing_keys) + print(f" Decoder {i}: Loaded {dec_loaded}/{dec_all} parameters.") else: - print(" Warning: 'decoder_state_dict' not found in checkpoint.") + print(" [ERROR] 'decoder_state_dict' not found in the file.") + From 99f7c89760ebe597c47a4c51f3b25953a90711c4 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Thu, 18 Dec 2025 13:02:14 +0900 Subject: [PATCH 091/258] Update utils.py --- utils.py | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/utils.py b/utils.py index b551b9d..6164991 100644 --- a/utils.py +++ b/utils.py @@ -312,14 +312,18 @@ def load_weights(encoder, decoders, filename): print(" [ERROR] 'encoder_state_dict' not found in the file.") # --- Decoders のロード --- +# --- Decoders のロード --- if 'decoder_state_dict' in state: print(f"--- Decoders Load Report ---") - # 元のコードのバグ(decodersリストの上書き)を修正 for i, (decoder, d_state) in enumerate(zip(decoders, state['decoder_state_dict'])): dec_msg = decoder.load_state_dict(d_state, strict=False) dec_all = len(decoder.state_dict()) dec_loaded = dec_all - len(dec_msg.missing_keys) print(f" Decoder {i}: Loaded {dec_loaded}/{dec_all} parameters.") + + # 読み込めなかったキーの名前を表示(原因特定のため) + if len(dec_msg.missing_keys) > 0: + print(f" Missing keys in Decoder {i}: {dec_msg.missing_keys}") else: print(" [ERROR] 'decoder_state_dict' not found in the file.") From eb60720e971088d48526458575c31f930e475d2a Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Thu, 18 Dec 2025 13:33:34 +0900 Subject: [PATCH 092/258] Update utils.py --- utils.py | 45 +++++++++++++++++++-------------------------- 1 file changed, 19 insertions(+), 26 deletions(-) diff --git a/utils.py b/utils.py index 6164991..9b3534c 100644 --- a/utils.py +++ b/utils.py @@ -285,45 +285,38 @@ def load_weights(encoder, decoders, filename):#12/16追加 print('Loading weights from {}'.format(filename)) ''' def load_weights(encoder, decoders, filename): - # デバイス不整合を防ぐため cpu でロードしてから配分 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 = set(encoder.state_dict().keys()) - loaded_keys = all_keys - set(enc_msg.missing_keys) - + all_keys = len(encoder.state_dict().keys()) + loaded_keys = all_keys - len(enc_msg.missing_keys) print(f"--- Encoder Load Report ---") - print(f" Total parameters in model: {len(all_keys)}") - print(f" Successfully loaded: {len(loaded_keys)}") - print(f" Missing (Not loaded): {len(enc_msg.missing_keys)}") - - # もし一つもロードされていなければ警告 - if len(loaded_keys) == 0: - print(" [WARNING] No weights were loaded into the Encoder! Check key names.") - elif len(enc_msg.missing_keys) > 0: - # 最初の5つだけ具体例を表示(ログが埋まるのを防ぐため) - print(f" Example of missing keys: {enc_msg.missing_keys[:5]}") + print(f" Successfully loaded: {loaded_keys}/{all_keys}") else: - print(" [ERROR] 'encoder_state_dict' not found in the file.") + print(" [ERROR] 'encoder_state_dict' not found.") # --- Decoders のロード --- -# --- Decoders のロード --- if 'decoder_state_dict' in state: - print(f"--- Decoders Load Report ---") + 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) - dec_all = len(decoder.state_dict()) - dec_loaded = dec_all - len(dec_msg.missing_keys) - print(f" Decoder {i}: Loaded {dec_loaded}/{dec_all} parameters.") - # 読み込めなかったキーの名前を表示(原因特定のため) - if len(dec_msg.missing_keys) > 0: - print(f" Missing keys in Decoder {i}: {dec_msg.missing_keys}") + 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.") + + 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 in the file.") + print(" [ERROR] 'decoder_state_dict' not found.") From bd37d102ea1f849ad662ce1053bc44c2cc3491be Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Fri, 19 Dec 2025 15:39:23 +0900 Subject: [PATCH 093/258] Update utils.py --- utils.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/utils.py b/utils.py index 9b3534c..375eac0 100644 --- a/utils.py +++ b/utils.py @@ -309,10 +309,12 @@ def load_weights(encoder, decoders, filename): 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: From 8dc2912b3223377dfdfa78c79fd9817ce2076b34 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Fri, 19 Dec 2025 15:44:48 +0900 Subject: [PATCH 094/258] Update utils.py --- utils.py | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) diff --git a/utils.py b/utils.py index 375eac0..4aeb6b1 100644 --- a/utils.py +++ b/utils.py @@ -309,12 +309,13 @@ def load_weights(encoder, decoders, filename): 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}") + # 修正箇所: インデントを削除しました + 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: From e7780db2677733306569789ec32baf99f1acd145 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Fri, 2 Jan 2026 21:06:47 +0900 Subject: [PATCH 095/258] Create main_ad.py --- main_ad.py | 324 +++++++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 324 insertions(+) create mode 100644 main_ad.py diff --git a/main_ad.py b/main_ad.py new file mode 100644 index 0000000..91ca7fb --- /dev/null +++ b/main_ad.py @@ -0,0 +1,324 @@ +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 + +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) + 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") + + # 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) From 66fb7779937a412890e644f4ef5e9133be40d369 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Thu, 8 Jan 2026 23:55:56 +0900 Subject: [PATCH 096/258] Update main_ad.py --- main_ad.py | 30 ++++++++++++++++++++++++++++-- 1 file changed, 28 insertions(+), 2 deletions(-) diff --git a/main_ad.py b/main_ad.py index 91ca7fb..894ba45 100644 --- a/main_ad.py +++ b/main_ad.py @@ -7,6 +7,7 @@ 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 @@ -96,6 +97,14 @@ def main(args): 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() + 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) @@ -118,6 +127,7 @@ def main(args): 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: @@ -138,10 +148,18 @@ def main(args): 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): @@ -149,7 +167,9 @@ def main(args): 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 @@ -170,10 +190,13 @@ def main(args): 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 @@ -185,6 +208,7 @@ def main(args): total_num += num scheduler_vq.step() + scheduler_ada.step() #追加1/8 scheduler0.step() scheduler1.step() @@ -258,6 +282,7 @@ def main(args): 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')) @@ -300,6 +325,7 @@ def load_mc_reference_features(root_dir: str, class_names, device: torch.device, 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('--eval_freq', type=int, default=1) parser.add_argument('--backbone', type=str, default="wide_resnet50_2") From 2f6c65624ccd4042e08709629c416800d85e34e3 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Fri, 9 Jan 2026 00:03:53 +0900 Subject: [PATCH 097/258] Update utils.py --- utils.py | 8 ++++++++ 1 file changed, 8 insertions(+) diff --git a/utils.py b/utils.py index 4aeb6b1..78adfd3 100644 --- a/utils.py +++ b/utils.py @@ -322,4 +322,12 @@ def load_weights(encoder, decoders, filename): 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)) From 7a23f3c310fd884aed0760db25b49753878aa3a9 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Fri, 9 Jan 2026 00:04:30 +0900 Subject: [PATCH 098/258] Update main_ib.py --- main_ib.py | 1 + 1 file changed, 1 insertion(+) diff --git a/main_ib.py b/main_ib.py index 25bbf8b..c808333 100644 --- a/main_ib.py +++ b/main_ib.py @@ -25,6 +25,7 @@ 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 42201b411c69c2163cef4ffdff2d71779d538b5d Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Fri, 9 Jan 2026 00:05:14 +0900 Subject: [PATCH 099/258] Update main_ad.py --- main_ad.py | 1 + 1 file changed, 1 insertion(+) diff --git a/main_ad.py b/main_ad.py index 894ba45..78d1b50 100644 --- a/main_ad.py +++ b/main_ad.py @@ -24,6 +24,7 @@ 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 9caa0c4290ae81cacabc5cbf0f3c34d527d6a5b0 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Fri, 9 Jan 2026 00:06:06 +0900 Subject: [PATCH 100/258] Update utils.py --- utils.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/utils.py b/utils.py index 78adfd3..766c358 100644 --- a/utils.py +++ b/utils.py @@ -323,7 +323,7 @@ def load_weights(encoder, decoders, filename): else: print(" [ERROR] 'decoder_state_dict' not found.") - def load_weights_ada(adapter, filename): +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) From ca9414dba12bb99183d32668d6ac26b6cb5641eb Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Fri, 9 Jan 2026 00:08:13 +0900 Subject: [PATCH 101/258] Update main_ad.py --- main_ad.py | 1 + 1 file changed, 1 insertion(+) diff --git a/main_ad.py b/main_ad.py index 78d1b50..83f5eda 100644 --- a/main_ad.py +++ b/main_ad.py @@ -327,6 +327,7 @@ def load_mc_reference_features(root_dir: str, class_names, device: torch.device, 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") From 0baafff8bab24755e73e998903f20a00c289c6d7 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Fri, 9 Jan 2026 00:14:59 +0900 Subject: [PATCH 102/258] Update main_ad.py --- main_ad.py | 1 + 1 file changed, 1 insertion(+) diff --git a/main_ad.py b/main_ad.py index 83f5eda..2843a2a 100644 --- a/main_ad.py +++ b/main_ad.py @@ -98,6 +98,7 @@ def main(args): 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 From fe6593b570baf5289d0c911ddb4f9b69bdda63c7 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sun, 11 Jan 2026 02:01:23 +0900 Subject: [PATCH 103/258] Update classes.py --- classes.py | 13 ------------- 1 file changed, 13 deletions(-) diff --git a/classes.py b/classes.py index 1130b26..529730f 100644 --- a/classes.py +++ b/classes.py @@ -39,18 +39,6 @@ 'unseen': ['brain']} MVTEC_TO_MVTEC = {'seen': ['bottle', 'cable', 'capsule', 'carpet', 'grid', -<<<<<<< HEAD - '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']} - -MVTECFEW_TO_MVTEC = {'seen': ['capsule','screw','transistor'], - 'unseen': ['bottle', 'cable', 'capsule', 'carpet', 'grid', - 'hazelnut', 'leather', 'metal_nut', 'pill', 'screw', - 'tile', 'toothbrush', 'transistor', 'wood', 'zipper']} -======= 'hazelnut', 'leather', 'metal_nut', 'pill', 'screw', 'tile', 'toothbrush', 'transistor', 'wood', 'zipper'], 'unseen': ['bottle', 'cable', 'capsule', 'carpet', 'grid', @@ -64,4 +52,3 @@ CAPSULES_TO_CAPSULES = {'seen': ['capsules'], 'unseen': ['capsules']} ->>>>>>> e7780db2677733306569789ec32baf99f1acd145 From e66c76f1afe7af7fc267ac16395575dc294f3d0d Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sun, 11 Jan 2026 02:02:51 +0900 Subject: [PATCH 104/258] Update extract_ref_features.py --- extract_ref_features.py | 11 ----------- 1 file changed, 11 deletions(-) diff --git a/extract_ref_features.py b/extract_ref_features.py index c666d70..ee614ea 100644 --- a/extract_ref_features.py +++ b/extract_ref_features.py @@ -237,7 +237,6 @@ def main2(args): 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") @@ -245,15 +244,6 @@ def main2(args): parser.add_argument('--flow_arch', type=str, default='conditional_flow_model') parser.add_argument('--bgadweight_dir', type=str, default="none")# 12/16追加 parser.add_argument('--save_dir', type=str, default="./ref_features/w50/mvtec_4shot") -<<<<<<< HEAD - parser.add_argument('--mode', type=str, default='main') - - args = parser.parse_args() - if args.mode == 'main': - main(args) - elif args.mode == 'main2': - main2(args) -======= 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) @@ -261,4 +251,3 @@ def main2(args): parser.add_argument('--device', type=str, default="cuda:0") args = parser.parse_args() main(args) ->>>>>>> e7780db2677733306569789ec32baf99f1acd145 From 90d422366793af556c75d725da62f5e50da01a74 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sun, 11 Jan 2026 02:04:12 +0900 Subject: [PATCH 105/258] Update main.py --- main.py | 166 ++++++-------------------------------------------------- 1 file changed, 18 insertions(+), 148 deletions(-) diff --git a/main.py b/main.py index a7c62af..b9e9b00 100644 --- a/main.py +++ b/main.py @@ -17,11 +17,7 @@ from datasets.mpdd import MPDD from datasets.mvtec_loco import MVTECLOCO from datasets.brats import BRATS -<<<<<<< HEAD -from datasets.mvtec_fewclass import MVTECFEWANO, MVTECFEW -======= from datasets.capsules import CAPSULES, CAPSULESANO ->>>>>>> e7780db2677733306569789ec32baf99f1acd145 from models.fc_flow import load_flow_model from models.modules import MultiScaleConv @@ -30,16 +26,10 @@ 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 -<<<<<<< HEAD -from classes import MVTEC_TO_MPDD, MVTEC_TO_MVTECLOCO, MVTEC_TO_BRATS, MVTEC_TO_MVTEC, MVTECFEW_TO_MVTEC -# visualizerのインポート -from visualizer import Visualizer, denormalization -======= 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 ->>>>>>> e7780db2677733306569789ec32baf99f1acd145 warnings.filterwarnings('ignore') TOTAL_SHOT = 4 # total few-shot reference samples @@ -47,11 +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, -<<<<<<< HEAD - 'mvtec_to_brats': MVTEC_TO_BRATS,'mvtec_to_mvtec': MVTEC_TO_MVTEC, 'mvtecfew_to_mvtec': MVTECFEW_TO_MVTEC} -======= 'mvtec_to_brats': MVTEC_TO_BRATS,'mvtec_to_mvtec':MVTEC_TO_MVTEC, 'visa_to_visa':VISA_TO_VISA, 'capsules_to_capsules': CAPSULES_TO_CAPSULES} ->>>>>>> e7780db2677733306569789ec32baf99f1acd145 def main(args): @@ -59,28 +45,14 @@ def main(args): CLASSES = SETTINGS[args.setting] else: raise ValueError(f"Dataset setting must be in {SETTINGS.keys()}, but got {args.setting}.") -<<<<<<< HEAD - # - if args.train_dataset == 'mvtec_few': - train_dataset1 = MVTECFEW(args.train_dataset_dir, class_name=CLASSES['seen'], train=True, -======= 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, ->>>>>>> e7780db2677733306569789ec32baf99f1acd145 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 ) -<<<<<<< HEAD - 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 - ) -======= 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) @@ -88,7 +60,6 @@ def main(args): train_dataset2, batch_size=args.batch_size, shuffle=True, num_workers=8, drop_last=True ) ->>>>>>> e7780db2677733306569789ec32baf99f1acd145 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", @@ -144,21 +115,8 @@ 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) -<<<<<<< HEAD - #best_pro = 0 #-1になってしまうので変更 -======= ->>>>>>> e7780db2677733306569789ec32baf99f1acd145 best_img_auc = 0 N_batch = 8192 - - # 可視化オブジェクトの初期化 - # 可視化結果を保存するディレクトリを指定 - 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() @@ -171,10 +129,9 @@ def main(args): train_loss_total, total_num = 0, 0 progress_bar = tqdm(total=len(train_loader)) progress_bar.set_description(f"Epoch[{epoch}/{args.epochs}]") -#data aug を適応するならここ? supervisedだからあまり意味ない? for step, batch in enumerate(train_loader): progress_bar.update(1) - images, _, masks, class_names, anomaly_types = batch + images, _, masks, class_names = batch images = images.to(args.device) masks = masks.to(args.device) @@ -185,12 +142,7 @@ 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, pos_flag=True) - #残差をつかうかどうか7/5 - if args.residual=='False': - rfeatures = features - else: - rfeatures = rfeatures - + lvl_masks = [] for l in range(args.feature_levels): _, _, h, w = rfeatures[l].size() @@ -242,14 +194,6 @@ def main(args): 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) -<<<<<<< HEAD - # 各クラスの評価結果とデータを一時的に保持する辞書 - 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, -======= for class_name in CLASSES['unseen']: if args.classes == 'capsules': test_dataset = CAPSULES(args.test_dataset_dir, class_name=class_name, train=False, @@ -257,74 +201,46 @@ def main(args): 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, ->>>>>>> e7780db2677733306569789ec32baf99f1acd145 normalize='w50', img_size=224, crp_size=224, msk_size=224, msk_crp_size=224) - elif class_name_eval in VISA.CLASS_NAMES: - test_dataset = VISA(args.test_dataset_dir, class_name=class_name_eval, train=False, + 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_eval in BTAD.CLASS_NAMES: - test_dataset = BTAD(args.test_dataset_dir, class_name=class_name_eval, train=False, + 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_eval in MVTEC3D.CLASS_NAMES: - test_dataset = MVTEC3D(args.test_dataset_dir, class_name=class_name_eval, train=False, + 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_eval in MPDD.CLASS_NAMES: - test_dataset = MPDD(args.test_dataset_dir, class_name=class_name_eval, train=False, + 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_eval in MVTECLOCO.CLASS_NAMES: - test_dataset = MVTECLOCO(args.test_dataset_dir, class_name=class_name_eval, train=False, + 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_eval in BRATS.CLASS_NAMES: - test_dataset = BRATS(args.test_dataset_dir, class_name=class_name_eval, train=False, + 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_eval)) + 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_eval], args.device, class_name_eval) - #metrics = validate(args, encoder, vq_ops, constraintor, estimators, test_loader, test_ref_features[class_name], args.device, class_name) - + 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_eval, img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro)) + 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']) - - # 可視化結果を保存するのは最終エポックのみ - 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) @@ -338,45 +254,6 @@ 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)) -<<<<<<< HEAD - if img_auc > best_img_auc: #pix_aupro > best_pro: - os.makedirs(args.checkpoint_path, exist_ok=True) - 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() -======= if img_auc > best_img_auc: os.makedirs(args.checkpoint_path, exist_ok=True) best_img_auc = img_auc @@ -385,7 +262,6 @@ def main(args): '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')) ->>>>>>> e7780db2677733306569789ec32baf99f1acd145 def load_mc_reference_features(root_dir: str, class_names, device: torch.device, num_shot=4): @@ -413,7 +289,6 @@ def load_mc_reference_features(root_dir: str, class_names, device: torch.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('--classes', type=str, default="none") parser.add_argument('--train_dataset_dir', type=str, default="") @@ -427,7 +302,6 @@ 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('--residual', type=str, default=True) # flow parameters parser.add_argument('--flow_arch', type=str, default='conditional_flow_model') @@ -448,8 +322,4 @@ def load_mc_reference_features(root_dir: str, class_names, device: torch.device, init_seeds(42) main(args) - - - - From 42b8c8e9cb89b2ac3308b99794e6dd8b62278160 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sun, 11 Jan 2026 02:05:12 +0900 Subject: [PATCH 106/258] Update validate.py --- validate.py | 51 +++++++-------------------------------------------- 1 file changed, 7 insertions(+), 44 deletions(-) diff --git a/validate.py b/validate.py index 9272b3b..dedf2e8 100644 --- a/validate.py +++ b/validate.py @@ -19,15 +19,7 @@ def validate(args, encoder, vq_ops, constraintor, estimators, test_loader, ref_f constraintor.eval() for estimator in estimators: estimator.eval() - - # UMAP可視化のために追加するリスト - all_features_to_return = [] - all_anomaly_types_to_return = [] - all_gts_to_return = [] # 0/1の画像レベルのラベル - # 可視化のために追加 - all_images_raw = [] # 生の画像データ - all_scores_map = [] # スコアマップ - + label_list, gt_mask_list = [], [] logps1_list = [list() for _ in range(args.feature_levels)] logps2_list = [list() for _ in range(args.feature_levels)] @@ -35,15 +27,8 @@ def validate(args, encoder, vq_ops, constraintor, estimators, test_loader, ref_f progress_bar.set_description(f"Evaluating") for idx, batch in enumerate(test_loader): progress_bar.update(1) - # データセットから返される値を修正したMVTEC/MVTECANOクラスの__getitem__を想定 - # image: 画像テンソル - # label: 画像レベルのGTラベル (0:正常, 1:異常) - # mask: ピクセルレベルのGTマスク - # class_name_batch: (元のコードの_に対応) その画像のクラス名 (str) - バッチ内の全画像で同じはず - # anomaly_type_batch: (追加) その画像の異常タイプ名 (str, 例: 'scratch', 'hole', 'good') - image, label, mask, class_name_batch, anomaly_type_batch = batch # ここを変更 - #image, label, mask, _ = batch - all_images_raw.append(image.cpu().numpy()) + + image, label, mask, _ = batch gt_mask_list.append(mask.squeeze(1).cpu().numpy().astype(bool)) label_list.append(label.cpu().numpy().astype(bool).ravel()) @@ -66,26 +51,11 @@ def validate(args, encoder, vq_ops, constraintor, estimators, test_loader, ref_f 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) - - if args.residual=='False': - rfeatures = features - else: - rfeatures = rfeatures - - # --- UMAP可視化のために特徴量と異常タイプ名を収集 --- - # ここで、UMAPに渡す特徴量を決定します。 - # 通常、最も深い層(最後の要素)の特徴量をフラットにして使います。 fdm_features = vq_ops(rfeatures, train=False) rfeatures = applying_EFDM(rfeatures, fdm_features, alpha=args.fdm_alpha) - rfeatures = constraintor(*rfeatures) - - current_features_flat = rfeatures[-1].cpu().numpy().reshape(image.shape[0], -1) - all_features_to_return.append(current_features_flat) - all_anomaly_types_to_return.extend(anomaly_type_batch) # リストのままextend - all_gts_to_return.extend(label.cpu().numpy()) # labelは0/1のGTラベル - - + rfeatures = constraintor(*rfeatures) + for l in range(args.feature_levels): e = rfeatures[l] # BxCxHxW bs, dim, h, w = e.size() @@ -122,19 +92,12 @@ def validate(args, encoder, vq_ops, constraintor, estimators, test_loader, ref_f 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) - #visualizerを使えるようにするためにtest_imgsを返す + 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] - # UMAP可視化のために追加した戻り値 - metrics['features'] = np.concatenate(all_features_to_return, axis=0) - metrics['anomaly_types'] = np.array(all_anomaly_types_to_return, dtype=object) # 文字列を含むのでobject型 - metrics['gts_labels'] = np.array(all_gts_to_return) # 0/1のGTラベル - # 可視化のために追加 - metrics['images_raw'] = np.concatenate(all_images_raw, axis=0) - metrics['scores_map'] = scores - metrics['gt_masks_raw'] = gt_masks + return metrics From 216b75954108c1f69b113e47ec3b71ab6c03850f Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sun, 11 Jan 2026 02:05:57 +0900 Subject: [PATCH 107/258] Update visualizer.py --- visualizer.py | 19 +++++-------------- 1 file changed, 5 insertions(+), 14 deletions(-) diff --git a/visualizer.py b/visualizer.py index 53c3bd8..a5a9a4e 100644 --- a/visualizer.py +++ b/visualizer.py @@ -29,19 +29,10 @@ def plot(self, test_imgs, scores, gt_masks): scores (ndarray): shape (N, h, w) gt_masks (ndarray): shape (N, 1, h, w) """ - #7/5 エラー修正 - if args.residual: - vmax = scores.max() * 255. - vmin = scores.min() * 255. + 80 - vmax = vmax - 20 - norm = matplotlib.colors.Normalize(vmin=vmin, vmax=vmax) - else: - vmax = scores.max() * 255. - vmin = scores.min() * 255. - if vmin == vmax: - vmax +=1e-6 - norm = matplotlib.colors.Normalize(vmin=vmin, vmax=vmax) - + vmax = scores.max() * 255. + 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) @@ -80,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() From 42ce5605fd63a76dcb726446fec054af85e5a620 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sun, 15 Feb 2026 00:32:38 +0900 Subject: [PATCH 108/258] Create validate1.py --- validate1.py | 148 +++++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 148 insertions(+) create mode 100644 validate1.py diff --git a/validate1.py b/validate1.py new file mode 100644 index 0000000..1d30daf --- /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 +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 + 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(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 From 96e81d0faf993c19d60a095b4e41b47118699679 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sun, 15 Feb 2026 00:38:00 +0900 Subject: [PATCH 109/258] Implement get_matched_ref_features_top function Added a function to get matched reference features based on rank. --- utils.py | 30 ++++++++++++++++++++++++++++++ 1 file changed, 30 insertions(+) diff --git a/utils.py b/utils.py index 766c358..30ac0c3 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 = [] From f9d7d01859cb9ddfb50d28a6d31d1152e7eae198 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sun, 15 Feb 2026 00:38:47 +0900 Subject: [PATCH 110/258] Add get_matched_ref_features_top import to validate1.py --- validate1.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/validate1.py b/validate1.py index 1d30daf..590bb86 100644 --- a/validate1.py +++ b/validate1.py @@ -9,7 +9,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_matched_ref_features_top from utils import calculate_metrics, applying_EFDM from losses.utils import get_logp_a From e7aa16575c247c8620f53873833393f84fa7ad5e Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sun, 15 Feb 2026 00:39:57 +0900 Subject: [PATCH 111/258] Update main.py --- main.py | 1 + 1 file changed, 1 insertion(+) diff --git a/main.py b/main.py index b9e9b00..f366589 100644 --- a/main.py +++ b/main.py @@ -302,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') From 5711032b287629b2bad1313edef50d1980aa94e7 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sun, 15 Feb 2026 00:40:28 +0900 Subject: [PATCH 112/258] Update validate1.py --- validate1.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/validate1.py b/validate1.py index 590bb86..8a7aab0 100644 --- a/validate1.py +++ b/validate1.py @@ -44,7 +44,7 @@ def validate(args, encoder, vq_ops, constraintor, estimators, test_loader, ref_f 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) + 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) From ffcd2c329a9b05cb0531f294aee98a92bab368b4 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sun, 15 Feb 2026 00:44:02 +0900 Subject: [PATCH 113/258] Update validate1.py --- validate1.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/validate1.py b/validate1.py index 8a7aab0..c5fd29d 100644 --- a/validate1.py +++ b/validate1.py @@ -16,7 +16,7 @@ warnings.filterwarnings('ignore') -def validate(args, encoder, vq_ops, constraintor, estimators, test_loader, ref_features, device, class_name): +def validate1(args, encoder, vq_ops, constraintor, estimators, test_loader, ref_features, device, class_name): vq_ops.eval() constraintor.eval() for estimator in estimators: From 322b129470d534583f91a714f7be826a6b91d12f Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sun, 15 Feb 2026 00:57:12 +0900 Subject: [PATCH 114/258] Update validate1.py --- validate1.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/validate1.py b/validate1.py index c5fd29d..270653e 100644 --- a/validate1.py +++ b/validate1.py @@ -30,7 +30,7 @@ def validate1(args, encoder, vq_ops, constraintor, estimators, test_loader, ref_ for idx, batch in enumerate(test_loader): progress_bar.update(1) - image, label, mask, _ = batch + 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()) From eeb729c703fc06d163db15a282d1f0d3065a8d20 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Tue, 17 Feb 2026 08:46:40 +0900 Subject: [PATCH 115/258] Update extract_ref_features.py --- extract_ref_features.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/extract_ref_features.py b/extract_ref_features.py index ee614ea..fbf503e 100644 --- a/extract_ref_features.py +++ b/extract_ref_features.py @@ -242,7 +242,7 @@ def main2(args): 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="none")# 12/16追加 + 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) From 3e169b70adad9ae9791ac98252ef94a42ce4c9a5 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Mon, 13 Apr 2026 02:46:02 +0900 Subject: [PATCH 116/258] Update validate.py --- validate.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/validate.py b/validate.py index dedf2e8..a0268c2 100644 --- a/validate.py +++ b/validate.py @@ -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 gt_mask_list.append(mask.squeeze(1).cpu().numpy().astype(bool)) label_list.append(label.cpu().numpy().astype(bool).ravel()) From 6822cd137e548662be7b86a4843b1ca91e2e6edd Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 15 Apr 2026 12:51:50 +0900 Subject: [PATCH 117/258] Create main_vit.py --- main_vit.py | 1 + 1 file changed, 1 insertion(+) create mode 100644 main_vit.py diff --git a/main_vit.py b/main_vit.py new file mode 100644 index 0000000..8b13789 --- /dev/null +++ b/main_vit.py @@ -0,0 +1 @@ + From a06c3cfdbc849bca805458314bcb5023f8a83190 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 15 Apr 2026 12:52:18 +0900 Subject: [PATCH 118/258] Update main_vit.py --- main_vit.py | 325 ++++++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 325 insertions(+) diff --git a/main_vit.py b/main_vit.py index 8b13789..f366589 100644 --- a/main_vit.py +++ b/main_vit.py @@ -1 +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 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} + + +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) + 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) + From 88e7ca88f6a4479304a0ce2cd0ede0dc1752cf36 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 15 Apr 2026 12:52:48 +0900 Subject: [PATCH 119/258] Create extract_ref_features_vit.py --- extract_ref_features_vit.py | 1 + 1 file changed, 1 insertion(+) create mode 100644 extract_ref_features_vit.py diff --git a/extract_ref_features_vit.py b/extract_ref_features_vit.py new file mode 100644 index 0000000..8b13789 --- /dev/null +++ b/extract_ref_features_vit.py @@ -0,0 +1 @@ + From 1d7c8836547d8f2374f1acc3d633766fee8a2e67 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 15 Apr 2026 12:59:00 +0900 Subject: [PATCH 120/258] Update extract_ref_features_vit.py --- extract_ref_features_vit.py | 252 ++++++++++++++++++++++++++++++++++++ 1 file changed, 252 insertions(+) diff --git a/extract_ref_features_vit.py b/extract_ref_features_vit.py index 8b13789..fbf503e 100644 --- a/extract_ref_features_vit.py +++ b/extract_ref_features_vit.py @@ -1 +1,253 @@ +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 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)) + 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) From 605ace38130ebe8c85bb0a6d1c7563dfd55f3dae Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 15 Apr 2026 13:07:27 +0900 Subject: [PATCH 121/258] Update main_vit.py --- main_vit.py | 27 +++++++++++++++++++++++++++ 1 file changed, 27 insertions(+) diff --git a/main_vit.py b/main_vit.py index f366589..e37e877 100644 --- a/main_vit.py +++ b/main_vit.py @@ -31,7 +31,34 @@ from classes import CAPSULES_TO_CAPSULES warnings.filterwarnings('ignore') +# --- import文の下、定数定義(TOTAL_SHOTなど)の上あたりに追加 --- +class ViTFeatureExtractor(nn.Module): + def __init__(self, model_name='vit_base_patch14_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, From 3b77ec6b617858c30dfe4098f550b77dab668ec0 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 15 Apr 2026 13:10:07 +0900 Subject: [PATCH 122/258] Update main_vit.py --- main_vit.py | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/main_vit.py b/main_vit.py index e37e877..9d71a24 100644 --- a/main_vit.py +++ b/main_vit.py @@ -123,6 +123,11 @@ def main(args): 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_224', 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) From dd9f41115d5b322124598733e7a5cd11e0665709 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 15 Apr 2026 13:12:04 +0900 Subject: [PATCH 123/258] Update extract_ref_features_vit.py --- extract_ref_features_vit.py | 32 +++++++++++++++++++++++++++++++- 1 file changed, 31 insertions(+), 1 deletion(-) diff --git a/extract_ref_features_vit.py b/extract_ref_features_vit.py index fbf503e..f3207ce 100644 --- a/extract_ref_features_vit.py +++ b/extract_ref_features_vit.py @@ -21,8 +21,34 @@ from models.imagebind import ImageBindModel from utils import load_weights +class ViTFeatureExtractor(nn.Module): + def __init__(self, model_name='vit_base_patch14_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', @@ -122,6 +148,10 @@ def main(args): 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) + elif args.backbone == 'vit_base_patch14': + encoder = ViTFeatureExtractor(model_name='vit_base_patch14_224', out_indices=(3, 7, 11)).eval() + encoder = encoder.to(device) + feat_dims = [encoder.embed_dim] * len(encoder.out_indices) 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] From 1ac2bb041653eb6c7ef05850fd28b333c025bded Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 15 Apr 2026 13:40:36 +0900 Subject: [PATCH 124/258] Update extract_ref_features_vit.py --- extract_ref_features_vit.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/extract_ref_features_vit.py b/extract_ref_features_vit.py index f3207ce..c97952b 100644 --- a/extract_ref_features_vit.py +++ b/extract_ref_features_vit.py @@ -22,7 +22,7 @@ from utils import load_weights class ViTFeatureExtractor(nn.Module): - def __init__(self, model_name='vit_base_patch14_224', out_indices=(3, 7, 11)): + 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 From 775a5d19519419e206741cb814df4bb492030269 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 15 Apr 2026 13:42:03 +0900 Subject: [PATCH 125/258] Update extract_ref_features_vit.py --- extract_ref_features_vit.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/extract_ref_features_vit.py b/extract_ref_features_vit.py index c97952b..68cd554 100644 --- a/extract_ref_features_vit.py +++ b/extract_ref_features_vit.py @@ -148,7 +148,7 @@ def main(args): 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) - elif args.backbone == 'vit_base_patch14': + elif args.backbone == 'deit_base_patch16_224': encoder = ViTFeatureExtractor(model_name='vit_base_patch14_224', out_indices=(3, 7, 11)).eval() encoder = encoder.to(device) feat_dims = [encoder.embed_dim] * len(encoder.out_indices) From 0a21193d780bb23d30197a4b307270ca9bc0b45d Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 15 Apr 2026 14:03:52 +0900 Subject: [PATCH 126/258] Update extract_ref_features_vit.py --- extract_ref_features_vit.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/extract_ref_features_vit.py b/extract_ref_features_vit.py index 68cd554..14a73f0 100644 --- a/extract_ref_features_vit.py +++ b/extract_ref_features_vit.py @@ -148,8 +148,8 @@ def main(args): 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) - elif args.backbone == 'deit_base_patch16_224': - encoder = ViTFeatureExtractor(model_name='vit_base_patch14_224', out_indices=(3, 7, 11)).eval() + elif args.backbone == 'vit_base_patch14': + encoder = ViTFeatureExtractor(model_name='vit_base_patch14_reg4_dinov2.lvd142m', out_indices=(3, 7, 11)).eval() encoder = encoder.to(device) feat_dims = [encoder.embed_dim] * len(encoder.out_indices) feat_dims = encoder.feature_info.channels() From a562903351c36f4fd6bb0f48704f53b660aaab64 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 15 Apr 2026 15:30:23 +0900 Subject: [PATCH 127/258] Update extract_ref_features_vit.py --- extract_ref_features_vit.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/extract_ref_features_vit.py b/extract_ref_features_vit.py index 14a73f0..a93458e 100644 --- a/extract_ref_features_vit.py +++ b/extract_ref_features_vit.py @@ -144,15 +144,16 @@ def main(args): 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_patch14_reg4_dinov2.lvd142m', out_indices=(3, 7, 11)).eval() encoder = encoder.to(device) - feat_dims = [encoder.embed_dim] * len(encoder.out_indices) - feat_dims = encoder.feature_info.channels() + 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] From 93bc864a09a7c21efee72bcef6eb039293c86e2a Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 15 Apr 2026 15:32:49 +0900 Subject: [PATCH 128/258] Update extract_ref_features_vit.py --- extract_ref_features_vit.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/extract_ref_features_vit.py b/extract_ref_features_vit.py index a93458e..f535b3c 100644 --- a/extract_ref_features_vit.py +++ b/extract_ref_features_vit.py @@ -24,7 +24,7 @@ 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.vit = timm.create_model(model_name, pretrained=True,dynamic_img_size=True) self.out_indices = out_indices self.patch_size = self.vit.patch_embed.patch_size[0] self.embed_dim = self.vit.embed_dim From d1e89ccda6fc1319a480f0824fa54930ddead456 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 15 Apr 2026 15:36:54 +0900 Subject: [PATCH 129/258] Update main_vit.py --- main_vit.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/main_vit.py b/main_vit.py index 9d71a24..5ed82cb 100644 --- a/main_vit.py +++ b/main_vit.py @@ -33,9 +33,9 @@ warnings.filterwarnings('ignore') # --- import文の下、定数定義(TOTAL_SHOTなど)の上あたりに追加 --- class ViTFeatureExtractor(nn.Module): - def __init__(self, model_name='vit_base_patch14_224', out_indices=(3, 7, 11)): + 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.vit = timm.create_model(model_name, pretrained=True,dynamic_img_size=True) self.out_indices = out_indices self.patch_size = self.vit.patch_embed.patch_size[0] self.embed_dim = self.vit.embed_dim From 2b4f74a3bf651767832563ce889b20a44880694c Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 15 Apr 2026 15:41:01 +0900 Subject: [PATCH 130/258] Update modules.py --- models/modules.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/models/modules.py b/models/modules.py index 6e3353b..3f0f1a5 100644 --- a/models/modules.py +++ b/models/modules.py @@ -4,7 +4,8 @@ import torch.nn as nn from timm.models.resnet import BasicBlock, create_aa, Bottleneck from timm.models.layers import create_attn -from timm.models.layers.create_act import create_act_layer +#from timm.models.layers.create_act import create_act_layer +from timm.layers import create_act_layer from timm.models.layers.helpers import make_divisible from einops import rearrange @@ -389,4 +390,4 @@ def forward(self, layer1_x, layer2_x, layer3_x, layer4_x): return out1, out2, out3, out4 - \ No newline at end of file + From 635ee1c584498ccbc1c0a20f80a606c122267935 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 15 Apr 2026 15:42:46 +0900 Subject: [PATCH 131/258] Update modules.py --- models/modules.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/models/modules.py b/models/modules.py index 3f0f1a5..783d6d9 100644 --- a/models/modules.py +++ b/models/modules.py @@ -3,10 +3,12 @@ import torch import torch.nn as nn from timm.models.resnet import BasicBlock, create_aa, Bottleneck -from timm.models.layers import create_attn +#from timm.models.layers import create_attn +from timm.layers import create_attn #from timm.models.layers.create_act import create_act_layer from timm.layers import create_act_layer -from timm.models.layers.helpers import make_divisible +#from timm.models.layers.helpers import make_divisible +from timm.layers import make_divisible from einops import rearrange From f53be479588cda1a50f642009f9008960133edc0 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 15 Apr 2026 15:43:32 +0900 Subject: [PATCH 132/258] Update main_vit.py --- main_vit.py | 1 + 1 file changed, 1 insertion(+) diff --git a/main_vit.py b/main_vit.py index 5ed82cb..c6b438e 100644 --- a/main_vit.py +++ b/main_vit.py @@ -29,6 +29,7 @@ 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など)の上あたりに追加 --- From f9657f9d8588d050f8855657be2596a9f2af79e6 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 15 Apr 2026 15:44:45 +0900 Subject: [PATCH 133/258] Update main_vit.py --- main_vit.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/main_vit.py b/main_vit.py index c6b438e..6d6e317 100644 --- a/main_vit.py +++ b/main_vit.py @@ -125,7 +125,7 @@ def main(args): 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_224', out_indices=(3, 7, 11)).eval() + encoder = ViTFeatureExtractor(model_name='vit_base_patch14_reg4_dinov2.lvd142m', out_indices=(3, 7, 11)).eval() encoder = encoder.to(args.device) feat_dims = [encoder.embed_dim] * len(encoder.out_indices) From 82ad67556f4c62e5a15de6e3d97b709dc05420b5 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 15 Apr 2026 16:03:30 +0900 Subject: [PATCH 134/258] Update validate.py --- validate.py | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/validate.py b/validate.py index a0268c2..e3725e1 100644 --- a/validate.py +++ b/validate.py @@ -44,6 +44,10 @@ 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 == '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)): From ab1f772ec27410607742e7089253a5879af5b70f Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Thu, 16 Apr 2026 08:39:30 +0900 Subject: [PATCH 135/258] Update main_vit.py --- main_vit.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/main_vit.py b/main_vit.py index 6d6e317..3c824ba 100644 --- a/main_vit.py +++ b/main_vit.py @@ -164,7 +164,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,_ = batch images = images.to(args.device) masks = masks.to(args.device) From c2f8fcda459225f7e0cbda1dc4774c84eae77f96 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Thu, 16 Apr 2026 08:46:56 +0900 Subject: [PATCH 136/258] Update validate.py --- validate.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/validate.py b/validate.py index e3725e1..4046c69 100644 --- a/validate.py +++ b/validate.py @@ -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()) From d0ecaf43b4d45914888307f2937c114ff8de1da2 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Thu, 16 Apr 2026 08:47:34 +0900 Subject: [PATCH 137/258] Update main_vit.py --- main_vit.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/main_vit.py b/main_vit.py index 3c824ba..11df83a 100644 --- a/main_vit.py +++ b/main_vit.py @@ -164,7 +164,8 @@ 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,_ = batch + images, _, masks, class_names = batch[0], batch[1], batch[2], batch[3] images = images.to(args.device) masks = masks.to(args.device) From fc8bcf631b3a81ee830fc3bbef79fdca9c7ad58f Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sun, 19 Apr 2026 21:26:13 +0900 Subject: [PATCH 138/258] Update modules.py --- models/modules.py | 9 +++------ 1 file changed, 3 insertions(+), 6 deletions(-) diff --git a/models/modules.py b/models/modules.py index 783d6d9..47cda41 100644 --- a/models/modules.py +++ b/models/modules.py @@ -3,12 +3,9 @@ import torch import torch.nn as nn from timm.models.resnet import BasicBlock, create_aa, Bottleneck -#from timm.models.layers import create_attn -from timm.layers import create_attn -#from timm.models.layers.create_act import create_act_layer -from timm.layers import create_act_layer -#from timm.models.layers.helpers import make_divisible -from timm.layers import make_divisible +from timm.models.layers import create_attn +from timm.models.layers.create_act import create_act_layer +from timm.models.layers.helpers import make_divisible from einops import rearrange From 32007cdb6b22a61af8123955811e9973ef143387 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 22 Apr 2026 00:00:09 +0900 Subject: [PATCH 139/258] Create extract_ref_features1.py --- extract_ref_features1.py | 254 +++++++++++++++++++++++++++++++++++++++ 1 file changed, 254 insertions(+) create mode 100644 extract_ref_features1.py diff --git a/extract_ref_features1.py b/extract_ref_features1.py new file mode 100644 index 0000000..1857da3 --- /dev/null +++ b/extract_ref_features1.py @@ -0,0 +1,254 @@ +#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 + +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)) + 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) From 1eb62a9896e8cb42a3690634039de726929bf5c5 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 22 Apr 2026 00:01:28 +0900 Subject: [PATCH 140/258] Create main_1.py --- main_1.py | 327 ++++++++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 327 insertions(+) create mode 100644 main_1.py diff --git a/main_1.py b/main_1.py new file mode 100644 index 0000000..e0a51b7 --- /dev/null +++ b/main_1.py @@ -0,0 +1,327 @@ +#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 + +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} + + +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) + 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) + From 44df1360be3eafd59c51516f45561573b9b702c9 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 22 Apr 2026 00:04:00 +0900 Subject: [PATCH 141/258] Update extract_ref_features1.py --- extract_ref_features1.py | 29 +++++++++++++++++++++++++++++ 1 file changed, 29 insertions(+) diff --git a/extract_ref_features1.py b/extract_ref_features1.py index 1857da3..cc8fe36 100644 --- a/extract_ref_features1.py +++ b/extract_ref_features1.py @@ -21,7 +21,36 @@ 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, From eedcaf8dfb6c4e2d69c60e4f83ce88c90bfee9fd Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 22 Apr 2026 00:04:24 +0900 Subject: [PATCH 142/258] Update main_1.py --- main_1.py | 29 +++++++++++++++++++++++++++++ 1 file changed, 29 insertions(+) diff --git a/main_1.py b/main_1.py index e0a51b7..49df4e0 100644 --- a/main_1.py +++ b/main_1.py @@ -40,7 +40,36 @@ '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] From 7d0f85dde86dc315e3ded279c6b5276ebcf69715 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 22 Apr 2026 00:06:25 +0900 Subject: [PATCH 143/258] Update extract_ref_features1.py --- extract_ref_features1.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/extract_ref_features1.py b/extract_ref_features1.py index cc8fe36..22aff0d 100644 --- a/extract_ref_features1.py +++ b/extract_ref_features1.py @@ -145,14 +145,14 @@ def main(args): 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 = 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() + #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] From 32c8de00ae3da5dda269617ddf76379ccbe00810 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 22 Apr 2026 00:07:46 +0900 Subject: [PATCH 144/258] Update main_1.py --- main_1.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/main_1.py b/main_1.py index 49df4e0..060b78c 100644 --- a/main_1.py +++ b/main_1.py @@ -117,10 +117,10 @@ def main(args): 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 = WideResNetFeatureExtractor(model_name='wide_resnet50_2', out_indices=(1, 2, 3)).eval() encoder = encoder.to(args.device) - feat_dims = encoder.feature_info.channels() + 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/ From bb45c3aa54f27c46b0ae0252b27f1abfd8dbdc65 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 22 Apr 2026 00:35:36 +0900 Subject: [PATCH 145/258] Update main_1.py --- main_1.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/main_1.py b/main_1.py index 060b78c..6cbcaa5 100644 --- a/main_1.py +++ b/main_1.py @@ -8,7 +8,7 @@ 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 4581a0ef30b0f213db6a6a4f2173ae972741c7b4 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 22 Apr 2026 00:37:52 +0900 Subject: [PATCH 146/258] Update main_1.py --- main_1.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/main_1.py b/main_1.py index 6cbcaa5..45922dd 100644 --- a/main_1.py +++ b/main_1.py @@ -161,7 +161,8 @@ 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 = batch + images, _, masks, class_names = batch[0], batch[1], batch[2], batch[3] images = images.to(args.device) masks = masks.to(args.device) From 55b8349a93290ea67f4dd005d05d7e9cc00b58ef Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Tue, 28 Apr 2026 00:46:07 +0900 Subject: [PATCH 147/258] Create main_Fourier --- main_Fourier | 326 +++++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 326 insertions(+) create mode 100644 main_Fourier diff --git a/main_Fourier b/main_Fourier new file mode 100644 index 0000000..f366589 --- /dev/null +++ b/main_Fourier @@ -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 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} + + +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) + 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) + From ac73eefa9d2a15ee588f4cd3d961f128cb16b4d3 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Tue, 28 Apr 2026 00:46:24 +0900 Subject: [PATCH 148/258] Rename main_Fourier to main_Fourier.py --- main_Fourier => main_Fourier.py | 0 1 file changed, 0 insertions(+), 0 deletions(-) rename main_Fourier => main_Fourier.py (100%) diff --git a/main_Fourier b/main_Fourier.py similarity index 100% rename from main_Fourier rename to main_Fourier.py From a3c6964487a1b61943d57030c588e814f0fdc760 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Tue, 28 Apr 2026 00:54:26 +0900 Subject: [PATCH 149/258] Implement Fourier residual feature calculation Add function to compute Fourier residual features from input feature lists using 2D FFT. --- utils.py | 36 ++++++++++++++++++++++++++++++++++++ 1 file changed, 36 insertions(+) diff --git a/utils.py b/utils.py index 30ac0c3..f06a557 100644 --- a/utils.py +++ b/utils.py @@ -99,7 +99,43 @@ def get_residual_features(features: List[Tensor], ref_features: List[Tensor], po residual_features.append(ri) return residual_features +import torch + +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 + + # 元の実装に合わせ、元の特徴量と残差を結合 (オプション) + if pos_flag: + res_f = torch.cat([f, res_f], dim=1) + + rfeatures.append(res_f) + return rfeatures def load_reference_features(root_dir: str, class_name: str, device: torch.device) -> List[Tensor]: """ From 435b8b25d31ddcd7ddddb64402b4b49bd4bc98ee Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Tue, 28 Apr 2026 00:55:55 +0900 Subject: [PATCH 150/258] Update main_Fourier.py --- main_Fourier.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/main_Fourier.py b/main_Fourier.py index f366589..245f300 100644 --- a/main_Fourier.py +++ b/main_Fourier.py @@ -22,7 +22,7 @@ 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 init_seeds, get_residual_features, get_mc_matched_ref_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 1f2edea0da1736b2816cf2f51de90d3c50b3990c Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Tue, 28 Apr 2026 00:56:32 +0900 Subject: [PATCH 151/258] Update main_Fourier.py --- main_Fourier.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/main_Fourier.py b/main_Fourier.py index 245f300..777e5dd 100644 --- a/main_Fourier.py +++ b/main_Fourier.py @@ -141,7 +141,7 @@ 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, pos_flag=True) + rfeatures = get_fourier_residual_features(features, mfeatures, pos_flag=True) lvl_masks = [] for l in range(args.feature_levels): From e035f3543c529c7065471d87dfd3eec399804372 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Tue, 28 Apr 2026 00:57:14 +0900 Subject: [PATCH 152/258] Update validate.py --- validate.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/validate.py b/validate.py index 4046c69..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 From d8a781b73c9b81f8f87222c7d4e67d4d4740cf3d Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Tue, 28 Apr 2026 00:57:57 +0900 Subject: [PATCH 153/258] Update validate.py --- validate.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/validate.py b/validate.py index ed3600a..661bdd1 100644 --- a/validate.py +++ b/validate.py @@ -39,15 +39,15 @@ def validate(args, encoder, vq_ops, constraintor, estimators, test_loader, ref_f 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) + 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) - rfeatures = get_residual_features(features, mfeatures, pos_flag=True) + 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) - rfeatures = get_residual_features(features, mfeatures, pos_flag=True) + rfeatures = get_fourier_residual_features(features, mfeatures, pos_flag=True) else: features = encoder.encode_image_from_tensors(image) for i in range(len(features)): From 2b6bb4884aebd7943bcea3e0432de0d54404aa09 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Tue, 28 Apr 2026 01:07:49 +0900 Subject: [PATCH 154/258] Create validate_Fourier.py --- validate_Fourier.py | 150 ++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 150 insertions(+) create mode 100644 validate_Fourier.py diff --git a/validate_Fourier.py b/validate_Fourier.py new file mode 100644 index 0000000..661bdd1 --- /dev/null +++ b/validate_Fourier.py @@ -0,0 +1,150 @@ +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, 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) + 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) + 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) + 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 From ecda4770d0d77591143bbed48b83cc8a3301233b Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Tue, 28 Apr 2026 01:08:23 +0900 Subject: [PATCH 155/258] Update validate.py --- validate.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/validate.py b/validate.py index 661bdd1..ed3600a 100644 --- a/validate.py +++ b/validate.py @@ -39,15 +39,15 @@ def validate(args, encoder, vq_ops, constraintor, estimators, test_loader, ref_f if args.backbone == 'wide_resnet50_2': features = encoder(image) mfeatures = get_matched_ref_features(features, ref_features) - rfeatures = get_fourier_residual_features(features, mfeatures, pos_flag=True) + 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_fourier_residual_features(features, mfeatures, pos_flag=True) + 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_fourier_residual_features(features, mfeatures, pos_flag=True) + rfeatures = get_residual_features(features, mfeatures, pos_flag=True) else: features = encoder.encode_image_from_tensors(image) for i in range(len(features)): From 65e6d3ea47fc4cff6b69c1c45e8d742fb4a4b519 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Tue, 28 Apr 2026 01:09:00 +0900 Subject: [PATCH 156/258] Update main_Fourier.py --- main_Fourier.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/main_Fourier.py b/main_Fourier.py index 777e5dd..61b7679 100644 --- a/main_Fourier.py +++ b/main_Fourier.py @@ -9,7 +9,7 @@ from torch.utils.data import DataLoader from train import train -from validate import validate +from validate_Fourier import validate from datasets.mvtec import MVTEC, MVTECANO from datasets.visa import VISA, VISAANO from datasets.btad import BTAD From dc628ac4eca8b1ec3f714d23a5e77c430f3bc5cb Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Tue, 28 Apr 2026 01:13:44 +0900 Subject: [PATCH 157/258] Update utils.py --- utils.py | 3 --- 1 file changed, 3 deletions(-) diff --git a/utils.py b/utils.py index f06a557..ce28400 100644 --- a/utils.py +++ b/utils.py @@ -129,9 +129,6 @@ def get_fourier_residual_features(features, mfeatures, pos_flag=True): # 5. 逆フーリエ変換 (IFFT) で空間領域の残差マップに戻す res_f = torch.fft.ifft2(res_fft_new, norm="ortho").real - # 元の実装に合わせ、元の特徴量と残差を結合 (オプション) - if pos_flag: - res_f = torch.cat([f, res_f], dim=1) rfeatures.append(res_f) From 9ebecc29b9cfca4dec09c9f7b553c6ef3487ada1 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Tue, 28 Apr 2026 01:15:39 +0900 Subject: [PATCH 158/258] Update utils.py --- utils.py | 8 ++------ 1 file changed, 2 insertions(+), 6 deletions(-) diff --git a/utils.py b/utils.py index ce28400..3b7aae5 100644 --- a/utils.py +++ b/utils.py @@ -266,12 +266,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 From f52f30afe62148869390c3d1fe81d81676ce325a Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Tue, 28 Apr 2026 23:26:27 +0900 Subject: [PATCH 159/258] Update utils.py --- utils.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/utils.py b/utils.py index 3b7aae5..38a749a 100644 --- a/utils.py +++ b/utils.py @@ -128,7 +128,9 @@ def get_fourier_residual_features(features, mfeatures, pos_flag=True): # 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) From 594a39ff7d3b036900cf31058782e34aa86467ec Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 29 Apr 2026 00:59:18 +0900 Subject: [PATCH 160/258] Add functions for image-level feature matching --- utils.py | 49 +++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 49 insertions(+) diff --git a/utils.py b/utils.py index 38a749a..8e0a937 100644 --- a/utils.py +++ b/utils.py @@ -101,6 +101,55 @@ def get_residual_features(features: List[Tensor], ref_features: List[Tensor], po 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) + + feat_flat = feature.view(B, 1, -1) + core_flat = coreset_spatial.view(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) + + feat_flat = feature.view(1, 1, -1) + core_flat = coreset_spatial.view(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): """ 周波数領域で残差を計算する関数 From d1f82ba497954bd0d8231069dc121882394dbb3d Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 29 Apr 2026 01:00:22 +0900 Subject: [PATCH 161/258] Update main_Fourier.py --- main_Fourier.py | 1 + 1 file changed, 1 insertion(+) diff --git a/main_Fourier.py b/main_Fourier.py index 61b7679..1aef6b7 100644 --- a/main_Fourier.py +++ b/main_Fourier.py @@ -23,6 +23,7 @@ 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 7d8871599ce0d22521609bcd20a9bb398c57c68a Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 29 Apr 2026 01:00:46 +0900 Subject: [PATCH 162/258] Update validate_Fourier.py --- validate_Fourier.py | 1 + 1 file changed, 1 insertion(+) diff --git a/validate_Fourier.py b/validate_Fourier.py index 661bdd1..66fd31d 100644 --- a/validate_Fourier.py +++ b/validate_Fourier.py @@ -8,6 +8,7 @@ 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 From 3c003dacb4e35d8db2f08a2d2600034a31561c4f Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 29 Apr 2026 01:01:44 +0900 Subject: [PATCH 163/258] Update main_Fourier.py --- main_Fourier.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/main_Fourier.py b/main_Fourier.py index 1aef6b7..7eed5d2 100644 --- a/main_Fourier.py +++ b/main_Fourier.py @@ -141,7 +141,8 @@ def main(args): 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_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 = [] From c8af31e5656cf4c35b51068658acf85b3f75e339 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 29 Apr 2026 01:02:37 +0900 Subject: [PATCH 164/258] Update validate_Fourier.py --- validate_Fourier.py | 9 ++++++--- 1 file changed, 6 insertions(+), 3 deletions(-) diff --git a/validate_Fourier.py b/validate_Fourier.py index 66fd31d..de2ebaa 100644 --- a/validate_Fourier.py +++ b/validate_Fourier.py @@ -39,15 +39,18 @@ def validate(args, encoder, vq_ops, constraintor, estimators, test_loader, ref_f with torch.no_grad(): if args.backbone == 'wide_resnet50_2': features = encoder(image) - mfeatures = get_matched_ref_features(features, ref_features) + #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_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_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) From a0a511c80d3ca201d237e6a8106a8e898e685d48 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 29 Apr 2026 01:07:06 +0900 Subject: [PATCH 165/258] Update utils.py --- utils.py | 12 ++++++------ 1 file changed, 6 insertions(+), 6 deletions(-) diff --git a/utils.py b/utils.py index 8e0a937..81bd0e2 100644 --- a/utils.py +++ b/utils.py @@ -109,10 +109,10 @@ def get_image_level_matched_features(features, ref_features): coreset = ref_features[layer_id] K = coreset.shape[0] // (H * W) - coreset_spatial = coreset.view(K, H, W, C).permute(0, 3, 1, 2) + coreset_spatial = coreset.view(K, H, W, C).permute(0, 3, 1, 2).contiguous() - feat_flat = feature.view(B, 1, -1) - core_flat = coreset_spatial.view(1, K, -1) + 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) @@ -135,10 +135,10 @@ def get_mc_image_level_matched_features(features, class_names, ref_features): 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) + coreset_spatial = coreset.view(K, H, W, C).permute(0, 3, 1, 2).contiguous() - feat_flat = feature.view(1, 1, -1) - core_flat = coreset_spatial.view(1, K, -1) + 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) From 4861524a26345c0ddb12b347a04af0aea5fa91c6 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 29 Apr 2026 01:25:30 +0900 Subject: [PATCH 166/258] Create main_attention.py --- main_attention.py | 326 ++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 326 insertions(+) create mode 100644 main_attention.py diff --git a/main_attention.py b/main_attention.py new file mode 100644 index 0000000..f366589 --- /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 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} + + +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) + 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) + From 1aaf485e9a4a475b0348a239d82d4f925b45bd21 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 29 Apr 2026 01:26:11 +0900 Subject: [PATCH 167/258] Update utils.py --- utils.py | 58 ++++++++++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 58 insertions(+) diff --git a/utils.py b/utils.py index 81bd0e2..c9a1030 100644 --- a/utils.py +++ b/utils.py @@ -440,4 +440,62 @@ def load_weights_ada(adapter, filename): #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 From dfdafb4a2a67998e01cdded6e0420151a9230654 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 29 Apr 2026 01:26:40 +0900 Subject: [PATCH 168/258] Update main_attention.py --- main_attention.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/main_attention.py b/main_attention.py index f366589..14aed25 100644 --- a/main_attention.py +++ b/main_attention.py @@ -22,7 +22,7 @@ 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 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 b59eb8d5f94c0ed6ddb14fd353f757b66c4febd1 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 29 Apr 2026 01:27:17 +0900 Subject: [PATCH 169/258] Update main_attention.py --- main_attention.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/main_attention.py b/main_attention.py index 14aed25..cb764a9 100644 --- a/main_attention.py +++ b/main_attention.py @@ -140,7 +140,7 @@ def main(args): 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_soft_matched_features(features, class_names, ref_features) rfeatures = get_residual_features(features, mfeatures, pos_flag=True) lvl_masks = [] From 67a89ad31e4204da49b484b19c29ec0f7430be64 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 29 Apr 2026 01:27:59 +0900 Subject: [PATCH 170/258] Add validation script for model evaluation This script implements a validation function for evaluating models using various estimators and metrics. It includes functions for converting log probabilities to anomaly scores and aggregating those scores across feature levels. --- validate_attention.py | 150 ++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 150 insertions(+) create mode 100644 validate_attention.py diff --git a/validate_attention.py b/validate_attention.py new file mode 100644 index 0000000..ed3600a --- /dev/null +++ b/validate_attention.py @@ -0,0 +1,150 @@ +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, 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) + 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)): + 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 From 5d8eff5773ace08f5f8713bf6e74f3e6dcebf73f Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 29 Apr 2026 01:29:26 +0900 Subject: [PATCH 171/258] Update validate_attention.py --- validate_attention.py | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) diff --git a/validate_attention.py b/validate_attention.py index ed3600a..8b43a6c 100644 --- a/validate_attention.py +++ b/validate_attention.py @@ -8,6 +8,7 @@ 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 @@ -38,15 +39,15 @@ def validate(args, encoder, vq_ops, constraintor, estimators, test_loader, ref_f with torch.no_grad(): if args.backbone == 'wide_resnet50_2': features = encoder(image) - mfeatures = get_matched_ref_features(features, ref_features) + 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_matched_ref_features(features, ref_features) + 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_matched_ref_features(features, ref_features) + 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) From e65b50405407900a23c00d1e98cc8f56ae55d946 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 29 Apr 2026 01:29:41 +0900 Subject: [PATCH 172/258] Update main_attention.py --- main_attention.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/main_attention.py b/main_attention.py index cb764a9..d2119b7 100644 --- a/main_attention.py +++ b/main_attention.py @@ -9,7 +9,7 @@ from torch.utils.data import DataLoader from train import train -from validate import validate +from validate_attention import validate from datasets.mvtec import MVTEC, MVTECANO from datasets.visa import VISA, VISAANO from datasets.btad import BTAD From aef1fbab69390e4940ebb591bd639220dbbe1413 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 29 Apr 2026 21:28:26 +0900 Subject: [PATCH 173/258] Update main_vit.py --- main_vit.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/main_vit.py b/main_vit.py index 11df83a..51f272d 100644 --- a/main_vit.py +++ b/main_vit.py @@ -125,7 +125,8 @@ def main(args): 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_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) From 01bc34f77bc420e6a7bf63bbbd543c5348bebe63 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 29 Apr 2026 21:28:53 +0900 Subject: [PATCH 174/258] Update extract_ref_features_vit.py --- extract_ref_features_vit.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/extract_ref_features_vit.py b/extract_ref_features_vit.py index f535b3c..e1d1ecf 100644 --- a/extract_ref_features_vit.py +++ b/extract_ref_features_vit.py @@ -151,7 +151,7 @@ def main(args): encoder = encoder.to(device) feat_dims = encoder.feature_info.channels() elif args.backbone == 'vit_base_patch14': - 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(device) feat_dims = [encoder.embed_dim] * len(encoder.out_indices) decoders = [load_flow_model(args, feat_dim) for feat_dim in feat_dims] From 6450b24c744d51bd1226afc6c52e3bed1a3668e8 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 29 Apr 2026 21:30:18 +0900 Subject: [PATCH 175/258] Update extract_ref_features_vit.py --- extract_ref_features_vit.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/extract_ref_features_vit.py b/extract_ref_features_vit.py index e1d1ecf..644d618 100644 --- a/extract_ref_features_vit.py +++ b/extract_ref_features_vit.py @@ -24,7 +24,7 @@ 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,dynamic_img_size=True) + 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 From f0d578f8bdea91bbadc4a7559de5cb7ea2028afd Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 29 Apr 2026 21:30:34 +0900 Subject: [PATCH 176/258] Update main_vit.py --- main_vit.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/main_vit.py b/main_vit.py index 51f272d..73c6c04 100644 --- a/main_vit.py +++ b/main_vit.py @@ -36,7 +36,7 @@ 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,dynamic_img_size=True) + 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 From 6ed7913cbc77323c87f151501ac3d0c3dd733f08 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Thu, 30 Apr 2026 22:44:13 +0900 Subject: [PATCH 177/258] Create extract_ref_features_filter.py --- extract_ref_features_filter.py | 253 +++++++++++++++++++++++++++++++++ 1 file changed, 253 insertions(+) create mode 100644 extract_ref_features_filter.py diff --git a/extract_ref_features_filter.py b/extract_ref_features_filter.py new file mode 100644 index 0000000..fbf503e --- /dev/null +++ b/extract_ref_features_filter.py @@ -0,0 +1,253 @@ +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 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)) + 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) From 777ed2c6e3e59caf37caa655fa653ddb262733cb Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Thu, 30 Apr 2026 22:48:30 +0900 Subject: [PATCH 178/258] Update extract_ref_features_filter.py --- extract_ref_features_filter.py | 17 ++++++++++++++++- 1 file changed, 16 insertions(+), 1 deletion(-) diff --git a/extract_ref_features_filter.py b/extract_ref_features_filter.py index fbf503e..e337d26 100644 --- a/extract_ref_features_filter.py +++ b/extract_ref_features_filter.py @@ -10,7 +10,7 @@ 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 @@ -21,6 +21,21 @@ 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, From 07433f04041dd7f2fd04ed73fdfa0b4d9f8d88a0 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Thu, 30 Apr 2026 22:50:37 +0900 Subject: [PATCH 179/258] Update extract_ref_features_filter.py --- extract_ref_features_filter.py | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) diff --git a/extract_ref_features_filter.py b/extract_ref_features_filter.py index e337d26..d643fe7 100644 --- a/extract_ref_features_filter.py +++ b/extract_ref_features_filter.py @@ -160,9 +160,10 @@ def main(args): 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]) + 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) From 18cd7e133d8911f15e42d6bfa9039dbfdc29eb93 Mon Sep 17 00:00:00 2001 From: KudanLabo <108343740+KudanLabo@users.noreply.github.com> Date: Fri, 1 May 2026 11:29:46 +0900 Subject: [PATCH 180/258] Update utils.py --- utils.py | 14 ++++++++++++++ 1 file changed, 14 insertions(+) diff --git a/utils.py b/utils.py index c9a1030..286a69a 100644 --- a/utils.py +++ b/utils.py @@ -289,6 +289,20 @@ 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). """ + #scaling + if False: + global_scaler = StandardScaler() + global_scaler.fit(scores.flatten()) + new_scores = [] + for _, s in zip(gt_masks, scores): + local_scaler.fit(s.flatten()) + local_scaler.sd = global_scaler.sd + local_scaler.transform(s.flatten()) + local_scaler.scale_ = global_scaler.scale_ + local_scaler.var_ = global_scaler.var_ + new_scores.append(local_scaler.transform(s)) + scores = np.array(new_scores) + # 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 From fcd6ca6638fe5b79da2d973b276a85086a199d0d Mon Sep 17 00:00:00 2001 From: KudanLabo <108343740+KudanLabo@users.noreply.github.com> Date: Fri, 1 May 2026 11:34:54 +0900 Subject: [PATCH 181/258] Update utils.py --- utils.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/utils.py b/utils.py index 286a69a..aa8c46b 100644 --- a/utils.py +++ b/utils.py @@ -292,12 +292,12 @@ def calculate_metrics(scores, labels, gt_masks, pro=True, only_max_value=True): #scaling if False: global_scaler = StandardScaler() - global_scaler.fit(scores.flatten()) + global_scaler.fit(np.maximum(scores.flatten(),0)) new_scores = [] for _, s in zip(gt_masks, scores): - local_scaler.fit(s.flatten()) + local_scaler.fit(np.maximum(s.flatten())) local_scaler.sd = global_scaler.sd - local_scaler.transform(s.flatten()) + local_scaler.transform(np.maximum(s.flatten())) local_scaler.scale_ = global_scaler.scale_ local_scaler.var_ = global_scaler.var_ new_scores.append(local_scaler.transform(s)) From 82ac83d83931dc3a043387990284e50ce2cdc973 Mon Sep 17 00:00:00 2001 From: KudanLabo <108343740+KudanLabo@users.noreply.github.com> Date: Fri, 8 May 2026 10:49:57 +0900 Subject: [PATCH 182/258] Update utils.py --- utils.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/utils.py b/utils.py index aa8c46b..f3cf996 100644 --- a/utils.py +++ b/utils.py @@ -295,9 +295,7 @@ def calculate_metrics(scores, labels, gt_masks, pro=True, only_max_value=True): global_scaler.fit(np.maximum(scores.flatten(),0)) new_scores = [] for _, s in zip(gt_masks, scores): - local_scaler.fit(np.maximum(s.flatten())) - local_scaler.sd = global_scaler.sd - local_scaler.transform(np.maximum(s.flatten())) + local_scaler.fit(np.maximum(s.flatten(),0)) local_scaler.scale_ = global_scaler.scale_ local_scaler.var_ = global_scaler.var_ new_scores.append(local_scaler.transform(s)) From 1a5ae16701aa941a9268735415d01797ff2d07e9 Mon Sep 17 00:00:00 2001 From: KudanLabo <108343740+KudanLabo@users.noreply.github.com> Date: Fri, 8 May 2026 10:51:23 +0900 Subject: [PATCH 183/258] Update utils.py --- utils.py | 1 + 1 file changed, 1 insertion(+) diff --git a/utils.py b/utils.py index f3cf996..aeee937 100644 --- a/utils.py +++ b/utils.py @@ -291,6 +291,7 @@ def calculate_metrics(scores, labels, gt_masks, pro=True, only_max_value=True): """ #scaling if False: + local_scaler = StandardScaler() global_scaler = StandardScaler() global_scaler.fit(np.maximum(scores.flatten(),0)) new_scores = [] From 884c29d70c24bf40ef4b194e85029a524d89175a Mon Sep 17 00:00:00 2001 From: KudanLabo <108343740+KudanLabo@users.noreply.github.com> Date: Fri, 8 May 2026 11:07:14 +0900 Subject: [PATCH 184/258] Update utils.py --- utils.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/utils.py b/utils.py index aeee937..904ddbb 100644 --- a/utils.py +++ b/utils.py @@ -293,10 +293,10 @@ def calculate_metrics(scores, labels, gt_masks, pro=True, only_max_value=True): if False: local_scaler = StandardScaler() global_scaler = StandardScaler() - global_scaler.fit(np.maximum(scores.flatten(),0)) + global_scaler.fit(np.maximum(scores.reshape(-1),0)) new_scores = [] for _, s in zip(gt_masks, scores): - local_scaler.fit(np.maximum(s.flatten(),0)) + local_scaler.fit(np.maximum(s.reshape(-1),0)) local_scaler.scale_ = global_scaler.scale_ local_scaler.var_ = global_scaler.var_ new_scores.append(local_scaler.transform(s)) From a05da01537e8c10578891f0e7cfbc0a3a08a5567 Mon Sep 17 00:00:00 2001 From: KudanLabo <108343740+KudanLabo@users.noreply.github.com> Date: Fri, 8 May 2026 11:30:35 +0900 Subject: [PATCH 185/258] Update utils.py --- utils.py | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/utils.py b/utils.py index 904ddbb..6b5308c 100644 --- a/utils.py +++ b/utils.py @@ -293,13 +293,13 @@ def calculate_metrics(scores, labels, gt_masks, pro=True, only_max_value=True): if False: local_scaler = StandardScaler() global_scaler = StandardScaler() - global_scaler.fit(np.maximum(scores.reshape(-1),0)) + global_scaler.fit(np.maximum(scores.reshape(-1,1),0)) new_scores = [] for _, s in zip(gt_masks, scores): - local_scaler.fit(np.maximum(s.reshape(-1),0)) - local_scaler.scale_ = global_scaler.scale_ - local_scaler.var_ = global_scaler.var_ - new_scores.append(local_scaler.transform(s)) + local_scaler.fit(np.maximum(s.reshape(-1,1),0)) + local_scaler.scale_ = global_scaler.mean_ + #local_scaler.var_ = global_scaler.var_ + new_scores.append(local_scaler.transform(s.reshape(-1,1)).reshape(H,W)) scores = np.array(new_scores) # average precision From 2e853d8fe2e445fcfb8180968bde01c5a7159963 Mon Sep 17 00:00:00 2001 From: KudanLabo <108343740+KudanLabo@users.noreply.github.com> Date: Fri, 8 May 2026 16:37:32 +0900 Subject: [PATCH 186/258] Update utils.py --- utils.py | 24 ++++++++++++------------ 1 file changed, 12 insertions(+), 12 deletions(-) diff --git a/utils.py b/utils.py index 6b5308c..f880424 100644 --- a/utils.py +++ b/utils.py @@ -289,18 +289,18 @@ 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). """ - #scaling - if False: - local_scaler = StandardScaler() - global_scaler = StandardScaler() - global_scaler.fit(np.maximum(scores.reshape(-1,1),0)) - new_scores = [] - for _, s in zip(gt_masks, scores): - local_scaler.fit(np.maximum(s.reshape(-1,1),0)) - local_scaler.scale_ = global_scaler.mean_ - #local_scaler.var_ = global_scaler.var_ - new_scores.append(local_scaler.transform(s.reshape(-1,1)).reshape(H,W)) - scores = np.array(new_scores) + #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)) + 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) From c9419dc9183ac54b47b92757d6e0664632fe54cc Mon Sep 17 00:00:00 2001 From: KudanLabo <108343740+KudanLabo@users.noreply.github.com> Date: Fri, 8 May 2026 16:41:02 +0900 Subject: [PATCH 187/258] Update utils.py --- utils.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/utils.py b/utils.py index f880424..8c4f7ca 100644 --- a/utils.py +++ b/utils.py @@ -298,6 +298,9 @@ def calculate_metrics(scores, labels, gt_masks, pro=True, only_max_value=True): 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:] + fn_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] From 27377e8d35f50006c0db9293c0f0f20bd8aa4142 Mon Sep 17 00:00:00 2001 From: KudanLabo <108343740+KudanLabo@users.noreply.github.com> Date: Fri, 8 May 2026 16:41:35 +0900 Subject: [PATCH 188/258] Update utils.py --- utils.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/utils.py b/utils.py index 8c4f7ca..011a0c8 100644 --- a/utils.py +++ b/utils.py @@ -299,7 +299,7 @@ def calculate_metrics(scores, labels, gt_masks, pro=True, only_max_value=True): 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:] - fn_unmask = (rankdata(rankdata(unmask_rank))[labels==0])[: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:] From fb5b9216008b4b13804a51bc517da2f23c741f98 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Mon, 11 May 2026 21:31:56 +0900 Subject: [PATCH 189/258] Update utils.py --- utils.py | 60 ++++++++++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 60 insertions(+) diff --git a/utils.py b/utils.py index 011a0c8..8033b57 100644 --- a/utils.py +++ b/utils.py @@ -515,3 +515,63 @@ def get_mc_soft_matched_features(features: List[Tensor], class_names: List[str], 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 From cfa552e5ecff834fa4edbd663909434d2362ba99 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Mon, 11 May 2026 21:34:33 +0900 Subject: [PATCH 190/258] Create main_osp.py --- main_osp.py | 326 ++++++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 326 insertions(+) create mode 100644 main_osp.py diff --git a/main_osp.py b/main_osp.py new file mode 100644 index 0000000..f366589 --- /dev/null +++ b/main_osp.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 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} + + +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) + 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) + From c0bb28c59bf989f325f06f6d51c7eaac04ac4074 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Mon, 11 May 2026 21:35:36 +0900 Subject: [PATCH 191/258] Create validate_osp.py --- validate_osp.py | 326 ++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 326 insertions(+) create mode 100644 validate_osp.py diff --git a/validate_osp.py b/validate_osp.py new file mode 100644 index 0000000..f366589 --- /dev/null +++ b/validate_osp.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 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} + + +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) + 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) + From 0fb2370b6f11310bec22b8739f8f09a9d479d44a Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Mon, 11 May 2026 21:37:11 +0900 Subject: [PATCH 192/258] Update validate_osp.py --- validate_osp.py | 428 ++++++++++++++---------------------------------- 1 file changed, 126 insertions(+), 302 deletions(-) diff --git a/validate_osp.py b/validate_osp.py index f366589..ed3600a 100644 --- a/validate_osp.py +++ b/validate_osp.py @@ -1,326 +1,150 @@ -import os import warnings -import argparse from tqdm import tqdm +from scipy.ndimage import gaussian_filter 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 +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, applying_EFDM +from losses.utils import get_logp_a 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) +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() - 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) - 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() + 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) + 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)): + 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) - loss = 0 - for l in range(args.feature_levels): # backward svdd loss - e = rfeatures[l] - t = rfeatures_t[l] + + 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) - 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) + # (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: - 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'] + 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)) - 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')) + 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 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) +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() - 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) + # 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 - 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) + #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() - args = parser.parse_args() - init_seeds(42) + # score aggregation + scores = np.zeros_like(abnormal_map[0]) + for l in range(feature_levels): + scores += abnormal_map[l] + scores /= feature_levels - main(args) - + for i in range(scores.shape[0]): + scores[i] = gaussian_filter(scores[i], sigma=4) + + return scores From 3360a44dcb1800872a67d397de290dbd9ec2de73 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Mon, 11 May 2026 21:57:46 +0900 Subject: [PATCH 193/258] Update main_osp.py --- main_osp.py | 211 +++++++++++++++++++++++++++------------------------- 1 file changed, 108 insertions(+), 103 deletions(-) diff --git a/main_osp.py b/main_osp.py index f366589..4c9c41d 100644 --- a/main_osp.py +++ b/main_osp.py @@ -8,8 +8,7 @@ import torch.nn.functional as F from torch.utils.data import DataLoader -from train import train -from validate import validate +from validate_osp import validate # validate_osp をインポート from datasets.mvtec import MVTEC, MVTECANO from datasets.visa import VISA, VISAANO from datasets.btad import BTAD @@ -21,7 +20,6 @@ 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 @@ -32,13 +30,49 @@ warnings.filterwarnings('ignore') -TOTAL_SHOT = 4 # total few-shot reference samples +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(): @@ -46,63 +80,42 @@ def main(args): 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 + 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 - ) + 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 - ) + 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 + 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 - ) + 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 + 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 - ) + 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 - ) + 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 = 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':#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/ + 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) - 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 のみ保持 (vq_ops は削除) 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) @@ -115,20 +128,27 @@ 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) + # SVD行列をキャッシュする辞書 + osp_cache = {} + + from train import train + 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 @@ -143,6 +163,23 @@ def main(args): 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() @@ -150,16 +187,10 @@ def main(args): 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() - + # constraintor適用 rfeatures = constraintor(*rfeatures) loss = 0 - for l in range(args.feature_levels): # backward svdd loss + for l in range(args.feature_levels): e = rfeatures[l] t = rfeatures_t[l] bs, dim, h, w = e.size() @@ -170,6 +201,7 @@ def main(args): loss_i, _, _ = calculate_log_barrier_bi_occ_loss(e, m, t) loss += loss_i + optimizer0.zero_grad() loss.backward() optimizer0.step() @@ -177,14 +209,11 @@ def main(args): 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() @@ -194,45 +223,36 @@ def main(args): 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) + 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) + 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) + 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) + 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) + 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) + 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) + 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) + 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) + + 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( @@ -247,22 +267,13 @@ def main(args): 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(), + 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 = {} @@ -283,10 +294,8 @@ def load_mc_reference_features(root_dir: str, class_names, device: torch.device, 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") @@ -294,7 +303,7 @@ def load_mc_reference_features(root_dir: str, class_names, device: torch.device, 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('--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) @@ -314,13 +323,9 @@ def load_mc_reference_features(root_dir: str, class_names, device: torch.device, 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) - From d73363fc0871fd3bdbfaf8da2b1278cfd171710e Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Mon, 11 May 2026 21:58:14 +0900 Subject: [PATCH 194/258] Integrate OSP application in validate function Refactor validate function to include OSP parameters and apply OSP to residual features. --- validate_osp.py | 61 ++++++++++++++++++++++++------------------------- 1 file changed, 30 insertions(+), 31 deletions(-) diff --git a/validate_osp.py b/validate_osp.py index ed3600a..f110192 100644 --- a/validate_osp.py +++ b/validate_osp.py @@ -7,15 +7,28 @@ 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, applying_EFDM +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, vq_ops, constraintor, estimators, test_loader, ref_features, device, class_name): - vq_ops.eval() +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() @@ -25,6 +38,7 @@ def validate(args, encoder, vq_ops, constraintor, estimators, test_loader, ref_f 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) @@ -36,15 +50,7 @@ def validate(args, encoder, vq_ops, constraintor, estimators, test_loader, ref_f 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(features, ref_features) - rfeatures = get_residual_features(features, mfeatures, pos_flag=True) - elif args.backbone == 'vit_base_patch14': + 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) @@ -56,16 +62,17 @@ def validate(args, encoder, vq_ops, constraintor, estimators, test_loader, ref_f 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) + # --- 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] # BxCxHxW + e = rfeatures[l] 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] @@ -79,8 +86,8 @@ def validate(args, encoder, vq_ops, constraintor, estimators, test_loader, ref_f 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) + 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)) @@ -109,21 +116,16 @@ def convert_to_anomaly_scores(logps_list, feature_levels=3, class_name=None, siz 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() + 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() - # 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) @@ -134,11 +136,8 @@ def aggregate_anomaly_scores(logps_list, feature_levels=3, class_name=None, size 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() + 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] From 47e01737e1787bece433d0698b81a1925cbd22c2 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Mon, 11 May 2026 22:41:24 +0900 Subject: [PATCH 195/258] Update main_osp.py --- main_osp.py | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) diff --git a/main_osp.py b/main_osp.py index 4c9c41d..8e1a87f 100644 --- a/main_osp.py +++ b/main_osp.py @@ -267,7 +267,12 @@ def main(args): 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 From a55e05714dce9602ac597b8e2b3249ed1565ae4b Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Tue, 12 May 2026 20:51:56 +0900 Subject: [PATCH 196/258] Update main_osp.py --- main_osp.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/main_osp.py b/main_osp.py index 8e1a87f..bc38aa7 100644 --- a/main_osp.py +++ b/main_osp.py @@ -189,6 +189,9 @@ def main(args): # 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] From 85675b47ed7d40995274861fa55ed6f21825616a Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 13 May 2026 20:53:53 +0900 Subject: [PATCH 197/258] Create main_wav.py --- main_wav.py | 326 ++++++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 326 insertions(+) create mode 100644 main_wav.py diff --git a/main_wav.py b/main_wav.py new file mode 100644 index 0000000..f366589 --- /dev/null +++ b/main_wav.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 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} + + +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) + 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) + From 1079e244c3361c1b9aecb4941c13838f27606200 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 13 May 2026 20:54:19 +0900 Subject: [PATCH 198/258] Create validate_wav.py --- validate_wav.py | 1 + 1 file changed, 1 insertion(+) create mode 100644 validate_wav.py diff --git a/validate_wav.py b/validate_wav.py new file mode 100644 index 0000000..8b13789 --- /dev/null +++ b/validate_wav.py @@ -0,0 +1 @@ + From 1df6887cf2d41e4b67fa172af16077a9d01beaaf Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 13 May 2026 20:54:31 +0900 Subject: [PATCH 199/258] Update validate_wav.py --- validate_wav.py | 148 ++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 148 insertions(+) diff --git a/validate_wav.py b/validate_wav.py index 8b13789..f110192 100644 --- a/validate_wav.py +++ b/validate_wav.py @@ -1 +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 From 8d24a41f287f3d1c6e07693ce202d70029b0ffe1 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 13 May 2026 20:55:08 +0900 Subject: [PATCH 200/258] Update utils.py --- utils.py | 66 ++++++++++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 66 insertions(+) diff --git a/utils.py b/utils.py index 8033b57..f419fcf 100644 --- a/utils.py +++ b/utils.py @@ -575,3 +575,69 @@ def apply_osp(rfeatures_list, proj_matrices, means): 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.0): + super().__init__() + self.low_freq_weight = low_freq_weight + self.high_freq_weight = high_freq_weight + + # Haarウェーブレットのカーネル定義 (2x2) + # LL: 低周波, HL: 垂直エッジ, LH: 水平エッジ, HH: 対角エッジ + 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('kernel_ll', ll.view(1, 1, 2, 2)) + self.register_buffer('kernel_hl', hl.view(1, 1, 2, 2)) + self.register_buffer('kernel_lh', lh.view(1, 1, 2, 2)) + self.register_buffer('kernel_hh', hh.view(1, 1, 2, 2)) + + def forward(self, x): + B, C, H, W = x.shape + + # チャネルごとに独立して畳み込むため、カーネルをチャネル数分コピー + weight_ll = self.kernel_ll.expand(C, 1, 2, 2) + weight_hl = self.kernel_hl.expand(C, 1, 2, 2) + weight_lh = self.kernel_lh.expand(C, 1, 2, 2) + weight_hh = self.kernel_hh.expand(C, 1, 2, 2) + + # DWT (離散ウェーブレット変換) - stride=2 でダウンサンプリング + x_ll = F.conv2d(x, weight_ll, stride=2, groups=C) + x_hl = F.conv2d(x, weight_hl, stride=2, groups=C) + x_lh = F.conv2d(x, weight_lh, stride=2, groups=C) + x_hh = F.conv2d(x, weight_hh, stride=2, groups=C) + + # --- フィルタリング処理 --- + # 位置ズレ(低周波)を弱め、キズ(高周波)を強調する + x_ll = x_ll * self.low_freq_weight + x_hl = x_hl * self.high_freq_weight + x_lh = x_lh * self.high_freq_weight + x_hh = x_hh * self.high_freq_weight + + # IDWT (逆離散ウェーブレット変換) - stride=2 で元の解像度に戻す + out_ll = F.conv_transpose2d(x_ll, weight_ll, stride=2, groups=C) + out_hl = F.conv_transpose2d(x_hl, weight_hl, stride=2, groups=C) + out_lh = F.conv_transpose2d(x_lh, weight_lh, stride=2, groups=C) + out_hh = F.conv_transpose2d(x_hh, weight_hh, stride=2, groups=C) + + # 全ての成分を足し合わせて再構成 + return out_ll + out_hl + out_lh + out_hh + +def apply_wavelet_filter(rfeatures): + """ + 残差特徴量のリストに対してウェーブレットフィルタを適用する + """ + device = rfeatures[0].device + # 低周波(位置ズレ)を10%に抑え、高周波(キズ)をそのまま通す + wavelet_filter = HaarWaveletFilter(low_freq_weight=0.1, high_freq_weight=1.0).to(device) + + filtered_rfeatures = [] + for rf in rfeatures: + filtered_rfeatures.append(wavelet_filter(rf)) + + return filtered_rfeatures From 10d5e9bc3db78820d3af979693d09e2ef916afb7 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 13 May 2026 21:06:29 +0900 Subject: [PATCH 201/258] Update main_wav.py --- main_wav.py | 249 +++++++++++++++++++++++----------------------------- 1 file changed, 109 insertions(+), 140 deletions(-) diff --git a/main_wav.py b/main_wav.py index f366589..a1ff419 100644 --- a/main_wav.py +++ b/main_wav.py @@ -4,12 +4,12 @@ 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 import validate +from validate_wav import validate # validate_wav をインポート from datasets.mvtec import MVTEC, MVTECANO from datasets.visa import VISA, VISAANO from datasets.btad import BTAD @@ -21,7 +21,6 @@ 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 @@ -32,13 +31,60 @@ warnings.filterwarnings('ignore') -TOTAL_SHOT = 4 # total few-shot reference samples +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 (S2SWCLIP / WMOE 思想) +# ========================================== +class HaarWaveletFilter(nn.Module): + def __init__(self, low_freq_weight=0.1, high_freq_weight=1.5): + """ + 低周波(LL: 位置ズレや全体構造)を抑制し、高周波(HL,LH,HH: キズやエッジ)を強調する + """ + super().__init__() + self.lf_w = low_freq_weight + self.hf_w = high_freq_weight + + # Haarウェーブレットのカーネル (2x2) + 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 + + # チャネルごとに独立してDWTを計算 + 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 + + # IDWTで元の解像度に再構成 + 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(): @@ -46,63 +92,36 @@ def main(args): 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.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() # the pretrained checkpoint will be in /home/.cache/torch/hub/checkpoints/ - encoder = encoder.to(args.device) + 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':#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) + 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) - 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) + # 1. ウェーブレットフィルタの初期化 (低周波を0.1に抑え、高周波を1.2倍に強調) + wav_filter = HaarWaveletFilter(low_freq_weight=0.1, high_freq_weight=1.2).to(args.device) + + # 2. constraintorのみ初期化 (vq_opsは削除) 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) @@ -115,20 +134,24 @@ 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) + from train import train 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 @@ -141,8 +164,13 @@ 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, pos_flag=True) + # --- 【重要】ウェーブレット処理の適用 --- + rfeatures = [wav_filter(rf) for rf in rfeatures] + lvl_masks = [] for l in range(args.feature_levels): _, _, h, w = rfeatures[l].size() @@ -150,16 +178,11 @@ def main(args): 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() - + # --- 【重要】constraintorで空間的に滑らかに整える --- rfeatures = constraintor(*rfeatures) + loss = 0 - for l in range(args.feature_levels): # backward svdd loss + for l in range(args.feature_levels): e = rfeatures[l] t = rfeatures_t[l] bs, dim, h, w = e.size() @@ -167,9 +190,9 @@ def main(args): 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() @@ -177,66 +200,40 @@ def main(args): 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 + # Normalizing Flowの学習 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) + 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) + 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) + # (中略 - 他のデータセットロード処理) 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)) + test_loader = DataLoader(test_dataset, batch_size=1, shuffle=False, num_workers=8, drop_last=False) + + # validate_wav の呼び出し (wav_filter と constraintor を渡す) + 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: {:.3f} | Pixel AUC: {:.3f} | AUPRO: {:.3f}".format( + epoch, class_name, img_auc, pix_auc, pix_aupro)) s1_res.append(metrics['scores1']) s2_res.append(metrics['scores2']) s_res.append(metrics['scores']) @@ -247,46 +244,29 @@ def main(args): 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(), + 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) - + 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 - 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) - + 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") @@ -294,7 +274,7 @@ def load_mc_reference_features(root_dir: str, class_names, device: torch.device, 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('--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) @@ -302,25 +282,14 @@ 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') 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) - From 1b5f1fd0ddd3e40bd964e52ec2e54796f112edfb Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 13 May 2026 21:06:50 +0900 Subject: [PATCH 202/258] Update validate_wav.py --- validate_wav.py | 52 +++++++++++++++---------------------------------- 1 file changed, 16 insertions(+), 36 deletions(-) diff --git a/validate_wav.py b/validate_wav.py index f110192..811c9f2 100644 --- a/validate_wav.py +++ b/validate_wav.py @@ -7,37 +7,25 @@ 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_matched_ref_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): +# validate関数に wav_filter を引数として追加 +def validate(args, encoder, constraintor, wav_filter, estimators, test_loader, ref_features, device, class_name): constraintor.eval() + wav_filter.eval() # ウェーブレットフィルタも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") + progress_bar.set_description(f"Evaluating {class_name}") for idx, batch in enumerate(test_loader): progress_bar.update(1) @@ -50,22 +38,16 @@ def validate(args, encoder, constraintor, estimators, test_loader, ref_features, 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) + features = encoder(image) + mfeatures = get_matched_ref_features(features, ref_features) - # --- VQとEFDMの代わりにOSPを適用 --- - rfeatures = apply_osp(rfeatures, osp_proj_matrices, osp_means) + # 生の残差を計算 + rfeatures = get_residual_features(features, mfeatures, pos_flag=True) - # --- constraintorで滑らかに補正 --- + # --- 【重要】学習時と同じウェーブレット処理で位置ズレノイズを抑制 --- + rfeatures = [wav_filter(rf) for rf in rfeatures] + + # --- 【重要】constraintorで空間的に滑らかに補正 --- rfeatures = constraintor(*rfeatures) for l in range(args.feature_levels): @@ -95,6 +77,7 @@ def validate(args, encoder, constraintor, estimators, test_loader, ref_features, 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) @@ -116,19 +99,17 @@ def convert_to_anomaly_scores(logps_list, feature_levels=3, class_name=None, siz 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) + 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 @@ -145,5 +126,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 From 713d3aed1bddc198730c88a98895dfc60e23184c Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 13 May 2026 21:20:41 +0900 Subject: [PATCH 203/258] Create extract_ref_feature_wav.py --- extract_ref_feature_wav.py | 253 +++++++++++++++++++++++++++++++++++++ 1 file changed, 253 insertions(+) create mode 100644 extract_ref_feature_wav.py diff --git a/extract_ref_feature_wav.py b/extract_ref_feature_wav.py new file mode 100644 index 0000000..fbf503e --- /dev/null +++ b/extract_ref_feature_wav.py @@ -0,0 +1,253 @@ +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 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)) + 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) From fd48a8cfd41f8d09131482d1a3427bfb78af60d2 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 13 May 2026 21:21:03 +0900 Subject: [PATCH 204/258] Rename extract_ref_feature_wav.py to extract_ref_features_wav.py --- extract_ref_feature_wav.py => extract_ref_features_wav.py | 0 1 file changed, 0 insertions(+), 0 deletions(-) rename extract_ref_feature_wav.py => extract_ref_features_wav.py (100%) diff --git a/extract_ref_feature_wav.py b/extract_ref_features_wav.py similarity index 100% rename from extract_ref_feature_wav.py rename to extract_ref_features_wav.py From b5da7e7efa57b4fe9ac82dcf4fc3b14ebae5aef2 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 13 May 2026 21:31:52 +0900 Subject: [PATCH 205/258] Update extract_ref_features_wav.py --- extract_ref_features_wav.py | 358 ++++++++++++++---------------------- 1 file changed, 136 insertions(+), 222 deletions(-) diff --git a/extract_ref_features_wav.py b/extract_ref_features_wav.py index fbf503e..37fbc1b 100644 --- a/extract_ref_features_wav.py +++ b/extract_ref_features_wav.py @@ -1,15 +1,12 @@ import os import argparse import numpy as np -from PIL import Image - +from tqdm import tqdm 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 +import timm +from torch.utils.data import DataLoader from datasets.mvtec import MVTEC from datasets.visa import VISA @@ -18,236 +15,153 @@ 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 +from datasets.capsules import CAPSULES -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) +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, CAPSULES_TO_CAPSULES +from utils import init_seeds - 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 +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 _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} +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): - 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] + init_seeds(42) + device = torch.device(args.device) + + if args.setting in SETTINGS.keys(): + CLASSES = SETTINGS[args.setting] 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()) - + raise ValueError(f"Dataset setting must be in {SETTINGS.keys()}, but got {args.setting}.") -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] + 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(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().to(device) else: - raise ValueError(f"Dataset setting must be in {SETTINGS.keys()}, but got {args.dataset}.") + raise ValueError("Unsupported backbone.") + + wav_filter = HaarWaveletFilter(low_freq_weight=args.lf_weight, high_freq_weight=args.hf_weight).to(device) + wav_filter.eval() + + all_classes = CLASSES['seen'] + CLASSES['unseen'] - 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 = [], [], [], [] + os.makedirs(args.save_dir, exist_ok=True) + print(f"Saving wavelet-filtered reference features to: {args.save_dir}") + print(f"Filter settings -> Low Freq (LL): {args.lf_weight}, High Freq (HL,LH,HH): {args.hf_weight}") + + for class_name in all_classes: + class_save_dir = os.path.join(args.save_dir, class_name) + os.makedirs(class_save_dir, exist_ok=True) - for batch in tqdm.tqdm(train_loader): - images, _, _, _ = batch + if os.path.exists(os.path.join(class_save_dir, 'layer1.npy')): + print(f"Features for {class_name} already exist. Skipping.") + continue + + if args.classes == 'capsules' or class_name in CAPSULES.CLASS_NAMES: + dataset = CAPSULES(args.dataset_dir, class_name=class_name, train=True, normalize='w50') + elif class_name in MVTEC.CLASS_NAMES: + dataset = MVTEC(args.dataset_dir, class_name=class_name, train=True, normalize='w50') + elif class_name in VISA.CLASS_NAMES: + dataset = VISA(args.dataset_dir, class_name=class_name, train=True, normalize='w50') + elif class_name in BTAD.CLASS_NAMES: + dataset = BTAD(args.dataset_dir, class_name=class_name, train=True, normalize='w50') + elif class_name in MVTEC3D.CLASS_NAMES: + dataset = MVTEC3D(args.dataset_dir, class_name=class_name, train=True, normalize='w50') + elif class_name in MPDD.CLASS_NAMES: + dataset = MPDD(args.dataset_dir, class_name=class_name, train=True, normalize='w50') + elif class_name in MVTECLOCO.CLASS_NAMES: + dataset = MVTECLOCO(args.dataset_dir, class_name=class_name, train=True, normalize='w50') + elif class_name in BRATS.CLASS_NAMES: + dataset = BRATS(args.dataset_dir, class_name=class_name, train=True, normalize='w50') + else: + raise ValueError(f"Unknown class {class_name}") + + loader = DataLoader(dataset, batch_size=1, shuffle=True, num_workers=4) + + extracted_features = {0: [], 1: [], 2: []} + count = 0 + + progress_bar = tqdm(loader, desc=f"Extracting {class_name}") + for batch in progress_bar: + if count >= args.train_ref_shot: + break + + images = batch[0].to(device) + 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()) - + features = encoder(images) + + features_wav = [wav_filter(f) for f in features] + + for i, feat in enumerate(features_wav): + flat_feat = feat.permute(0, 2, 3, 1).reshape(-1, feat.shape[1]).cpu().numpy() + extracted_features[i].append(flat_feat) + + count += 1 + + for i in range(3): + layer_feats = np.concatenate(extracted_features[i], axis=0) + save_path = os.path.join(class_save_dir, f'layer{i+1}.npy') + np.save(save_path, layer_feats) -if __name__ == '__main__': + print(f"Saved {count} shots for {class_name}") + + +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('--setting', type=str, default="mvtec_to_mvtec") + parser.add_argument('--classes', type=str, default="none") + parser.add_argument('--dataset_dir', type=str, required=True, help="Path to the dataset directory") + parser.add_argument('--save_dir', type=str, default="./ref_features/w50/mvtec_4shot_wav", help="Directory to save filtered features") + parser.add_argument('--backbone', type=str, default="wide_resnet50_2") parser.add_argument('--device', type=str, default="cuda:0") + parser.add_argument("--train_ref_shot", type=int, default=4, help="Number of reference images to extract") + + 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) From 9e1242b77183a65426f1eea38ca0485fadff24e8 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 13 May 2026 21:41:06 +0900 Subject: [PATCH 206/258] Update extract_ref_features_wav.py --- extract_ref_features_wav.py | 320 +++++++++++++++++++++++++----------- 1 file changed, 226 insertions(+), 94 deletions(-) diff --git a/extract_ref_features_wav.py b/extract_ref_features_wav.py index 37fbc1b..9a41841 100644 --- a/extract_ref_features_wav.py +++ b/extract_ref_features_wav.py @@ -1,12 +1,16 @@ import os import argparse import numpy as np -from tqdm import tqdm +from PIL import Image + import torch +import tqdm +import timm import torch.nn as nn import torch.nn.functional as F -import timm -from torch.utils.data import DataLoader +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 @@ -15,22 +19,12 @@ from datasets.mpdd import MPDD from datasets.mvtec_loco import MVTECLOCO from datasets.brats import BRATS -from datasets.capsules import CAPSULES - -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, CAPSULES_TO_CAPSULES -from utils import init_seeds - -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 -} - +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__() @@ -49,7 +43,6 @@ def __init__(self, low_freq_weight=0.1, high_freq_weight=1.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) @@ -65,101 +58,240 @@ def forward(self, x): 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 main(args): - init_seeds(42) - device = torch.device(args.device) - - 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.backbone == 'wide_resnet50_2': - encoder = timm.create_model('wide_resnet50_2', features_only=True, out_indices=(1, 2, 3), pretrained=True).eval().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().to(device) - else: - raise ValueError("Unsupported backbone.") - - wav_filter = HaarWaveletFilter(low_freq_weight=args.lf_weight, high_freq_weight=args.hf_weight).to(device) - wav_filter.eval() - - all_classes = CLASSES['seen'] + CLASSES['unseen'] + 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 - os.makedirs(args.save_dir, exist_ok=True) - print(f"Saving wavelet-filtered reference features to: {args.save_dir}") - print(f"Filter settings -> Low Freq (LL): {args.lf_weight}, High Freq (HL,LH,HH): {args.hf_weight}") + 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 - for class_name in all_classes: - class_save_dir = os.path.join(args.save_dir, class_name) - os.makedirs(class_save_dir, exist_ok=True) + def _load_data(self, class_name): + image_paths, labels, mask_paths = [], [], [] + phase = 'train' if self.train else 'test' - if os.path.exists(os.path.join(class_save_dir, 'layer1.npy')): - print(f"Features for {class_name} already exist. Skipping.") - continue + image_dir = os.path.join(self.root, class_name, phase) + mask_dir = os.path.join(self.root, class_name, 'ground_truth') - if args.classes == 'capsules' or class_name in CAPSULES.CLASS_NAMES: - dataset = CAPSULES(args.dataset_dir, class_name=class_name, train=True, normalize='w50') - elif class_name in MVTEC.CLASS_NAMES: - dataset = MVTEC(args.dataset_dir, class_name=class_name, train=True, normalize='w50') - elif class_name in VISA.CLASS_NAMES: - dataset = VISA(args.dataset_dir, class_name=class_name, train=True, normalize='w50') - elif class_name in BTAD.CLASS_NAMES: - dataset = BTAD(args.dataset_dir, class_name=class_name, train=True, normalize='w50') - elif class_name in MVTEC3D.CLASS_NAMES: - dataset = MVTEC3D(args.dataset_dir, class_name=class_name, train=True, normalize='w50') - elif class_name in MPDD.CLASS_NAMES: - dataset = MPDD(args.dataset_dir, class_name=class_name, train=True, normalize='w50') - elif class_name in MVTECLOCO.CLASS_NAMES: - dataset = MVTECLOCO(args.dataset_dir, class_name=class_name, train=True, normalize='w50') - elif class_name in BRATS.CLASS_NAMES: - dataset = BRATS(args.dataset_dir, class_name=class_name, train=True, normalize='w50') - else: - raise ValueError(f"Unknown class {class_name}") + 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) - loader = DataLoader(dataset, batch_size=1, shuffle=True, num_workers=4) + 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 + - extracted_features = {0: [], 1: [], 2: []} - count = 0 +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} - progress_bar = tqdm(loader, desc=f"Extracting {class_name}") - for batch in progress_bar: - if count >= args.train_ref_shot: - break - - images = batch[0].to(device) +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(): - features = encoder(images) + patch_tokens = encoder(images.to(device)) + # 【追加】保存用に変換する前に、ウェーブレット変換をかける + patch_tokens = [wav_filter(f) for f in patch_tokens] - features_wav = [wav_filter(f) for f in features] - - for i, feat in enumerate(features_wav): - flat_feat = feat.permute(0, 2, 3, 1).reshape(-1, feat.shape[1]).cpu().numpy() - extracted_features[i].append(flat_feat) + layer1_features.append(patch_tokens[0]) + layer2_features.append(patch_tokens[1]) + layer3_features.append(patch_tokens[2]) - count += 1 - - for i in range(3): - layer_feats = np.concatenate(extracted_features[i], axis=0) - save_path = os.path.join(class_save_dir, f'layer{i+1}.npy') - np.save(save_path, layer_feats) + 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) - print(f"Saved {count} shots for {class_name}") + 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()) + -if __name__ == "__main__": +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('--setting', type=str, default="mvtec_to_mvtec") - parser.add_argument('--classes', type=str, default="none") - parser.add_argument('--dataset_dir', type=str, required=True, help="Path to the dataset directory") - parser.add_argument('--save_dir', type=str, default="./ref_features/w50/mvtec_4shot_wav", help="Directory to save filtered features") + 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("--train_ref_shot", type=int, default=4, help="Number of reference images to extract") + # 追加: ウェーブレット変換のパラメータ 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") From 77fa6342b19b89c3e784a595a1347c792f107ba1 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 13 May 2026 21:45:08 +0900 Subject: [PATCH 207/258] Update main_wav.py --- main_wav.py | 129 +++++++++++++++++++++++++--------------------------- 1 file changed, 62 insertions(+), 67 deletions(-) diff --git a/main_wav.py b/main_wav.py index a1ff419..9a33938 100644 --- a/main_wav.py +++ b/main_wav.py @@ -9,7 +9,7 @@ import torch.nn.functional as F from torch.utils.data import DataLoader -from validate_wav import validate # validate_wav をインポート +from validate_wav import validate from datasets.mvtec import MVTEC, MVTECANO from datasets.visa import VISA, VISAANO from datasets.btad import BTAD @@ -39,18 +39,14 @@ '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 (S2SWCLIP / WMOE 思想) +# Haar Wavelet Filter # ========================================== class HaarWaveletFilter(nn.Module): - def __init__(self, low_freq_weight=0.1, high_freq_weight=1.5): - """ - 低周波(LL: 位置ズレや全体構造)を抑制し、高周波(HL,LH,HH: キズやエッジ)を強調する - """ + 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 - # Haarウェーブレットのカーネル (2x2) 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]]) @@ -63,27 +59,21 @@ def __init__(self, low_freq_weight=0.1, high_freq_weight=1.5): def forward(self, x): B, C, H, W = x.shape - - # チャネルごとに独立してDWTを計算 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 - # IDWTで元の解像度に再構成 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): @@ -110,29 +100,25 @@ def main(args): 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 = encoder.feature_info.channels() boundary_ops = BoundaryAverager(num_levels=args.feature_levels) - # 1. ウェーブレットフィルタの初期化 (低周波を0.1に抑え、高周波を1.2倍に強調) - wav_filter = HaarWaveletFilter(low_freq_weight=0.1, high_freq_weight=1.2).to(args.device) + wav_filter = HaarWaveletFilter(low_freq_weight=args.lf_weight, high_freq_weight=args.hf_weight).to(args.device) + wav_filter.eval() - # 2. 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) + optimizer0 = torch.optim.Adam(constraintor.parameters(), lr=args.lr, weight_decay=0.005) # Weight decay少し強め + scheduler0 = torch.optim.lr_scheduler.MultiStepLR(optimizer0, milestones=[30, 50], 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] + estimators = [load_flow_model(args, feat_dim).to(args.device) for feat_dim in feat_dims] 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) + optimizer1 = torch.optim.Adam(params, lr=args.lr, weight_decay=0.005) + scheduler1 = torch.optim.lr_scheduler.MultiStepLR(optimizer1, milestones=[30, 50], gamma=0.1) from train import train best_img_auc = 0 @@ -143,53 +129,54 @@ def main(args): for estimator in estimators: estimator.train() - if epoch < FIRST_STAGE_EPOCH: - train_loader = train_loader1 - else: - train_loader = train_loader2 - + 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}]") + + progress_bar = tqdm(total=len(train_loader), desc=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) + images, masks = images.to(args.device), masks.to(args.device) with torch.no_grad(): + # 1. 画像から特徴抽出 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) - - # --- 【重要】ウェーブレット処理の適用 --- - rfeatures = [wav_filter(rf) for rf in rfeatures] + + # 2. テスト画像(訓練バッチ)の特徴量をウェーブレット変換 (Pre-filter) + features = [wav_filter(f) for f in features] + + # 3. 訓練用の参照特徴量を抽出し、それらもウェーブレット変換 + ref_features = get_mc_reference_features(encoder, args.train_dataset_dir, class_names, images.device, args.train_ref_shot) + for c_name in ref_features.keys(): + ref_features[c_name] = [wav_filter(rf) for rf in ref_features[c_name]] + + # 4. 「エッジ強調された特徴量同士」でマッチング + mfeatures = get_mc_matched_ref_features(features, class_names, ref_features) + + # 5. 残差の計算 + 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) + lvl_masks.append(F.interpolate(masks, size=(h, w), mode='nearest').squeeze(1)) rfeatures_t = [rfeature.detach().clone() for rfeature in rfeatures] - # --- 【重要】constraintorで空間的に滑らかに整える --- + # 6. Constraintorによる空間補正 rfeatures = constraintor(*rfeatures) + # (任意) 特徴量への微小ノイズ付加による過学習防止 + noise_std = 0.01 + 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] + e = rfeatures_noisy[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) + e, t = e.permute(0, 2, 3, 1).reshape(-1, dim), t.permute(0, 2, 3, 1).reshape(-1, dim) + m = lvl_masks[l].reshape(-1) loss_i, _, _ = calculate_log_barrier_bi_occ_loss(e, m, t) loss += loss_i @@ -201,20 +188,20 @@ def main(args): total_num += 1 rfeatures = [rfeature.detach().clone() for rfeature in rfeatures] - # Normalizing Flowの学習 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 = [], [], [] + + # 既にextract_ref_features_wav.pyでウェーブレット変換済みのカンペをロード 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']: @@ -222,13 +209,24 @@ def main(args): 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 の呼び出し (wav_filter と constraintor を渡す) + # validate_wavに 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'] @@ -238,13 +236,7 @@ def main(args): 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) - + img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro = np.mean(np.array(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)) @@ -273,23 +265,26 @@ def load_mc_reference_features(root_dir: str, class_names, device: torch.device, 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('--test_ref_feature_dir', type=str, default="./ref_features/w50/mvtec_4shot_wav") # デフォルトを _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('--epochs', type=int, default=60) # 過学習防止のためエポック数を削減 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") - # 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('--pos_embed_dim', type=int, default=256) 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, help="Weight for low frequency (LL)") + parser.add_argument("--hf_weight", type=float, default=1.2, help="Weight for high frequency (LH, HL, HH)") + args = parser.parse_args() init_seeds(42) main(args) From b9cab4ac9ba8caa4e8ea268bcd81e7393e94eac0 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 13 May 2026 21:45:38 +0900 Subject: [PATCH 208/258] Update validate_wav.py --- validate_wav.py | 36 ++++++++++++++++++------------------ 1 file changed, 18 insertions(+), 18 deletions(-) diff --git a/validate_wav.py b/validate_wav.py index 811c9f2..92b0b80 100644 --- a/validate_wav.py +++ b/validate_wav.py @@ -13,10 +13,9 @@ warnings.filterwarnings('ignore') -# validate関数に wav_filter を引数として追加 def validate(args, encoder, constraintor, wav_filter, estimators, test_loader, ref_features, device, class_name): constraintor.eval() - wav_filter.eval() # ウェーブレットフィルタもevalモードに + wav_filter.eval() for estimator in estimators: estimator.eval() @@ -24,8 +23,7 @@ def validate(args, encoder, constraintor, wav_filter, estimators, test_loader, r 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}") + progress_bar = tqdm(total=len(test_loader), desc=f"Evaluating {class_name}") for idx, batch in enumerate(test_loader): progress_bar.update(1) @@ -39,15 +37,17 @@ def validate(args, encoder, constraintor, wav_filter, estimators, test_loader, r with torch.no_grad(): features = encoder(image) + + # --- 【重要】テスト画像側の特徴量をウェーブレット変換 (Pre-filter) --- + features = [wav_filter(f) for f in features] + + # --- 【重要】カンペ側は既に _wav ファイルとして保存・ロードされているためそのままマッチング --- mfeatures = get_matched_ref_features(features, ref_features) - # 生の残差を計算 + # マッチング後の残差を計算 rfeatures = get_residual_features(features, mfeatures, pos_flag=True) - # --- 【重要】学習時と同じウェーブレット処理で位置ズレノイズを抑制 --- - rfeatures = [wav_filter(rf) for rf in rfeatures] - - # --- 【重要】constraintorで空間的に滑らかに補正 --- + # Constraintorで空間的に滑らかに補正 rfeatures = constraintor(*rfeatures) for l in range(args.feature_levels): @@ -78,8 +78,8 @@ def validate(args, encoder, constraintor, wav_filter, estimators, test_loader, r 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) + 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) @@ -87,15 +87,15 @@ def validate(args, encoder, constraintor, wav_filter, estimators, test_loader, r 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] - + metrics = { + 'scores1': [img_auc1, img_ap1, img_f1_score1, pix_auc1, pix_ap1, pix_f1_score1, pix_aupro1], + 'scores2': [img_auc2, img_ap2, img_f1_score2, pix_auc2, pix_ap2, pix_f1_score2, pix_aupro2], + '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): +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) @@ -113,7 +113,7 @@ def convert_to_anomaly_scores(logps_list, feature_levels=3, class_name=None, siz return scores -def aggregate_anomaly_scores(logps_list, feature_levels=3, class_name=None, size=224): +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) From 5cdffe07794a9fe4110d8033d4ed2b617858609c Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 13 May 2026 21:49:26 +0900 Subject: [PATCH 209/258] Update main_wav.py --- main_wav.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/main_wav.py b/main_wav.py index 9a33938..8a3c133 100644 --- a/main_wav.py +++ b/main_wav.py @@ -110,8 +110,8 @@ def main(args): wav_filter.eval() constraintor = MultiScaleConv(feat_dims).to(args.device) - optimizer0 = torch.optim.Adam(constraintor.parameters(), lr=args.lr, weight_decay=0.005) # Weight decay少し強め - scheduler0 = torch.optim.lr_scheduler.MultiStepLR(optimizer0, milestones=[30, 50], gamma=0.1) # スケジューラ前倒し + optimizer0 = torch.optim.Adam(constraintor.parameters(), lr=args.lr, weight_decay=0.0005) # Weight decay + scheduler0 = torch.optim.lr_scheduler.MultiStepLR(optimizer0, milestones=[40, 50], gamma=0.1) # スケジューラ前倒し estimators = [load_flow_model(args, feat_dim).to(args.device) for feat_dim in feat_dims] params = list(estimators[0].parameters()) @@ -167,8 +167,8 @@ def main(args): rfeatures = constraintor(*rfeatures) # (任意) 特徴量への微小ノイズ付加による過学習防止 - noise_std = 0.01 - rfeatures_noisy = [rf + torch.randn_like(rf) * noise_std for rf in rfeatures] + #noise_std = 0.01 + #rfeatures_noisy = [rf + torch.randn_like(rf) * noise_std for rf in rfeatures] loss = 0 for l in range(args.feature_levels): From d20bd9d792eae88fc2f85e9bcf6f4c13ab138204 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 13 May 2026 21:54:04 +0900 Subject: [PATCH 210/258] Update main_wav.py --- main_wav.py | 9 ++++++++- 1 file changed, 8 insertions(+), 1 deletion(-) diff --git a/main_wav.py b/main_wav.py index 8a3c133..25adae7 100644 --- a/main_wav.py +++ b/main_wav.py @@ -280,7 +280,14 @@ def load_mc_reference_features(root_dir: str, class_names, device: torch.device, parser.add_argument('--pos_embed_dim', type=int, default=256) parser.add_argument("--train_ref_shot", type=int, default=4) parser.add_argument("--num_ref_shot", type=int, default=4) - + parser.add_argument('--pos_beta', type=float, default=0.05) + parser.add_argument('--coupling_layers', type=int, default=10) + parser.add_argument('--clamp_alpha', type=float, default=1.9) + 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("--lf_weight", type=float, default=0.1, help="Weight for low frequency (LL)") parser.add_argument("--hf_weight", type=float, default=1.2, help="Weight for high frequency (LH, HL, HH)") From d0bf8c9c9c7cb68efec5fb850324ed3f06358a98 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 13 May 2026 21:54:34 +0900 Subject: [PATCH 211/258] Update main_wav.py --- main_wav.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/main_wav.py b/main_wav.py index 25adae7..80c2a5a 100644 --- a/main_wav.py +++ b/main_wav.py @@ -285,7 +285,7 @@ def load_mc_reference_features(root_dir: str, class_names, device: torch.device, parser.add_argument('--clamp_alpha', type=float, default=1.9) 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('--fdm_alpha', type=float, default=0.4) parser.add_argument('--num_embeddings', type=int, default=1536) # ウェーブレット用パラメータ From 3684d37c062fcf1d2ab3b1f84629b190bce9d58e Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 13 May 2026 21:59:45 +0900 Subject: [PATCH 212/258] Update utils.py --- utils.py | 74 +++++++++++++++++--------------------------------------- 1 file changed, 22 insertions(+), 52 deletions(-) diff --git a/utils.py b/utils.py index f419fcf..ced8055 100644 --- a/utils.py +++ b/utils.py @@ -578,66 +578,36 @@ def apply_osp(rfeatures_list, proj_matrices, means): 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.0): + def __init__(self, low_freq_weight=0.1, high_freq_weight=1.2): super().__init__() - self.low_freq_weight = low_freq_weight - self.high_freq_weight = high_freq_weight + self.lf_w = low_freq_weight + self.hf_w = high_freq_weight - # Haarウェーブレットのカーネル定義 (2x2) - # LL: 低周波, HL: 垂直エッジ, LH: 水平エッジ, HH: 対角エッジ 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('kernel_ll', ll.view(1, 1, 2, 2)) - self.register_buffer('kernel_hl', hl.view(1, 1, 2, 2)) - self.register_buffer('kernel_lh', lh.view(1, 1, 2, 2)) - self.register_buffer('kernel_hh', hh.view(1, 1, 2, 2)) + 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 - - # チャネルごとに独立して畳み込むため、カーネルをチャネル数分コピー - weight_ll = self.kernel_ll.expand(C, 1, 2, 2) - weight_hl = self.kernel_hl.expand(C, 1, 2, 2) - weight_lh = self.kernel_lh.expand(C, 1, 2, 2) - weight_hh = self.kernel_hh.expand(C, 1, 2, 2) - - # DWT (離散ウェーブレット変換) - stride=2 でダウンサンプリング - x_ll = F.conv2d(x, weight_ll, stride=2, groups=C) - x_hl = F.conv2d(x, weight_hl, stride=2, groups=C) - x_lh = F.conv2d(x, weight_lh, stride=2, groups=C) - x_hh = F.conv2d(x, weight_hh, stride=2, groups=C) - - # --- フィルタリング処理 --- - # 位置ズレ(低周波)を弱め、キズ(高周波)を強調する - x_ll = x_ll * self.low_freq_weight - x_hl = x_hl * self.high_freq_weight - x_lh = x_lh * self.high_freq_weight - x_hh = x_hh * self.high_freq_weight - - # IDWT (逆離散ウェーブレット変換) - stride=2 で元の解像度に戻す - out_ll = F.conv_transpose2d(x_ll, weight_ll, stride=2, groups=C) - out_hl = F.conv_transpose2d(x_hl, weight_hl, stride=2, groups=C) - out_lh = F.conv_transpose2d(x_lh, weight_lh, stride=2, groups=C) - out_hh = F.conv_transpose2d(x_hh, weight_hh, stride=2, groups=C) - - # 全ての成分を足し合わせて再構成 - return out_ll + out_hl + out_lh + out_hh - -def apply_wavelet_filter(rfeatures): - """ - 残差特徴量のリストに対してウェーブレットフィルタを適用する - """ - device = rfeatures[0].device - # 低周波(位置ズレ)を10%に抑え、高周波(キズ)をそのまま通す - wavelet_filter = HaarWaveletFilter(low_freq_weight=0.1, high_freq_weight=1.0).to(device) - - filtered_rfeatures = [] - for rf in rfeatures: - filtered_rfeatures.append(wavelet_filter(rf)) - - return filtered_rfeatures + 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 From 4f063545fea9636bb12e73201f88cac911a6701e Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 13 May 2026 22:03:53 +0900 Subject: [PATCH 213/258] Update utils.py --- utils.py | 24 +++++++++++++++++++++++- 1 file changed, 23 insertions(+), 1 deletion(-) diff --git a/utils.py b/utils.py index ced8055..d272a4f 100644 --- a/utils.py +++ b/utils.py @@ -233,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: From 588c2d88172da338349c22da1c0c4758a37327f1 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 13 May 2026 22:05:03 +0900 Subject: [PATCH 214/258] Update main_wav.py --- main_wav.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/main_wav.py b/main_wav.py index 80c2a5a..47a8a05 100644 --- a/main_wav.py +++ b/main_wav.py @@ -21,7 +21,7 @@ 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 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 @@ -147,7 +147,7 @@ def main(args): features = [wav_filter(f) for f in features] # 3. 訓練用の参照特徴量を抽出し、それらもウェーブレット変換 - ref_features = get_mc_reference_features(encoder, args.train_dataset_dir, class_names, images.device, args.train_ref_shot) + ref_features = get_mc_reference_features_wav(encoder, args.train_dataset_dir, class_names, images.device, args.train_ref_shot,wav_filter=wav_filter) for c_name in ref_features.keys(): ref_features[c_name] = [wav_filter(rf) for rf in ref_features[c_name]] From 563260e131a2572738b35a376a778cef521b911e Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 13 May 2026 22:09:47 +0900 Subject: [PATCH 215/258] Update main_wav.py --- main_wav.py | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/main_wav.py b/main_wav.py index 47a8a05..ff5e18c 100644 --- a/main_wav.py +++ b/main_wav.py @@ -148,8 +148,7 @@ def main(args): # 3. 訓練用の参照特徴量を抽出し、それらもウェーブレット変換 ref_features = get_mc_reference_features_wav(encoder, args.train_dataset_dir, class_names, images.device, args.train_ref_shot,wav_filter=wav_filter) - for c_name in ref_features.keys(): - ref_features[c_name] = [wav_filter(rf) for rf in ref_features[c_name]] + # 4. 「エッジ強調された特徴量同士」でマッチング mfeatures = get_mc_matched_ref_features(features, class_names, ref_features) From 9a2635d7dbeaf84002611618f7cc8d39fb2203e7 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 13 May 2026 22:10:53 +0900 Subject: [PATCH 216/258] Update main_wav.py --- main_wav.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/main_wav.py b/main_wav.py index ff5e18c..29dbe82 100644 --- a/main_wav.py +++ b/main_wav.py @@ -171,7 +171,7 @@ def main(args): loss = 0 for l in range(args.feature_levels): - e = rfeatures_noisy[l] + e = rfeatures[l] t = rfeatures_t[l] bs, dim, h, w = e.size() e, t = e.permute(0, 2, 3, 1).reshape(-1, dim), t.permute(0, 2, 3, 1).reshape(-1, dim) From bbc546e4e54873b7978d2d04819448afd78d61c2 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 13 May 2026 22:43:44 +0900 Subject: [PATCH 217/258] Create main_wav1.py --- main_wav1.py | 296 +++++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 296 insertions(+) create mode 100644 main_wav1.py diff --git a/main_wav1.py b/main_wav1.py new file mode 100644 index 0000000..29dbe82 --- /dev/null +++ b/main_wav1.py @@ -0,0 +1,296 @@ +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 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().to(args.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().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 = MultiScaleConv(feat_dims).to(args.device) + optimizer0 = torch.optim.Adam(constraintor.parameters(), lr=args.lr, weight_decay=0.0005) # Weight decay + scheduler0 = torch.optim.lr_scheduler.MultiStepLR(optimizer0, milestones=[40, 50], gamma=0.1) # スケジューラ前倒し + + estimators = [load_flow_model(args, feat_dim).to(args.device) for feat_dim in feat_dims] + 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.005) + scheduler1 = torch.optim.lr_scheduler.MultiStepLR(optimizer1, milestones=[30, 50], gamma=0.1) + + 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() + + 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), desc=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 = encoder(images) + + # 2. テスト画像(訓練バッチ)の特徴量をウェーブレット変換 (Pre-filter) + features = [wav_filter(f) for f in features] + + # 3. 訓練用の参照特徴量を抽出し、それらもウェーブレット変換 + ref_features = get_mc_reference_features_wav(encoder, args.train_dataset_dir, class_names, images.device, args.train_ref_shot,wav_filter=wav_filter) + + + # 4. 「エッジ強調された特徴量同士」でマッチング + mfeatures = get_mc_matched_ref_features(features, class_names, ref_features) + + # 5. 残差の計算 + rfeatures = get_residual_features(features, mfeatures, pos_flag=True) + + lvl_masks = [] + for l in range(args.feature_levels): + _, _, h, w = rfeatures[l].size() + lvl_masks.append(F.interpolate(masks, size=(h, w), mode='nearest').squeeze(1)) + rfeatures_t = [rfeature.detach().clone() for rfeature in rfeatures] + + # 6. Constraintorによる空間補正 + rfeatures = constraintor(*rfeatures) + + # (任意) 特徴量への微小ノイズ付加による過学習防止 + #noise_std = 0.01 + #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, t = e.permute(0, 2, 3, 1).reshape(-1, dim), t.permute(0, 2, 3, 1).reshape(-1, dim) + m = lvl_masks[l].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 = [], [], [] + + # 既にextract_ref_features_wav.pyでウェーブレット変換済みのカンペをロード + 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に 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'] + print("Epoch: {}, Class Name: {}, Image AUC: {:.3f} | Pixel AUC: {:.3f} | AUPRO: {:.3f}".format( + epoch, class_name, img_auc, pix_auc, pix_aupro)) + s1_res.append(metrics['scores1']) + s2_res.append(metrics['scores2']) + s_res.append(metrics['scores']) + + img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro = np.mean(np.array(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_wav") # デフォルトを _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=60) # 過学習防止のためエポック数を削減 + 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('--flow_arch', type=str, default='conditional_flow_model') + parser.add_argument('--feature_levels', default=3, type=int) + parser.add_argument('--pos_embed_dim', type=int, default=256) + parser.add_argument("--train_ref_shot", type=int, default=4) + parser.add_argument("--num_ref_shot", type=int, default=4) + parser.add_argument('--pos_beta', type=float, default=0.05) + parser.add_argument('--coupling_layers', type=int, default=10) + parser.add_argument('--clamp_alpha', type=float, default=1.9) + 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("--lf_weight", type=float, default=0.1, help="Weight for low frequency (LL)") + parser.add_argument("--hf_weight", type=float, default=1.2, help="Weight for high frequency (LH, HL, HH)") + + args = parser.parse_args() + init_seeds(42) + main(args) From 26ceb31d13cc0b2309c121efbc397310d22208fe Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 13 May 2026 22:44:09 +0900 Subject: [PATCH 218/258] Create validate_wav1.py --- validate_wav1.py | 1 + 1 file changed, 1 insertion(+) create mode 100644 validate_wav1.py diff --git a/validate_wav1.py b/validate_wav1.py new file mode 100644 index 0000000..8b13789 --- /dev/null +++ b/validate_wav1.py @@ -0,0 +1 @@ + From 0e291d825fddeebde96b9c425496010350bddeec Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 13 May 2026 22:44:34 +0900 Subject: [PATCH 219/258] Implement model validation and anomaly scoring Added validation function for model evaluation with anomaly scoring and metrics calculation. --- validate_wav1.py | 129 +++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 129 insertions(+) diff --git a/validate_wav1.py b/validate_wav1.py index 8b13789..336ab41 100644 --- a/validate_wav1.py +++ b/validate_wav1.py @@ -1 +1,130 @@ +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), desc=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] + + # --- 【重要】カンペ側は既に _wav ファイルとして保存・ロードされているためそのままマッチング --- + mfeatures = get_matched_ref_features(features, ref_features) + + # マッチング後の残差を計算 + rfeatures = get_residual_features(features, mfeatures, pos_flag=True) + + # 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, 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 = { + 'scores1': [img_auc1, img_ap1, img_f1_score1, pix_auc1, pix_ap1, pix_f1_score1, pix_aupro1], + 'scores2': [img_auc2, img_ap2, img_f1_score2, pix_auc2, pix_ap2, pix_f1_score2, pix_aupro2], + '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 From 0d7e0a9e565b27b3beb2a66a51209bda1cc2000a Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 13 May 2026 22:48:33 +0900 Subject: [PATCH 220/258] Update main_wav1.py --- main_wav1.py | 60 +++++++++++++++++++++++++++------------------------- 1 file changed, 31 insertions(+), 29 deletions(-) diff --git a/main_wav1.py b/main_wav1.py index 29dbe82..02cffca 100644 --- a/main_wav1.py +++ b/main_wav1.py @@ -9,7 +9,7 @@ import torch.nn.functional as F from torch.utils.data import DataLoader -from validate_wav import validate +from validate_wav1 import validate from datasets.mvtec import MVTEC, MVTECANO from datasets.visa import VISA, VISAANO from datasets.btad import BTAD @@ -21,6 +21,7 @@ from models.fc_flow import load_flow_model from models.modules import MultiScaleConv +from models.vq import VectorQuantize 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 @@ -38,9 +39,6 @@ '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__() @@ -109,9 +107,14 @@ def main(args): wav_filter = HaarWaveletFilter(low_freq_weight=args.lf_weight, high_freq_weight=args.hf_weight).to(args.device) wav_filter.eval() + vqs = [VectorQuantize(dim=feat_dim, n_embed=args.num_embeddings).to(args.device) for feat_dim in feat_dims] + params_vq = [p for vq in vqs for p in vq.parameters()] + optimizer2 = torch.optim.Adam(params_vq, lr=args.lr, weight_decay=0.005) + scheduler2 = torch.optim.lr_scheduler.MultiStepLR(optimizer2, milestones=[30, 50], gamma=0.1) + constraintor = MultiScaleConv(feat_dims).to(args.device) - optimizer0 = torch.optim.Adam(constraintor.parameters(), lr=args.lr, weight_decay=0.0005) # Weight decay - scheduler0 = torch.optim.lr_scheduler.MultiStepLR(optimizer0, milestones=[40, 50], gamma=0.1) # スケジューラ前倒し + optimizer0 = torch.optim.Adam(constraintor.parameters(), lr=args.lr, weight_decay=0.005) + scheduler0 = torch.optim.lr_scheduler.MultiStepLR(optimizer0, milestones=[30, 50], gamma=0.1) estimators = [load_flow_model(args, feat_dim).to(args.device) for feat_dim in feat_dims] params = list(estimators[0].parameters()) @@ -126,6 +129,8 @@ def main(args): for epoch in range(args.epochs): constraintor.train() + for vq in vqs: + vq.train() for estimator in estimators: estimator.train() @@ -140,38 +145,34 @@ def main(args): images, masks = images.to(args.device), masks.to(args.device) with torch.no_grad(): - # 1. 画像から特徴抽出 features = encoder(images) - - # 2. テスト画像(訓練バッチ)の特徴量をウェーブレット変換 (Pre-filter) features = [wav_filter(f) for f in features] - # 3. 訓練用の参照特徴量を抽出し、それらもウェーブレット変換 - ref_features = get_mc_reference_features_wav(encoder, args.train_dataset_dir, class_names, images.device, args.train_ref_shot,wav_filter=wav_filter) - + ref_features = get_mc_reference_features_wav(encoder, args.train_dataset_dir, class_names, images.device, args.train_ref_shot, wav_filter=wav_filter) - # 4. 「エッジ強調された特徴量同士」でマッチング mfeatures = get_mc_matched_ref_features(features, class_names, ref_features) - - # 5. 残差の計算 rfeatures = get_residual_features(features, mfeatures, pos_flag=True) + vq_loss_total = 0 + for l in range(args.feature_levels): + out = vqs[l](rfeatures[l]) + rfeatures[l] = out[0] + vq_loss_total += out[-1].mean() + lvl_masks = [] for l in range(args.feature_levels): _, _, h, w = rfeatures[l].size() lvl_masks.append(F.interpolate(masks, size=(h, w), mode='nearest').squeeze(1)) rfeatures_t = [rfeature.detach().clone() for rfeature in rfeatures] - # 6. Constraintorによる空間補正 rfeatures = constraintor(*rfeatures) - # (任意) 特徴量への微小ノイズ付加による過学習防止 - #noise_std = 0.01 - #rfeatures_noisy = [rf + torch.randn_like(rf) * noise_std for rf in rfeatures] + noise_std = 0.01 + 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] + e = rfeatures_noisy[l] t = rfeatures_t[l] bs, dim, h, w = e.size() e, t = e.permute(0, 2, 3, 1).reshape(-1, dim), t.permute(0, 2, 3, 1).reshape(-1, dim) @@ -179,9 +180,12 @@ def main(args): loss_i, _, _ = calculate_log_barrier_bi_occ_loss(e, m, t) loss += loss_i + loss += vq_loss_total optimizer0.zero_grad() + optimizer2.zero_grad() loss.backward() optimizer0.step() + optimizer2.step() train_loss_total += loss.item() total_num += 1 @@ -193,14 +197,13 @@ def main(args): scheduler0.step() scheduler1.step() + scheduler2.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 = [], [], [] - # 既にextract_ref_features_wav.pyでウェーブレット変換済みのカンペをロード 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']: @@ -225,8 +228,7 @@ def main(args): test_loader = DataLoader(test_dataset, batch_size=1, shuffle=False, num_workers=8, drop_last=False) - # validate_wavに wav_filter を渡す - metrics = validate(args, encoder, constraintor, wav_filter, estimators, test_loader, test_ref_features[class_name], args.device, class_name) + metrics = validate(args, encoder, constraintor, vqs, 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: {:.3f} | Pixel AUC: {:.3f} | AUPRO: {:.3f}".format( @@ -243,6 +245,7 @@ def main(args): os.makedirs(args.checkpoint_path, exist_ok=True) best_img_auc = img_auc state_dict = {'constraintor': constraintor.state_dict(), + 'vqs': [vq.state_dict() for vq in vqs], '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')) @@ -264,11 +267,11 @@ def load_mc_reference_features(root_dir: str, class_names, device: torch.device, 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") # デフォルトを _wav に変更 + 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=60) # 過学習防止のためエポック数を削減 + parser.add_argument('--epochs', type=int, default=60) 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) @@ -287,9 +290,8 @@ def load_mc_reference_features(root_dir: str, class_names, device: torch.device, parser.add_argument('--fdm_alpha', type=float, default=0.4) parser.add_argument('--num_embeddings', type=int, default=1536) - # ウェーブレット用パラメータ - parser.add_argument("--lf_weight", type=float, default=0.1, help="Weight for low frequency (LL)") - parser.add_argument("--hf_weight", type=float, default=1.2, help="Weight for high frequency (LH, HL, HH)") + 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) From 6df71bab6f9d156de0b286c2d7cc565c39897ace Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 13 May 2026 22:48:57 +0900 Subject: [PATCH 221/258] Update validate_wav1.py --- validate_wav1.py | 15 +++++++-------- 1 file changed, 7 insertions(+), 8 deletions(-) diff --git a/validate_wav1.py b/validate_wav1.py index 336ab41..0e23916 100644 --- a/validate_wav1.py +++ b/validate_wav1.py @@ -1,4 +1,3 @@ - import warnings from tqdm import tqdm from scipy.ndimage import gaussian_filter @@ -14,9 +13,11 @@ warnings.filterwarnings('ignore') -def validate(args, encoder, constraintor, wav_filter, estimators, test_loader, ref_features, device, class_name): +def validate(args, encoder, constraintor, vqs, wav_filter, estimators, test_loader, ref_features, device, class_name): constraintor.eval() wav_filter.eval() + for vq in vqs: + vq.eval() for estimator in estimators: estimator.eval() @@ -38,17 +39,15 @@ def validate(args, encoder, constraintor, wav_filter, estimators, test_loader, r with torch.no_grad(): features = encoder(image) - - # --- 【重要】テスト画像側の特徴量をウェーブレット変換 (Pre-filter) --- features = [wav_filter(f) for f in features] - # --- 【重要】カンペ側は既に _wav ファイルとして保存・ロードされているためそのままマッチング --- mfeatures = get_matched_ref_features(features, ref_features) - - # マッチング後の残差を計算 rfeatures = get_residual_features(features, mfeatures, pos_flag=True) - # Constraintorで空間的に滑らかに補正 + for l in range(args.feature_levels): + out = vqs[l](rfeatures[l]) + rfeatures[l] = out[0] + rfeatures = constraintor(*rfeatures) for l in range(args.feature_levels): From 56282e60ea2ce3f7d9398e2495995a2e2730a5b6 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 13 May 2026 22:51:23 +0900 Subject: [PATCH 222/258] Update main_wav1.py --- main_wav1.py | 25 +++++++++++++++++-------- 1 file changed, 17 insertions(+), 8 deletions(-) diff --git a/main_wav1.py b/main_wav1.py index 02cffca..7744a76 100644 --- a/main_wav1.py +++ b/main_wav1.py @@ -21,7 +21,7 @@ from models.fc_flow import load_flow_model from models.modules import MultiScaleConv -from models.vq import VectorQuantize +from models.vq import VectorQuantizer # 正しいクラス名に変更 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 @@ -39,6 +39,9 @@ '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__() @@ -107,7 +110,8 @@ def main(args): wav_filter = HaarWaveletFilter(low_freq_weight=args.lf_weight, high_freq_weight=args.hf_weight).to(args.device) wav_filter.eval() - vqs = [VectorQuantize(dim=feat_dim, n_embed=args.num_embeddings).to(args.device) for feat_dim in feat_dims] + # VQモデルの初期化 (引数を VectorQuantizer の仕様に合わせる) + vqs = [VectorQuantizer(n_e=args.num_embeddings, vq_embed_dim=feat_dim, beta=0.25).to(args.device) for feat_dim in feat_dims] params_vq = [p for vq in vqs for p in vq.parameters()] optimizer2 = torch.optim.Adam(params_vq, lr=args.lr, weight_decay=0.005) scheduler2 = torch.optim.lr_scheduler.MultiStepLR(optimizer2, milestones=[30, 50], gamma=0.1) @@ -153,18 +157,23 @@ def main(args): mfeatures = get_mc_matched_ref_features(features, class_names, ref_features) rfeatures = get_residual_features(features, mfeatures, pos_flag=True) - vq_loss_total = 0 - for l in range(args.feature_levels): - out = vqs[l](rfeatures[l]) - rfeatures[l] = out[0] - vq_loss_total += out[-1].mean() - + # 先にマスクのリサイズ処理を行ってVQに渡せるようにする lvl_masks = [] for l in range(args.feature_levels): _, _, h, w = rfeatures[l].size() lvl_masks.append(F.interpolate(masks, size=(h, w), mode='nearest').squeeze(1)) + + # VQの適用 (マスク情報を渡す) + vq_loss_total = 0 + for l in range(args.feature_levels): + z_q, vq_loss, _ = vqs[l](rfeatures[l], lvl_masks[l]) + rfeatures[l] = z_q + vq_loss_total += vq_loss + + # 制約対象となる量子化後の特徴量を保存 rfeatures_t = [rfeature.detach().clone() for rfeature in rfeatures] + # Constraintorで空間補正 rfeatures = constraintor(*rfeatures) noise_std = 0.01 From 3dae178b1cbd7a9043c28433c8da1056b45c5942 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 13 May 2026 22:51:45 +0900 Subject: [PATCH 223/258] Update validate_wav1.py --- validate_wav1.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/validate_wav1.py b/validate_wav1.py index 0e23916..d07b04b 100644 --- a/validate_wav1.py +++ b/validate_wav1.py @@ -44,9 +44,9 @@ def validate(args, encoder, constraintor, vqs, wav_filter, estimators, test_load mfeatures = get_matched_ref_features(features, ref_features) rfeatures = get_residual_features(features, mfeatures, pos_flag=True) + # 推論時のVQ適用 (get_condebook_entryを使用) for l in range(args.feature_levels): - out = vqs[l](rfeatures[l]) - rfeatures[l] = out[0] + rfeatures[l] = vqs[l].get_condebook_entry(rfeatures[l]) rfeatures = constraintor(*rfeatures) From 7e3e404bab372e2e3fe2d21660bd8c293b82dc01 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 13 May 2026 22:55:29 +0900 Subject: [PATCH 224/258] Update validate_wav1.py --- validate_wav1.py | 376 ++++++++++++++++++++++++++++++++++------------- 1 file changed, 274 insertions(+), 102 deletions(-) diff --git a/validate_wav1.py b/validate_wav1.py index d07b04b..194c49f 100644 --- a/validate_wav1.py +++ b/validate_wav1.py @@ -1,129 +1,301 @@ +import os import warnings +import argparse from tqdm import tqdm -from scipy.ndimage import gaussian_filter 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 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 +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 +from models.vq import MultiScaleVQ # 元の MultiScaleVQ に戻す +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') -def validate(args, encoder, constraintor, vqs, wav_filter, estimators, test_loader, ref_features, device, class_name): - constraintor.eval() +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().to(args.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().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() - for vq in vqs: - vq.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)] + # 元の main.py と同じ MultiScaleVQ の初期化 + 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=[30, 50], gamma=0.1) + + 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=[30, 50], gamma=0.1) - progress_bar = tqdm(total=len(test_loader), desc=f"Evaluating {class_name}") + estimators = [load_flow_model(args, feat_dim).to(args.device) for feat_dim in feat_dims] + 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=[30, 50], gamma=0.1) - 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()) + from train import train + best_img_auc = 0 + N_batch = 8192 + + for epoch in range(args.epochs): + constraintor.train() + vq_ops.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 - image = image.to(device) - size = image.shape[-1] + progress_bar = tqdm(total=len(train_loader), desc=f"Epoch[{epoch}/{args.epochs}]") - with torch.no_grad(): - features = encoder(image) - features = [wav_filter(f) for f in features] + 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) - mfeatures = get_matched_ref_features(features, ref_features) - rfeatures = get_residual_features(features, mfeatures, pos_flag=True) + with torch.no_grad(): + features = encoder(images) + features = [wav_filter(f) for f in features] + + 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) - # 推論時のVQ適用 (get_condebook_entryを使用) + lvl_masks = [] for l in range(args.feature_levels): - rfeatures[l] = vqs[l].get_condebook_entry(rfeatures[l]) + _, _, h, w = rfeatures[l].size() + lvl_masks.append(F.interpolate(masks, size=(h, w), mode='nearest').squeeze(1)) + rfeatures_t = [rfeature.detach().clone() for rfeature in rfeatures] + + # 元の main.py に準拠: vq_ops でロスだけを計算 (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 は4次元のまま constraintor に入る rfeatures = constraintor(*rfeatures) - - for l in range(args.feature_levels): - e = rfeatures[l] + + # 過学習防止ノイズ (不要な場合はコメントアウト) + noise_std = 0.01 + rfeatures_noisy = [rf + torch.randn_like(rf) * noise_std for rf in rfeatures] + + loss = 0 + for l in range(args.feature_levels): + e = rfeatures_noisy[l] + t = rfeatures_t[l] bs, dim, h, w = e.size() - e = e.permute(0, 2, 3, 1).reshape(-1, dim) + e, t = e.permute(0, 2, 3, 1).reshape(-1, dim), t.permute(0, 2, 3, 1).reshape(-1, dim) + m = lvl_masks[l].reshape(-1) + loss_i, _, _ = calculate_log_barrier_bi_occ_loss(e, m, t) + loss += loss_i - 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) + 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 + + 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: - 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)) + raise ValueError('Unrecognized class name: {}'.format(class_name)) - 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 = { - 'scores1': [img_auc1, img_ap1, img_f1_score1, pix_auc1, pix_ap1, pix_f1_score1, pix_aupro1], - 'scores2': [img_auc2, img_ap2, img_f1_score2, pix_auc2, pix_ap2, pix_f1_score2, pix_aupro2], - '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 - + test_loader = DataLoader(test_dataset, batch_size=1, shuffle=False, num_workers=8, drop_last=False) + + metrics = validate(args, encoder, vq_ops, 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: {:.3f} | Pixel AUC: {:.3f} | AUPRO: {:.3f}".format( + epoch, class_name, img_auc, pix_auc, pix_aupro)) + s1_res.append(metrics['scores1']) + s2_res.append(metrics['scores2']) + s_res.append(metrics['scores']) + + img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro = np.mean(np.array(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 = {'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')) -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() +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=60) + 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") - scores = np.zeros_like(abnormal_map[0]) - for l in range(feature_levels): - scores += abnormal_map[l] - scores /= feature_levels + parser.add_argument('--flow_arch', type=str, default='conditional_flow_model') + parser.add_argument('--feature_levels', default=3, type=int) + parser.add_argument('--pos_embed_dim', type=int, default=256) + parser.add_argument("--train_ref_shot", type=int, default=4) + parser.add_argument("--num_ref_shot", type=int, default=4) + parser.add_argument('--pos_beta', type=float, default=0.05) + parser.add_argument('--coupling_layers', type=int, default=10) + parser.add_argument('--clamp_alpha', type=float, default=1.9) + 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("--lf_weight", type=float, default=0.1) + parser.add_argument("--hf_weight", type=float, default=1.2) - for i in range(scores.shape[0]): - scores[i] = gaussian_filter(scores[i], sigma=4) - return scores + args = parser.parse_args() + init_seeds(42) + main(args) From ad10a517aa430e2ae8b46d9773321ba99bab77f8 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 13 May 2026 22:56:04 +0900 Subject: [PATCH 225/258] Update main_wav1.py --- main_wav1.py | 50 ++++++++++++++++++++++---------------------------- 1 file changed, 22 insertions(+), 28 deletions(-) diff --git a/main_wav1.py b/main_wav1.py index 7744a76..194c49f 100644 --- a/main_wav1.py +++ b/main_wav1.py @@ -21,7 +21,7 @@ from models.fc_flow import load_flow_model from models.modules import MultiScaleConv -from models.vq import VectorQuantizer # 正しいクラス名に変更 +from models.vq import MultiScaleVQ # 元の MultiScaleVQ に戻す 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 @@ -110,21 +110,20 @@ def main(args): wav_filter = HaarWaveletFilter(low_freq_weight=args.lf_weight, high_freq_weight=args.hf_weight).to(args.device) wav_filter.eval() - # VQモデルの初期化 (引数を VectorQuantizer の仕様に合わせる) - vqs = [VectorQuantizer(n_e=args.num_embeddings, vq_embed_dim=feat_dim, beta=0.25).to(args.device) for feat_dim in feat_dims] - params_vq = [p for vq in vqs for p in vq.parameters()] - optimizer2 = torch.optim.Adam(params_vq, lr=args.lr, weight_decay=0.005) - scheduler2 = torch.optim.lr_scheduler.MultiStepLR(optimizer2, milestones=[30, 50], gamma=0.1) + # 元の main.py と同じ MultiScaleVQ の初期化 + 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=[30, 50], gamma=0.1) constraintor = MultiScaleConv(feat_dims).to(args.device) - optimizer0 = torch.optim.Adam(constraintor.parameters(), lr=args.lr, weight_decay=0.005) + optimizer0 = torch.optim.Adam(constraintor.parameters(), lr=args.lr, weight_decay=0.0005) scheduler0 = torch.optim.lr_scheduler.MultiStepLR(optimizer0, milestones=[30, 50], gamma=0.1) estimators = [load_flow_model(args, feat_dim).to(args.device) for feat_dim in feat_dims] 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.005) + optimizer1 = torch.optim.Adam(params, lr=args.lr, weight_decay=0.0005) scheduler1 = torch.optim.lr_scheduler.MultiStepLR(optimizer1, milestones=[30, 50], gamma=0.1) from train import train @@ -133,8 +132,7 @@ def main(args): for epoch in range(args.epochs): constraintor.train() - for vq in vqs: - vq.train() + vq_ops.train() for estimator in estimators: estimator.train() @@ -157,25 +155,24 @@ def main(args): mfeatures = get_mc_matched_ref_features(features, class_names, ref_features) rfeatures = get_residual_features(features, mfeatures, pos_flag=True) - # 先にマスクのリサイズ処理を行ってVQに渡せるようにする lvl_masks = [] for l in range(args.feature_levels): _, _, h, w = rfeatures[l].size() lvl_masks.append(F.interpolate(masks, size=(h, w), mode='nearest').squeeze(1)) - - # VQの適用 (マスク情報を渡す) - vq_loss_total = 0 - for l in range(args.feature_levels): - z_q, vq_loss, _ = vqs[l](rfeatures[l], lvl_masks[l]) - rfeatures[l] = z_q - vq_loss_total += vq_loss - - # 制約対象となる量子化後の特徴量を保存 rfeatures_t = [rfeature.detach().clone() for rfeature in rfeatures] - # Constraintorで空間補正 + # 元の main.py に準拠: vq_ops でロスだけを計算 (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 は4次元のまま constraintor に入る rfeatures = constraintor(*rfeatures) + # 過学習防止ノイズ (不要な場合はコメントアウト) noise_std = 0.01 rfeatures_noisy = [rf + torch.randn_like(rf) * noise_std for rf in rfeatures] @@ -189,12 +186,9 @@ def main(args): loss_i, _, _ = calculate_log_barrier_bi_occ_loss(e, m, t) loss += loss_i - loss += vq_loss_total optimizer0.zero_grad() - optimizer2.zero_grad() loss.backward() optimizer0.step() - optimizer2.step() train_loss_total += loss.item() total_num += 1 @@ -204,9 +198,9 @@ def main(args): train_loss_total += loss total_num += num + scheduler_vq.step() scheduler0.step() scheduler1.step() - scheduler2.step() progress_bar.close() print(f"Epoch[{epoch}/{args.epochs}]: train_loss: {train_loss_total / total_num}") @@ -237,7 +231,7 @@ def main(args): test_loader = DataLoader(test_dataset, batch_size=1, shuffle=False, num_workers=8, drop_last=False) - metrics = validate(args, encoder, constraintor, vqs, wav_filter, estimators, test_loader, test_ref_features[class_name], args.device, class_name) + metrics = validate(args, encoder, vq_ops, 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: {:.3f} | Pixel AUC: {:.3f} | AUPRO: {:.3f}".format( @@ -253,8 +247,8 @@ def main(args): 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(), - 'vqs': [vq.state_dict() for vq in vqs], + 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')) From dba422b9ad939f1d70f7814cb883247d8135f77e Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 13 May 2026 22:56:27 +0900 Subject: [PATCH 226/258] Update validate_wav1.py --- validate_wav1.py | 376 +++++++++++++---------------------------------- 1 file changed, 101 insertions(+), 275 deletions(-) diff --git a/validate_wav1.py b/validate_wav1.py index 194c49f..da3bf6e 100644 --- a/validate_wav1.py +++ b/validate_wav1.py @@ -1,301 +1,127 @@ -import os import warnings -import argparse from tqdm import tqdm +from scipy.ndimage import gaussian_filter 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 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 -from models.vq import MultiScaleVQ # 元の MultiScaleVQ に戻す -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 +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') -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().to(args.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().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) +def validate(args, encoder, vq_ops, constraintor, wav_filter, estimators, test_loader, ref_features, device, class_name): + vq_ops.eval() + constraintor.eval() wav_filter.eval() + for estimator in estimators: + estimator.eval() - # 元の main.py と同じ MultiScaleVQ の初期化 - 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=[30, 50], gamma=0.1) - - 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=[30, 50], gamma=0.1) - - estimators = [load_flow_model(args, feat_dim).to(args.device) for feat_dim in feat_dims] - 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=[30, 50], gamma=0.1) + label_list, gt_mask_list = [], [] + logps1_list = [list() for _ in range(args.feature_levels)] + logps2_list = [list() for _ in range(args.feature_levels)] - from train import train - best_img_auc = 0 - N_batch = 8192 + progress_bar = tqdm(total=len(test_loader), desc=f"Evaluating {class_name}") - for epoch in range(args.epochs): - constraintor.train() - vq_ops.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 + for idx, batch in enumerate(test_loader): + progress_bar.update(1) - progress_bar = tqdm(total=len(train_loader), desc=f"Epoch[{epoch}/{args.epochs}]") + 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()) - 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(): - features = encoder(images) - features = [wav_filter(f) for f in features] - - 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) + image = image.to(device) + size = image.shape[-1] + + with torch.no_grad(): + features = encoder(image) + features = [wav_filter(f) for f in features] - lvl_masks = [] - for l in range(args.feature_levels): - _, _, h, w = rfeatures[l].size() - lvl_masks.append(F.interpolate(masks, size=(h, w), mode='nearest').squeeze(1)) - rfeatures_t = [rfeature.detach().clone() for rfeature in rfeatures] + mfeatures = get_matched_ref_features(features, ref_features) + rfeatures = get_residual_features(features, mfeatures, pos_flag=True) - # 元の main.py に準拠: vq_ops でロスだけを計算 (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() + # 推論時のみ VQ で特徴量を量子化 (tupleで返ってくるのでlist化) + rfeatures = list(vq_ops(rfeatures, train=False)) - # rfeatures は4次元のまま constraintor に入る rfeatures = constraintor(*rfeatures) - - # 過学習防止ノイズ (不要な場合はコメントアウト) - noise_std = 0.01 - rfeatures_noisy = [rf + torch.randn_like(rf) * noise_std for rf in rfeatures] - - loss = 0 - for l in range(args.feature_levels): - e = rfeatures_noisy[l] - t = rfeatures_t[l] + + for l in range(args.feature_levels): + e = rfeatures[l] bs, dim, h, w = e.size() - e, t = e.permute(0, 2, 3, 1).reshape(-1, dim), t.permute(0, 2, 3, 1).reshape(-1, dim) - m = lvl_masks[l].reshape(-1) - loss_i, _, _ = calculate_log_barrier_bi_occ_loss(e, m, t) - loss += loss_i + e = e.permute(0, 2, 3, 1).reshape(-1, dim) - 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 - - 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) + 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: - 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, wav_filter, estimators, test_loader, test_ref_features[class_name], args.device, class_name) + 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)) - img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro = metrics['scores'] - print("Epoch: {}, Class Name: {}, Image AUC: {:.3f} | Pixel AUC: {:.3f} | AUPRO: {:.3f}".format( - epoch, class_name, img_auc, pix_auc, pix_aupro)) - s1_res.append(metrics['scores1']) - s2_res.append(metrics['scores2']) - s_res.append(metrics['scores']) - - img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro = np.mean(np.array(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 = {'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')) + 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 = { + 'scores1': [img_auc1, img_ap1, img_f1_score1, pix_auc1, pix_ap1, pix_f1_score1, pix_aupro1], + 'scores2': [img_auc2, img_ap2, img_f1_score2, pix_auc2, pix_ap2, pix_f1_score2, pix_aupro2], + 'scores': [img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro] + } + return metrics -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=60) - 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") + +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() - parser.add_argument('--flow_arch', type=str, default='conditional_flow_model') - parser.add_argument('--feature_levels', default=3, type=int) - parser.add_argument('--pos_embed_dim', type=int, default=256) - parser.add_argument("--train_ref_shot", type=int, default=4) - parser.add_argument("--num_ref_shot", type=int, default=4) - parser.add_argument('--pos_beta', type=float, default=0.05) - parser.add_argument('--coupling_layers', type=int, default=10) - parser.add_argument('--clamp_alpha', type=float, default=1.9) - 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("--lf_weight", type=float, default=0.1) - parser.add_argument("--hf_weight", type=float, default=1.2) + 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 - args = parser.parse_args() - init_seeds(42) - main(args) + for i in range(scores.shape[0]): + scores[i] = gaussian_filter(scores[i], sigma=4) + return scores From 5f2ef35d91377a795b44558a6e6cc0284bd2ca13 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 13 May 2026 23:25:55 +0900 Subject: [PATCH 227/258] Update validate_wav1.py --- validate_wav1.py | 388 +++++++++++++++++++++++++++++++++++------------ 1 file changed, 289 insertions(+), 99 deletions(-) diff --git a/validate_wav1.py b/validate_wav1.py index da3bf6e..72f8035 100644 --- a/validate_wav1.py +++ b/validate_wav1.py @@ -1,127 +1,317 @@ +import os import warnings +import argparse from tqdm import tqdm -from scipy.ndimage import gaussian_filter 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 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 +from validate_wav1 import validate # validate_wav1 をインポート +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 # 元の MultiScaleVQ を使用 +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') -def validate(args, encoder, vq_ops, constraintor, wav_filter, estimators, test_loader, ref_features, device, class_name): - vq_ops.eval() - constraintor.eval() +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().to(args.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().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() - 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)] + # VQの初期化 (元のmain.pyと同じ) + 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=[30, 50], gamma=0.1) + + 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=[30, 50], gamma=0.1) - progress_bar = tqdm(total=len(test_loader), desc=f"Evaluating {class_name}") + 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=[30, 50], gamma=0.1) - 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()) + from train import train + best_img_auc = 0 + N_batch = 8192 + + for epoch in range(args.epochs): + vq_ops.train() + 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 - image = image.to(device) - size = image.shape[-1] + progress_bar = tqdm(total=len(train_loader), desc=f"Epoch[{epoch}/{args.epochs}]") - with torch.no_grad(): - features = encoder(image) - features = [wav_filter(f) for f in features] + 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(): + # 画像から特徴抽出 + features = encoder(images) + + # --- 追加: ウェーブレットフィルタを適用 (Pre-filter) --- + features = [wav_filter(f) for f in features] + + # 訓練用の参照特徴量を抽出し、内部でウェーブレット変換させる + 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) - mfeatures = get_matched_ref_features(features, 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() + lvl_masks.append(F.interpolate(masks, size=(h, w), mode='nearest').squeeze(1)) + rfeatures_t = [rfeature.detach().clone() for rfeature in rfeatures] - # 推論時のみ VQ で特徴量を量子化 (tupleで返ってくるのでlist化) - rfeatures = list(vq_ops(rfeatures, train=False)) + # --- 元の main.py と同じ VQ の学習 --- + # VQは rfeatures を変更せず、loss_vq だけを返す + 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() + # --- 元の main.py と同じ constraintor の適用 --- rfeatures = constraintor(*rfeatures) - - for l in range(args.feature_levels): + + 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) - 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) + 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 + + 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 = [], [], [] + + # 抽出済みの _wav 特徴量をロード + 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: - 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)) + raise ValueError('Unrecognized class name: {}'.format(class_name)) - 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 = { - 'scores1': [img_auc1, img_ap1, img_f1_score1, pix_auc1, pix_ap1, pix_f1_score1, pix_aupro1], - 'scores2': [img_auc2, img_ap2, img_f1_score2, pix_auc2, pix_ap2, pix_f1_score2, pix_aupro2], - '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 - + test_loader = DataLoader(test_dataset, batch_size=1, shuffle=False, num_workers=8, drop_last=False) + + # validate関数に vq_ops と wav_filter を渡す + metrics = validate(args, encoder, vq_ops, 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: {:.3f} | Pixel AUC: {:.3f} | AUPRO: {:.3f}".format( + epoch, class_name, img_auc, pix_auc, 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 = {'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')) -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() +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=60) + 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") - scores = np.zeros_like(abnormal_map[0]) - for l in range(feature_levels): - scores += abnormal_map[l] - scores /= feature_levels + parser.add_argument('--flow_arch', type=str, default='conditional_flow_model') + parser.add_argument('--feature_levels', default=3, type=int) + parser.add_argument('--pos_embed_dim', type=int, default=256) + parser.add_argument("--train_ref_shot", type=int, default=4) + parser.add_argument("--num_ref_shot", type=int, default=4) + parser.add_argument('--pos_beta', type=float, default=0.05) + parser.add_argument('--coupling_layers', type=int, default=10) + parser.add_argument('--clamp_alpha', type=float, default=1.9) + 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("--lf_weight", type=float, default=0.1) + parser.add_argument("--hf_weight", type=float, default=1.2) - for i in range(scores.shape[0]): - scores[i] = gaussian_filter(scores[i], sigma=4) - return scores + args = parser.parse_args() + init_seeds(42) + main(args) From 178eb1b2a1d35483aa1d625bf0037e64a6f91cb0 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 13 May 2026 23:26:08 +0900 Subject: [PATCH 228/258] Update main_wav1.py --- main_wav1.py | 46 +++++++++++++++++++++++++++++++--------------- 1 file changed, 31 insertions(+), 15 deletions(-) diff --git a/main_wav1.py b/main_wav1.py index 194c49f..72f8035 100644 --- a/main_wav1.py +++ b/main_wav1.py @@ -9,7 +9,7 @@ import torch.nn.functional as F from torch.utils.data import DataLoader -from validate_wav1 import validate +from validate_wav1 import validate # validate_wav1 をインポート from datasets.mvtec import MVTEC, MVTECANO from datasets.visa import VISA, VISAANO from datasets.btad import BTAD @@ -21,7 +21,7 @@ from models.fc_flow import load_flow_model from models.modules import MultiScaleConv -from models.vq import MultiScaleVQ # 元の MultiScaleVQ に戻す +from models.vq import MultiScaleVQ # 元の MultiScaleVQ を使用 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 @@ -107,10 +107,11 @@ def main(args): 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() - # 元の main.py と同じ MultiScaleVQ の初期化 + # VQの初期化 (元のmain.pyと同じ) 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=[30, 50], gamma=0.1) @@ -119,7 +120,8 @@ def main(args): optimizer0 = torch.optim.Adam(constraintor.parameters(), lr=args.lr, weight_decay=0.0005) scheduler0 = torch.optim.lr_scheduler.MultiStepLR(optimizer0, milestones=[30, 50], gamma=0.1) - estimators = [load_flow_model(args, feat_dim).to(args.device) for feat_dim in feat_dims] + 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()) @@ -131,8 +133,8 @@ def main(args): N_batch = 8192 for epoch in range(args.epochs): - constraintor.train() vq_ops.train() + constraintor.train() for estimator in estimators: estimator.train() @@ -147,11 +149,16 @@ def main(args): images, masks = images.to(args.device), masks.to(args.device) with torch.no_grad(): + # 画像から特徴抽出 features = encoder(images) + + # --- 追加: ウェーブレットフィルタを適用 (Pre-filter) --- features = [wav_filter(f) for f in features] + # 訓練用の参照特徴量を抽出し、内部でウェーブレット変換させる 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) @@ -161,7 +168,8 @@ def main(args): lvl_masks.append(F.interpolate(masks, size=(h, w), mode='nearest').squeeze(1)) rfeatures_t = [rfeature.detach().clone() for rfeature in rfeatures] - # 元の main.py に準拠: vq_ops でロスだけを計算 (rfeatures は無傷) + # --- 元の main.py と同じ VQ の学習 --- + # VQは rfeatures を変更せず、loss_vq だけを返す loss_vq = vq_ops(rfeatures, lvl_masks, train=True) train_loss_total += loss_vq.item() total_num += 1 @@ -169,20 +177,19 @@ def main(args): loss_vq.backward() optimizer_vq.step() - # rfeatures は4次元のまま constraintor に入る + # --- 元の main.py と同じ constraintor の適用 --- rfeatures = constraintor(*rfeatures) - # 過学習防止ノイズ (不要な場合はコメントアウト) - noise_std = 0.01 - rfeatures_noisy = [rf + torch.randn_like(rf) * noise_std for rf in rfeatures] - loss = 0 for l in range(args.feature_levels): - e = rfeatures_noisy[l] + e = rfeatures[l] t = rfeatures_t[l] bs, dim, h, w = e.size() - e, t = e.permute(0, 2, 3, 1).reshape(-1, dim), t.permute(0, 2, 3, 1).reshape(-1, dim) - m = lvl_masks[l].reshape(-1) + 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 @@ -204,9 +211,11 @@ def main(args): 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 = [], [], [] + # 抽出済みの _wav 特徴量をロード 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']: @@ -231,6 +240,7 @@ def main(args): test_loader = DataLoader(test_dataset, batch_size=1, shuffle=False, num_workers=8, drop_last=False) + # validate関数に vq_ops と wav_filter を渡す metrics = validate(args, encoder, vq_ops, 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'] @@ -240,7 +250,12 @@ def main(args): s2_res.append(metrics['scores2']) s_res.append(metrics['scores']) - img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro = np.mean(np.array(s_res), axis=0) + 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)) @@ -279,6 +294,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") parser.add_argument('--flow_arch', type=str, default='conditional_flow_model') parser.add_argument('--feature_levels', default=3, type=int) From 3dac92bc8dc10cfb0a17dcbf109effc7406fabc3 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 13 May 2026 23:26:27 +0900 Subject: [PATCH 229/258] Update validate_wav1.py --- validate_wav1.py | 392 +++++++++++++---------------------------------- 1 file changed, 104 insertions(+), 288 deletions(-) diff --git a/validate_wav1.py b/validate_wav1.py index 72f8035..1f22dad 100644 --- a/validate_wav1.py +++ b/validate_wav1.py @@ -1,317 +1,133 @@ -import os import warnings -import argparse from tqdm import tqdm +from scipy.ndimage import gaussian_filter 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 validate_wav1 import validate # validate_wav1 をインポート -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 # 元の MultiScaleVQ を使用 -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 +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') -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().to(args.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().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) +def validate(args, encoder, vq_ops, constraintor, wav_filter, estimators, test_loader, ref_features, device, class_name): + vq_ops.eval() + constraintor.eval() wav_filter.eval() + for estimator in estimators: + estimator.eval() - # VQの初期化 (元のmain.pyと同じ) - 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=[30, 50], gamma=0.1) - - 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=[30, 50], 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()) - optimizer1 = torch.optim.Adam(params, lr=args.lr, weight_decay=0.0005) - scheduler1 = torch.optim.lr_scheduler.MultiStepLR(optimizer1, milestones=[30, 50], gamma=0.1) + label_list, gt_mask_list = [], [] + logps1_list = [list() for _ in range(args.feature_levels)] + logps2_list = [list() for _ in range(args.feature_levels)] - from train import train - best_img_auc = 0 - N_batch = 8192 + progress_bar = tqdm(total=len(test_loader), desc=f"Evaluating {class_name}") - for epoch in range(args.epochs): - vq_ops.train() - 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 + for idx, batch in enumerate(test_loader): + progress_bar.update(1) - progress_bar = tqdm(total=len(train_loader), desc=f"Epoch[{epoch}/{args.epochs}]") + 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()) - 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) + image = image.to(device) + size = image.shape[-1] + + with torch.no_grad(): + # 画像から特徴抽出 + features = encoder(image) - with torch.no_grad(): - # 画像から特徴抽出 - features = encoder(images) - - # --- 追加: ウェーブレットフィルタを適用 (Pre-filter) --- - features = [wav_filter(f) for f in features] - - # 訓練用の参照特徴量を抽出し、内部でウェーブレット変換させる - 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) + # --- 追加: テスト画像をウェーブレット変換 (Pre-filter) --- + features = [wav_filter(f) for f in features] - lvl_masks = [] - for l in range(args.feature_levels): - _, _, h, w = rfeatures[l].size() - lvl_masks.append(F.interpolate(masks, size=(h, w), mode='nearest').squeeze(1)) - rfeatures_t = [rfeature.detach().clone() for rfeature in rfeatures] + # ウェーブレット済みのカンペとマッチング + mfeatures = get_matched_ref_features(features, ref_features) + rfeatures = get_residual_features(features, mfeatures, pos_flag=True) - # --- 元の main.py と同じ VQ の学習 --- - # VQは rfeatures を変更せず、loss_vq だけを返す - 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() + # --- 元の validate.py と同じ VQ の適用 --- + # VQを通すことで、特徴量が離散的なコードブックにマッピングされる + rfeatures = list(vq_ops(rfeatures, train=False)) - # --- 元の main.py と同じ constraintor の適用 --- + # Constraintorで空間を滑らかにする rfeatures = constraintor(*rfeatures) - - loss = 0 - for l in range(args.feature_levels): + + 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 - - 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 = [], [], [] - - # 抽出済みの _wav 特徴量をロード - 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) + 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: - 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関数に vq_ops と wav_filter を渡す - metrics = validate(args, encoder, vq_ops, constraintor, wav_filter, estimators, test_loader, test_ref_features[class_name], args.device, class_name) + 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)) - img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro = metrics['scores'] - print("Epoch: {}, Class Name: {}, Image AUC: {:.3f} | Pixel AUC: {:.3f} | AUPRO: {:.3f}".format( - epoch, class_name, img_auc, pix_auc, 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 = {'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')) + 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 = { + 'scores1': [img_auc1, img_ap1, img_f1_score1, pix_auc1, pix_ap1, pix_f1_score1, pix_aupro1], + 'scores2': [img_auc2, img_ap2, img_f1_score2, pix_auc2, pix_ap2, pix_f1_score2, pix_aupro2], + 'scores': [img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro] + } + return metrics -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=60) - 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") + +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() - parser.add_argument('--flow_arch', type=str, default='conditional_flow_model') - parser.add_argument('--feature_levels', default=3, type=int) - parser.add_argument('--pos_embed_dim', type=int, default=256) - parser.add_argument("--train_ref_shot", type=int, default=4) - parser.add_argument("--num_ref_shot", type=int, default=4) - parser.add_argument('--pos_beta', type=float, default=0.05) - parser.add_argument('--coupling_layers', type=int, default=10) - parser.add_argument('--clamp_alpha', type=float, default=1.9) - 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("--lf_weight", type=float, default=0.1) - parser.add_argument("--hf_weight", type=float, default=1.2) + 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 - args = parser.parse_args() - init_seeds(42) - main(args) + for i in range(scores.shape[0]): + scores[i] = gaussian_filter(scores[i], sigma=4) + return scores From f4571e71eb43c30004c16781f05efbf633b35ed2 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 13 May 2026 23:43:27 +0900 Subject: [PATCH 230/258] Update main_wav1.py --- main_wav1.py | 217 +++++++++++++++++++++++++++++++++------------------ 1 file changed, 140 insertions(+), 77 deletions(-) diff --git a/main_wav1.py b/main_wav1.py index 72f8035..408b3b9 100644 --- a/main_wav1.py +++ b/main_wav1.py @@ -9,7 +9,8 @@ import torch.nn.functional as F from torch.utils.data import DataLoader -from validate_wav1 import validate # validate_wav1 をインポート +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 @@ -21,7 +22,7 @@ from models.fc_flow import load_flow_model from models.modules import MultiScaleConv -from models.vq import MultiScaleVQ # 元の MultiScaleVQ を使用 +from models.vq import MultiScaleVQ 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 @@ -32,7 +33,7 @@ warnings.filterwarnings('ignore') -TOTAL_SHOT = 4 +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, @@ -83,42 +84,70 @@ def main(args): 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.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) + 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().to(args.device) - feat_dims = encoder.feature_info.channels() + 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() - # VQの初期化 (元のmain.pyと同じ) 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=[30, 50], gamma=0.1) - + scheduler_vq = torch.optim.lr_scheduler.MultiStepLR(optimizer_vq, milestones=[70, 90], gamma=0.1) + 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=[30, 50], gamma=0.1) + 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] @@ -126,9 +155,8 @@ def main(args): 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=[30, 50], gamma=0.1) + scheduler1 = torch.optim.lr_scheduler.MultiStepLR(optimizer1, milestones=[70, 90], gamma=0.1) - from train import train best_img_auc = 0 N_batch = 8192 @@ -138,38 +166,39 @@ def main(args): for estimator in estimators: estimator.train() - train_loader = train_loader1 if epoch < FIRST_STAGE_EPOCH else train_loader2 + 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), desc=f"Epoch[{epoch}/{args.epochs}]") + 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) + + 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] - - # 訓練用の参照特徴量を抽出し、内部でウェーブレット変換させる - 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) + + # --- 変更: get_mc_reference_features_wav を呼び出し、wav_filterを渡す --- + 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() - lvl_masks.append(F.interpolate(masks, size=(h, w), mode='nearest').squeeze(1)) + m = F.interpolate(masks, size=(h, w), mode='nearest').squeeze(1) + lvl_masks.append(m) rfeatures_t = [rfeature.detach().clone() for rfeature in rfeatures] - # --- 元の main.py と同じ VQ の学習 --- - # VQは rfeatures を変更せず、loss_vq だけを返す loss_vq = vq_ops(rfeatures, lvl_masks, train=True) train_loss_total += loss_vq.item() total_num += 1 @@ -177,9 +206,7 @@ def main(args): loss_vq.backward() optimizer_vq.step() - # --- 元の main.py と同じ constraintor の適用 --- rfeatures = constraintor(*rfeatures) - loss = 0 for l in range(args.feature_levels): e = rfeatures[l] @@ -192,7 +219,6 @@ def main(args): loss_i, _, _ = calculate_log_barrier_bi_occ_loss(e, m, t) loss += loss_i - optimizer0.zero_grad() loss.backward() optimizer0.step() @@ -208,44 +234,60 @@ def main(args): 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 = [], [], [] - - # 抽出済みの _wav 特徴量をロード 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) + 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) + 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) + 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) + 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) + 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) + 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) + 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) + 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) + test_loader = DataLoader( + test_dataset, batch_size=1, shuffle=False, num_workers=8, drop_last=False + ) - # validate関数に vq_ops と wav_filter を渡す + # --- 変更: validate に wav_filter を渡す --- metrics = validate(args, encoder, vq_ops, 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: {:.3f} | Pixel AUC: {:.3f} | AUPRO: {:.3f}".format( - epoch, class_name, img_auc, pix_auc, pix_aupro)) + + 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']) @@ -256,6 +298,11 @@ def main(args): 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)) @@ -267,18 +314,30 @@ def main(args): '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) + 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 - refs[class_name] = (layer1_refs[:K1, :], layer2_refs[:K2, :], layer3_refs[:K3, :]) + 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") @@ -289,29 +348,33 @@ def load_mc_reference_features(root_dir: str, class_names, device: torch.device, 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=60) + 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('--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('--pos_embed_dim', type=int, default=256) - parser.add_argument("--train_ref_shot", type=int, default=4) - parser.add_argument("--num_ref_shot", type=int, default=4) - parser.add_argument('--pos_beta', type=float, default=0.05) parser.add_argument('--coupling_layers', type=int, default=10) - parser.add_argument('--clamp_alpha', type=float, default=1.9) + 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('--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) From 4e6c3058f9b3410705f7b88191c181334b33e0ab Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 13 May 2026 23:44:02 +0900 Subject: [PATCH 231/258] Update validate_wav1.py --- validate_wav1.py | 26 ++++++++++++-------------- 1 file changed, 12 insertions(+), 14 deletions(-) diff --git a/validate_wav1.py b/validate_wav1.py index 1f22dad..701597f 100644 --- a/validate_wav1.py +++ b/validate_wav1.py @@ -16,7 +16,7 @@ def validate(args, encoder, vq_ops, constraintor, wav_filter, estimators, test_loader, ref_features, device, class_name): vq_ops.eval() constraintor.eval() - wav_filter.eval() + wav_filter.eval() # 追加 for estimator in estimators: estimator.eval() @@ -24,7 +24,8 @@ def validate(args, encoder, vq_ops, constraintor, wav_filter, estimators, test_l 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), desc=f"Evaluating {class_name}") + 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) @@ -37,21 +38,17 @@ def validate(args, encoder, vq_ops, constraintor, wav_filter, estimators, test_l 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) - # --- 元の validate.py と同じ VQ の適用 --- - # VQを通すことで、特徴量が離散的なコードブックにマッピングされる - rfeatures = list(vq_ops(rfeatures, train=False)) - - # Constraintorで空間を滑らかにする + # 元の validate.py と全く同じ VQ と constraintor の処理 + qx1, qx2, qx3 = vq_ops(rfeatures, train=False) + rfeatures = [qx1, qx2, qx3] rfeatures = constraintor(*rfeatures) for l in range(args.feature_levels): @@ -91,11 +88,12 @@ def validate(args, encoder, vq_ops, constraintor, wav_filter, estimators, test_l 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 = { - 'scores1': [img_auc1, img_ap1, img_f1_score1, pix_auc1, pix_ap1, pix_f1_score1, pix_aupro1], - 'scores2': [img_auc2, img_ap2, img_f1_score2, pix_auc2, pix_ap2, pix_f1_score2, pix_aupro2], - 'scores': [img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro] - } + # 元の validate.py と全く同じ出力を返す + 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 From 9415e723717b73a5d7a8c1f257f11ad9b9689a99 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 13 May 2026 23:46:07 +0900 Subject: [PATCH 232/258] Update validate_wav.py --- validate_wav.py | 19 ++++++++----------- 1 file changed, 8 insertions(+), 11 deletions(-) diff --git a/validate_wav.py b/validate_wav.py index 92b0b80..2b732e0 100644 --- a/validate_wav.py +++ b/validate_wav.py @@ -23,7 +23,8 @@ def validate(args, encoder, constraintor, wav_filter, estimators, test_loader, r 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), desc=f"Evaluating {class_name}") + 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) @@ -38,16 +39,12 @@ def validate(args, encoder, constraintor, wav_filter, estimators, test_loader, r with torch.no_grad(): features = encoder(image) - # --- 【重要】テスト画像側の特徴量をウェーブレット変換 (Pre-filter) --- + # --- 追加: テスト画像をウェーブレット変換 (Pre-filter) --- features = [wav_filter(f) for f in features] - # --- 【重要】カンペ側は既に _wav ファイルとして保存・ロードされているためそのままマッチング --- mfeatures = get_matched_ref_features(features, ref_features) - - # マッチング後の残差を計算 rfeatures = get_residual_features(features, mfeatures, pos_flag=True) - # Constraintorで空間的に滑らかに補正 rfeatures = constraintor(*rfeatures) for l in range(args.feature_levels): @@ -87,11 +84,11 @@ def validate(args, encoder, constraintor, wav_filter, estimators, test_loader, r 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 = { - 'scores1': [img_auc1, img_ap1, img_f1_score1, pix_auc1, pix_ap1, pix_f1_score1, pix_aupro1], - 'scores2': [img_auc2, img_ap2, img_f1_score2, pix_auc2, pix_ap2, pix_f1_score2, pix_aupro2], - 'scores': [img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro] - } + 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 From 5426a3d49a39b6d68e8db46905433c25985015f8 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Thu, 14 May 2026 00:12:01 +0900 Subject: [PATCH 233/258] Update main_wav.py --- main_wav.py | 186 +++++++++++++++++++++++++++++++++------------------- 1 file changed, 117 insertions(+), 69 deletions(-) diff --git a/main_wav.py b/main_wav.py index 29dbe82..6836f1d 100644 --- a/main_wav.py +++ b/main_wav.py @@ -9,6 +9,7 @@ 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 @@ -82,45 +83,77 @@ def main(args): 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.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) + 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().to(args.device) - feat_dims = encoder.feature_info.channels() + 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) # Weight decay - scheduler0 = torch.optim.lr_scheduler.MultiStepLR(optimizer0, milestones=[40, 50], gamma=0.1) # スケジューラ前倒し + 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] + # 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.005) - scheduler1 = torch.optim.lr_scheduler.MultiStepLR(optimizer1, milestones=[30, 50], gamma=0.1) + 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) - from train import train best_img_auc = 0 N_batch = 8192 @@ -129,56 +162,53 @@ def main(args): for estimator in estimators: estimator.train() - train_loader = train_loader1 if epoch < FIRST_STAGE_EPOCH else train_loader2 + 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), desc=f"Epoch[{epoch}/{args.epochs}]") + 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) + + images = images.to(args.device) + masks = masks.to(args.device) with torch.no_grad(): - # 1. 画像から特徴抽出 features = encoder(images) - - # 2. テスト画像(訓練バッチ)の特徴量をウェーブレット変換 (Pre-filter) + # --- ウェーブレット変換 (Pre-filter) --- features = [wav_filter(f) for f in features] - - # 3. 訓練用の参照特徴量を抽出し、それらもウェーブレット変換 - ref_features = get_mc_reference_features_wav(encoder, args.train_dataset_dir, class_names, images.device, args.train_ref_shot,wav_filter=wav_filter) - - - # 4. 「エッジ強調された特徴量同士」でマッチング - mfeatures = get_mc_matched_ref_features(features, class_names, ref_features) - - # 5. 残差の計算 - rfeatures = get_residual_features(features, mfeatures, pos_flag=True) + + # --- 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() - lvl_masks.append(F.interpolate(masks, size=(h, w), mode='nearest').squeeze(1)) + m = F.interpolate(masks, size=(h, w), mode='nearest').squeeze(1) + lvl_masks.append(m) rfeatures_t = [rfeature.detach().clone() for rfeature in rfeatures] - # 6. Constraintorによる空間補正 + # constraintor 適用 rfeatures = constraintor(*rfeatures) - - # (任意) 特徴量への微小ノイズ付加による過学習防止 - #noise_std = 0.01 - #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, t = e.permute(0, 2, 3, 1).reshape(-1, dim), t.permute(0, 2, 3, 1).reshape(-1, dim) - m = lvl_masks[l].reshape(-1) + 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() @@ -193,14 +223,13 @@ def main(args): 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 = [], [], [] - - # 既にextract_ref_features_wav.pyでウェーブレット変換済みのカンペをロード 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']: @@ -225,17 +254,30 @@ def main(args): test_loader = DataLoader(test_dataset, batch_size=1, shuffle=False, num_workers=8, drop_last=False) - # validate_wavに wav_filter を渡す + # 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'] - print("Epoch: {}, Class Name: {}, Image AUC: {:.3f} | Pixel AUC: {:.3f} | AUPRO: {:.3f}".format( - epoch, class_name, img_auc, pix_auc, pix_aupro)) + + # 元の 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']) - img_auc, img_ap, img_f1_score, pix_auc, pix_ap, pix_f1_score, pix_aupro = np.mean(np.array(s_res), axis=0) + 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)) @@ -246,6 +288,7 @@ def main(args): '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: @@ -258,39 +301,44 @@ def load_mc_reference_features(root_dir: str, class_names, device: torch.device, 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") # デフォルトを _wav に変更 + 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=60) # 過学習防止のためエポック数を削減 + 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('--pos_embed_dim', type=int, default=256) - parser.add_argument("--train_ref_shot", type=int, default=4) - parser.add_argument("--num_ref_shot", type=int, default=4) - parser.add_argument('--pos_beta', type=float, default=0.05) parser.add_argument('--coupling_layers', type=int, default=10) - parser.add_argument('--clamp_alpha', type=float, default=1.9) + 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('--fdm_alpha', type=float, default=0.4) parser.add_argument('--num_embeddings', type=int, default=1536) - - # ウェーブレット用パラメータ - parser.add_argument("--lf_weight", type=float, default=0.1, help="Weight for low frequency (LL)") - parser.add_argument("--hf_weight", type=float, default=1.2, help="Weight for high frequency (LH, HL, HH)") + 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) From 0d4f893b674af77061e2f4a0020e7f2671a31ab9 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Thu, 14 May 2026 18:37:47 +0900 Subject: [PATCH 234/258] Update main_wav1.py --- main_wav1.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/main_wav1.py b/main_wav1.py index 408b3b9..f299656 100644 --- a/main_wav1.py +++ b/main_wav1.py @@ -366,7 +366,7 @@ def load_mc_reference_features(root_dir: str, class_names, device: torch.device, 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('--num_embeddings', type=int, default=2048) parser.add_argument("--train_ref_shot", type=int, default=4) parser.add_argument("--num_ref_shot", type=int, default=4) From 0dd3d5452b8b6bc5573cf0935a7328eeeb7ff12a Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sat, 16 May 2026 01:04:46 +0900 Subject: [PATCH 235/258] Update validate_wav1.py --- validate_wav1.py | 413 +++++++++++++++++++++++++++++++++++------------ 1 file changed, 309 insertions(+), 104 deletions(-) diff --git a/validate_wav1.py b/validate_wav1.py index 701597f..44d7ccb 100644 --- a/validate_wav1.py +++ b/validate_wav1.py @@ -1,131 +1,336 @@ +import os import warnings +import argparse from tqdm import tqdm -from scipy.ndimage import gaussian_filter 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 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 +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') -def validate(args, encoder, vq_ops, constraintor, wav_filter, estimators, test_loader, ref_features, device, class_name): - vq_ops.eval() - constraintor.eval() - wav_filter.eval() # 追加 - for estimator in estimators: - estimator.eval() +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) - label_list, gt_mask_list = [], [] - logps1_list = [list() for _ in range(args.feature_levels)] - logps2_list = [list() for _ in range(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) - progress_bar = tqdm(total=len(test_loader)) - progress_bar.set_description(f"Evaluating {class_name}") + 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) - 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()) + 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()) - image = image.to(device) - size = image.shape[-1] + # ★ ゲーティングネットワークのパラメータをオプティマイザに追加 + 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}]") - with torch.no_grad(): - features = encoder(image) + 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) - # --- 追加: テスト画像をウェーブレット変換 (Pre-filter) --- - features = [wav_filter(f) for f in features] + with torch.no_grad(): + # 1. 生の特徴量を抽出 (まだウェーブレットはかけない) + features_raw = encoder(images) - mfeatures = get_matched_ref_features(features, ref_features) - rfeatures = get_residual_features(features, mfeatures, pos_flag=True) + # 2. 画像の深い特徴量から、層ごとの重みを予測 [w1, w2, w3] + w1, w2, w3 = gating_net(features_raw[-1].detach()) + weights = [w1, w2, w3] - # 元の validate.py と全く同じ VQ と constraintor の処理 - qx1, qx2, qx3 = vq_ops(rfeatures, train=False) - rfeatures = [qx1, qx2, qx3] - rfeatures = constraintor(*rfeatures) - + # 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) - 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)) + 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)) - 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) - - # 元の validate.py と全く同じ出力を返す - 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 - + 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 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() +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") - 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() + # 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) - scores = np.zeros_like(abnormal_map[0]) - for l in range(feature_levels): - scores += abnormal_map[l] - scores /= feature_levels + args = parser.parse_args() + init_seeds(42) - for i in range(scores.shape[0]): - scores[i] = gaussian_filter(scores[i], sigma=4) - return scores + main(args) From 0cf405e06658a2cd4a511405ee5e4a752a63f6c7 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sat, 16 May 2026 01:08:41 +0900 Subject: [PATCH 236/258] Update main_wav1.py --- main_wav1.py | 240 +++++++++++++++++++++------------------------------ 1 file changed, 98 insertions(+), 142 deletions(-) diff --git a/main_wav1.py b/main_wav1.py index f299656..44d7ccb 100644 --- a/main_wav1.py +++ b/main_wav1.py @@ -22,8 +22,8 @@ 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_wav +# 注意: 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 @@ -33,7 +33,7 @@ warnings.filterwarnings('ignore') -TOTAL_SHOT = 4 # total few-shot reference samples +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, @@ -41,14 +41,36 @@ '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 +# 1. Frequency Gating Network (層独立) # ========================================== -class HaarWaveletFilter(nn.Module): - def __init__(self, low_freq_weight=0.1, high_freq_weight=1.2): +class FrequencyGatingNetwork(nn.Module): + def __init__(self, in_channels): super().__init__() - self.lf_w = low_freq_weight - self.hf_w = high_freq_weight + 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]]) @@ -59,17 +81,17 @@ def __init__(self, low_freq_weight=0.1, high_freq_weight=1.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): + 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 * self.lf_w - hl = hl * self.hf_w - lh = lh * self.hf_w - hh = hh * self.hf_w + 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) + \ @@ -85,65 +107,36 @@ def main(args): 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 - ) + 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 - ) + 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 - ) + 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 = 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 = 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() - - 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) + # --- モジュールの初期化 --- + 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) @@ -154,6 +147,10 @@ def main(args): 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) @@ -161,7 +158,7 @@ def main(args): N_batch = 8192 for epoch in range(args.epochs): - vq_ops.train() + gating_net.train() # 学習モード constraintor.train() for estimator in estimators: estimator.train() @@ -178,19 +175,26 @@ def main(args): 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) + images, masks = images.to(args.device), masks.to(args.device) with torch.no_grad(): - features = encoder(images) - # --- 追加: ウェーブレット変換の適用 --- - features = [wav_filter(f) for f in features] + # 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)] - # --- 変更: get_mc_reference_features_wav を呼び出し、wav_filterを渡す --- - 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) + # 5. フィルタリング後の特徴量同士で残差を計算 + rfeatures = get_residual_features(features_wav, mfeatures_wav, pos_flag=True) lvl_masks = [] for l in range(args.feature_levels): @@ -199,13 +203,6 @@ def main(args): 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): @@ -231,58 +228,33 @@ def main(args): 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 = [], [], [] + # 注意: テスト用のカンペも、ウェーブレット変換をしていない「元の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)) + 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 - ) + test_loader = DataLoader(test_dataset, batch_size=1, shuffle=False, num_workers=8, drop_last=False) - # --- 変更: validate に wav_filter を渡す --- - metrics = validate(args, encoder, vq_ops, constraintor, wav_filter, estimators, test_loader, test_ref_features[class_name], args.device, class_name) + # 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'] @@ -309,42 +281,32 @@ def main(args): 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(), + # ★ ここに 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 = 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) - + 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 - 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) - + 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('--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) @@ -364,16 +326,10 @@ def load_mc_reference_features(root_dir: str, class_names, device: torch.device, 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=2048) 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) From e734251aac2d55f42442c42626c29bb9c250bee3 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Sat, 16 May 2026 01:09:00 +0900 Subject: [PATCH 237/258] Update validate_wav1.py --- validate_wav1.py | 404 ++++++++++++----------------------------------- 1 file changed, 101 insertions(+), 303 deletions(-) diff --git a/validate_wav1.py b/validate_wav1.py index 44d7ccb..9c62ba9 100644 --- a/validate_wav1.py +++ b/validate_wav1.py @@ -1,336 +1,134 @@ -import os import warnings -import argparse from tqdm import tqdm +from scipy.ndimage import gaussian_filter 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 +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') -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_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() - # --- モジュールの初期化 --- - gating_in_channels = feat_dims[-1] - gating_net = FrequencyGatingNetwork(in_channels=gating_in_channels).to(args.device) - wav_filter = HaarWaveletFilterDynamic().to(args.device) + label_list, gt_mask_list = [], [] + logps1_list = [list() for _ in range(args.feature_levels)] + logps2_list = [list() for _ in range(args.feature_levels)] - 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) + progress_bar = tqdm(total=len(test_loader)) + progress_bar.set_description(f"Evaluating {class_name}") - 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()) + for idx, batch in enumerate(test_loader): + progress_bar.update(1) - # ★ ゲーティングネットワークのパラメータをオプティマイザに追加 - 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}]") + 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()) - 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) + image = image.to(device) + size = image.shape[-1] + + with torch.no_grad(): + # 1. 生の特徴量を抽出 + features_raw = encoder(image) - # 2. 画像の深い特徴量から、層ごとの重みを予測 [w1, w2, w3] + # 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) + # 3. 生のカンペと生の特徴量をマッチング + mfeatures_raw = get_matched_ref_features(features_raw, ref_features) - # 4. テスト画像とカンペの両方に、予測した「層ごとの同じ重み」でフィルタをかける + # 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. フィルタリング後の特徴量同士で残差を計算 + # 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): + + 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'] + 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)) - 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')) + 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 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") +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 - # 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) + 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() - args = parser.parse_args() - init_seeds(42) + scores = np.zeros_like(abnormal_map[0]) + for l in range(feature_levels): + scores += abnormal_map[l] + scores /= feature_levels - main(args) + for i in range(scores.shape[0]): + scores[i] = gaussian_filter(scores[i], sigma=4) + return scores From 6bd8d22e86e16293df5b3df77d228317c93819dc Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Mon, 18 May 2026 23:32:14 +0900 Subject: [PATCH 238/258] Create main_wav_cf.py --- main_wav_cf.py | 341 +++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 341 insertions(+) create mode 100644 main_wav_cf.py diff --git a/main_wav_cf.py b/main_wav_cf.py new file mode 100644 index 0000000..2ebe3d6 --- /dev/null +++ b/main_wav_cf.py @@ -0,0 +1,341 @@ +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_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 +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. Haar Wavelet Filter (LFとHFの2成分に分離) +# ========================================== +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) + + # 高周波成分は3方向のエッジをすべて足し合わせて1つのテンソルにする + 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 + +# ========================================== +# 2. Cross Frequency (CF) Module +# ========================================== +class CrossFrequencyModule(nn.Module): + def __init__(self, in_channels): + super().__init__() + # 論文に沿って LF と HF を結合してから畳み込む + 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): + # 1. 結合 (Concat) + x = torch.cat([lf, hf], dim=1) + # 2. 畳み込み + x_out = self.conv_block(x) + # 3. 分割 (Split) + lf_out, hf_out = torch.chunk(x_out, 2, dim=1) + # 4. 残差加算 (Residual Add) + 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() + + # 特徴量を LF と HF に分けて連結するため、後段に渡す次元数は元の 2 倍になります + feat_dims_cat = [dim * 2 for dim in feat_dims] + + boundary_ops = BoundaryAverager(num_levels=args.feature_levels) + wav_filter = HaarWaveletFilter2Component().to(args.device) + + # 各階層(Layer1, 2, 3)用の CFモジュール を作成 + cf_modules = nn.ModuleList([CrossFrequencyModule(dim) for dim in feat_dims]).to(args.device) + + # 後段ネットワークは 2倍の次元(feat_dims_cat)で初期化 + constraintor = MultiScaleConv(feat_dims_cat).to(args.device) + + # オプティマイザに CFモジュール のパラメータも追加 + 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_raw = get_mc_reference_features(encoder, args.train_dataset_dir, class_names, images.device, args.train_ref_shot) + + features_cat = [] + ref_features_cat = [] + + # --- アイデア1 + CFモジュール 処理 --- + for l in range(args.feature_levels): + # 1. テスト画像とカンペを LF と HF に分離 + test_lf, test_hf = wav_filter.get_LF_HF(features_raw[l]) + + # リスト内のカンペ全てに対して分離処理を行う + ref_lfs, ref_hfs = [], [] + for ref_raw in ref_features_raw[l]: + r_lf, r_hf = wav_filter.get_LF_HF(ref_raw.unsqueeze(0)) + ref_lfs.append(r_lf.squeeze(0)) + ref_hfs.append(r_hf.squeeze(0)) + ref_lf = torch.stack(ref_lfs) + ref_hf = torch.stack(ref_hfs) + + # 2. CFモジュールを通す (周波数間の相互作用) + test_lf, test_hf = cf_modules[l](test_lf, test_hf) + + B_ref, C_ref, H_ref, W_ref = ref_lf.shape + # バッチ処理としてCFモジュールに通す + ref_lf, ref_hf = cf_modules[l](ref_lf, ref_hf) + + # 3. チャネル方向に連結 (Concat) して2倍の次元にする + features_cat.append(torch.cat([test_lf, test_hf], dim=1)) + + # カンペ側も連結し、元のカンペリスト構造に戻す + ref_cat = torch.cat([ref_lf, ref_hf], dim=1) + ref_features_cat.append(ref_cat) + + # 連結された特徴量を使ってカンペとマッチング + mfeatures = get_mc_matched_ref_features(features_cat, class_names, ref_features_cat) + 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 = [], [], [] + # テスト用のカンペも、ウェーブレット変換をしていない「元の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) + + # CFモジュールを評価関数に渡す + 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="") + # ★ ここは CNN で抽出したオリジナルの生特徴量カンペを指定します + 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) From e1a28dd45008bfc999b49986d9cd403ccf911932 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Mon, 18 May 2026 23:32:57 +0900 Subject: [PATCH 239/258] Create validate_wav_cf.py --- validate_wav_cf.py | 144 +++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 144 insertions(+) create mode 100644 validate_wav_cf.py diff --git a/validate_wav_cf.py b/validate_wav_cf.py new file mode 100644 index 0000000..37b31c5 --- /dev/null +++ b/validate_wav_cf.py @@ -0,0 +1,144 @@ +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') + +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() + + 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 = [] + ref_features_cat = [] + + # --- 推論時の アイデア1 + CFモジュール 処理 --- + for l in range(args.feature_levels): + test_lf, test_hf = wav_filter.get_LF_HF(features_raw[l]) + + # カンペ側もLFとHFを分離 + ref_lf, ref_hf = wav_filter.get_LF_HF(ref_features[l]) + + # CFモジュールを通す + test_lf, test_hf = cf_modules[l](test_lf, test_hf) + ref_lf, ref_hf = cf_modules[l](ref_lf, ref_hf) + + # チャネル方向に連結 + features_cat.append(torch.cat([test_lf, test_hf], dim=1)) + ref_features_cat.append(torch.cat([ref_lf, ref_hf], dim=1)) + + # 2倍次元のテンソル同士でマッチング + 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 From fdc81137447b87e9b478e1b477d14030b08bfb58 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Mon, 18 May 2026 23:41:10 +0900 Subject: [PATCH 240/258] Update main_wav_cf.py --- main_wav_cf.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/main_wav_cf.py b/main_wav_cf.py index 2ebe3d6..060d999 100644 --- a/main_wav_cf.py +++ b/main_wav_cf.py @@ -10,7 +10,7 @@ from torch.utils.data import DataLoader from train import train -from validate_wav1_cf import validate +from validate_wav_cf import validate from datasets.mvtec import MVTEC, MVTECANO from datasets.visa import VISA, VISAANO from datasets.btad import BTAD From 0840f40fd461e28bc2a7dd9938eb2b5fd8d61ef8 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Mon, 18 May 2026 23:57:08 +0900 Subject: [PATCH 241/258] Update main_wav_cf.py --- main_wav_cf.py | 64 ++++++++++---------------------------------------- 1 file changed, 12 insertions(+), 52 deletions(-) diff --git a/main_wav_cf.py b/main_wav_cf.py index 060d999..1e4d072 100644 --- a/main_wav_cf.py +++ b/main_wav_cf.py @@ -1,3 +1,4 @@ +# main_wav1_cf.py import os import warnings import argparse @@ -10,7 +11,7 @@ from torch.utils.data import DataLoader from train import train -from validate_wav_cf import validate +from validate_wav1_cf import validate from datasets.mvtec import MVTEC, MVTECANO from datasets.visa import VISA, VISAANO from datasets.btad import BTAD @@ -39,9 +40,6 @@ '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. Haar Wavelet Filter (LFとHFの2成分に分離) -# ========================================== class HaarWaveletFilter2Component(nn.Module): def __init__(self): super().__init__() @@ -64,7 +62,6 @@ def get_LF_HF(self, x): lf = F.conv_transpose2d(ll, self.k_ll.expand(C, 1, 2, 2), stride=2, groups=C) - # 高周波成分は3方向のエッジをすべて足し合わせて1つのテンソルにする 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) @@ -72,13 +69,9 @@ def get_LF_HF(self, x): return lf, hf -# ========================================== -# 2. Cross Frequency (CF) Module -# ========================================== class CrossFrequencyModule(nn.Module): def __init__(self, in_channels): super().__init__() - # 論文に沿って LF と HF を結合してから畳み込む 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), @@ -86,13 +79,9 @@ def __init__(self, in_channels): ) def forward(self, lf, hf): - # 1. 結合 (Concat) x = torch.cat([lf, hf], dim=1) - # 2. 畳み込み x_out = self.conv_block(x) - # 3. 分割 (Split) lf_out, hf_out = torch.chunk(x_out, 2, dim=1) - # 4. 残差加算 (Residual Add) lf_refined = lf + lf_out hf_refined = hf + hf_out return lf_refined, hf_refined @@ -100,7 +89,6 @@ def forward(self, lf, hf): 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) @@ -124,19 +112,13 @@ def main(args): 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() - # 特徴量を LF と HF に分けて連結するため、後段に渡す次元数は元の 2 倍になります feat_dims_cat = [dim * 2 for dim in feat_dims] boundary_ops = BoundaryAverager(num_levels=args.feature_levels) wav_filter = HaarWaveletFilter2Component().to(args.device) - - # 各階層(Layer1, 2, 3)用の CFモジュール を作成 cf_modules = nn.ModuleList([CrossFrequencyModule(dim) for dim in feat_dims]).to(args.device) - # 後段ネットワークは 2倍の次元(feat_dims_cat)で初期化 constraintor = MultiScaleConv(feat_dims_cat).to(args.device) - - # オプティマイザに CFモジュール のパラメータも追加 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) @@ -169,44 +151,27 @@ def main(args): with torch.no_grad(): features_raw = encoder(images) - # 生のカンペを取得 ref_features_raw = get_mc_reference_features(encoder, args.train_dataset_dir, class_names, images.device, args.train_ref_shot) features_cat = [] - ref_features_cat = [] - # --- アイデア1 + CFモジュール 処理 --- for l in range(args.feature_levels): - # 1. テスト画像とカンペを LF と HF に分離 test_lf, test_hf = wav_filter.get_LF_HF(features_raw[l]) - - # リスト内のカンペ全てに対して分離処理を行う - ref_lfs, ref_hfs = [], [] - for ref_raw in ref_features_raw[l]: - r_lf, r_hf = wav_filter.get_LF_HF(ref_raw.unsqueeze(0)) - ref_lfs.append(r_lf.squeeze(0)) - ref_hfs.append(r_hf.squeeze(0)) - ref_lf = torch.stack(ref_lfs) - ref_hf = torch.stack(ref_hfs) - - # 2. CFモジュールを通す (周波数間の相互作用) test_lf, test_hf = cf_modules[l](test_lf, test_hf) - - B_ref, C_ref, H_ref, W_ref = ref_lf.shape - # バッチ処理としてCFモジュールに通す - ref_lf, ref_hf = cf_modules[l](ref_lf, ref_hf) - - # 3. チャネル方向に連結 (Concat) して2倍の次元にする features_cat.append(torch.cat([test_lf, test_hf], dim=1)) - # カンペ側も連結し、元のカンペリスト構造に戻す - ref_cat = torch.cat([ref_lf, ref_hf], dim=1) - ref_features_cat.append(ref_cat) + ref_features_cat_dict = {} + for c_name, refs_tuple in ref_features_raw.items(): + refs_cat_list = [] + for l in range(args.feature_levels): + ref_l = refs_tuple[l] + r_lf, r_hf = wav_filter.get_LF_HF(ref_l) + r_lf, r_hf = cf_modules[l](r_lf, r_hf) + refs_cat_list.append(torch.cat([r_lf, r_hf], dim=1)) + ref_features_cat_dict[c_name] = tuple(refs_cat_list) - # 連結された特徴量を使ってカンペとマッチング - mfeatures = get_mc_matched_ref_features(features_cat, class_names, ref_features_cat) + 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): @@ -246,10 +211,8 @@ def main(args): 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']: @@ -265,7 +228,6 @@ def main(args): test_loader = DataLoader(test_dataset, batch_size=1, shuffle=False, num_workers=8, drop_last=False) - # CFモジュールを評価関数に渡す 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'] @@ -311,7 +273,6 @@ def load_mc_reference_features(root_dir: str, class_names, device: torch.device, 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="") - # ★ ここは CNN で抽出したオリジナルの生特徴量カンペを指定します 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) @@ -323,7 +284,6 @@ def load_mc_reference_features(root_dir: str, class_names, device: torch.device, 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) From 603b0c30457a26b84bba7d4c1224981c391629f9 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Mon, 18 May 2026 23:57:32 +0900 Subject: [PATCH 242/258] Update validate_wav_cf.py --- validate_wav_cf.py | 25 ++++++++----------------- 1 file changed, 8 insertions(+), 17 deletions(-) diff --git a/validate_wav_cf.py b/validate_wav_cf.py index 37b31c5..97aef2a 100644 --- a/validate_wav_cf.py +++ b/validate_wav_cf.py @@ -1,3 +1,4 @@ +# validate_wav1_cf.py import warnings from tqdm import tqdm from scipy.ndimage import gaussian_filter @@ -7,7 +8,6 @@ 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 @@ -19,6 +19,13 @@ def validate(args, encoder, constraintor, wav_filter, cf_modules, estimators, te cf_modules.eval() for estimator in estimators: estimator.eval() + + ref_features_cat = [] + with torch.no_grad(): + for l in range(args.feature_levels): + ref_lf, ref_hf = wav_filter.get_LF_HF(ref_features[l]) + ref_lf, ref_hf = cf_modules[l](ref_lf, ref_hf) + ref_features_cat.append(torch.cat([ref_lf, ref_hf], dim=1)) label_list, gt_mask_list = [], [] logps1_list = [list() for _ in range(args.feature_levels)] @@ -39,31 +46,15 @@ def validate(args, encoder, constraintor, wav_filter, cf_modules, estimators, te with torch.no_grad(): features_raw = encoder(image) - features_cat = [] - ref_features_cat = [] - # --- 推論時の アイデア1 + CFモジュール 処理 --- for l in range(args.feature_levels): test_lf, test_hf = wav_filter.get_LF_HF(features_raw[l]) - - # カンペ側もLFとHFを分離 - ref_lf, ref_hf = wav_filter.get_LF_HF(ref_features[l]) - - # CFモジュールを通す test_lf, test_hf = cf_modules[l](test_lf, test_hf) - ref_lf, ref_hf = cf_modules[l](ref_lf, ref_hf) - - # チャネル方向に連結 features_cat.append(torch.cat([test_lf, test_hf], dim=1)) - ref_features_cat.append(torch.cat([ref_lf, ref_hf], dim=1)) - # 2倍次元のテンソル同士でマッチング 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) From 138603bad160033727e8a4c799944b8f3a5e96a3 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Mon, 18 May 2026 23:58:01 +0900 Subject: [PATCH 243/258] Update main_wav_cf.py --- main_wav_cf.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/main_wav_cf.py b/main_wav_cf.py index 1e4d072..e3ccd63 100644 --- a/main_wav_cf.py +++ b/main_wav_cf.py @@ -11,7 +11,7 @@ from torch.utils.data import DataLoader from train import train -from validate_wav1_cf import validate +from validate_wav_cf import validate from datasets.mvtec import MVTEC, MVTECANO from datasets.visa import VISA, VISAANO from datasets.btad import BTAD From 53860f360b502d22d4421bd23f7e7bb4273a5d3b Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Tue, 19 May 2026 00:03:45 +0900 Subject: [PATCH 244/258] Update main_wav_cf.py --- main_wav_cf.py | 49 +++++++++++++++++++++++++++++++++---------------- 1 file changed, 33 insertions(+), 16 deletions(-) diff --git a/main_wav_cf.py b/main_wav_cf.py index e3ccd63..36596b7 100644 --- a/main_wav_cf.py +++ b/main_wav_cf.py @@ -1,4 +1,3 @@ -# main_wav1_cf.py import os import warnings import argparse @@ -23,7 +22,9 @@ 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 +# ★ 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 @@ -40,6 +41,31 @@ '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__() @@ -61,12 +87,10 @@ def get_LF_HF(self, x): 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): @@ -151,27 +175,20 @@ def main(args): with torch.no_grad(): features_raw = encoder(images) - ref_features_raw = get_mc_reference_features(encoder, args.train_dataset_dir, class_names, images.device, args.train_ref_shot) + # --- 修正箇所:専用関数で安全にカンペを抽出 --- + 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)) - ref_features_cat_dict = {} - for c_name, refs_tuple in ref_features_raw.items(): - refs_cat_list = [] - for l in range(args.feature_levels): - ref_l = refs_tuple[l] - r_lf, r_hf = wav_filter.get_LF_HF(ref_l) - r_lf, r_hf = cf_modules[l](r_lf, r_hf) - refs_cat_list.append(torch.cat([r_lf, r_hf], dim=1)) - ref_features_cat_dict[c_name] = tuple(refs_cat_list) - 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): From 60f54e340a2b28e37b1a0f26db81c8599134af2d Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Tue, 19 May 2026 00:04:05 +0900 Subject: [PATCH 245/258] Update validate_wav_cf.py --- validate_wav_cf.py | 27 ++++++++++++++++++++++++--- 1 file changed, 24 insertions(+), 3 deletions(-) diff --git a/validate_wav_cf.py b/validate_wav_cf.py index 97aef2a..087d0d2 100644 --- a/validate_wav_cf.py +++ b/validate_wav_cf.py @@ -1,4 +1,3 @@ -# validate_wav1_cf.py import warnings from tqdm import tqdm from scipy.ndimage import gaussian_filter @@ -21,11 +20,31 @@ def validate(args, encoder, constraintor, wav_filter, cf_modules, estimators, te 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): - ref_lf, ref_hf = wav_filter.get_LF_HF(ref_features[l]) + # エンコーダの出力から正しい解像度(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) - ref_features_cat.append(torch.cat([ref_lf, ref_hf], dim=1)) + + # 連結して再びマッチング用の平坦な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)] @@ -51,6 +70,8 @@ def validate(args, encoder, constraintor, wav_filter, cf_modules, estimators, te 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) From f054ecca552a123f05a8bd6d6ceb856e16ddb25d Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 20 May 2026 00:58:47 +0900 Subject: [PATCH 246/258] Create main_global.py --- main_global.py | 319 +++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 319 insertions(+) create mode 100644 main_global.py diff --git a/main_global.py b/main_global.py new file mode 100644 index 0000000..86f92dc --- /dev/null +++ b/main_global.py @@ -0,0 +1,319 @@ +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) From cf24930bdf70b433a2530ce802ad9cfd6f40d81d Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 20 May 2026 00:59:17 +0900 Subject: [PATCH 247/258] Create validate_global.py --- validate_global.py | 1 + 1 file changed, 1 insertion(+) create mode 100644 validate_global.py diff --git a/validate_global.py b/validate_global.py new file mode 100644 index 0000000..8b13789 --- /dev/null +++ b/validate_global.py @@ -0,0 +1 @@ + From 343e323fca7be5058738d14a28b7ec16130851c8 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 20 May 2026 01:07:54 +0900 Subject: [PATCH 248/258] Update validate_global.py --- validate_global.py | 136 +++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 136 insertions(+) diff --git a/validate_global.py b/validate_global.py index 8b13789..916c15c 100644 --- a/validate_global.py +++ b/validate_global.py @@ -1 +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 From db3b1022380ccf2115207c70083f111c73537603 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 20 May 2026 01:35:56 +0900 Subject: [PATCH 249/258] Update main_global.py --- main_global.py | 80 ++++++++++++++++++++++++++++++++++++-------------- 1 file changed, 58 insertions(+), 22 deletions(-) diff --git a/main_global.py b/main_global.py index 86f92dc..ded6db5 100644 --- a/main_global.py +++ b/main_global.py @@ -4,7 +4,6 @@ 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 @@ -21,6 +20,7 @@ from datasets.capsules import CAPSULES, CAPSULESANO from models.fc_flow import load_flow_model +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 @@ -35,19 +35,28 @@ 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_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} +# ========================================== +# アイデア3: Global Context Constraintor(情報ボトルネック内蔵) +# ========================================== class GlobalContextConstraintor(nn.Module): - def __init__(self, feat_dims, num_heads=4, num_layers=1): + def __init__(self, feat_dims, num_heads=4, num_layers=1, bottleneck_ratio=4): super().__init__() self.num_levels = len(feat_dims) self.local_convs = nn.ModuleList() self.transformers = nn.ModuleList() + # 丸暗記(恒等写像)を防止するための情報ボトルネック + self.proj_down = nn.ModuleList() + self.proj_up = nn.ModuleList() + + # Layer 1(浅い層)はテクスチャ重視のためCNNのみ、Layer 2,3(深い層)は論理構造重視のためTransformerを適用 self.apply_transformer = [l > 0 for l in range(self.num_levels)] for i, dim in enumerate(feat_dims): + # 1. 局所ノイズ平滑化用CNN conv = nn.Sequential( nn.Conv2d(dim, dim, kernel_size=3, padding=1, bias=False), nn.BatchNorm2d(dim), @@ -58,12 +67,17 @@ def __init__(self, feat_dims, num_heads=4, num_layers=1): ) self.local_convs.append(conv) + # 2. 大域論理検証用Transformer + ボトルネック if self.apply_transformer[i]: - heads = num_heads if dim % num_heads == 0 else 1 + compressed_dim = max(dim // bottleneck_ratio, 16) + self.proj_down.append(nn.Linear(dim, compressed_dim)) + self.proj_up.append(nn.Linear(compressed_dim, dim)) + + heads = num_heads if compressed_dim % num_heads == 0 else 1 encoder_layer = nn.TransformerEncoderLayer( - d_model=dim, + d_model=compressed_dim, nhead=heads, - dim_feedforward=dim * 2, + dim_feedforward=compressed_dim * 2, activation='relu', batch_first=True, dropout=0.1 @@ -71,6 +85,8 @@ def __init__(self, feat_dims, num_heads=4, num_layers=1): self.transformers.append(nn.TransformerEncoder(encoder_layer, num_layers=num_layers)) else: self.transformers.append(nn.Identity()) + self.proj_down.append(nn.Identity()) + self.proj_up.append(nn.Identity()) def forward(self, *features): out_features = [] @@ -78,18 +94,27 @@ def forward(self, *features): x = features[i] B, C, H, W = x.shape + # CNNによる局所的な平滑化 x_local = self.local_convs[i](x) + x if self.apply_transformer[i]: + # 2D Sinusoidal 位置エンコーディング 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_local.view(B, C, -1).permute(0, 2, 1) # (B, H*W, C) x_flat = x_flat + pos_embed - x_global = self.transformers[i](x_flat) + # 情報ボトルネックによる次元圧縮 + x_compressed = self.proj_down[i](x_flat) + + # Transformerによる大域コンテキスト補正 + x_global = self.transformers[i](x_compressed) + + # 次元復元と空間次元への再構成 + x_restored = self.proj_up[i](x_global) + x_out = x_restored.permute(0, 2, 1).view(B, C, H, W) - 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) @@ -100,9 +125,7 @@ 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) + grid_h, grid_w = grid_h.reshape(-1), 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) @@ -116,12 +139,8 @@ 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) + emb = torch.cat([torch.sin(out), torch.cos(out)], dim=1) if embed_dim % 2 != 0: emb = F.pad(emb, (0, 1)) return emb @@ -154,8 +173,13 @@ def main(args): boundary_ops = BoundaryAverager(num_levels=args.feature_levels) - constraintor = GlobalContextConstraintor(feat_dims).to(args.device) + # 元の VQモジュール を組み戻し + 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 = 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) @@ -170,6 +194,7 @@ def main(args): N_batch = 8192 for epoch in range(args.epochs): + vq_ops.train() constraintor.train() for estimator in estimators: estimator.train() @@ -187,9 +212,7 @@ def main(args): 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 = [] @@ -199,6 +222,15 @@ def main(args): lvl_masks.append(m) rfeatures_t = [rfeature.detach().clone() for rfeature in rfeatures] + # --- VQモジュールの最適化 (元の処理) --- + 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() + + # --- 新Constraintor (大域ハイブリッド) の最適化 --- rfeatures = constraintor(*rfeatures) loss = 0 for l in range(args.feature_levels): @@ -224,6 +256,7 @@ def main(args): train_loss_total += loss total_num += num + scheduler_vq.step() scheduler0.step() scheduler1.step() @@ -247,7 +280,8 @@ def main(args): 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) + # validateに関数を正しく引き渡す + 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( @@ -269,7 +303,8 @@ def main(args): 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(), + 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')) @@ -311,6 +346,7 @@ def load_mc_reference_features(root_dir: str, class_names, device: torch.device, 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) From 8681c51549274020984bf3592bdedce7f162430e Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 20 May 2026 01:36:06 +0900 Subject: [PATCH 250/258] Update validate_global.py --- validate_global.py | 10 +++++++--- 1 file changed, 7 insertions(+), 3 deletions(-) diff --git a/validate_global.py b/validate_global.py index 916c15c..4233906 100644 --- a/validate_global.py +++ b/validate_global.py @@ -29,8 +29,10 @@ def get_matched_ref_features(features, ref_features): matched_ref_features.append(index_feats) return matched_ref_features -def validate(args, encoder, constraintor, estimators, test_loader, ref_features, device, class_name): +# 引数に vq_ops を追加 +def validate(args, encoder, vq_ops, constraintor, estimators, test_loader, ref_features, device, class_name): constraintor.eval() + vq_ops.eval() for estimator in estimators: estimator.eval() @@ -53,11 +55,13 @@ def validate(args, encoder, constraintor, estimators, test_loader, ref_features, 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 = vq_ops(rfeatures, train=False) + + # 補正モジュール(大域ハイブリッド)に通過 rfeatures = constraintor(*rfeatures) for l in range(args.feature_levels): From 99f83178be03151e2c4b1e8951ba2ba52e3fa0b7 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 20 May 2026 01:37:56 +0900 Subject: [PATCH 251/258] Update main_global.py --- main_global.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/main_global.py b/main_global.py index ded6db5..4ee36d1 100644 --- a/main_global.py +++ b/main_global.py @@ -7,7 +7,7 @@ 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_global import validate from datasets.mvtec import MVTEC, MVTECANO From 6c4de6e5a17e842684c39ea3b5a57ea53335807e Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 20 May 2026 02:07:49 +0900 Subject: [PATCH 252/258] Update main_global.py --- main_global.py | 85 +++++++++++++++----------------------------------- 1 file changed, 25 insertions(+), 60 deletions(-) diff --git a/main_global.py b/main_global.py index 4ee36d1..2963e46 100644 --- a/main_global.py +++ b/main_global.py @@ -4,10 +4,11 @@ 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 -import torch.nn as nn + from train import train from validate_global import validate from datasets.mvtec import MVTEC, MVTECANO @@ -20,7 +21,6 @@ from datasets.capsules import CAPSULES, CAPSULESANO from models.fc_flow import load_flow_model -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 @@ -35,28 +35,19 @@ 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_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} -# ========================================== -# アイデア3: Global Context Constraintor(情報ボトルネック内蔵) -# ========================================== class GlobalContextConstraintor(nn.Module): - def __init__(self, feat_dims, num_heads=4, num_layers=1, bottleneck_ratio=4): + 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.proj_down = nn.ModuleList() - self.proj_up = nn.ModuleList() - - # Layer 1(浅い層)はテクスチャ重視のためCNNのみ、Layer 2,3(深い層)は論理構造重視のためTransformerを適用 self.apply_transformer = [l > 0 for l in range(self.num_levels)] for i, dim in enumerate(feat_dims): - # 1. 局所ノイズ平滑化用CNN conv = nn.Sequential( nn.Conv2d(dim, dim, kernel_size=3, padding=1, bias=False), nn.BatchNorm2d(dim), @@ -67,17 +58,12 @@ def __init__(self, feat_dims, num_heads=4, num_layers=1, bottleneck_ratio=4): ) self.local_convs.append(conv) - # 2. 大域論理検証用Transformer + ボトルネック if self.apply_transformer[i]: - compressed_dim = max(dim // bottleneck_ratio, 16) - self.proj_down.append(nn.Linear(dim, compressed_dim)) - self.proj_up.append(nn.Linear(compressed_dim, dim)) - - heads = num_heads if compressed_dim % num_heads == 0 else 1 + heads = num_heads if dim % num_heads == 0 else 1 encoder_layer = nn.TransformerEncoderLayer( - d_model=compressed_dim, + d_model=dim, nhead=heads, - dim_feedforward=compressed_dim * 2, + dim_feedforward=dim * 2, activation='relu', batch_first=True, dropout=0.1 @@ -85,8 +71,6 @@ def __init__(self, feat_dims, num_heads=4, num_layers=1, bottleneck_ratio=4): self.transformers.append(nn.TransformerEncoder(encoder_layer, num_layers=num_layers)) else: self.transformers.append(nn.Identity()) - self.proj_down.append(nn.Identity()) - self.proj_up.append(nn.Identity()) def forward(self, *features): out_features = [] @@ -94,27 +78,18 @@ def forward(self, *features): x = features[i] B, C, H, W = x.shape - # CNNによる局所的な平滑化 x_local = self.local_convs[i](x) + x if self.apply_transformer[i]: - # 2D Sinusoidal 位置エンコーディング 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) # (B, H*W, C) + x_flat = x_local.view(B, C, -1).permute(0, 2, 1) x_flat = x_flat + pos_embed - # 情報ボトルネックによる次元圧縮 - x_compressed = self.proj_down[i](x_flat) - - # Transformerによる大域コンテキスト補正 - x_global = self.transformers[i](x_compressed) - - # 次元復元と空間次元への再構成 - x_restored = self.proj_up[i](x_global) - x_out = x_restored.permute(0, 2, 1).view(B, C, H, W) + 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) @@ -125,7 +100,9 @@ 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_w = grid_h.reshape(-1), grid_w.reshape(-1) + + 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) @@ -139,8 +116,12 @@ 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 = torch.cat([torch.sin(out), torch.cos(out)], dim=1) + 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 @@ -173,13 +154,8 @@ def main(args): boundary_ops = BoundaryAverager(num_levels=args.feature_levels) - # 元の VQモジュール を組み戻し - 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 = 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) @@ -194,7 +170,6 @@ def main(args): N_batch = 8192 for epoch in range(args.epochs): - vq_ops.train() constraintor.train() for estimator in estimators: estimator.train() @@ -212,7 +187,9 @@ def main(args): 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 = [] @@ -222,15 +199,6 @@ def main(args): lvl_masks.append(m) rfeatures_t = [rfeature.detach().clone() for rfeature in rfeatures] - # --- VQモジュールの最適化 (元の処理) --- - 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() - - # --- 新Constraintor (大域ハイブリッド) の最適化 --- rfeatures = constraintor(*rfeatures) loss = 0 for l in range(args.feature_levels): @@ -256,7 +224,6 @@ def main(args): train_loss_total += loss total_num += num - scheduler_vq.step() scheduler0.step() scheduler1.step() @@ -280,8 +247,7 @@ def main(args): test_loader = DataLoader(test_dataset, batch_size=1, shuffle=False, num_workers=8, drop_last=False) - # validateに関数を正しく引き渡す - metrics = validate(args, encoder, vq_ops, constraintor, estimators, test_loader, test_ref_features[class_name], args.device, class_name) + 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( @@ -303,8 +269,7 @@ def main(args): 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(), + 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')) @@ -346,10 +311,10 @@ def load_mc_reference_features(root_dir: str, class_names, device: torch.device, 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) args = parser.parse_args() init_seeds(42) - main(args) + main(args)s + From 82a88109642730f5d3b2d8e4cd9196ed3784fcd0 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 20 May 2026 02:08:11 +0900 Subject: [PATCH 253/258] Refactor validate function to remove vq_ops Removed vq_ops parameter from validate function and its usage. --- validate_global.py | 10 +++------- 1 file changed, 3 insertions(+), 7 deletions(-) diff --git a/validate_global.py b/validate_global.py index 4233906..916c15c 100644 --- a/validate_global.py +++ b/validate_global.py @@ -29,10 +29,8 @@ def get_matched_ref_features(features, ref_features): matched_ref_features.append(index_feats) return matched_ref_features -# 引数に vq_ops を追加 -def validate(args, encoder, vq_ops, constraintor, estimators, test_loader, ref_features, device, class_name): +def validate(args, encoder, constraintor, estimators, test_loader, ref_features, device, class_name): constraintor.eval() - vq_ops.eval() for estimator in estimators: estimator.eval() @@ -55,13 +53,11 @@ def validate(args, encoder, vq_ops, constraintor, estimators, test_loader, ref_f 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 = vq_ops(rfeatures, train=False) + rfeatures = get_residual_features(features, mfeatures, pos_flag=True) - # 補正モジュール(大域ハイブリッド)に通過 rfeatures = constraintor(*rfeatures) for l in range(args.feature_levels): From 15ace0226ee3abf46d276e4af563dac22d295849 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Wed, 20 May 2026 02:09:35 +0900 Subject: [PATCH 254/258] Update main_global.py --- main_global.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/main_global.py b/main_global.py index 2963e46..5fe0e9e 100644 --- a/main_global.py +++ b/main_global.py @@ -316,5 +316,5 @@ def load_mc_reference_features(root_dir: str, class_names, device: torch.device, args = parser.parse_args() init_seeds(42) - main(args)s + main(args) From 7fe7f21f0d4b9e05443e9778108e0f800977a63c Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Thu, 21 May 2026 00:33:00 +0900 Subject: [PATCH 255/258] Create main_freq_blend.py --- main_freq_blend.py | 352 +++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 352 insertions(+) create mode 100644 main_freq_blend.py 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) From 261dbcdf1013be31b1add959e4a05e94c6a84095 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Thu, 21 May 2026 00:35:22 +0900 Subject: [PATCH 256/258] Create validate_freq_blend.py --- validate_freq_blend.py | 181 +++++++++++++++++++++++++++++++++++++++++ 1 file changed, 181 insertions(+) create mode 100644 validate_freq_blend.py diff --git a/validate_freq_blend.py b/validate_freq_blend.py new file mode 100644 index 0000000..03facd9 --- /dev/null +++ b/validate_freq_blend.py @@ -0,0 +1,181 @@ +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 = [] + + # 事前計算: 既存の未分離カンペ(.npy)を空間次元に復元してLFとHFに分離する + 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) + + pos_embed = get_position_encoding(dim, h, w).to(device).unsqueeze(0).repeat(bs, 1, 1, 1) + pos_embed = pos_embed.permute(0, 2, 3, 1).reshape(-1, 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 From eb29e35cee46ce166f3a318e4c5769a56e28fc72 Mon Sep 17 00:00:00 2001 From: tomo082 <131239927+tomo082@users.noreply.github.com> Date: Thu, 21 May 2026 00:56:09 +0900 Subject: [PATCH 257/258] Update validate_freq_blend.py --- validate_freq_blend.py | 7 +++---- 1 file changed, 3 insertions(+), 4 deletions(-) diff --git a/validate_freq_blend.py b/validate_freq_blend.py index 03facd9..80948e4 100644 --- a/validate_freq_blend.py +++ b/validate_freq_blend.py @@ -54,7 +54,6 @@ def validate(args, encoder, constraintor, wav_filter, estimators, test_loader, r ref_lf_cat = [] ref_hf_cat = [] - # 事前計算: 既存の未分離カンペ(.npy)を空間次元に復元してLFとHFに分離する with torch.no_grad(): dummy_img = torch.zeros(1, 3, 224, 224).to(device) dummy_feats = encoder(dummy_img) @@ -96,7 +95,6 @@ def validate(args, encoder, constraintor, wav_filter, estimators, test_loader, r 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 @@ -109,8 +107,9 @@ def validate(args, encoder, constraintor, wav_filter, estimators, test_loader, r bs, dim, h, w = e.size() e = e.permute(0, 2, 3, 1).reshape(-1, dim) - pos_embed = get_position_encoding(dim, h, w).to(device).unsqueeze(0).repeat(bs, 1, 1, 1) - pos_embed = pos_embed.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': From 49eb0f71a945cfeb6d71357b3e217b117096bb73 Mon Sep 17 00:00:00 2001 From: Tomoya Ueno Date: Mon, 29 Jun 2026 05:35:15 +0900 Subject: [PATCH 258/258] Add AdaCLIP text-reference map evaluation --- scripts/eval_text_ref_map.py | 846 +++++++++++++++++++++++++++++++++++ 1 file changed, 846 insertions(+) create mode 100644 scripts/eval_text_ref_map.py 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()