Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
261 commits
Select commit Hold shift + click to select a range
dfef067
Update main_ib.py
tomo082 Jun 15, 2025
3fc5b09
Update main.py
tomo082 Jun 19, 2025
6d86c70
Update main.py
tomo082 Jun 19, 2025
42848f0
Update mvtec.py
tomo082 Jun 19, 2025
0e738b0
Update validate.py
tomo082 Jun 19, 2025
95858c0
Update main.py
tomo082 Jun 19, 2025
7eca7ce
紛らわしいのでclass_nameをclass_name_evalに変更
tomo082 Jun 19, 2025
32162cf
class_namesのエラーについて
tomo082 Jun 19, 2025
c96043b
unpackのエラー
tomo082 Jun 19, 2025
54cd05a
anomaly_typesを追加
tomo082 Jun 19, 2025
7fdc037
anomaly_typesをテキストで保存
tomo082 Jun 19, 2025
dd4da74
特徴量を上書き保存するように変更
tomo082 Jun 19, 2025
1979bd1
取ってくる特徴量を変更
tomo082 Jun 19, 2025
1d049a8
visualizer適応
tomo082 Jun 20, 2025
447af86
可視化は最後のエポックのみ適応
tomo082 Jun 20, 2025
2700503
タイプミス修正
tomo082 Jun 20, 2025
586bf9d
Create mvtec_fewclass.py
tomo082 Jun 21, 2025
7ac66ff
MVTECFEW
tomo082 Jun 21, 2025
be8ead5
Update mvtec_fewclass.py
tomo082 Jun 21, 2025
382d9e8
設定追加
tomo082 Jun 21, 2025
1e7f616
設定追加
tomo082 Jun 21, 2025
01a8beb
Update classes.py
tomo082 Jun 21, 2025
894bbed
Merge branch 'main' of https://github.com/tomo082/ResAD_SSL
tomo082 Jun 21, 2025
9c226d1
Update classes.py
tomo082 Jun 21, 2025
605f470
a
tomo082 Jun 21, 2025
e87c40f
Merge branch 'main' of https://github.com/tomo082/ResAD_SSL
tomo082 Jun 21, 2025
f1746b3
fewclass setting
tomo082 Jun 21, 2025
2ca6c7a
Update classes.py
tomo082 Jun 21, 2025
41c6cbc
6/27ゼミ
tomo082 Jun 27, 2025
4e7eed0
Update main_ib.p
tomo082 Jun 30, 2025
29e9796
main_ibをmainと揃える
tomo082 Jun 30, 2025
3b06c32
Update main_ib.py
tomo082 Jun 30, 2025
54d8c93
Update main_ib.py
tomo082 Jun 30, 2025
cff925f
Update main_ib.py
tomo082 Jun 30, 2025
1230492
Update main_ib.py
tomo082 Jun 30, 2025
6142cd9
Update main_ib.py
tomo082 Jun 30, 2025
ba46e84
Update main_ib.py
tomo082 Jun 30, 2025
cf98326
Update main_ib.py
tomo082 Jun 30, 2025
53100e7
Update main_ib.py
tomo082 Jun 30, 2025
791f82d
残差を使わない設定を追加
tomo082 Jul 4, 2025
c8ef4ed
残差を使わないように変更
tomo082 Jul 4, 2025
203ace5
Update main.py
tomo082 Oct 18, 2025
550c029
Update classes.py
tomo082 Oct 18, 2025
7cd89ba
Update main.py
tomo082 Oct 18, 2025
e5db4c9
Update main.py
tomo082 Oct 18, 2025
340352b
Update main.py
tomo082 Oct 18, 2025
d12ed77
Update classes.py
tomo082 Oct 18, 2025
772f017
Update validate.py
tomo082 Oct 26, 2025
c21f2c3
Update main.py
tomo082 Oct 26, 2025
daded88
Update main.py
tomo082 Oct 26, 2025
cf5b41a
Update main.py
tomo082 Oct 26, 2025
75359f7
Update extract_ref_features.py
tomo082 Oct 26, 2025
b143c32
Update extract_ref_features.py
tomo082 Oct 26, 2025
5b24649
Update extract_ref_features.py
tomo082 Oct 26, 2025
317689c
Update extract_ref_features.py
tomo082 Oct 26, 2025
26e598e
Update extract_ref_features.py
tomo082 Oct 26, 2025
afbff24
Update extract_ref_features.py
tomo082 Oct 26, 2025
8787658
Add VISA and VISAANO dataset classes
tomo082 Dec 13, 2025
1a8a0a9
Update capsule_visa.py
tomo082 Dec 13, 2025
3b7d1ed
Add VISACAPSULES_TO_VISACAPSULES mapping
tomo082 Dec 13, 2025
5920e0a
Rename class VISACAPSULESANO to CAPSULESANO
tomo082 Dec 13, 2025
49ae25f
Update classes.py
tomo082 Dec 13, 2025
7d41316
Update main.py
tomo082 Dec 13, 2025
69a75ab
Update main.py
tomo082 Dec 13, 2025
deabd80
Update main.py
tomo082 Dec 13, 2025
101a7b1
Update main.py
tomo082 Dec 13, 2025
398aa06
Update classes.py
tomo082 Dec 13, 2025
e170345
Update capsules.py
tomo082 Dec 13, 2025
feef1c2
Update extract_ref_features.py
tomo082 Dec 16, 2025
b334fe9
Add load_weights function to utils.py
tomo082 Dec 16, 2025
905b1c5
Update extract_ref_features.py
tomo082 Dec 16, 2025
1023ba3
Update extract_ref_features.py
tomo082 Dec 16, 2025
5b8a30e
Update extract_ref_features.py
tomo082 Dec 16, 2025
02f3ad0
Update extract_ref_features.py
tomo082 Dec 16, 2025
41fb21c
Update extract_ref_features.py
tomo082 Dec 16, 2025
c57e409
Update main.py
tomo082 Dec 16, 2025
77e2b17
Update extract_ref_features.py
tomo082 Dec 16, 2025
97b1d0d
Update utils.py
tomo082 Dec 16, 2025
c2b1012
Update utils.py
tomo082 Dec 18, 2025
3ac708e
Update utils.py
tomo082 Dec 18, 2025
1bea5ca
Update utils.py
tomo082 Dec 18, 2025
99f7c89
Update utils.py
tomo082 Dec 18, 2025
eb60720
Update utils.py
tomo082 Dec 18, 2025
bd37d10
Update utils.py
tomo082 Dec 19, 2025
8dc2912
Update utils.py
tomo082 Dec 19, 2025
e7780db
Create main_ad.py
tomo082 Jan 2, 2026
66ec0ad
Merge branch 'main' of https://github.com/tomo082/ResAD_SSL
tomo082 Jan 8, 2026
66fb777
Update main_ad.py
tomo082 Jan 8, 2026
2f6c656
Update utils.py
tomo082 Jan 8, 2026
7a23f3c
Update main_ib.py
tomo082 Jan 8, 2026
42201b4
Update main_ad.py
tomo082 Jan 8, 2026
9caa0c4
Update utils.py
tomo082 Jan 8, 2026
ca9414d
Update main_ad.py
tomo082 Jan 8, 2026
0baafff
Update main_ad.py
tomo082 Jan 8, 2026
fe6593b
Update classes.py
tomo082 Jan 10, 2026
e66c76f
Update extract_ref_features.py
tomo082 Jan 10, 2026
90d4223
Update main.py
tomo082 Jan 10, 2026
42b8c8e
Update validate.py
tomo082 Jan 10, 2026
216b759
Update visualizer.py
tomo082 Jan 10, 2026
42ce560
Create validate1.py
tomo082 Feb 14, 2026
96e81d0
Implement get_matched_ref_features_top function
tomo082 Feb 14, 2026
f9d7d01
Add get_matched_ref_features_top import to validate1.py
tomo082 Feb 14, 2026
e7aa165
Update main.py
tomo082 Feb 14, 2026
5711032
Update validate1.py
tomo082 Feb 14, 2026
ffcd2c3
Update validate1.py
tomo082 Feb 14, 2026
322b129
Update validate1.py
tomo082 Feb 14, 2026
eeb729c
Update extract_ref_features.py
tomo082 Feb 16, 2026
3e169b7
Update validate.py
tomo082 Apr 12, 2026
6822cd1
Create main_vit.py
tomo082 Apr 15, 2026
a06c3cf
Update main_vit.py
tomo082 Apr 15, 2026
88e7ca8
Create extract_ref_features_vit.py
tomo082 Apr 15, 2026
1d7c883
Update extract_ref_features_vit.py
tomo082 Apr 15, 2026
605ace3
Update main_vit.py
tomo082 Apr 15, 2026
3b77ec6
Update main_vit.py
tomo082 Apr 15, 2026
dd9f411
Update extract_ref_features_vit.py
tomo082 Apr 15, 2026
1ac2bb0
Update extract_ref_features_vit.py
tomo082 Apr 15, 2026
775a5d1
Update extract_ref_features_vit.py
tomo082 Apr 15, 2026
0a21193
Update extract_ref_features_vit.py
tomo082 Apr 15, 2026
a562903
Update extract_ref_features_vit.py
tomo082 Apr 15, 2026
93bc864
Update extract_ref_features_vit.py
tomo082 Apr 15, 2026
d1e89cc
Update main_vit.py
tomo082 Apr 15, 2026
2b4f74a
Update modules.py
tomo082 Apr 15, 2026
635ee1c
Update modules.py
tomo082 Apr 15, 2026
f53be47
Update main_vit.py
tomo082 Apr 15, 2026
f9657f9
Update main_vit.py
tomo082 Apr 15, 2026
82ad675
Update validate.py
tomo082 Apr 15, 2026
ab1f772
Update main_vit.py
tomo082 Apr 15, 2026
c2f8fcd
Update validate.py
tomo082 Apr 15, 2026
d0ecaf4
Update main_vit.py
tomo082 Apr 15, 2026
fc8bcf6
Update modules.py
tomo082 Apr 19, 2026
32007cd
Create extract_ref_features1.py
tomo082 Apr 21, 2026
1eb62a9
Create main_1.py
tomo082 Apr 21, 2026
44df136
Update extract_ref_features1.py
tomo082 Apr 21, 2026
eedcaf8
Update main_1.py
tomo082 Apr 21, 2026
7d0f85d
Update extract_ref_features1.py
tomo082 Apr 21, 2026
32c8de0
Update main_1.py
tomo082 Apr 21, 2026
bb45c3a
Update main_1.py
tomo082 Apr 21, 2026
4581a0e
Update main_1.py
tomo082 Apr 21, 2026
55b8349
Create main_Fourier
tomo082 Apr 27, 2026
ac73eef
Rename main_Fourier to main_Fourier.py
tomo082 Apr 27, 2026
a3c6964
Implement Fourier residual feature calculation
tomo082 Apr 27, 2026
435b8b2
Update main_Fourier.py
tomo082 Apr 27, 2026
1f2edea
Update main_Fourier.py
tomo082 Apr 27, 2026
e035f35
Update validate.py
tomo082 Apr 27, 2026
d8a781b
Update validate.py
tomo082 Apr 27, 2026
2b6bb48
Create validate_Fourier.py
tomo082 Apr 27, 2026
ecda477
Update validate.py
tomo082 Apr 27, 2026
65e6d3e
Update main_Fourier.py
tomo082 Apr 27, 2026
dc628ac
Update utils.py
tomo082 Apr 27, 2026
9ebecc2
Update utils.py
tomo082 Apr 27, 2026
f52f30a
Update utils.py
tomo082 Apr 28, 2026
594a39f
Add functions for image-level feature matching
tomo082 Apr 28, 2026
d1f82ba
Update main_Fourier.py
tomo082 Apr 28, 2026
7d88715
Update validate_Fourier.py
tomo082 Apr 28, 2026
3c003da
Update main_Fourier.py
tomo082 Apr 28, 2026
c8af31e
Update validate_Fourier.py
tomo082 Apr 28, 2026
a0a511c
Update utils.py
tomo082 Apr 28, 2026
4861524
Create main_attention.py
tomo082 Apr 28, 2026
1aaf485
Update utils.py
tomo082 Apr 28, 2026
dfdafb4
Update main_attention.py
tomo082 Apr 28, 2026
b59eb8d
Update main_attention.py
tomo082 Apr 28, 2026
67a89ad
Add validation script for model evaluation
tomo082 Apr 28, 2026
5d8eff5
Update validate_attention.py
tomo082 Apr 28, 2026
e65b504
Update main_attention.py
tomo082 Apr 28, 2026
aef1fba
Update main_vit.py
tomo082 Apr 29, 2026
01bc34f
Update extract_ref_features_vit.py
tomo082 Apr 29, 2026
6450b24
Update extract_ref_features_vit.py
tomo082 Apr 29, 2026
f0d578f
Update main_vit.py
tomo082 Apr 29, 2026
6ed7913
Create extract_ref_features_filter.py
tomo082 Apr 30, 2026
777ed2c
Update extract_ref_features_filter.py
tomo082 Apr 30, 2026
07433f0
Update extract_ref_features_filter.py
tomo082 Apr 30, 2026
18cd7e1
Update utils.py
KudanLabo May 1, 2026
fcd6ca6
Update utils.py
KudanLabo May 1, 2026
82ac83d
Update utils.py
KudanLabo May 8, 2026
1a5ae16
Update utils.py
KudanLabo May 8, 2026
884c29d
Update utils.py
KudanLabo May 8, 2026
a05da01
Update utils.py
KudanLabo May 8, 2026
2e853d8
Update utils.py
KudanLabo May 8, 2026
c9419dc
Update utils.py
KudanLabo May 8, 2026
27377e8
Update utils.py
KudanLabo May 8, 2026
fb5b921
Update utils.py
tomo082 May 11, 2026
cfa552e
Create main_osp.py
tomo082 May 11, 2026
c0bb28c
Create validate_osp.py
tomo082 May 11, 2026
0fb2370
Update validate_osp.py
tomo082 May 11, 2026
3360a44
Update main_osp.py
tomo082 May 11, 2026
d73363f
Integrate OSP application in validate function
tomo082 May 11, 2026
47e0173
Update main_osp.py
tomo082 May 11, 2026
a55e057
Update main_osp.py
tomo082 May 12, 2026
85675b4
Create main_wav.py
tomo082 May 13, 2026
1079e24
Create validate_wav.py
tomo082 May 13, 2026
1df6887
Update validate_wav.py
tomo082 May 13, 2026
8d24a41
Update utils.py
tomo082 May 13, 2026
10d5e9b
Update main_wav.py
tomo082 May 13, 2026
1b5f1fd
Update validate_wav.py
tomo082 May 13, 2026
713d3ae
Create extract_ref_feature_wav.py
tomo082 May 13, 2026
fd48a8c
Rename extract_ref_feature_wav.py to extract_ref_features_wav.py
tomo082 May 13, 2026
b5da7e7
Update extract_ref_features_wav.py
tomo082 May 13, 2026
9e1242b
Update extract_ref_features_wav.py
tomo082 May 13, 2026
77fa634
Update main_wav.py
tomo082 May 13, 2026
b9cab4a
Update validate_wav.py
tomo082 May 13, 2026
5cdffe0
Update main_wav.py
tomo082 May 13, 2026
d20bd9d
Update main_wav.py
tomo082 May 13, 2026
d0bf8c9
Update main_wav.py
tomo082 May 13, 2026
3684d37
Update utils.py
tomo082 May 13, 2026
4f06354
Update utils.py
tomo082 May 13, 2026
588c2d8
Update main_wav.py
tomo082 May 13, 2026
563260e
Update main_wav.py
tomo082 May 13, 2026
9a2635d
Update main_wav.py
tomo082 May 13, 2026
bbc546e
Create main_wav1.py
tomo082 May 13, 2026
26ceb31
Create validate_wav1.py
tomo082 May 13, 2026
0e291d8
Implement model validation and anomaly scoring
tomo082 May 13, 2026
0d7e0a9
Update main_wav1.py
tomo082 May 13, 2026
6df71ba
Update validate_wav1.py
tomo082 May 13, 2026
56282e6
Update main_wav1.py
tomo082 May 13, 2026
3dae178
Update validate_wav1.py
tomo082 May 13, 2026
7e3e404
Update validate_wav1.py
tomo082 May 13, 2026
ad10a51
Update main_wav1.py
tomo082 May 13, 2026
dba422b
Update validate_wav1.py
tomo082 May 13, 2026
5f2ef35
Update validate_wav1.py
tomo082 May 13, 2026
178eb1b
Update main_wav1.py
tomo082 May 13, 2026
3dac92b
Update validate_wav1.py
tomo082 May 13, 2026
f4571e7
Update main_wav1.py
tomo082 May 13, 2026
4e6c305
Update validate_wav1.py
tomo082 May 13, 2026
9415e72
Update validate_wav.py
tomo082 May 13, 2026
5426a3d
Update main_wav.py
tomo082 May 13, 2026
0d4f893
Update main_wav1.py
tomo082 May 14, 2026
0dd3d54
Update validate_wav1.py
tomo082 May 15, 2026
0cf405e
Update main_wav1.py
tomo082 May 15, 2026
e734251
Update validate_wav1.py
tomo082 May 15, 2026
6bd8d22
Create main_wav_cf.py
tomo082 May 18, 2026
e1a28dd
Create validate_wav_cf.py
tomo082 May 18, 2026
fdc8113
Update main_wav_cf.py
tomo082 May 18, 2026
0840f40
Update main_wav_cf.py
tomo082 May 18, 2026
603b0c3
Update validate_wav_cf.py
tomo082 May 18, 2026
138603b
Update main_wav_cf.py
tomo082 May 18, 2026
53860f3
Update main_wav_cf.py
tomo082 May 18, 2026
60f54e3
Update validate_wav_cf.py
tomo082 May 18, 2026
f054ecc
Create main_global.py
tomo082 May 19, 2026
cf24930
Create validate_global.py
tomo082 May 19, 2026
343e323
Update validate_global.py
tomo082 May 19, 2026
db3b102
Update main_global.py
tomo082 May 19, 2026
8681c51
Update validate_global.py
tomo082 May 19, 2026
99f8317
Update main_global.py
tomo082 May 19, 2026
6c4de6e
Update main_global.py
tomo082 May 19, 2026
82a8810
Refactor validate function to remove vq_ops
tomo082 May 19, 2026
15ace02
Update main_global.py
tomo082 May 19, 2026
7fe7f21
Create main_freq_blend.py
tomo082 May 20, 2026
261dbcd
Create validate_freq_blend.py
tomo082 May 20, 2026
eb29e35
Update validate_freq_blend.py
tomo082 May 20, 2026
49eb0f7
Add AdaCLIP text-reference map evaluation
tomo082 Jun 28, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 5 additions & 0 deletions change_plan
Original file line number Diff line number Diff line change
@@ -0,0 +1,5 @@
ResADの場合、例えばmvtecでtestする場合、visaの全データを用いてtrainしている。このときtrainをnormal+dataaug anomalt or normal + anomaly + dataaug anomalyにする?
fewshotなので他のデータセットのクラスの異常にも対応できるようにしたいcutpasteよりはdream?dreamを物体があるに場所しか適応できないように変える。maskのデータはあるからperlinnoiseを生成してそこから絞り込み。
NSA: cutpasteの応用、コピー元が正常画像、切り取った画像をノイズやブレンドなどを行い異常風に生成、それをほかの正常画像にはりつける。
SPADEの設定を適応するにはどうするか?
新しいテスト用のデータセットを作る
17 changes: 16 additions & 1 deletion classes.py
Original file line number Diff line number Diff line change
Expand Up @@ -36,4 +36,19 @@
MVTEC_TO_BRATS = {'seen': ['bottle', 'cable', 'capsule', 'carpet', 'grid',
'hazelnut', 'leather', 'metal_nut', 'pill', 'screw',
'tile', 'toothbrush', 'transistor', 'wood', 'zipper'],
'unseen': ['brain']}
'unseen': ['brain']}

MVTEC_TO_MVTEC = {'seen': ['bottle', 'cable', 'capsule', 'carpet', 'grid',
'hazelnut', 'leather', 'metal_nut', 'pill', 'screw',
'tile', 'toothbrush', 'transistor', 'wood', 'zipper'],
'unseen': ['bottle', 'cable', 'capsule', 'carpet', 'grid',
'hazelnut', 'leather', 'metal_nut', 'pill', 'screw',
'tile', 'toothbrush', 'transistor', 'wood', 'zipper']}

VISA_TO_VISA = {'seen': ['candle', 'capsules', 'cashew', 'chewinggum', 'fryum',
'macaroni1', 'macaroni2', 'pcb1', 'pcb2', 'pcb3', 'pcb4', 'pipe_fryum'],
'unseen': ['candle', 'capsules', 'cashew', 'chewinggum', 'fryum',
'macaroni1', 'macaroni2', 'pcb1', 'pcb2', 'pcb3', 'pcb4', 'pipe_fryum']}

CAPSULES_TO_CAPSULES = {'seen': ['capsules'],
'unseen': ['capsules']}
309 changes: 309 additions & 0 deletions datasets/capsules.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,309 @@
import os
import pandas
import torch
import numpy as np
from PIL import Image
from typing import Callable, Optional
from torch.utils.data import Dataset
from torchvision import transforms as T
from torchvision.transforms.transforms import RandomHorizontalFlip


