| 71 | |
| 72 | |
| 73 | def flatten(gt, pred, types = None): |
| 74 | assert len(gt) == len(pred) |
| 75 | |
| 76 | gt_flat = [] |
| 77 | pred_flat = [] |
| 78 | |
| 79 | for (sample_gt, sample_pred) in zip(gt, pred): |
| 80 | union = set() |
| 81 | |
| 82 | union.update(sample_gt) |
| 83 | union.update(sample_pred) |
| 84 | |
| 85 | for s in union: |
| 86 | if types: |
| 87 | if s in types: |
| 88 | if s in sample_gt: |
| 89 | gt_flat.append(types.index(s)+1) |
| 90 | else: |
| 91 | gt_flat.append(0) |
| 92 | |
| 93 | if s in sample_pred: |
| 94 | pred_flat.append(types.index(s)+1) |
| 95 | else: |
| 96 | pred_flat.append(0) |
| 97 | else: |
| 98 | gt_flat.append(0) |
| 99 | pred_flat.append(0) |
| 100 | else: |
| 101 | if s in sample_gt: |
| 102 | gt_flat.append(1) |
| 103 | else: |
| 104 | gt_flat.append(0) |
| 105 | |
| 106 | if s in sample_pred: |
| 107 | pred_flat.append(1) |
| 108 | else: |
| 109 | pred_flat.append(0) |
| 110 | return gt_flat, pred_flat |
| 111 | |
| 112 | def print_results(per_type, micro, macro, types, result_dict = None): |
| 113 | columns = ('type', 'precision', 'recall', 'f1-score', 'support') |