CD-Models / utils /metrics.py
Dineth Perera
Publish tested dataset winners and benchmark rankings
ce209f5
Raw
History Blame Contribute Delete
5.92 kB
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)