MCPcopy Create free account
hub / github.com/chengsen/PyTorch_TextGCN / macro_f1

Function macro_f1

utils.py:10–40  ·  view source on GitHub ↗
(pred, targ, num_classes=None)

Source from the content-addressed store, hash-verified

8
9
10def macro_f1(pred, targ, num_classes=None):
11 pred = th.max(pred, 1)[1]
12 tp_out = []
13 fp_out = []
14 fn_out = []
15 if num_classes is None:
16 num_classes = sorted(set(targ.cpu().numpy().tolist()))
17 else:
18 num_classes = range(num_classes)
19 for i in num_classes:
20 tp = ((pred == i) & (targ == i)).sum().item() # 预测为i,且标签的确为i的
21 fp = ((pred == i) & (targ != i)).sum().item() # 预测为i,但标签不是为i的
22 fn = ((pred != i) & (targ == i)).sum().item() # 预测不是i,但标签是i的
23 tp_out.append(tp)
24 fp_out.append(fp)
25 fn_out.append(fn)
26
27 eval_tp = np.array(tp_out)
28 eval_fp = np.array(fp_out)
29 eval_fn = np.array(fn_out)
30
31 precision = eval_tp / (eval_tp + eval_fp)
32 precision[np.isnan(precision)] = 0
33 precision = np.mean(precision)
34
35 recall = eval_tp / (eval_tp + eval_fn)
36 recall[np.isnan(recall)] = 0
37 recall = np.mean(recall)
38
39 f1 = 2 * (precision * recall) / (precision + recall)
40 return f1, precision, recall
41
42
43def accuracy(pred, targ):

Callers 1

valMethod · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected