MCPcopy Create free account
hub / github.com/Vegetebird/GraphMLP / test

Function test

main.py:62–93  ·  view source on GitHub ↗
(actions, dataloader, model, model_refine)

Source from the content-addressed store, hash-verified

60
61
62def test(actions, dataloader, model, model_refine):
63 model.eval()
64
65 action_error = define_error_list(actions)
66
67 for i, data in enumerate(tqdm(dataloader, 0)):
68 batch_cam, gt_3D, input_2D, input_2D_GT, action, subject, cam_ind = data
69 [input_2D, input_2D_GT, gt_3D, batch_cam] = get_varialbe('test', [input_2D, input_2D_GT, gt_3D, batch_cam])
70
71 output_3D_non_flip = model(input_2D[:, 0])
72 output_3D_flip = model(input_2D[:, 1])
73
74 output_3D_flip[:, :, :, 0] *= -1
75 output_3D_flip[:, :, args.joints_left + args.joints_right, :] = output_3D_flip[:, :, args.joints_right + args.joints_left, :]
76
77 output_3D = (output_3D_non_flip + output_3D_flip) / 2
78
79 out_target = gt_3D.clone()
80 out_target = out_target[:, args.pad].unsqueeze(1)
81
82 if args.refine:
83 model_refine.eval()
84 output_3D = refine_model(model_refine, output_3D, input_2D[:, 0], gt_3D, batch_cam, args.pad, args.root_joint)
85
86 output_3D[:, :, args.root_joint] = 0
87 out_target[:, :, args.root_joint] = 0
88
89 action_error = eval_cal.test_calculation(output_3D, out_target, action, action_error, args.dataset, subject)
90
91 p1, p2, pck, auc = print_error(args.dataset, action_error, args.train)
92
93 return p1, p2, pck, auc
94
95if __name__ == '__main__':
96 seed = 1

Callers 1

main.pyFile · 0.85

Calls 4

refine_modelFunction · 0.90
define_error_listFunction · 0.85
get_varialbeFunction · 0.85
print_errorFunction · 0.85

Tested by

no test coverage detected