MCPcopy Create free account
hub / github.com/VITA-Group/GNT / eval

Function eval

eval.py:41–121  ·  view source on GitHub ↗
(args)

Source from the content-addressed store, hash-verified

39
40@torch.no_grad()
41def 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)

Callers 1

eval.pyFile · 0.85

Calls 5

create_training_datasetFunction · 0.90
GNTModelClass · 0.90
ProjectorClass · 0.90
log_viewFunction · 0.70

Tested by

no test coverage detected