| 303 | |
| 304 | class ChairsSDHom(object): |
| 305 | def __init__(self, args, is_cropped, root='/path/to/chairssdhom/data', dstype='train', replicates=1): |
| 306 | self.args = args |
| 307 | self.is_cropped = is_cropped |
| 308 | self.crop_size = args.crop_size |
| 309 | self.render_size = args.inference_size |
| 310 | self.replicates = replicates |
| 311 | |
| 312 | image1 = sorted(glob(join(root, dstype, 't0/*.png'))) |
| 313 | image2 = sorted(glob(join(root, dstype, 't1/*.png'))) |
| 314 | self.flow_list = sorted(glob(join(root, dstype, 'flow/*.flo'))) |
| 315 | |
| 316 | assert (len(image1) == len(self.flow_list)) |
| 317 | |
| 318 | self.image_list = [] |
| 319 | for i in range(len(self.flow_list)): |
| 320 | im1 = image1[i] |
| 321 | im2 = image2[i] |
| 322 | self.image_list += [[im1, im2]] |
| 323 | |
| 324 | assert len(self.image_list) == len(self.flow_list) |
| 325 | |
| 326 | self.size = len(self.image_list) |
| 327 | |
| 328 | self.frame_size = frame_utils.read_gen(self.image_list[0][0]).shape |
| 329 | |
| 330 | if (self.render_size[0] < 0) or (self.render_size[1] < 0) or (self.frame_size[0] % 64) or ( |
| 331 | self.frame_size[1] % 64): |
| 332 | self.render_size[0] = ((self.frame_size[0]) // 64) * 64 |
| 333 | self.render_size[1] = ((self.frame_size[1]) // 64) * 64 |
| 334 | |
| 335 | args.inference_size = self.render_size |
| 336 | |
| 337 | def __getitem__(self, index): |
| 338 | index = index % self.size |