Download training_code/utils/metrics.py from ODELIA-AI/Pimed: direct link, hf CLI and curl.
- Browser
- Download file 3.61 kB
-
https://huggingface.co/ODELIA-AI/Pimed/resolve/main/training_code/utils/metrics.py
- Command line
-
hf download hf://ODELIA-AI/Pimed/training_code/utils/metrics.py
-
curl -L -o metrics.py https://huggingface.co/ODELIA-AI/Pimed/resolve/main/training_code/utils/metrics.py
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 | |