232 lines
7.6 KiB
Python
232 lines
7.6 KiB
Python
"""Metric computation, calibration, and threshold/bias tuning for V2."""
|
|
from __future__ import annotations
|
|
|
|
from typing import Optional
|
|
|
|
import numpy as np
|
|
import torch
|
|
import torch.nn.functional as F
|
|
from sklearn.metrics import (
|
|
cohen_kappa_score,
|
|
f1_score,
|
|
matthews_corrcoef,
|
|
recall_score,
|
|
roc_auc_score,
|
|
)
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Loss
|
|
# ---------------------------------------------------------------------------
|
|
|
|
def focal_loss(
|
|
logits: torch.Tensor,
|
|
targets: torch.Tensor,
|
|
gamma: float = 0.0,
|
|
weight: Optional[torch.Tensor] = None,
|
|
reduction: str = "mean",
|
|
) -> torch.Tensor:
|
|
"""
|
|
Standard focal loss wrapper. When gamma=0 it reduces to cross entropy.
|
|
weight should be per-class weights (same semantics as CrossEntropyLoss).
|
|
"""
|
|
if gamma <= 0:
|
|
return F.cross_entropy(logits, targets, weight=weight, reduction=reduction)
|
|
|
|
log_probs = F.log_softmax(logits, dim=1)
|
|
probs = log_probs.exp()
|
|
|
|
targets = targets.long().view(-1, 1)
|
|
logpt = log_probs.gather(1, targets)
|
|
pt = probs.gather(1, targets)
|
|
|
|
focal_factor = (1.0 - pt).clamp_min(0.0) ** gamma
|
|
loss = -focal_factor * logpt
|
|
|
|
if weight is not None:
|
|
class_weight = weight.gather(0, targets.view(-1))
|
|
loss = loss * class_weight.view(-1, 1)
|
|
|
|
loss = loss.view(-1)
|
|
if reduction == "sum":
|
|
return loss.sum()
|
|
if reduction == "mean":
|
|
return loss.mean()
|
|
return loss
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Basic array scoring
|
|
# ---------------------------------------------------------------------------
|
|
|
|
def _score_arrays(y_true: np.ndarray, probs: np.ndarray, num_classes: int):
|
|
"""Returns (acc, auc, n)."""
|
|
if y_true.size == 0:
|
|
return float("nan"), float("nan"), 0
|
|
acc = float((probs.argmax(1) == y_true).mean())
|
|
try:
|
|
auc = (
|
|
float(roc_auc_score(y_true, probs[:, 1]))
|
|
if num_classes == 2
|
|
else float(roc_auc_score(y_true, probs, multi_class="ovr", average="macro"))
|
|
)
|
|
except Exception:
|
|
auc = float("nan")
|
|
return acc, auc, int(len(y_true))
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Calibration
|
|
# ---------------------------------------------------------------------------
|
|
|
|
def compute_ece(y_true: np.ndarray, probs: np.ndarray, n_bins: int = 10) -> float:
|
|
"""Expected Calibration Error: weighted mean of |confidence - accuracy| per bin."""
|
|
if y_true.size == 0:
|
|
return float("nan")
|
|
confidences = probs.max(axis=1)
|
|
predictions = probs.argmax(axis=1)
|
|
bin_edges = np.linspace(0.0, 1.0, n_bins + 1)
|
|
ece = 0.0
|
|
n = len(y_true)
|
|
for i, (lo, hi) in enumerate(zip(bin_edges[:-1], bin_edges[1:])):
|
|
mask = (confidences >= lo) & (
|
|
confidences <= hi if i == n_bins - 1 else confidences < hi
|
|
)
|
|
if not mask.any():
|
|
continue
|
|
bin_acc = float((predictions[mask] == y_true[mask]).mean())
|
|
bin_conf = float(confidences[mask].mean())
|
|
ece += float(mask.sum()) / n * abs(bin_conf - bin_acc)
|
|
return float(ece)
|
|
|
|
|
|
def compute_extended_metrics(
|
|
y_true: np.ndarray,
|
|
probs: np.ndarray,
|
|
num_classes: int,
|
|
n_bins: int = 10,
|
|
preds_override: Optional[np.ndarray] = None,
|
|
) -> dict:
|
|
nan = float("nan")
|
|
if y_true.size == 0:
|
|
return dict(
|
|
kappa=nan, mcc=nan, macro_f1=nan,
|
|
per_class_recall=np.full(num_classes, nan), ece=nan,
|
|
)
|
|
preds = preds_override if preds_override is not None else probs.argmax(axis=1)
|
|
try:
|
|
kappa = float(cohen_kappa_score(y_true, preds))
|
|
except Exception:
|
|
kappa = nan
|
|
try:
|
|
mcc = float(matthews_corrcoef(y_true, preds))
|
|
except Exception:
|
|
mcc = nan
|
|
try:
|
|
macro_f1 = float(f1_score(y_true, preds, average="macro", zero_division=0))
|
|
except Exception:
|
|
macro_f1 = nan
|
|
try:
|
|
pcr = recall_score(
|
|
y_true, preds, average=None,
|
|
labels=list(range(num_classes)), zero_division=0,
|
|
).astype(float)
|
|
except Exception:
|
|
pcr = np.full(num_classes, nan)
|
|
ece = compute_ece(y_true, probs, n_bins=n_bins)
|
|
return dict(kappa=kappa, mcc=mcc, macro_f1=macro_f1, per_class_recall=pcr, ece=ece)
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Threshold / bias tuning
|
|
# ---------------------------------------------------------------------------
|
|
|
|
def tune_binary_threshold(y_true: np.ndarray, p1: np.ndarray) -> float:
|
|
if y_true.size == 0:
|
|
return 0.5
|
|
grid = np.linspace(0.0, 1.0, 1001)
|
|
best_t, best_acc = 0.5, -1.0
|
|
for t in grid:
|
|
pred = (p1 >= t).astype(int)
|
|
acc = float((pred == y_true).mean())
|
|
if acc > best_acc or (acc == best_acc and abs(t - 0.5) < abs(best_t - 0.5)):
|
|
best_acc, best_t = acc, float(t)
|
|
return best_t
|
|
|
|
|
|
def multiclass_acc_with_bias(y_true: np.ndarray, probs: np.ndarray, bias: np.ndarray) -> float:
|
|
if y_true.size == 0:
|
|
return float("nan")
|
|
logits = np.log(np.clip(probs, 1e-8, 1.0)) + bias.reshape(1, -1)
|
|
return float((np.argmax(logits, axis=1) == y_true).mean())
|
|
|
|
|
|
def tune_multiclass_bias(y_true: np.ndarray, probs: np.ndarray, *, iters: int = 2) -> np.ndarray:
|
|
if y_true.size == 0 or probs.size == 0:
|
|
return np.zeros((0,), dtype=float)
|
|
c = probs.shape[1]
|
|
bias = np.zeros((c,), dtype=float)
|
|
grid = np.linspace(-1.0, 1.0, 41)
|
|
for _ in range(iters):
|
|
for k in range(c):
|
|
best_v = bias[k]
|
|
best_acc = multiclass_acc_with_bias(y_true, probs, bias)
|
|
old = bias[k]
|
|
for v in grid:
|
|
bias[k] = float(v)
|
|
acc = multiclass_acc_with_bias(y_true, probs, bias)
|
|
if acc > best_acc or (acc == best_acc and abs(v) < abs(best_v)):
|
|
best_acc, best_v = acc, float(v)
|
|
bias[k] = best_v
|
|
if np.isnan(best_acc):
|
|
bias[k] = old
|
|
return bias
|
|
|
|
|
|
def _svf(vec) -> Optional[str]:
|
|
"""Serialise a float vector to pipe-separated string, or None if empty."""
|
|
if vec is None:
|
|
return None
|
|
arr = np.asarray(vec, dtype=float)
|
|
if arr.size == 0:
|
|
return None
|
|
return "|".join(f"{float(v):.4f}" for v in arr.tolist())
|
|
|
|
|
|
def _tune_and_snap(
|
|
y: np.ndarray,
|
|
p: np.ndarray,
|
|
acc: float,
|
|
num_classes: int,
|
|
args,
|
|
n_bins: int,
|
|
) -> tuple[dict, float, Optional[np.ndarray], Optional[np.ndarray]]:
|
|
"""
|
|
Apply threshold/bias tuning and compute extended metrics.
|
|
Returns (snap_dict, tuned_auc, threshold, bias).
|
|
"""
|
|
thr = 0.5 if num_classes == 2 else float("nan")
|
|
bias = None
|
|
ext_preds = None
|
|
|
|
if args.tune_binary_threshold and num_classes == 2 and y.size > 0:
|
|
thr = tune_binary_threshold(y, p[:, 1])
|
|
ext_preds = (p[:, 1] >= thr).astype(int)
|
|
acc = float((ext_preds == y).mean())
|
|
elif args.tune_multiclass_bias and num_classes > 2 and y.size > 0:
|
|
bias = tune_multiclass_bias(y, p)
|
|
logits = np.log(np.clip(p, 1e-8, 1.0)) + bias.reshape(1, -1)
|
|
ext_preds = np.argmax(logits, axis=1)
|
|
acc = float((ext_preds == y).mean())
|
|
|
|
ext = compute_extended_metrics(y, p, num_classes, n_bins=n_bins, preds_override=ext_preds)
|
|
_, auc, n = _score_arrays(y, p, num_classes)
|
|
|
|
snap = dict(
|
|
auc=auc, acc=acc, n=n,
|
|
kappa=ext["kappa"], mcc=ext["mcc"], macro_f1=ext["macro_f1"],
|
|
per_class_recall=ext["per_class_recall"], ece=ext["ece"],
|
|
threshold=thr, bias=bias,
|
|
)
|
|
return snap, auc, thr, bias
|