File size: 3,608 Bytes
64fd08f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
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