import os import torch from torchvision.io import [read_image](https://docs.pytorch.org/vision/stable/generated/torchvision.io.read_image.html#torchvision.io.read_image "torchvision.io.read_image") from torchvision.ops.boxes import [masks_to_boxes](https://docs.pytorch.org/vision/stable/generated/torchvision.ops.masks_to_boxes.html#torchvision.ops.masks_to_boxes "torchvision.ops.masks_to_boxes") from torchvision import tv_tensors from torchvision.transforms.v2 import functional as F class PennFudanDataset([torch.utils.data.Dataset](https://docs.pytorch.org/docs/stable/data.html#torch.utils.data.Dataset "torch.utils.data.Dataset")): def __init__(self, root, transforms): self.root = root self.transforms = transforms # load all image files, sorting them to # ensure that they are aligned self.imgs = list(sorted(os.listdir(os.path.join(root, "PNGImages")))) self.[masks](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor") = list(sorted(os.listdir(os.path.join(root, "PedMasks")))) def __getitem__(self, idx): # load images and masks img_path = os.path.join(self.root, "PNGImages", self.imgs[idx]) mask_path = os.path.join(self.root, "PedMasks", self.[masks](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor")[idx]) img = [read_image](https://docs.pytorch.org/vision/stable/generated/torchvision.io.read_image.html#torchvision.io.read_image "torchvision.io.read_image")(img_path) [mask](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor") = [read_image](https://docs.pytorch.org/vision/stable/generated/torchvision.io.read_image.html#torchvision.io.read_image "torchvision.io.read_image")(mask_path) # instances are encoded as different colors obj_ids = [torch.unique](https://docs.pytorch.org/docs/stable/generated/torch.unique.html#torch.unique "torch.unique")([mask](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor")) # first id is the background, so remove it obj_ids = obj_ids[1:] num_objs = len(obj_ids) # split the color-encoded mask into a set # of binary masks [masks](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor") = ([mask](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor") == obj_ids[:, None, None]).to(dtype=[torch.uint8](https://docs.pytorch.org/docs/stable/tensor_attributes.html#torch.dtype "torch.dtype")) # get bounding box coordinates for each mask boxes = [masks_to_boxes](https://docs.pytorch.org/vision/stable/generated/torchvision.ops.masks_to_boxes.html#torchvision.ops.masks_to_boxes "torchvision.ops.masks_to_boxes")([masks](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor")) # there is only one class labels = [torch.ones](https://docs.pytorch.org/docs/stable/generated/torch.ones.html#torch.ones "torch.ones")((num_objs,), dtype=[torch.int64](https://docs.pytorch.org/docs/stable/tensor_attributes.html#torch.dtype "torch.dtype")) image_id = idx area = (boxes[:, 3] - boxes[:, 1]) * (boxes[:, 2] - boxes[:, 0]) # suppose all instances are not crowd iscrowd = [torch.zeros](https://docs.pytorch.org/docs/stable/generated/torch.zeros.html#torch.zeros "torch.zeros")((num_objs,), dtype=[torch.int64](https://docs.pytorch.org/docs/stable/tensor_attributes.html#torch.dtype "torch.dtype")) # Wrap sample and targets into torchvision tv_tensors: img = [tv_tensors.Image](https://docs.pytorch.org/vision/stable/generated/torchvision.tv_tensors.Image.html#torchvision.tv_tensors.Image "torchvision.tv_tensors.Image")(img) target = {} target["boxes"] = [tv_tensors.BoundingBoxes](https://docs.pytorch.org/vision/stable/generated/torchvision.tv_tensors.BoundingBoxes.html#torchvision.tv_tensors.BoundingBoxes "torchvision.tv_tensors.BoundingBoxes")(boxes, format="XYXY", canvas_size=F.get_size(img)) target["masks"] = [tv_tensors.Mask](https://docs.pytorch.org/vision/stable/generated/torchvision.tv_tensors.Mask.html#torchvision.tv_tensors.Mask "torchvision.tv_tensors.Mask")([masks](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor")) target["labels"] = labels target["image_id"] = image_id target["area"] = area target["iscrowd"] = iscrowd if self.transforms is not None: img, target = self.transforms(img, target) return img, target def __len__(self): return len(self.imgs)