MCPcopy Create free account
hub / github.com/Gorilla-Lab-SCUT/tango / test

Function test

eval.py:33–160  ·  view source on GitHub ↗
(args)

Source from the content-addressed store, hash-verified

31 return mesh
32
33def test(args):
34 Path(args.output_dir).mkdir(parents=True, exist_ok=True)
35 torch.set_default_dtype(torch.float32)
36 # torch.set_num_threads(8)
37 # Constrain all sources of randomness
38 torch.manual_seed(args.seed)
39 torch.cuda.manual_seed(args.seed)
40 torch.cuda.manual_seed_all(args.seed)
41 random.seed(args.seed)
42 np.random.seed(args.seed)
43 torch.backends.cudnn.benchmark = False
44 torch.backends.cudnn.deterministic = True
45
46 objbase, extension = os.path.splitext(os.path.basename(args.obj_path))
47 # Check that isn't already done
48 if (not args.overwrite) and os.path.exists(os.path.join(args.output_dir, "loss.png")) and \
49 os.path.exists(os.path.join(args.output_dir, f"{objbase}_final.obj")):
50 print(f"Already done with {args.output_dir}")
51 exit()
52 elif args.overwrite and os.path.exists(os.path.join(args.output_dir, "loss.png")) and \
53 os.path.exists(os.path.join(args.output_dir, f"{objbase}_final.obj")):
54 import shutil
55 for filename in os.listdir(args.output_dir):
56 file_path = os.path.join(args.output_dir, filename)
57 try:
58 if os.path.isfile(file_path) or os.path.islink(file_path):
59 os.unlink(file_path)
60 elif os.path.isdir(file_path):
61 shutil.rmtree(file_path)
62 except Exception as e:
63 print('Failed to delete %s. Reason: %s' % (file_path, e))
64
65 n_augs = args.n_augs
66 dir = args.output_dir
67
68 model = NeuralStyleField(args.material_random_pe_numfreq,
69 args.material_random_pe_sigma,
70 args.num_lgt_sgs,
71 args.max_delta_theta,
72 args.max_delta_phi,
73 args.normal_nerf_pe_numfreq,
74 args.normal_random_pe_numfreq,
75 args.symmetry,
76 args.radius,
77 args.background,
78 args.init_r_and_s,
79 args.width,
80 args.init_roughness,
81 args.init_specular,
82 args.material_nerf_pe_numfreq,
83 args.normal_random_pe_sigma,
84 args.if_normal_clamp
85 )
86 state_dict = torch.load(args.model_dir)
87 model.load_state_dict(state_dict['model'])
88 model.eval()
89 envmap = compute_envmap(lgtSGs=model.svbrdf_network.get_light(), H=256, W=512, upper_hemi=model.svbrdf_network.upper_hemi)
90 envmap = envmap.cpu().numpy()

Callers 1

eval.pyFile · 0.85

Calls 6

render_single_imageMethod · 0.95
NeuralStyleFieldClass · 0.90
compute_envmapFunction · 0.90
save_gifFunction · 0.85
get_lightMethod · 0.80
get_normalize_meshFunction · 0.70

Tested by

no test coverage detected