dataloader for polyp segmentation tasks
| 8 | |
| 9 | |
| 10 | class PolypDataset(data.Dataset): |
| 11 | """ |
| 12 | dataloader for polyp segmentation tasks |
| 13 | """ |
| 14 | def __init__(self, image_root, gt_root, trainsize, augmentations): |
| 15 | self.trainsize = trainsize |
| 16 | self.augmentations = augmentations |
| 17 | print(self.augmentations) |
| 18 | self.images = [image_root + f for f in os.listdir(image_root) if f.endswith('.jpg') or f.endswith('.png')] |
| 19 | self.gts = [gt_root + f for f in os.listdir(gt_root) if f.endswith('.png') or f.endswith('.jpg')] |
| 20 | self.images = sorted(self.images) |
| 21 | self.gts = sorted(self.gts) |
| 22 | self.filter_files() |
| 23 | self.size = len(self.images) |
| 24 | if self.augmentations == 'True': |
| 25 | print('Using RandomRotation, RandomFlip') |
| 26 | self.img_transform = transforms.Compose([ |
| 27 | transforms.RandomRotation(90, resample=False, expand=False, center=None, fill=None), |
| 28 | transforms.RandomVerticalFlip(p=0.5), |
| 29 | transforms.RandomHorizontalFlip(p=0.5), |
| 30 | transforms.Resize((self.trainsize, self.trainsize)), |
| 31 | transforms.ToTensor(), |
| 32 | transforms.Normalize([0.485, 0.456, 0.406], |
| 33 | [0.229, 0.224, 0.225])]) |
| 34 | self.gt_transform = transforms.Compose([ |
| 35 | transforms.RandomRotation(90, resample=False, expand=False, center=None, fill=None), |
| 36 | transforms.RandomVerticalFlip(p=0.5), |
| 37 | transforms.RandomHorizontalFlip(p=0.5), |
| 38 | transforms.Resize((self.trainsize, self.trainsize)), |
| 39 | transforms.ToTensor()]) |
| 40 | |
| 41 | else: |
| 42 | print('no augmentation') |
| 43 | self.img_transform = transforms.Compose([ |
| 44 | transforms.Resize((self.trainsize, self.trainsize)), |
| 45 | transforms.ToTensor(), |
| 46 | transforms.Normalize([0.485, 0.456, 0.406], |
| 47 | [0.229, 0.224, 0.225])]) |
| 48 | |
| 49 | self.gt_transform = transforms.Compose([ |
| 50 | transforms.Resize((self.trainsize, self.trainsize)), |
| 51 | transforms.ToTensor()]) |
| 52 | |
| 53 | |
| 54 | def __getitem__(self, index): |
| 55 | |
| 56 | image = self.rgb_loader(self.images[index]) |
| 57 | gt = self.binary_loader(self.gts[index]) |
| 58 | |
| 59 | seed = np.random.randint(2147483647) # make a seed with numpy generator |
| 60 | random.seed(seed) # apply this seed to img tranfsorms |
| 61 | torch.manual_seed(seed) # needed for torchvision 0.7 |
| 62 | if self.img_transform is not None: |
| 63 | image = self.img_transform(image) |
| 64 | |
| 65 | random.seed(seed) # apply this seed to img tranfsorms |
| 66 | torch.manual_seed(seed) # needed for torchvision 0.7 |
| 67 | if self.gt_transform is not None: |