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)