import os import json from PIL import Image from torch.utils.data import Dataset from torchvision.transforms import ToTensor TRAIN_ANN_PATH = 'annotations/instances_train2017.json' VAL_ANN_PATH = 'annotations/instances_val2017.json' TEST_ANN_PATH = 'annotations/instances_test2017.json' TRAIN_IMG_DIR = 'images/instances_train2017' VAL_IMG_DIR = 'images/instances_val2017' TEST_IMG_DIR = 'images/instances_test2017' class TrainDataset(Dataset): def __init__(self, dataset_root, is_test = False, transform=None): super(Dataset, self).__init__() self.dataset_root = dataset_root img_dir = TEST_IMG_DIR if is_test else TRAIN_IMG_DIR ann_path = TEST_ANN_PATH if is_test else TRAIN_ANN_PATH self.img_dir = os.path.join(dataset_root, img_dir) ann_path = os.path.join(dataset_root, ann_path) with open(ann_path, 'r') as f: self.ann_data = json.load(f) self.transform = transform self.id_file_dict = {img['id']: img['file_name'] for img in self.ann_data['images']} def __len__(self): return len(self.ann_data['annotations']) def __getitem__(self, idx): ann = self.ann_data['annotations'][idx] bbox = ann['bbox'] img = Image.open(os.path.join(self.img_dir, self.id_file_dict[ann['image_id']])) img = img.crop((bbox[0], bbox[1], bbox[0] + bbox[2], bbox[1] + bbox[3])) if self.transform: img = self.transform(img) else: img = ToTensor()(img) return img, ann['category_id'] - 1 class PredDataset(Dataset): def __init__(self, dataset_root, region_path, transform=None): super(Dataset, self).__init__() self.dataset_root = dataset_root self.img_dir = os.path.join(dataset_root, TEST_IMG_DIR) ann_path = os.path.join(dataset_root, TEST_ANN_PATH) with open(ann_path, 'r') as f: self.ann_data = json.load(f) with open(region_path, 'r') as f: self.regions = json.load(f) self.transform = transform self.id_file_dict = {ann['id']: ann['file_name'] for ann in self.ann_data['images']} def __len__(self): return len(self.regions) def __getitem__(self, idx): region = self.regions[idx] bbox = region['bbox'] img = Image.open(os.path.join(self.img_dir, self.id_file_dict[region['image_id']])) img = img.crop((bbox[0], bbox[1], bbox[0] + bbox[2], bbox[1] + bbox[3])) if self.transform: img = self.transform(img) else: img = ToTensor()(img) return img