MCPcopy Create free account
hub / github.com/ChunmingHe/WS-SAM / PolypObjDataset_noEdge

Class PolypObjDataset_noEdge

utils/data_val.py:224–322  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

222
223# dataset for training
224class PolypObjDataset_noEdge(data.Dataset):
225 def __init__(self, image_root, gt_root, trainsize=384):
226 self.trainsize = trainsize
227 # get filenames
228 self.images = [image_root + f for f in os.listdir(image_root) if f.endswith('.jpg') or f.endswith('.png')]
229 self.gts = [gt_root + f for f in os.listdir(gt_root) if f.endswith('.jpg') or f.endswith('.png')]
230 # self.edges = [edge_root + f for f in os.listdir(edge_root) if f.endswith('.jpg') or f.endswith('.png')]
231 # self.grads = [grad_root + f for f in os.listdir(grad_root) if f.endswith('.jpg')
232 # or f.endswith('.png')]
233 # self.depths = [depth_root + f for f in os.listdir(depth_root) if f.endswith('.bmp')
234 # or f.endswith('.png')]
235 # 将图像输入大小 -> 边缘转成 // 8
236 self.edgesize = self.trainsize
237 # sorted files
238 self.images = sorted(self.images)
239 self.gts = sorted(self.gts)
240 # self.edges = sorted(self.edges)
241 # self.grads = sorted(self.grads)
242 # self.depths = sorted(self.depths)
243 # filter mathcing degrees of files
244 self.filter_files()
245 # transforms
246 self.img_transform = transforms.Compose([
247 transforms.Resize((self.trainsize, self.trainsize)),
248 transforms.ToTensor(),
249 transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])])
250 self.gt_transform = transforms.Compose([
251 transforms.Resize((self.trainsize, self.trainsize)),
252 transforms.ToTensor()])
253
254 self.edge_transform = transforms.Compose([
255 transforms.Resize((self.edgesize, self.edgesize)),
256 transforms.ToTensor()])
257
258 self.small_transform = transforms.Compose([
259 transforms.Resize((self.edgesize//32, self.edgesize//32)),
260 transforms.ToTensor()])
261
262 self.kernel = np.ones((3, 3), np.uint8)
263 # get size of dataset
264 self.size = len(self.images)
265
266 def __getitem__(self, index):
267 # read imgs/gts/grads/depths
268 image = self.rgb_loader(self.images[index])
269 gt = self.binary_loader(self.gts[index])
270
271 # edge = cv2.imread(self.edges[index], cv2.IMREAD_GRAYSCALE)
272 # edge = cv2.dilate(edge, self.kernel, iterations=1)
273 # edge = Image.fromarray(edge)
274
275 # data augumentation
276 image, gt = cv_random_flip_noEdge(image, gt)
277 image, gt = randomCrop_noEdge(image, gt)
278 image, gt = randomRotation_noEdge(image, gt)
279 gt_small = self.small_transform(gt)
280
281 image = colorEnhance(image)

Callers 1

get_loader_noEdgeFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected