(action_error_sum, is_train, data_type)
| 43 | return mean_error_p1, mean_error_p2, pck, auc |
| 44 | |
| 45 | def print_error_action(action_error_sum, is_train, data_type): |
| 46 | mean_error_each = {'p1': 0.0, 'p2': 0.0, 'pck': 0.0, 'auc': 0.0} |
| 47 | mean_error_all = {'p1': AccumLoss(), 'p2': AccumLoss(), 'pck': AccumLoss(), 'auc': AccumLoss()} |
| 48 | |
| 49 | if not is_train: |
| 50 | if data_type.startswith('3dhp'): |
| 51 | print("{0:=^12} {1:=^10} {2:=^8} {3:=^8} {4:=^8}".format("Action", "p#1 mm", "p#2 mm", "PCK", "AUC")) |
| 52 | else: |
| 53 | print("{0:=^12} {1:=^10} {2:=^8}".format("Action", "p#1 mm", "p#2 mm")) |
| 54 | |
| 55 | for action, value in action_error_sum.items(): |
| 56 | if not is_train: |
| 57 | print("{0:<12} ".format(action), end="") |
| 58 | |
| 59 | mean_error_each['p1'] = action_error_sum[action]['p1'].avg * 1000.0 |
| 60 | mean_error_all['p1'].update(mean_error_each['p1'], 1) |
| 61 | |
| 62 | mean_error_each['p2'] = action_error_sum[action]['p2'].avg * 1000.0 |
| 63 | mean_error_all['p2'].update(mean_error_each['p2'], 1) |
| 64 | |
| 65 | mean_error_each['pck'] = action_error_sum[action]['pck'].avg * 100.0 |
| 66 | mean_error_all['pck'].update(mean_error_each['pck'], 1) |
| 67 | |
| 68 | mean_error_each['auc'] = action_error_sum[action]['auc'].avg * 100.0 |
| 69 | mean_error_all['auc'].update(mean_error_each['auc'], 1) |
| 70 | |
| 71 | if not is_train: |
| 72 | if data_type.startswith('3dhp'): |
| 73 | print("{0:>6.2f} {1:>10.2f} {2:>10.2f} {3:>10.2f}".format( |
| 74 | mean_error_each['p1'], mean_error_each['p2'], |
| 75 | mean_error_each['pck'], mean_error_each['auc'])) |
| 76 | else: |
| 77 | print("{0:>6.2f} {1:>10.2f}".format(mean_error_each['p1'], mean_error_each['p2'])) |
| 78 | |
| 79 | if not is_train: |
| 80 | if data_type.startswith('3dhp'): |
| 81 | print("{0:<12} {1:>6.2f} {2:>10.2f} {3:>10.2f} {4:>10.2f}".format("Average", |
| 82 | mean_error_all['p1'].avg, mean_error_all['p2'].avg, |
| 83 | mean_error_all['pck'].avg, mean_error_all['auc'].avg)) |
| 84 | else: |
| 85 | print("{0:<12} {1:>6.2f} {2:>10.2f}".format("Average", mean_error_all['p1'].avg, \ |
| 86 | mean_error_all['p2'].avg)) |
| 87 | |
| 88 | if data_type.startswith('3dhp'): |
| 89 | return mean_error_all['p1'].avg, mean_error_all['p2'].avg, \ |
| 90 | mean_error_all['pck'].avg, mean_error_all['auc'].avg |
| 91 | else: |
| 92 | return mean_error_all['p1'].avg, mean_error_all['p2'].avg, 0, 0 |
| 93 | |
| 94 | |
| 95 | def save_model(args, epoch, mpjpe, model, model_name): |
no test coverage detected