| import os |
| import torch.nn as nn |
| import torch |
| import sys |
|
|
| from MTIL_datasets.caltech101 import Caltech101 |
| from MTIL_datasets.cifar100 import CIFAR100 |
| from MTIL_datasets.dtd import DescribableTextures as DTD |
| from MTIL_datasets.eurosat import EuroSAT |
| from MTIL_datasets.fgvc_aircraft import FGVCAircraft as Aircraft |
| from MTIL_datasets.food101 import Food101 as Food |
| from MTIL_datasets.mnist import MNIST |
| from MTIL_datasets.oxford_flowers import OxfordFlowers as Flowers |
| from MTIL_datasets.oxford_pets import OxfordPets as OxfordPet |
| from MTIL_datasets.stanford_cars import StanfordCars |
| from MTIL_datasets.sun397 import SUN397 |
| from MTIL_datasets.ucf101 import UCF101 |
| from MTIL_datasets.country211 import Country211 |
| from MTIL_datasets.sst2 import SST2 |
| from MTIL_datasets.hatefulmemes import HatefulMemes |
| from MTIL_datasets.gtsrb import GTSRB |
| from MTIL_datasets.resisc import RESISC45 |
| from MTIL_datasets.fer2013 import FER2013 |
| from MTIL_datasets.cifar10 import CIFAR10 |
| from MTIL_datasets.stl10 import STL10 |
| from MTIL_datasets.voc2007 import VOC2007 |
| from MTIL_datasets.imagenet_r import ImageNetR |
| from MTIL_datasets.kitti_distance import KittiDistance |
| from MTIL_datasets.pcam import PCam |
| from MTIL_datasets.clevr_count import CLEVRCount |
| from MTIL_datasets.utils import DatasetWrapper |
|
|
|
|
| def get_dataset(cfg, split, transforms=None): |
| if split == 'val' and (not cfg.use_validation): |
| return None, None, None |
|
|
| is_train = (split == 'train') |
| templates = None |
| dataset_names = None |
|
|
| if cfg.dataset == "MTIL": |
| |
| |
| all_sets = [ |
| Aircraft, |
| Caltech101, |
| CIFAR100, |
| DTD, |
| EuroSAT, |
| Flowers, |
| Food, |
| MNIST, |
| OxfordPet, |
| StanfordCars, |
| SUN397, |
| Country211, |
| SST2, |
| HatefulMemes, |
| GTSRB, |
| RESISC45, |
| FER2013, |
| UCF101, |
| CIFAR10, |
| STL10, |
| VOC2007, |
| ImageNetR, |
| KittiDistance, |
| PCam, |
| CLEVRCount, |
| ] |
| |
| |
| if getattr(cfg, 'MTIL_order_2', False): |
| legacy_first11 = [StanfordCars, Food, MNIST, OxfordPet, Flowers, SUN397, Aircraft, Caltech101, DTD, EuroSAT, CIFAR100] |
| all_sets = legacy_first11 + all_sets[11:] |
| single_mode = int(getattr(cfg, 'train_one_dataset', -1)) |
| selected_indices = [single_mode] if single_mode >= 0 else list(range(len(all_sets))) |
| dataset = [] |
| classes_names = [] |
| templates = [] |
| dataset_names = [] |
| for idx in selected_indices: |
| base_ctor = all_sets[idx] if idx < len(all_sets) else None |
| if base_ctor is None: |
| if single_mode >= 0: |
| raise ValueError(f"MTIL dataset index {idx} is not available (module missing).") |
| |
| continue |
| base = base_ctor(cfg.dataset_root, seed=getattr(cfg, 'seed', 1)) |
| classes_names.append(base.classnames) |
| |
| templates.append(getattr(base, 'templates', None)) |
| dataset_names.append(base_ctor.__name__) |
| if split == 'train': |
| dataset.append(DatasetWrapper(base.train_x, transform=transforms, is_train=is_train)) |
| elif split == 'val': |
| dataset.append(DatasetWrapper(base.val, transform=transforms, is_train=is_train)) |
| elif split == 'test': |
| dataset.append(DatasetWrapper(base.test, transform=transforms, is_train=is_train)) |
| else: |
| ValueError(f"'{cfg.dataset}' is a invalid dataset.") |
| return dataset, classes_names, templates, dataset_names |
|
|
|
|
|
|
| def parse_sample(sample, is_train, task_id, cfg): |
| return sample[0], sample[1], torch.IntTensor([task_id]).repeat(sample[0].size(0)) |