MCPcopy Create free account
hub / github.com/Walter0807/MotionBERT / evaluate

Function evaluate

train.py:56–153  ·  view source on GitHub ↗
(args, model_pos, test_loader, datareader)

Source from the content-addressed store, hash-verified

54 }, chk_path)
55
56def evaluate(args, model_pos, test_loader, datareader):
57 print('INFO: Testing')
58 results_all = []
59 model_pos.eval()
60 with torch.no_grad():
61 for batch_input, batch_gt in tqdm(test_loader):
62 N, T = batch_gt.shape[:2]
63 if torch.cuda.is_available():
64 batch_input = batch_input.cuda()
65 if args.no_conf:
66 batch_input = batch_input[:, :, :, :2]
67 if args.flip:
68 batch_input_flip = flip_data(batch_input)
69 predicted_3d_pos_1 = model_pos(batch_input)
70 predicted_3d_pos_flip = model_pos(batch_input_flip)
71 predicted_3d_pos_2 = flip_data(predicted_3d_pos_flip) # Flip back
72 predicted_3d_pos = (predicted_3d_pos_1+predicted_3d_pos_2) / 2
73 else:
74 predicted_3d_pos = model_pos(batch_input)
75 if args.rootrel:
76 predicted_3d_pos[:,:,0,:] = 0 # [N,T,17,3]
77 else:
78 batch_gt[:,0,0,2] = 0
79
80 if args.gt_2d:
81 predicted_3d_pos[...,:2] = batch_input[...,:2]
82 results_all.append(predicted_3d_pos.cpu().numpy())
83 results_all = np.concatenate(results_all)
84 results_all = datareader.denormalize(results_all)
85 _, split_id_test = datareader.get_split_id()
86 actions = np.array(datareader.dt_dataset['test']['action'])
87 factors = np.array(datareader.dt_dataset['test']['2.5d_factor'])
88 gts = np.array(datareader.dt_dataset['test']['joints_2.5d_image'])
89 sources = np.array(datareader.dt_dataset['test']['source'])
90
91 num_test_frames = len(actions)
92 frames = np.array(range(num_test_frames))
93 action_clips = actions[split_id_test]
94 factor_clips = factors[split_id_test]
95 source_clips = sources[split_id_test]
96 frame_clips = frames[split_id_test]
97 gt_clips = gts[split_id_test]
98 assert len(results_all)==len(action_clips)
99
100 e1_all = np.zeros(num_test_frames)
101 e2_all = np.zeros(num_test_frames)
102 oc = np.zeros(num_test_frames)
103 results = {}
104 results_procrustes = {}
105 action_names = sorted(set(datareader.dt_dataset['test']['action']))
106 for action in action_names:
107 results[action] = []
108 results_procrustes[action] = []
109 block_list = ['s_09_act_05_subact_02',
110 's_09_act_10_subact_02',
111 's_09_act_13_subact_01']
112 for idx in range(len(action_clips)):
113 source = source_clips[idx][0][:-6]

Callers 1

train_with_configFunction · 0.85

Calls 5

flip_dataFunction · 0.90
mpjpeFunction · 0.85
p_mpjpeFunction · 0.85
denormalizeMethod · 0.80
get_split_idMethod · 0.45

Tested by

no test coverage detected