(args, model_pos, test_loader, datareader)
| 54 | }, chk_path) |
| 55 | |
| 56 | def 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] |
no test coverage detected