"""Calibration and selective-prediction metrics not provided by sklearn."""
from __future__ import annotations
import numpy as np
import pandas as pd
from sklearn.metrics import f1_score, matthews_corrcoef
_trapz = getattr(np, "trapezoid", None) or np.trapz
[docs]
def expected_calibration_error(
y_true: np.ndarray,
proba: np.ndarray,
n_bins: int = 15,
strategy: str = "uniform",
) -> tuple[float, pd.DataFrame]:
"""Top-label ECE plus the per-bin table used for reliability diagrams."""
confidence = proba.max(axis=1)
correct = (proba.argmax(axis=1) == np.asarray(y_true)).astype(float)
if strategy == "quantile":
edges = np.quantile(confidence, np.linspace(0.0, 1.0, n_bins + 1))
edges = np.unique(edges)
else:
edges = np.linspace(0.0, 1.0, n_bins + 1)
idx = np.clip(np.searchsorted(edges, confidence, side="right") - 1, 0, len(edges) - 2)
rows = []
ece = 0.0
n = len(confidence)
for b in range(len(edges) - 1):
mask = idx == b
count = int(mask.sum())
if count == 0:
continue
conf_b = float(confidence[mask].mean())
acc_b = float(correct[mask].mean())
ece += (count / n) * abs(acc_b - conf_b)
rows.append({"bin_lo": edges[b], "bin_hi": edges[b + 1],
"confidence": conf_b, "accuracy": acc_b, "count": count})
return float(ece), pd.DataFrame(rows)
[docs]
def multiclass_brier(y_true: np.ndarray, proba: np.ndarray) -> float:
"""Mean squared error between one-hot labels and predicted probabilities."""
y_true = np.asarray(y_true)
onehot = np.zeros_like(proba)
onehot[np.arange(len(y_true)), y_true] = 1.0
return float(np.mean(np.sum((proba - onehot) ** 2, axis=1)))
[docs]
def risk_coverage(
y_true: np.ndarray,
y_pred: np.ndarray,
confidence: np.ndarray,
) -> pd.DataFrame:
"""Selective-risk curve: abstain on the least confident samples first."""
order = np.argsort(-np.asarray(confidence), kind="stable")
errors = (np.asarray(y_true)[order] != np.asarray(y_pred)[order]).astype(float)
n = len(errors)
kept = np.arange(1, n + 1)
return pd.DataFrame({
"coverage": kept / n,
"risk": np.cumsum(errors) / kept,
"threshold": np.asarray(confidence)[order],
})
[docs]
def aurc(y_true: np.ndarray, y_pred: np.ndarray, confidence: np.ndarray) -> float:
"""Area under the risk-coverage curve (lower is better)."""
curve = risk_coverage(y_true, y_pred, confidence)
return float(_trapz(curve["risk"].to_numpy(), curve["coverage"].to_numpy()))
[docs]
def selective_metrics_at_coverage(
y_true: np.ndarray,
y_pred: np.ndarray,
confidence: np.ndarray,
coverage: float,
) -> dict[str, float]:
"""Accuracy/macro-F1/MCC on the retained fraction at a target coverage."""
n_keep = max(1, int(round(coverage * len(y_true))))
keep = np.argsort(-np.asarray(confidence), kind="stable")[:n_keep]
yt, yp = np.asarray(y_true)[keep], np.asarray(y_pred)[keep]
return {
"coverage": n_keep / len(y_true),
"selective_accuracy": float(np.mean(yt == yp)),
"selective_macro_f1": float(f1_score(yt, yp, average="macro", zero_division=0)),
"selective_mcc": float(matthews_corrcoef(yt, yp)) if len(np.unique(yt)) > 1 else float("nan"),
}