CharlesCNorton
Image-level person classification on EUPE-ViT-B features with no free parameters
e8b8483 | """Binary classification metrics, single-sourced so every stage scores identically.""" | |
| from typing import NamedTuple | |
| import torch | |
| class Metrics(NamedTuple): | |
| """F1, precision, recall, and the threshold they were measured at.""" | |
| f1: float | |
| precision: float | |
| recall: float | |
| threshold: float = float('nan') | |
| def asdict(self) -> dict: | |
| d = {'F1': self.f1, 'precision': self.precision, 'recall': self.recall} | |
| if self.threshold == self.threshold: # excludes NaN | |
| d['threshold'] = self.threshold | |
| return d | |
| def prf1(pred: torch.Tensor, labels: torch.Tensor) -> Metrics: | |
| """Metrics for boolean prediction and label tensors.""" | |
| tp = (pred & labels).sum().float() | |
| fp = (pred & ~labels).sum().float() | |
| fn = (~pred & labels).sum().float() | |
| precision = tp / (tp + fp).clamp(min=1) | |
| recall = tp / (tp + fn).clamp(min=1) | |
| f1 = 2 * precision * recall / (precision + recall).clamp(min=1e-9) | |
| return Metrics(float(f1), float(precision), float(recall)) | |
| def f1_at(scores: torch.Tensor, labels: torch.Tensor, threshold: float) -> Metrics: | |
| """Metrics at a fixed threshold.""" | |
| return prf1(scores > threshold, labels)._replace(threshold=float(threshold)) | |
| def f1_sweep(scores: torch.Tensor, labels: torch.Tensor, n_candidates: int = 500) -> Metrics: | |
| """Best metrics over candidate thresholds drawn evenly from the sorted unique scores.""" | |
| uniq = torch.unique(scores).sort().values | |
| stride = max(1, len(uniq) // n_candidates) | |
| best = Metrics(0.0, 0.0, 0.0, 0.0) | |
| for t in uniq.tolist()[::stride]: | |
| m = prf1(scores > t, labels) | |
| if m.f1 > best.f1: | |
| best = m._replace(threshold=float(t)) | |
| return best | |