| from __future__ import annotations |
|
|
| from dataclasses import dataclass |
| from typing import Dict, Tuple |
|
|
| import numpy as np |
| import torch |
| import torch.nn.functional as F |
|
|
|
|
| def normalize_binary_prediction(preds: torch.Tensor, threshold: float = 0.5) -> Tuple[torch.Tensor, torch.Tensor | None]: |
| """Return binary predictions and optional changed-class score maps. |
| |
| Two-channel tensors are treated as class logits/log-probabilities and are |
| converted with argmax. One-channel floating tensors outside [0, 1] are |
| treated as logits and passed through sigmoid. Integer tensors are treated as |
| labels. |
| """ |
| if preds.ndim == 4 and preds.shape[1] == 2: |
| pred = torch.argmax(preds, dim=1, keepdim=True).float() |
| scores = torch.softmax(preds.float(), dim=1)[:, 1:2] |
| return pred, scores |
| if preds.ndim == 3: |
| preds = preds.unsqueeze(1) |
| if preds.dtype.is_floating_point: |
| if float(preds.min()) < 0.0 or float(preds.max()) > 1.0: |
| preds = torch.sigmoid(preds) |
| return (preds >= threshold).float(), preds.float() |
| return (preds > 0).float(), None |
|
|
|
|
| def _as_binary(preds: torch.Tensor, threshold: float) -> torch.Tensor: |
| pred, _ = normalize_binary_prediction(preds, threshold=threshold) |
| return pred |
|
|
|
|
| def _target(mask: torch.Tensor) -> torch.Tensor: |
| if mask.ndim == 3: |
| mask = mask.unsqueeze(1) |
| return (mask > 0).float() |
|
|
|
|
| @dataclass |
| class BinaryMetrics: |
| threshold: float = 0.5 |
| tp: int = 0 |
| fp: int = 0 |
| fn: int = 0 |
| tn: int = 0 |
|
|
| def update(self, preds: torch.Tensor, targets: torch.Tensor) -> None: |
| pred = _as_binary(preds.detach().cpu(), self.threshold).bool() |
| tgt = _target(targets.detach().cpu()).bool() |
| self.tp += int((pred & tgt).sum().item()) |
| self.fp += int((pred & ~tgt).sum().item()) |
| self.fn += int((~pred & tgt).sum().item()) |
| self.tn += int((~pred & ~tgt).sum().item()) |
|
|
| def compute(self) -> Dict[str, float]: |
| eps = 1e-8 |
| precision = self.tp / (self.tp + self.fp + eps) |
| recall = self.tp / (self.tp + self.fn + eps) |
| f1 = 2 * precision * recall / (precision + recall + eps) |
| iou = self.tp / (self.tp + self.fp + self.fn + eps) |
| oa = (self.tp + self.tn) / (self.tp + self.fp + self.fn + self.tn + eps) |
| total = self.tp + self.fp + self.fn + self.tn |
| pe = ( |
| ((self.tp + self.fp) * (self.tp + self.fn)) |
| + ((self.fn + self.tn) * (self.fp + self.tn)) |
| ) / ((total * total) + eps) |
| kappa = (oa - pe) / (1.0 - pe + eps) |
| precision_0 = self.tn / (self.tn + self.fn + eps) |
| recall_0 = self.tn / (self.tn + self.fp + eps) |
| f1_0 = 2 * precision_0 * recall_0 / (precision_0 + recall_0 + eps) |
| iou_0 = self.tn / (self.tn + self.fp + self.fn + eps) |
| miou = (iou + iou_0) / 2.0 |
| return { |
| "threshold": self.threshold, |
| "f1": f1, |
| "iou": iou, |
| "miou": miou, |
| "precision": precision, |
| "recall": recall, |
| "oa": oa, |
| "kappa": kappa, |
| "f1_0": f1_0, |
| "iou_0": iou_0, |
| "tp": float(self.tp), |
| "fp": float(self.fp), |
| "fn": float(self.fn), |
| "tn": float(self.tn), |
| } |
|
|
|
|
| def _boundary_map(mask: torch.Tensor, radius: int) -> torch.Tensor: |
| mask = _target(mask).float() |
| if radius < 1: |
| radius = 1 |
| k = 2 * radius + 1 |
| eroded = -F.max_pool2d(-mask, kernel_size=k, stride=1, padding=radius) |
| return (mask - eroded).clamp(min=0.0, max=1.0).bool() |
|
|
|
|
| @dataclass |
| class BoundaryMetrics: |
| tolerance: int = 2 |
| boundary_tp_pred: int = 0 |
| boundary_total_pred: int = 0 |
| boundary_tp_gt: int = 0 |
| boundary_total_gt: int = 0 |
|
|
| def update(self, preds: torch.Tensor, targets: torch.Tensor) -> None: |
| pred = _as_binary(preds.detach().cpu(), 0.5) |
| tgt = _target(targets.detach().cpu()) |
| pred_b = _boundary_map(pred, radius=1) |
| tgt_b = _boundary_map(tgt, radius=1) |
| tol = max(int(self.tolerance), 1) |
| k = 2 * tol + 1 |
| pred_match = F.max_pool2d(pred_b.float(), kernel_size=k, stride=1, padding=tol).bool() |
| tgt_match = F.max_pool2d(tgt_b.float(), kernel_size=k, stride=1, padding=tol).bool() |
| self.boundary_tp_pred += int((pred_b & tgt_match).sum().item()) |
| self.boundary_total_pred += int(pred_b.sum().item()) |
| self.boundary_tp_gt += int((tgt_b & pred_match).sum().item()) |
| self.boundary_total_gt += int(tgt_b.sum().item()) |
|
|
| def compute(self) -> Dict[str, float]: |
| eps = 1e-8 |
| precision = self.boundary_tp_pred / (self.boundary_total_pred + eps) |
| recall = self.boundary_tp_gt / (self.boundary_total_gt + eps) |
| bf1 = 2 * precision * recall / (precision + recall + eps) |
| return { |
| "bf1": bf1, |
| "boundary_precision": precision, |
| "boundary_recall": recall, |
| "boundary_tolerance": int(self.tolerance), |
| } |
|
|
|
|
| def compute_binary_metrics(preds: torch.Tensor, targets: torch.Tensor, threshold: float = 0.5) -> Dict[str, float]: |
| metrics = BinaryMetrics(threshold=threshold) |
| metrics.update(preds, targets) |
| return metrics.compute() |
|
|
|
|
| def compute_binary_cd_metrics(pred, target, threshold: float = 0.5, logits: bool = False) -> Dict[str, float]: |
| if isinstance(pred, np.ndarray): |
| pred_t = torch.from_numpy(pred) |
| else: |
| pred_t = pred.detach().cpu() if hasattr(pred, "detach") else torch.as_tensor(pred) |
| if isinstance(target, np.ndarray): |
| target_t = torch.from_numpy(target) |
| else: |
| target_t = target.detach().cpu() if hasattr(target, "detach") else torch.as_tensor(target) |
| if logits and pred_t.dtype.is_floating_point: |
| pred_t = torch.sigmoid(pred_t) |
| return compute_binary_metrics(pred_t, target_t, threshold=threshold) |
|
|