MCPcopy Create free account
hub / github.com/SLDGroup/EMCAD / PolypDataset

Class PolypDataset

utils/dataloader.py:10–125  ·  view source on GitHub ↗

dataloader for polyp segmentation tasks

Source from the content-addressed store, hash-verified

8
9
10class 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:

Callers 1

get_loaderFunction · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected