Pimed / training_code /utils /metrics.py
deboraJ23's picture
upload training_code
64fd08f verified
Raw History Blame Contribute Delete
3.61 kB
import numpy as np
from typing import Dict, Any
from sklearn.metrics import balanced_accuracy_score, roc_auc_score, roc_curve, confusion_matrix
from sklearn.preprocessing import label_binarize
def compute_metrics_bc(labels: np.ndarray, preds: np.ndarray, probs: np.ndarray) -> Dict[str, Any]:
"""
Compute all relevant metrics for binary classification.
Args:
labels: Ground truth binary labels (0 or 1)
preds: Binary predictions (0 or 1)
probs: Prediction probabilities (0-1)
Returns:
Dictionary containing all calculated metrics
"""
metrics = {}
# Balanced accuracy
metrics['balanced_accuracy'] = balanced_accuracy_score(labels, preds)
# AUC
metrics['auc'] = roc_auc_score(labels, probs)
# Confusion matrix for sensitivity and specificity
cm = confusion_matrix(labels, preds)
metrics['confusion_matrix'] = cm
tn, fp, fn, tp = cm.ravel()
# Sensitivity and specificity
metrics['sensitivity'] = tp / (tp + fn) if (tp + fn) > 0 else 0
metrics['specificity'] = tn / (tn + fp) if (tn + fp) > 0 else 0
# ROC curve data
metrics['roc_curve'] = roc_curve(labels, probs)
return metrics
def compute_metrics_mc(labels: np.ndarray, preds: np.ndarray, probs: np.ndarray, num_classes: int = 3) -> Dict[str, Any]:
"""
Compute relevant metrics for multiclass classification (default: 3 classes).
Args:
labels: Ground-truth class indices with shape (N,), values in \{0, ..., num_classes-1\}
preds: Predicted class indices with shape (N,)
probs: Predicted class probabilities with shape (N, num_classes)
num_classes: Number of classes present in the problem (default 3)
Returns:
Dictionary mapping metric names to their computed values.
"""
metrics = {}
# Normalize probabilities to handle floating point precision errors from mixed precision training
probs = np.array(probs)
probs = probs / probs.sum(axis=1, keepdims=True)
labels = np.array(labels)
# Overall balanced accuracy (macro-averaged recall)
metrics["balanced_accuracy"] = balanced_accuracy_score(labels, preds)
# Macro-averaged ROC-AUC using the one-vs-rest strategy
labels_auc = label_binarize(labels, classes=list(range(num_classes)))
metrics["auc"] = roc_auc_score(labels_auc, probs, average="micro")
# Confusion matrix (shape: num_classes × num_classes)
cm = confusion_matrix(labels, preds, labels=list(range(num_classes)))
metrics["confusion_matrix"] = cm
# Per-class sensitivity (recall) and specificity
sensitivities = []
specificities = []
total = cm.sum()
for c in range(num_classes):
tp = cm[c, c]
fn = cm[c, :].sum() - tp
fp = cm[:, c].sum() - tp
tn = total - (tp + fp + fn)
sens = tp / (tp + fn) if (tp + fn) > 0 else 0.0
spec = tn / (tn + fp) if (tn + fp) > 0 else 0.0
sensitivities.append(sens)
specificities.append(spec)
metrics["sensitivity"] = np.array(sensitivities)
metrics["specificity"] = np.array(specificities)
# ROC curve per class (fpr, tpr, thresholds for each class)
roc_curves = []
for c in range(num_classes):
try:
fpr, tpr, thresholds = roc_curve((np.array(labels) == c).astype(int), np.array(probs)[:, c])
roc_curves.append((fpr, tpr, thresholds))
except ValueError:
roc_curves.append((np.array([]), np.array([]), np.array([])))
metrics["roc_curve"] = roc_curves
return metrics