Data loader for Pose Regression Network
(args)
| 347 | return train_set, val_set, bounds |
| 348 | |
| 349 | def load_Cambridge_dataloader(args): |
| 350 | ''' Data loader for Pose Regression Network ''' |
| 351 | if args.pose_only: # if train posenet is true |
| 352 | pass |
| 353 | else: |
| 354 | raise Exception('load_Cambridge_dataloader() currently only support PoseNet Training, not NeRF training') |
| 355 | data_dir, scene = osp.split(args.datadir) # ../data/7Scenes, chess |
| 356 | dataset_folder, dataset = osp.split(data_dir) # ../data, 7Scenes |
| 357 | |
| 358 | # transformer |
| 359 | data_transform = transforms.Compose([ |
| 360 | transforms.ToTensor(), |
| 361 | ]) |
| 362 | target_transform = transforms.Lambda(lambda x: torch.Tensor(x)) |
| 363 | |
| 364 | ret_idx = False # return frame index |
| 365 | fix_idx = False # return frame index=0 in training |
| 366 | ret_hist = False |
| 367 | |
| 368 | if 'NeRFH' in args: |
| 369 | if args.NeRFH == True: |
| 370 | ret_idx = True |
| 371 | if args.fix_index: |
| 372 | fix_idx = True |
| 373 | |
| 374 | # encode hist experiment |
| 375 | if args.encode_hist: |
| 376 | ret_idx = False |
| 377 | fix_idx = False |
| 378 | ret_hist = True |
| 379 | |
| 380 | kwargs = dict(scene=scene, data_path=data_dir, |
| 381 | transform=data_transform, target_transform=target_transform, |
| 382 | df=args.df, ret_idx=ret_idx, fix_idx=fix_idx, |
| 383 | ret_hist=ret_hist, hist_bin=args.hist_bin) |
| 384 | |
| 385 | if args.finetune_unlabel: # direct-pn + unlabel |
| 386 | train_set = Cambridge2(train=False, testskip=args.trainskip, **kwargs) |
| 387 | val_set = Cambridge2(train=False, testskip=args.testskip, **kwargs) |
| 388 | |
| 389 | # if not args.eval: |
| 390 | # # remove overlap data in val_set that was already in train_set, |
| 391 | # train_set, val_set = remove_overlap_data(train_set, val_set) |
| 392 | else: |
| 393 | train_set = Cambridge2(train=True, trainskip=args.trainskip, **kwargs) |
| 394 | val_set = Cambridge2(train=False, testskip=args.testskip, **kwargs) |
| 395 | L = len(train_set) |
| 396 | |
| 397 | i_train = train_set.gt_idx |
| 398 | i_val = val_set.gt_idx |
| 399 | i_test = val_set.gt_idx |
| 400 | # use a pose average stats computed earlier to unify posenet and nerf training |
| 401 | if args.save_pose_avg_stats or args.load_pose_avg_stats: |
| 402 | pose_avg_stats_file = osp.join(args.datadir, 'pose_avg_stats.txt') |
| 403 | train_set, val_set, bounds = fix_coord(args, train_set, val_set, pose_avg_stats_file, rescale_coord=False) # only adjust coord. systems, rescale are done at training |
| 404 | else: |
| 405 | train_set, val_set, bounds = fix_coord(args, train_set, val_set, rescale_coord=False) |
| 406 |
no test coverage detected