(args)
| 39 | |
| 40 | @torch.no_grad() |
| 41 | def eval(args): |
| 42 | |
| 43 | device = "cuda:{}".format(args.local_rank) |
| 44 | out_folder = os.path.join(args.rootdir, "out", args.expname) |
| 45 | print("outputs will be saved to {}".format(out_folder)) |
| 46 | os.makedirs(out_folder, exist_ok=True) |
| 47 | |
| 48 | # save the args and config files |
| 49 | f = os.path.join(out_folder, "args.txt") |
| 50 | with open(f, "w") as file: |
| 51 | for arg in sorted(vars(args)): |
| 52 | attr = getattr(args, arg) |
| 53 | file.write("{} = {}\n".format(arg, attr)) |
| 54 | |
| 55 | if args.config is not None: |
| 56 | f = os.path.join(out_folder, "config.txt") |
| 57 | if not os.path.isfile(f): |
| 58 | shutil.copy(args.config, f) |
| 59 | |
| 60 | if args.run_val == False: |
| 61 | # create training dataset |
| 62 | dataset, sampler = create_training_dataset(args) |
| 63 | # currently only support batch_size=1 (i.e., one set of target and source views) for each GPU node |
| 64 | # please use distributed parallel on multiple GPUs to train multiple target views per batch |
| 65 | loader = torch.utils.data.DataLoader( |
| 66 | dataset, |
| 67 | batch_size=1, |
| 68 | worker_init_fn=lambda _: np.random.seed(), |
| 69 | num_workers=args.workers, |
| 70 | pin_memory=True, |
| 71 | sampler=sampler, |
| 72 | shuffle=True if sampler is None else False, |
| 73 | ) |
| 74 | iterator = iter(loader) |
| 75 | else: |
| 76 | # create validation dataset |
| 77 | dataset = dataset_dict[args.eval_dataset](args, "validation", scenes=args.eval_scenes) |
| 78 | loader = DataLoader(dataset, batch_size=1) |
| 79 | iterator = iter(loader) |
| 80 | |
| 81 | # Create GNT model |
| 82 | model = GNTModel( |
| 83 | args, load_opt=not args.no_load_opt, load_scheduler=not args.no_load_scheduler |
| 84 | ) |
| 85 | # create projector |
| 86 | projector = Projector(device=device) |
| 87 | |
| 88 | indx = 0 |
| 89 | psnr_scores = [] |
| 90 | lpips_scores = [] |
| 91 | ssim_scores = [] |
| 92 | while True: |
| 93 | try: |
| 94 | data = next(iterator) |
| 95 | except: |
| 96 | break |
| 97 | if args.local_rank == 0: |
| 98 | tmp_ray_sampler = RaySamplerSingleImage(data, device, render_stride=args.render_stride) |
no test coverage detected