IMAGENET_MEAN = [0.485, 0.456, 0.406]
IMAGENET_STD = [0.229, 0.224, 0.225]


class CAPSULES(Dataset):

CLASS_NAMES = ['capsules']

def __init__(self,
root: str,
class_name: str,
train: bool = True,
normalize: str = 'imagebind',
transform: Optional[Callable] = None,
target_transform: Optional[Callable] = None,
**kwargs):

self.root = root
self.class_name = class_name
self.train = train
self.cropsize = [kwargs.get('crp_size'), kwargs.get('crp_size')]

# load dataset
if isinstance(self.class_name, str):
self.image_paths, self.labels, self.mask_paths, self.class_names = self._load_data(self.class_name)
elif self.class_name is None: # load all classes
self.image_paths, self.labels, self.mask_paths, self.class_names = self._load_all_data()
else:
self.image_paths, self.labels, self.mask_paths, self.class_names = self._load_all_data(self.class_name)

# set transforms
if normalize == "imagebind":
self.transform = T.Compose( # for imagebind
[
T.Resize(
224, interpolation=T.InterpolationMode.BICUBIC
),
T.CenterCrop(224),
T.ToTensor(),
T.Normalize(
mean=(0.48145466, 0.4578275, 0.40821073),
std=(0.26862954, 0.26130258, 0.27577711),
),
]
)
else:
self.transform = T.Compose([
T.Resize(kwargs.get('img_size', 224), T.InterpolationMode.BICUBIC),
T.CenterCrop(kwargs.get('crp_size', 224)),
T.ToTensor(),
T.Compose([T.Normalize(IMAGENET_MEAN, IMAGENET_STD)])])

# mask
self.target_transform = T.Compose([
T.Resize(kwargs.get('img_size'), Image.NEAREST),
T.CenterCrop(kwargs.get('crp_size')),
T.ToTensor()])

self.class_to_idx = {'capsules': 0 }
self.idx_to_class = {0:'capsules' }

def __getitem__(self, idx):
image_path, label, mask, class_name = self.image_paths[idx], self.labels[idx], self.mask_paths[idx], self.class_names[idx]

image = Image.open(image_path).convert('RGB')
image = self.transform(image)

if label == 0:
mask = torch.zeros([1, self.cropsize[0], self.cropsize[1]])
else:
mask = Image.open(mask)
mask = np.array(mask)
mask[mask != 0] = 255
mask = Image.fromarray(mask)
mask = self.target_transform(mask)

if self.train:
label = self.class_to_idx[class_name]

return image, label, mask, class_name

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

def _load_data(self, class_name):
split_csv_file = os.path.join(self.root, 'split_csv', '1cls.csv')
csv_data = pandas.read_csv(split_csv_file)

class_data = csv_data.loc[csv_data['object'] == class_name]

if self.train:
train_data = class_data.loc[class_data['split'] == 'train']
image_paths = train_data['image'].to_list()
image_paths = [os.path.join(self.root, file_name) for file_name in image_paths]
labels = [0] * len(image_paths)
mask_paths = [None] * len(image_paths)
else:
image_paths, labels, mask_paths = [], [], []

