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

Class PolypObjDataset

utils/data_val.py:117–220  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

115
116# dataset for training
117class 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)

Callers 1

get_loaderFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected