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

Function print_error_action

common/utils.py:45–92  ·  view source on GitHub ↗
(action_error_sum, is_train, data_type)

Source from the content-addressed store, hash-verified

43 return mean_error_p1, mean_error_p2, pck, auc
44
45def 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
95def save_model(args, epoch, mpjpe, model, model_name):

Callers 1

print_errorFunction · 0.85

Calls 2

AccumLossClass · 0.85
updateMethod · 0.45

Tested by

no test coverage detected