test_data = class_data.loc[class_data['split'] == 'test']
test_normal_data = test_data.loc[test_data['label'] == 'normal']
test_anomaly_data = test_data.loc[test_data['label'] == 'anomaly']

normal_image_paths = test_normal_data['image'].to_list()
normal_image_paths = [os.path.join(self.root, file_name) for file_name in normal_image_paths]
image_paths.extend(normal_image_paths)
labels.extend([0] * len(normal_image_paths))
mask_paths.extend([None] * len(normal_image_paths))

anomaly_image_paths = test_anomaly_data['image'].to_list()
anomaly_mask_paths = test_anomaly_data['mask'].to_list()
anomaly_image_paths = [os.path.join(self.root, file_name) for file_name in anomaly_image_paths]
anomaly_mask_paths = [os.path.join(self.root, file_name) for file_name in anomaly_mask_paths]
image_paths.extend(anomaly_image_paths)
labels.extend([1] * len(anomaly_image_paths))
mask_paths.extend(anomaly_mask_paths)

class_names = [class_name] * len(image_paths)
return image_paths, labels, mask_paths, class_names

def _load_all_data(self, class_names=None):
all_image_paths = []
all_labels = []
all_mask_paths = []
all_class_names = []
CLASS_NAMES = class_names if class_names is not None else self.CLASS_NAMES
for class_name in CLASS_NAMES:
image_paths, labels, mask_paths, class_names = self._load_data(class_name)
all_image_paths.extend(image_paths)
all_labels.extend(labels)
all_mask_paths.extend(mask_paths)
all_class_names.extend(class_names)
return all_image_paths, all_labels, all_mask_paths, all_class_names

