| 115 | |
| 116 | # dataset for training |
| 117 | class PolypObjDataset(data.Dataset): |
| 118 | def __init__(self, image_root, gt_root, edge_root, trainsize): |
| 119 | self.trainsize = trainsize |
| 120 | # get filenames |
| 121 | self.images = [image_root + f for f in os.listdir(image_root) if f.endswith('.jpg')] |
| 122 | self.gts = [gt_root + f for f in os.listdir(gt_root) if f.endswith('.jpg') or f.endswith('.png')] |
| 123 | self.edges = [edge_root + f for f in os.listdir(edge_root) if f.endswith('.jpg') or f.endswith('.png')] |
| 124 | # self.grads = [grad_root + f for f in os.listdir(grad_root) if f.endswith('.jpg') |
| 125 | # or f.endswith('.png')] |
| 126 | # self.depths = [depth_root + f for f in os.listdir(depth_root) if f.endswith('.bmp') |
| 127 | # or f.endswith('.png')] |
| 128 | # 将图像输入大小 -> 边缘转成 // 8 |
| 129 | self.edgesize = self.trainsize |
| 130 | # sorted files |
| 131 | self.images = sorted(self.images) |
| 132 | self.gts = sorted(self.gts) |
| 133 | self.edges = sorted(self.edges) |
| 134 | # self.grads = sorted(self.grads) |
| 135 | # self.depths = sorted(self.depths) |
| 136 | # filter mathcing degrees of files |
| 137 | self.filter_files() |
| 138 | # transforms |
| 139 | self.img_transform = transforms.Compose([ |
| 140 | transforms.Resize((self.trainsize, self.trainsize)), |
| 141 | transforms.ToTensor(), |
| 142 | transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])]) |
| 143 | self.gt_transform = transforms.Compose([ |
| 144 | transforms.Resize((self.trainsize, self.trainsize)), |
| 145 | transforms.ToTensor()]) |
| 146 | |
| 147 | self.edge_transform = transforms.Compose([ |
| 148 | transforms.Resize((self.edgesize, self.edgesize)), |
| 149 | transforms.ToTensor()]) |
| 150 | |
| 151 | self.small_transform = transforms.Compose([ |
| 152 | transforms.Resize((self.edgesize//32, self.edgesize//32)), |
| 153 | transforms.ToTensor()]) |
| 154 | |
| 155 | self.kernel = np.ones((3, 3), np.uint8) |
| 156 | # get size of dataset |
| 157 | self.size = len(self.images) |
| 158 | |
| 159 | def __getitem__(self, index): |
| 160 | # read imgs/gts/grads/depths |
| 161 | image = self.rgb_loader(self.images[index]) |
| 162 | gt = self.binary_loader(self.gts[index]) |
| 163 | |
| 164 | edge = cv2.imread(self.edges[index], cv2.IMREAD_GRAYSCALE) |
| 165 | edge = cv2.dilate(edge, self.kernel, iterations=1) |
| 166 | edge = Image.fromarray(edge) |
| 167 | |
| 168 | # data augumentation |
| 169 | image, gt, edge = cv_random_flip(image, gt, edge) |
| 170 | image, gt, edge = randomCrop(image, gt, edge) |
| 171 | image, gt, edge = randomRotation(image, gt, edge) |
| 172 | gt_small = self.small_transform(gt) |
| 173 | |
| 174 | image = colorEnhance(image) |