MCPcopy Create free account
hub / github.com/PaddlePaddle/Research / FlyingChairs

Class FlyingChairs

CV/PWCNet/data/datasets.py:137–214  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

135
136
137class FlyingChairs(object):
138 def __init__(self, train_val, args, is_cropped, txt_file, root='/path/to/FlyingChairs_release/data', replicates=1):
139 self.args = args
140 self.is_cropped = is_cropped
141 self.crop_size = args.crop_size
142 self.render_size = args.inference_size
143 self.replicates = replicates
144
145 images = sorted(glob(join(root, '*.ppm')))
146
147 flow_list = sorted(glob(join(root, '*.flo')))
148
149 assert (len(images) // 2 == len(flow_list))
150
151 image_list = []
152 for i in range(len(flow_list)):
153 im1 = images[2 * i]
154 im2 = images[2 * i + 1]
155 image_list += [[im1, im2]]
156
157 assert len(image_list) == len(flow_list)
158 if train_val == 'train':
159 intindex = np.array(read_txt_to_index(txt_file))
160 image_list = np.array(image_list)
161 image_list = image_list[intindex == 1]
162 image_list = image_list.tolist()
163 flow_list = np.array(flow_list)
164 flow_list = flow_list[intindex == 1]
165 flow_list = flow_list.tolist()
166 assert len(image_list) == len(flow_list)
167 elif train_val == 'val':
168 intindex = np.array(read_txt_to_index(txt_file))
169 image_list = np.array(image_list)
170 image_list = image_list[intindex == 2]
171 image_list = image_list.tolist()
172 flow_list = np.array(flow_list)
173 flow_list = flow_list[intindex == 2]
174 flow_list = flow_list.tolist()
175 assert len(image_list) == len(flow_list)
176 else:
177 raise ValueError('FlyingChairs_train_val.txt not found for txt_file ......')
178 self.flow_list = flow_list
179 self.image_list = image_list
180
181 self.size = len(self.image_list)
182
183 self.frame_size = frame_utils.read_gen(self.image_list[0][0]).shape
184
185 if (self.render_size[0] < 0) or (self.render_size[1] < 0) or (self.frame_size[0] % 64) or (
186 self.frame_size[1] % 64):
187 self.render_size[0] = ((self.frame_size[0]) // 64) * 64
188 self.render_size[1] = ((self.frame_size[1]) // 64) * 64
189
190 args.inference_size = self.render_size
191
192 def __getitem__(self, index):
193 index = index % self.size
194

Callers 2

mainFunction · 0.90
datasets.pyFile · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected