| 226 | |
| 227 | class FlyingThings(object): |
| 228 | def __init__(self, args, is_cropped, root='/path/to/flyingthings3d', dstype='frames_cleanpass', replicates=1): |
| 229 | self.args = args |
| 230 | self.is_cropped = is_cropped |
| 231 | self.crop_size = args.crop_size |
| 232 | self.render_size = args.inference_size |
| 233 | self.replicates = replicates |
| 234 | |
| 235 | image_dirs = sorted(glob(join(root, dstype, 'TRAIN/*/*'))) |
| 236 | image_dirs = sorted([join(f, 'left') for f in image_dirs] + [join(f, 'right') for f in image_dirs]) |
| 237 | |
| 238 | flow_dirs = sorted(glob(join(root, 'optical_flow_flo_format/TRAIN/*/*'))) |
| 239 | flow_dirs = sorted( |
| 240 | [join(f, 'into_future/left') for f in flow_dirs] + [join(f, 'into_future/right') for f in flow_dirs]) |
| 241 | |
| 242 | assert (len(image_dirs) == len(flow_dirs)) |
| 243 | |
| 244 | self.image_list = [] |
| 245 | self.flow_list = [] |
| 246 | |
| 247 | for idir, fdir in zip(image_dirs, flow_dirs): |
| 248 | images = sorted(glob(join(idir, '*.png'))) |
| 249 | flows = sorted(glob(join(fdir, '*.flo'))) |
| 250 | for i in range(len(flows)): |
| 251 | self.image_list += [[images[i], images[i + 1]]] |
| 252 | self.flow_list += [flows[i]] |
| 253 | |
| 254 | assert len(self.image_list) == len(self.flow_list) |
| 255 | |
| 256 | self.size = len(self.image_list) |
| 257 | |
| 258 | self.frame_size = frame_utils.read_gen(self.image_list[0][0]).shape |
| 259 | |
| 260 | if (self.render_size[0] < 0) or (self.render_size[1] < 0) or (self.frame_size[0] % 64) or ( |
| 261 | self.frame_size[1] % 64): |
| 262 | self.render_size[0] = ((self.frame_size[0]) // 64) * 64 |
| 263 | self.render_size[1] = ((self.frame_size[1]) // 64) * 64 |
| 264 | |
| 265 | args.inference_size = self.render_size |
| 266 | |
| 267 | def __getitem__(self, index): |
| 268 | index = index % self.size |