File size: 5,915 Bytes
ce209f5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
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)