diff --git a/api/CocoDetectionDataset.py b/api/CocoDetectionDataset.py new file mode 100644 index 0000000..6f27871 --- /dev/null +++ b/api/CocoDetectionDataset.py @@ -0,0 +1,65 @@ +# Custom PyTorch Dataset to load COCO-format annotations and images +import matplotlib.pyplot as plt +import cv2 +import os +import torch +from PIL import Image +import numpy as np +from torch.utils.data import Dataset +from pycocotools.coco import COCO + +# website: https://visionbrick.com/pipeline-for-training-custom-faster-rcnn-object-detection-models-with-pytorch/ + +class CocoDetectionDataset(Dataset): + # Init function: loads annotation file and prepares list of image IDs + def __init__(self, image_dir, annotation_path, transforms=None): + self.image_dir = image_dir + self.coco = COCO(annotation_path) + self.image_ids = list(self.coco.imgs.keys()) + self.transforms = transforms + + # Returns total number of images + def __len__(self): + return len(self.image_ids) + + # Fetches a single image and its annotations + def __getitem__(self, idx): + image_id = self.image_ids[idx] + image_info = self.coco.loadImgs(image_id)[0] + image_path = os.path.join(self.image_dir, image_info['file_name']) + image = Image.open(image_path).convert("RGB") + + # Load all annotations for this image + annotation_ids = self.coco.getAnnIds(imgIds=image_id) + annotations = self.coco.loadAnns(annotation_ids) + + # Extract bounding boxes and labels from annotations + boxes = [] + labels = [] + for obj in annotations: + xmin, ymin, width, height = obj['bbox'] + xmax = xmin + width + ymax = ymin + height + boxes.append([xmin, ymin, xmax, ymax]) + labels.append(obj['category_id']) + + # Convert annotations to PyTorch tensors + boxes = torch.as_tensor(boxes, dtype=torch.float32) + labels = torch.as_tensor(labels, dtype=torch.int64) + area = torch.as_tensor([obj['area'] for obj in annotations], dtype=torch.float32) + iscrowd = torch.as_tensor([obj.get('iscrowd', 0) for obj in annotations], dtype=torch.int64) + + # Package everything into a target dictionary + target = { + "boxes": boxes, + "labels": labels, + "image_id": image_id, + "area": area, + "iscrowd": iscrowd + } + + # Apply transforms if any were passed + if self.transforms: + image = self.transforms(image) + + return image, target \ No newline at end of file