Evaluate clustering metrics on two subsets of data, as defined by the mask 'mask' (Mask usually corresponding to `Old' and `New' classes in GCD setting) :param targets: All ground truth labels :param preds: All predictions :param mask: Mask defining two subsets :return:
(targets, preds, mask)
| 73 | # Mixed Eval Function |
| 74 | # ------------------------------- |
| 75 | def mixed_eval(targets, preds, mask): |
| 76 | |
| 77 | """ |
| 78 | Evaluate clustering metrics on two subsets of data, as defined by the mask 'mask' |
| 79 | (Mask usually corresponding to `Old' and `New' classes in GCD setting) |
| 80 | :param targets: All ground truth labels |
| 81 | :param preds: All predictions |
| 82 | :param mask: Mask defining two subsets |
| 83 | :return: |
| 84 | """ |
| 85 | |
| 86 | mask = mask.astype(bool) |
| 87 | |
| 88 | # Labelled examples |
| 89 | if mask.sum() == 0: # All examples come from unlabelled classes |
| 90 | |
| 91 | unlabelled_acc, unlabelled_nmi, unlabelled_ari = cluster_acc(targets.astype(int), preds.astype(int)), \ |
| 92 | nmi_score(targets, preds), \ |
| 93 | ari_score(targets, preds) |
| 94 | |
| 95 | print('Unlabelled Classes Test acc {:.4f}, nmi {:.4f}, ari {:.4f}' |
| 96 | .format(unlabelled_acc, unlabelled_nmi, unlabelled_ari)) |
| 97 | |
| 98 | # Also return ratio between labelled and unlabelled examples |
| 99 | return (unlabelled_acc, unlabelled_nmi, unlabelled_ari), mask.mean() |
| 100 | |
| 101 | else: |
| 102 | |
| 103 | labelled_acc, labelled_nmi, labelled_ari = cluster_acc(targets.astype(int)[mask], |
| 104 | preds.astype(int)[mask]), \ |
| 105 | nmi_score(targets[mask], preds[mask]), \ |
| 106 | ari_score(targets[mask], preds[mask]) |
| 107 | |
| 108 | unlabelled_acc, unlabelled_nmi, unlabelled_ari = cluster_acc(targets.astype(int)[~mask], |
| 109 | preds.astype(int)[~mask]), \ |
| 110 | nmi_score(targets[~mask], preds[~mask]), \ |
| 111 | ari_score(targets[~mask], preds[~mask]) |
| 112 | |
| 113 | # Also return ratio between labelled and unlabelled examples |
| 114 | return (labelled_acc, labelled_nmi, labelled_ari), ( |
| 115 | unlabelled_acc, unlabelled_nmi, unlabelled_ari), mask.mean() |
| 116 | |
| 117 | |
| 118 | class AverageMeter(object): |
nothing calls this directly
no test coverage detected