stereoid's picture
Add files using upload-large-folder tool
af46737 verified
Raw
History Blame Contribute Delete
2.74 kB
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