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