def update_class_to_idx(self, class_to_idx):
for class_name in self.class_to_idx.keys():
self.class_to_idx[class_name] = class_to_idx[class_name]
class_names = self.class_to_idx.keys()
idxs = self.class_to_idx.values()
self.idx_to_class = dict(zip(idxs, class_names))


class CAPSULESANO(Dataset):

CLASS_NAMES = ['capsules']

def __init__(self,
root: str,
class_name: str,
train: bool = True,
normalize: str = 'imagebind',
transform: Optional[Callable] = None,
target_transform: Optional[Callable] = None,
**kwargs):

self.root = root
self.class_name = class_name
self.train = train
self.cropsize = [kwargs.get('crp_size'), kwargs.get('crp_size')]

# load dataset
if isinstance(self.class_name, str):
self.image_paths, self.labels, self.mask_paths, self.class_names = self._load_data(self.class_name)
elif self.class_name is None:
self.image_paths, self.labels, self.mask_paths, self.class_names = self._load_all_data()
else:
self.image_paths, self.labels, self.mask_paths, self.class_names = self._load_all_data(self.class_name)

if normalize == "imagebind":
self.transform = T.Compose( # for imagebind
[
T.Resize(
224, interpolation=T.InterpolationMode.BICUBIC
),
T.CenterCrop(224),
T.ToTensor(),
T.Normalize(
mean=(0.48145466, 0.4578275, 0.40821073),
std=(0.26862954, 0.26130258, 0.27577711),
),
]
)
else:
self.transform = T.Compose([
T.Resize(kwargs.get('img_size', 224), T.InterpolationMode.BICUBIC),
T.CenterCrop(kwargs.get('crp_size', 224)),
T.ToTensor(),
T.Compose([T.Normalize(IMAGENET_MEAN, IMAGENET_STD)])])

# mask
self.target_transform = T.Compose([
T.Resize(kwargs.get('img_size'), Image.NEAREST),
T.CenterCrop(kwargs.get('crp_size')),
T.ToTensor()])

self.class_to_idx = {'capsules':0}
self.idx_to_class = {0:'capsules'}

def __getitem__(self, idx):
image_path, label, mask, class_name = self.image_paths[idx], self.labels[idx], self.mask_paths[idx], self.class_names[idx]

image = Image.open(image_path).convert('RGB')
image = self.transform(image)

if label == 0:
mask = torch.zeros([1, self.cropsize[0], self.cropsize[1]])
else:
mask = Image.open(mask)
mask = np.array(mask)
mask[mask != 0] = 255
mask = Image.fromarray(mask)
mask = self.target_transform(mask)

if self.train:
label = self.class_to_idx[class_name]

return image, label, mask, class_name

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

def _load_data(self, class_name):
split_csv_file = os.path.join(self.root, 'split_csv', '1cls.csv')
csv_data = pandas.read_csv(split_csv_file)

class_data = csv_data.loc[csv_data['object'] == class_name]
all_image_paths, all_labels, all_mask_paths = [], [], []

# train
train_data = class_data.loc[class_data['split'] == 'train']
image_paths = train_data['image'].to_list()
image_paths = [os.path.join(self.root, file_name) for file_name in image_paths]
labels = [0] * len(image_paths)
mask_paths = [None] * len(image_paths)
all_image_paths.extend(image_paths)
all_labels.extend(labels)
all_mask_paths.extend(mask_paths)

