AnonymousUser20's picture
Upload 1314 files
178d33b verified
Raw
History Blame Contribute Delete
23.9 kB
import os
import gdown
import zipfile
from torch.utils.data import DataLoader
import torchvision as tvs
if tvs.__version__ >= '0.13':
tvs_new = True
else:
tvs_new = False
from openood.datasets.imglist_dataset import ImglistDataset
from openood.datasets.bronze_dataset import BronzeWare_ID_Dataset
from openood.preprocessors import BasePreprocessor
from .preprocessor import get_default_preprocessor, ImageNetCPreProcessor
DATA_INFO = {
'cifar10': {
'num_classes': 10,
'id': {
'train': {
'data_dir': 'images_classic/',
'imglist_path': 'benchmark_imglist/cifar10/train_cifar10.txt'
},
'val': {
'data_dir': 'images_classic/',
'imglist_path': 'benchmark_imglist/cifar10/val_cifar10.txt'
},
'test': {
'data_dir': 'images_classic/',
'imglist_path': 'benchmark_imglist/cifar10/test_cifar10.txt'
}
},
'csid': {
'datasets': ['cifar10c'],
'cinic10': {
'data_dir': 'images_classic/',
'imglist_path': 'benchmark_imglist/cifar10/val_cinic10.txt'
},
'cifar10c': {
'data_dir': 'images_classic/',
'imglist_path': 'benchmark_imglist/cifar10/test_cifar10c.txt'
}
},
'ood': {
'val': {
'data_dir': 'images_classic/',
'imglist_path': 'benchmark_imglist/cifar10/val_tin.txt'
},
'near': {
'datasets': ['cifar100', 'tin'],
'cifar100': {
'data_dir': 'images_classic/',
'imglist_path':
'benchmark_imglist/cifar10/test_cifar100.txt'
},
'tin': {
'data_dir': 'images_classic/',
'imglist_path': 'benchmark_imglist/cifar10/test_tin.txt'
}
},
'far': {
'datasets': ['mnist', 'svhn', 'texture', 'places365'],
'mnist': {
'data_dir': 'images_classic/',
'imglist_path': 'benchmark_imglist/cifar10/test_mnist.txt'
},
'svhn': {
'data_dir': 'images_classic/',
'imglist_path': 'benchmark_imglist/cifar10/test_svhn.txt'
},
'texture': {
'data_dir': 'images_classic/',
'imglist_path':
'benchmark_imglist/cifar10/test_texture.txt'
},
'places365': {
'data_dir': 'images_classic/',
'imglist_path':
'benchmark_imglist/cifar10/test_places365.txt'
},
}
}
},
'cifar100': {
'num_classes': 100,
'id': {
'train': {
'data_dir': 'images_classic/',
'imglist_path': 'benchmark_imglist/cifar100/train_cifar100.txt'
},
'val': {
'data_dir': 'images_classic/',
'imglist_path': 'benchmark_imglist/cifar100/val_cifar100.txt'
},
'test': {
'data_dir': 'images_classic/',
'imglist_path': 'benchmark_imglist/cifar100/test_cifar100.txt'
}
},
'csid': {
'datasets': [],
},
'ood': {
'val': {
'data_dir': 'images_classic/',
'imglist_path': 'benchmark_imglist/cifar100/val_tin.txt'
},
'near': {
'datasets': ['cifar10', 'tin'],
'cifar10': {
'data_dir': 'images_classic/',
'imglist_path':
'benchmark_imglist/cifar100/test_cifar10.txt'
},
'tin': {
'data_dir': 'images_classic/',
'imglist_path': 'benchmark_imglist/cifar100/test_tin.txt'
}
},
'far': {
'datasets': ['mnist', 'svhn', 'texture', 'places365'],
'mnist': {
'data_dir': 'images_classic/',
'imglist_path': 'benchmark_imglist/cifar100/test_mnist.txt'
},
'svhn': {
'data_dir': 'images_classic/',
'imglist_path': 'benchmark_imglist/cifar100/test_svhn.txt'
},
'texture': {
'data_dir': 'images_classic/',
'imglist_path':
'benchmark_imglist/cifar100/test_texture.txt'
},
'places365': {
'data_dir': 'images_classic/',
'imglist_path':
'benchmark_imglist/cifar100/test_places365.txt'
}
},
}
},
'imagenet200': {
'num_classes': 200,
'id': {
'train': {
'data_dir':
'images_largescale/',
'imglist_path':
'benchmark_imglist/imagenet200/train_imagenet200.txt'
},
'val': {
'data_dir': 'images_largescale/',
'imglist_path':
'benchmark_imglist/imagenet200/val_imagenet200.txt'
},
'test': {
'data_dir':
'images_largescale/',
'imglist_path':
'benchmark_imglist/imagenet200/test_imagenet200.txt'
}
},
'csid': {
'datasets': ['imagenet_v2', 'imagenet_c', 'imagenet_r'],
'imagenet_v2': {
'data_dir':
'images_largescale/',
'imglist_path':
'benchmark_imglist/imagenet200/test_imagenet200_v2.txt'
},
'imagenet_c': {
'data_dir':
'images_largescale/',
'imglist_path':
'benchmark_imglist/imagenet200/test_imagenet200_c.txt'
},
'imagenet_r': {
'data_dir':
'images_largescale/',
'imglist_path':
'benchmark_imglist/imagenet200/test_imagenet200_r.txt'
},
},
'ood': {
'val': {
'data_dir': 'images_largescale/',
'imglist_path':
'benchmark_imglist/imagenet200/val_openimage_o.txt'
},
'near': {
'datasets': ['ssb_hard', 'ninco'],
'ssb_hard': {
'data_dir':
'images_largescale/',
'imglist_path':
'benchmark_imglist/imagenet200/test_ssb_hard.txt'
},
'ninco': {
'data_dir': 'images_largescale/',
'imglist_path':
'benchmark_imglist/imagenet200/test_ninco.txt'
}
},
'far': {
'datasets': ['inaturalist', 'textures', 'openimage_o'],
'inaturalist': {
'data_dir':
'images_largescale/',
'imglist_path':
'benchmark_imglist/imagenet200/test_inaturalist.txt'
},
'textures': {
'data_dir':
'images_classic/',
'imglist_path':
'benchmark_imglist/imagenet200/test_textures.txt'
},
'openimage_o': {
'data_dir':
'images_largescale/',
'imglist_path':
'benchmark_imglist/imagenet200/test_openimage_o.txt'
},
},
}
},
'imagenet': {
'num_classes': 1000,
'id': {
'train': {
'data_dir': 'images_largescale/',
'imglist_path': 'benchmark_imglist/imagenet/train_imagenet.txt'
},
'val': {
'data_dir': 'images_largescale/',
'imglist_path': 'benchmark_imglist/imagenet/val_imagenet.txt'
},
'test': {
'data_dir': 'images_largescale/',
'imglist_path': 'benchmark_imglist/imagenet/test_imagenet.txt'
}
},
'csid': {
'datasets': [],
},
# 'csid': {
# 'datasets': ['imagenet_v2', 'imagenet_c', 'imagenet_r'],
# 'imagenet_v2': {
# 'data_dir': 'images_largescale/',
# 'imglist_path':
# 'benchmark_imglist/imagenet/test_imagenet_v2.txt'
# },
# 'imagenet_c': {
# 'data_dir': 'images_largescale/',
# 'imglist_path':
# 'benchmark_imglist/imagenet/test_imagenet_c.txt'
# },
# 'imagenet_r': {
# 'data_dir': 'images_largescale/',
# 'imglist_path':
# 'benchmark_imglist/imagenet/test_imagenet_r.txt'
# },
# },
'ood': {
'val': {
'data_dir': 'images_largescale/',
'imglist_path':
'benchmark_imglist/imagenet/val_openimage_o.txt'
},
'near': {
'datasets': ['imagenet22k_container', 'ssb_hard', 'ninco'],
'imagenet22k_container': {
'data_dir': 'images_largescale/',
'imglist_path':
'benchmark_imglist/imagenet21k_container/imagenet21k_container_file-list.txt'
},
'ssb_hard': {
'data_dir': 'images_largescale/',
'imglist_path':
'benchmark_imglist/imagenet/test_ssb_hard.txt'
},
'ninco': {
'data_dir': 'images_largescale/',
'imglist_path': 'benchmark_imglist/imagenet/test_ninco.txt'
}
},
'far': {
'datasets': ['inaturalist', 'textures', 'openimage_o'],
'inaturalist': {
'data_dir':
'images_largescale/',
'imglist_path':
'benchmark_imglist/imagenet/test_inaturalist.txt'
},
'textures': {
'data_dir': 'images_classic/',
'imglist_path':
'benchmark_imglist/imagenet/test_textures.txt'
},
'openimage_o': {
'data_dir':
'images_largescale/',
'imglist_path':
'benchmark_imglist/imagenet/test_openimage_o.txt'
},
},
}
},
'bronze2': {
'num_classes': 11,
'id': {
'train': {
'data_dir': 'bronze_ID_and_OOD/bronze2/images',
'imglist_path': 'benchmark_imglist/bronzeData/ding_and_gui_train.xlsx',
'xml_path': 'bronze_ID_and_OOD/bronze2/xmls'
},
'val': {
'data_dir': 'bronze_ID_and_OOD/bronze2/images',
'imglist_path': 'benchmark_imglist/bronzeData/ding_and_gui_val.xlsx',
'xml_path': 'bronze_ID_and_OOD/bronze2/xmls'
},
'test': {
'data_dir': 'bronze_ID_and_OOD/bronze2/images',
'imglist_path': 'benchmark_imglist/bronzeData/ding_and_gui_test.xlsx',
'xml_path': 'bronze_ID_and_OOD/bronze2/xmls'
}
},
'csid': {
'datasets': [],
},
'ood': {
'val': {
'data_dir': 'images_largescale/',
'imglist_path':
'benchmark_imglist/imagenet/val_openimage_o.txt'
},
'near': {
'datasets': ['imagenet22k_container', 'imagenet22k_container_refine', 'bronzeS_containerM', 'bronzeM_containerS', 'bronze_Line'],
'imagenet22k_container': {
'data_dir': 'images_largescale/',
'imglist_path':
'benchmark_imglist/imagenet21k_container/imagenet21k_container_file-list.txt'
},
'imagenet22k_container_refine': {
'data_dir': 'images_largescale/',
'imglist_path':
'benchmark_imglist/imagenet21k_container_refine/imagenet21k_container_file-list-refine.txt'
},
'bronzeS_containerM': {
'data_dir': 'images_largescale/',
'imglist_path': 'images_largescale/transfer_dataset/bronze_structure_container_material/bronze_structure_container_material_test.txt'
},
'bronzeM_containerS': {
'data_dir': 'images_largescale/',
'imglist_path': 'images_largescale/transfer_dataset/container_structure_bronze_material/container_structure_bronze_material_test.txt'
},
'bronze_Line': {
'data_dir': 'images_largescale/',
'imglist_path': 'images_largescale/bronze_line/bronze2_Line_OOD_list.txt'
}
},
'mid': {
'datasets': ['ssb_hard', 'ninco'],
'ssb_hard': {
'data_dir': 'images_largescale/',
'imglist_path':
'benchmark_imglist/imagenet/test_ssb_hard.txt'
},
'ninco': {
'data_dir': 'images_largescale/',
'imglist_path': 'benchmark_imglist/imagenet/test_ninco.txt'
}
},
'far': {
'datasets': ['inaturalist', 'textures', 'openimage_o'],
'inaturalist': {
'data_dir':
'images_largescale/',
'imglist_path':
'benchmark_imglist/imagenet/test_inaturalist.txt'
},
'textures': {
'data_dir': 'images_classic/',
'imglist_path':
'benchmark_imglist/imagenet/test_textures.txt'
},
'openimage_o': {
'data_dir':
'images_largescale/',
'imglist_path':
'benchmark_imglist/imagenet/test_openimage_o.txt'
},
},
}
},
}
download_id_dict = {
'cifar10': '1Co32RiiWe16lTaiOU6JMMnyUYS41IlO1',
'cifar100': '1PGKheHUsf29leJPPGuXqzLBMwl8qMF8_',
'tin': '1PZ-ixyx52U989IKsMA2OT-24fToTrelC',
'mnist': '1CCHAGWqA1KJTFFswuF9cbhmB-j98Y1Sb',
'svhn': '1DQfc11HOtB1nEwqS4pWUFp8vtQ3DczvI',
'texture': '1OSz1m3hHfVWbRdmMwKbUzoU8Hg9UKcam',
'places365': '1Ec-LRSTf6u5vEctKX9vRp9OA6tqnJ0Ay',
'imagenet_1k': '1i1ipLDFARR-JZ9argXd2-0a6DXwVhXEj',
'species_sub': '1-JCxDx__iFMExkYRMylnGJYTPvyuX6aq',
'ssb_hard': '1PzkA-WGG8Z18h0ooL_pDdz9cO-DCIouE',
'ninco': '1Z82cmvIB0eghTehxOGP5VTdLt7OD3nk6',
'inaturalist': '1zfLfMvoUD0CUlKNnkk7LgxZZBnTBipdj',
'places': '1fZ8TbPC4JGqUCm-VtvrmkYxqRNp2PoB3',
'sun': '1ISK0STxWzWmg-_uUr4RQ8GSLFW7TZiKp',
'openimage_o': '1VUFXnB_z70uHfdgJG2E_pjYOcEgqM7tE',
'imagenet_v2': '1akg2IiE22HcbvTBpwXQoD7tgfPCdkoho',
'imagenet_r': '1EzjMN2gq-bVV7lg-MEAdeuBuz-7jbGYU',
'imagenet_c': '1JeXL9YH4BO8gCJ631c5BHbaSsl-lekHt',
'benchmark_imglist': '1XKzBdWCqg3vPoj-D32YixJyJJ0hL63gP'
}
dir_dict = {
'images_classic/': [
'cifar100', 'tin', 'tin597', 'svhn', 'cinic10', 'imagenet10', 'mnist',
'fashionmnist', 'cifar10', 'cifar100c', 'places365', 'cifar10c',
'fractals_and_fvis', 'usps', 'texture', 'notmnist'
],
'images_largescale/': [
'imagenet_1k',
'ssb_hard',
'ninco',
'inaturalist',
'places',
'sun',
'openimage_o',
'imagenet_v2',
'imagenet_c',
'imagenet_r',
'imagenet-21k-container',
],
'images_medical/': ['actmed', 'bimcv', 'ct', 'hannover', 'xraybone'],
'bronze_ID_and_OOD/': ['bronze2'],
}
benchmarks_dict = {
'cifar10':
['cifar10', 'cifar100', 'tin', 'mnist', 'svhn', 'texture', 'places365'],
'cifar100':
['cifar100', 'cifar10', 'tin', 'mnist', 'svhn', 'texture', 'places365'],
'imagenet200': [
'imagenet_1k', 'ssb_hard', 'ninco', 'inaturalist', 'texture',
'openimage_o', 'imagenet_v2', 'imagenet_c', 'imagenet_r'
],
'imagenet': [
'imagenet_1k', 'ssb_hard', 'ninco', 'inaturalist', 'texture',
'openimage_o', 'imagenet_v2', 'imagenet_c', 'imagenet_r'
],
'bronze2':
['bronze2', 'cifar10'],
}
def require_download(filename, path):
for item in os.listdir(path):
if item.startswith(filename) or filename.startswith(
item) or path.endswith(filename):
return False
else:
print(filename + ' needs download:')
return True
def download_dataset(dataset, data_root):
for key in dir_dict.keys():
if dataset in dir_dict[key]:
store_path = os.path.join(data_root, key, dataset)
if not os.path.exists(store_path):
os.makedirs(store_path)
break
else:
print('Invalid dataset detected {}'.format(dataset))
return
if require_download(dataset, store_path):
print(store_path)
if not store_path.endswith('/'):
store_path = store_path + '/'
gdown.download(id=download_id_dict[dataset], output=store_path)
file_path = os.path.join(store_path, dataset + '.zip')
with zipfile.ZipFile(file_path, 'r') as zip_file:
zip_file.extractall(store_path)
os.remove(file_path)
def data_setup(data_root, id_data_name):
if not data_root.endswith('/'):
data_root = data_root + '/'
if not os.path.exists(os.path.join(data_root, 'benchmark_imglist')):
gdown.download(id=download_id_dict['benchmark_imglist'],
output=data_root)
file_path = os.path.join(data_root, 'benchmark_imglist.zip')
with zipfile.ZipFile(file_path, 'r') as zip_file:
zip_file.extractall(data_root)
os.remove(file_path)
for dataset in benchmarks_dict[id_data_name]:
download_dataset(dataset, data_root)
def get_id_ood_dataloader(id_name, data_root, preprocessor, **loader_kwargs):
if 'imagenet' in id_name:
if tvs_new:
if isinstance(preprocessor,
tvs.transforms._presets.ImageClassification):
mean, std = preprocessor.mean, preprocessor.std
elif isinstance(preprocessor, tvs.transforms.Compose):
temp = preprocessor.transforms[-1]
mean, std = temp.mean, temp.std
elif isinstance(preprocessor, BasePreprocessor):
temp = preprocessor.transform.transforms[-1]
mean, std = temp.mean, temp.std
else:
raise TypeError
else:
if isinstance(preprocessor, tvs.transforms.Compose):
temp = preprocessor.transforms[-1]
mean, std = temp.mean, temp.std
elif isinstance(preprocessor, BasePreprocessor):
temp = preprocessor.transform.transforms[-1]
mean, std = temp.mean, temp.std
else:
raise TypeError
imagenet_c_preprocessor = ImageNetCPreProcessor(mean, std)
# weak augmentation for data_aux
test_standard_preprocessor = get_default_preprocessor(id_name)
dataloader_dict = {}
data_info = DATA_INFO[id_name]
# id
sub_dataloader_dict = {}
for split in data_info['id'].keys():
if id_name == "bronze2":
dataset = BronzeWare_ID_Dataset(os.path.join(data_root,
data_info['id'][split]['data_dir']),
os.path.join(data_root,
data_info['id'][split]['xml_path']),
os.path.join(data_root,
data_info['id'][split]['imglist_path']),
input_transform=preprocessor)
else:
dataset = ImglistDataset(
name='_'.join((id_name, split)),
imglist_pth=os.path.join(data_root,
data_info['id'][split]['imglist_path']),
data_dir=os.path.join(data_root,
data_info['id'][split]['data_dir']),
num_classes=data_info['num_classes'],
preprocessor=preprocessor,
data_aux_preprocessor=test_standard_preprocessor)
dataloader = DataLoader(dataset, **loader_kwargs)
sub_dataloader_dict[split] = dataloader
dataloader_dict['id'] = sub_dataloader_dict
# csid
sub_dataloader_dict = {}
for dataset_name in data_info['csid']['datasets']:
dataset = ImglistDataset(
name='_'.join((id_name, 'csid', dataset_name)),
imglist_pth=os.path.join(
data_root, data_info['csid'][dataset_name]['imglist_path']),
data_dir=os.path.join(data_root,
data_info['csid'][dataset_name]['data_dir']),
num_classes=data_info['num_classes'],
preprocessor=preprocessor
if dataset_name != 'imagenet_c' else imagenet_c_preprocessor,
data_aux_preprocessor=test_standard_preprocessor)
dataloader = DataLoader(dataset, **loader_kwargs)
sub_dataloader_dict[dataset_name] = dataloader
dataloader_dict['csid'] = sub_dataloader_dict
# ood
dataloader_dict['ood'] = {}
for split in data_info['ood'].keys():
split_config = data_info['ood'][split]
if split == 'val':
# validation set
dataset = ImglistDataset(
name='_'.join((id_name, 'ood', split)),
imglist_pth=os.path.join(data_root,
split_config['imglist_path']),
data_dir=os.path.join(data_root, split_config['data_dir']),
num_classes=data_info['num_classes'],
preprocessor=preprocessor,
data_aux_preprocessor=test_standard_preprocessor)
dataloader = DataLoader(dataset, **loader_kwargs)
dataloader_dict['ood'][split] = dataloader
else:
# dataloaders for nearood, farood
sub_dataloader_dict = {}
for dataset_name in split_config['datasets']:
dataset_config = split_config[dataset_name]
dataset = ImglistDataset(
name='_'.join((id_name, 'ood', dataset_name)),
imglist_pth=os.path.join(data_root,
dataset_config['imglist_path']),
data_dir=os.path.join(data_root,
dataset_config['data_dir']),
num_classes=data_info['num_classes'],
preprocessor=preprocessor,
data_aux_preprocessor=test_standard_preprocessor)
dataloader = DataLoader(dataset, **loader_kwargs)
sub_dataloader_dict[dataset_name] = dataloader
dataloader_dict['ood'][split] = sub_dataloader_dict
return dataloader_dict