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

Class PolypObjDataset_scribble_noEdge

utils/data_val.py:427–539  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

425
426
427class PolypObjDataset_scribble_noEdge(data.Dataset):
428 def __init__(self, image_root, gt_root, scribble_root, trainsize=384):
429 self.trainsize = trainsize
430 # get filenames
431 self.images = [image_root + f for f in os.listdir(image_root) if f.endswith('.jpg')]
432 self.gts = [gt_root + f for f in os.listdir(gt_root) if f.endswith('.jpg') or f.endswith('.png')]
433 self.scribbles = [scribble_root + f for f in os.listdir(scribble_root) if f.endswith('.jpg') or f.endswith('.png')]
434 # self.edges = [edge_root + f for f in os.listdir(edge_root) if f.endswith('.jpg') or f.endswith('.png')]
435 # self.grads = [grad_root + f for f in os.listdir(grad_root) if f.endswith('.jpg')
436 # or f.endswith('.png')]
437 # self.depths = [depth_root + f for f in os.listdir(depth_root) if f.endswith('.bmp')
438 # or f.endswith('.png')]
439 # 将图像输入大小 -> 边缘转成 // 8
440 self.edgesize = self.trainsize
441 # sorted files
442 self.images = sorted(self.images)
443 self.gts = sorted(self.gts)
444 self.scribbles = sorted(self.scribbles)
445
446 # self.edges = sorted(self.edges)
447 # self.grads = sorted(self.grads)
448 # self.depths = sorted(self.depths)
449 # filter mathcing degrees of files
450 self.filter_files()
451 # transforms
452 self.img_transform = transforms.Compose([
453 transforms.Resize((self.trainsize, self.trainsize)),
454 transforms.ToTensor(),
455 transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])])
456 self.gt_transform = transforms.Compose([
457 transforms.Resize((self.trainsize, self.trainsize)),
458 transforms.ToTensor()])
459
460 self.edge_transform = transforms.Compose([
461 transforms.Resize((self.edgesize, self.edgesize)),
462 transforms.ToTensor()])
463
464 self.small_transform = transforms.Compose([
465 transforms.Resize((self.edgesize//32, self.edgesize//32)),
466 transforms.ToTensor()])
467
468 self.kernel = np.ones((3, 3), np.uint8)
469 # get size of dataset
470 self.size = len(self.images)
471
472 def __getitem__(self, index):
473 # read imgs/gts/grads/depths
474 image = self.rgb_loader(self.images[index])
475 gt = self.binary_loader(self.gts[index])
476 scribble = self.binary_loader(self.scribbles[index])
477
478
479 # edge = cv2.imread(self.edges[index], cv2.IMREAD_GRAYSCALE)
480 # edge = cv2.dilate(edge, self.kernel, iterations=1)
481 # edge = Image.fromarray(edge)
482
483 # data augumentation
484 image, gt, scribble = cv_random_flip(image, gt, scribble)

Callers 1

Calls

no outgoing calls

Tested by

no test coverage detected