# test
image_paths, labels, mask_paths = [], [], []
test_data = class_data.loc[class_data['split'] == 'test']
test_normal_data = test_data.loc[test_data['label'] == 'normal']
test_anomaly_data = test_data.loc[test_data['label'] == 'anomaly']

normal_image_paths = test_normal_data['image'].to_list()
normal_image_paths = [os.path.join(self.root, file_name) for file_name in normal_image_paths]
image_paths.extend(normal_image_paths)
labels.extend([0] * len(normal_image_paths))
mask_paths.extend([None] * len(normal_image_paths))

anomaly_image_paths = test_anomaly_data['image'].to_list()
anomaly_mask_paths = test_anomaly_data['mask'].to_list()
anomaly_image_paths = [os.path.join(self.root, file_name) for file_name in anomaly_image_paths]
anomaly_mask_paths = [os.path.join(self.root, file_name) for file_name in anomaly_mask_paths]
image_paths.extend(anomaly_image_paths)
labels.extend([1] * len(anomaly_image_paths))
mask_paths.extend(anomaly_mask_paths)

all_image_paths.extend(image_paths)
all_labels.extend(labels)
all_mask_paths.extend(mask_paths)

class_names = [class_name] * len(all_image_paths)
return all_image_paths, all_labels, all_mask_paths, class_names

def _load_all_data(self, class_names=None):
all_image_paths = []
all_labels = []
all_mask_paths = []
all_class_names = []
CLASS_NAMES = class_names if class_names is not None else self.CLASS_NAMES
for class_name in CLASS_NAMES:
image_paths, labels, mask_paths, class_names = self._load_data(class_name)
all_image_paths.extend(image_paths)
all_labels.extend(labels)
all_mask_paths.extend(mask_paths)
all_class_names.extend(class_names)
return all_image_paths, all_labels, all_mask_paths, all_class_names

def update_class_to_idx(self, class_to_idx):
for class_name in self.class_to_idx.keys():
self.class_to_idx[class_name] = class_to_idx[class_name]
class_names = self.class_to_idx.keys()
idxs = self.class_to_idx.values()
self.idx_to_class = dict(zip(idxs, class_names))


def get_normal_image_paths_visa(root, class_name):
split_csv_file = os.path.join(root, 'split_csv', '1cls.csv')
csv_data = pandas.read_csv(split_csv_file)

class_data = csv_data.loc[csv_data['object'] == class_name]

train_data = class_data.loc[class_data['split'] == 'train']
image_paths = train_data['image'].to_list()
image_paths = [os.path.join(root, file_name) for file_name in image_paths]

return image_paths
Loading