Download training_code/utils/plots.py from ODELIA-AI/Pimed: direct link, hf CLI and curl.
- Browser
- Download file 10.3 kB
-
https://huggingface.co/ODELIA-AI/Pimed/resolve/main/training_code/utils/plots.py
- Command line
-
hf download hf://ODELIA-AI/Pimed/training_code/utils/plots.py
-
curl -L -o plots.py https://huggingface.co/ODELIA-AI/Pimed/resolve/main/training_code/utils/plots.py
10.3 kB
| import numpy as np | |
| import matplotlib.pyplot as plt | |
| from sklearn.metrics import balanced_accuracy_score, roc_auc_score, roc_curve, confusion_matrix, auc | |
| def moving_average(data, window_size=5): | |
| """ | |
| Compute moving average over specified window size. | |
| """ | |
| if len(data) < window_size: | |
| return data | |
| return np.convolve(data, np.ones(window_size)/window_size, mode='valid') | |
| def plot_training_progress_classification(all_train_loss, all_val_loss, all_train_acc, all_val_acc, | |
| all_train_auc, all_val_auc, save_path, window_size=5): | |
| """ | |
| Plot training progress with both raw epoch-by-epoch data and moving averages. | |
| Args: | |
| all_train_loss: List of training losses per epoch | |
| all_val_loss: List of validation losses per epoch | |
| all_train_acc: List of training accuracies per epoch | |
| all_val_acc: List of validation accuracies per epoch | |
| all_train_auc: List of training AUCs per epoch | |
| all_val_auc: List of validation AUCs per epoch | |
| save_path: Path to save the plot | |
| window_size: Window size for moving average (default: 5) | |
| """ | |
| plt.figure(figsize=(10, 5)) | |
| epochs = np.arange(1, len(all_train_loss) + 1) | |
| # Loss subplot | |
| plt.subplot(1, 3, 1) | |
| train_loss_line = plt.plot(epochs, all_train_loss, '--', alpha=0.6, label="Train Loss (raw)")[0] | |
| val_loss_line = plt.plot(epochs, all_val_loss, '--', alpha=0.6, label="Val Loss (raw)")[0] | |
| if len(all_train_loss) >= window_size: | |
| ma_epochs = np.arange(window_size, len(all_train_loss) + 1) | |
| plt.plot(ma_epochs, moving_average(all_train_loss, window_size), '-', | |
| color=train_loss_line.get_color(), label=f"Train Loss (MA-{window_size})") | |
| plt.plot(ma_epochs, moving_average(all_val_loss, window_size), '-', | |
| color=val_loss_line.get_color(), label=f"Val Loss (MA-{window_size})") | |
| plt.legend() | |
| plt.xlabel("Epoch") | |
| plt.ylabel("Loss") | |
| # Accuracy subplot | |
| plt.subplot(1, 3, 2) | |
| train_acc_line = plt.plot(epochs, all_train_acc, '--', alpha=0.6, label="Train Acc (raw)")[0] | |
| val_acc_line = plt.plot(epochs, all_val_acc, '--', alpha=0.6, label="Val Acc (raw)")[0] | |
| if len(all_train_acc) >= window_size: | |
| ma_epochs = np.arange(window_size, len(all_train_acc) + 1) | |
| plt.plot(ma_epochs, moving_average(all_train_acc, window_size), '-', | |
| color=train_acc_line.get_color(), label=f"Train Acc (MA-{window_size})") | |
| plt.plot(ma_epochs, moving_average(all_val_acc, window_size), '-', | |
| color=val_acc_line.get_color(), label=f"Val Acc (MA-{window_size})") | |
| plt.legend() | |
| plt.xlabel("Epoch") | |
| plt.ylabel("Accuracy") | |
| # AUC subplot | |
| plt.subplot(1, 3, 3) | |
| train_auc_line = plt.plot(epochs, all_train_auc, '--', alpha=0.6, label="Train AUC (raw)")[0] | |
| val_auc_line = plt.plot(epochs, all_val_auc, '--', alpha=0.6, label="Val AUC (raw)")[0] | |
| if len(all_train_auc) >= window_size: | |
| ma_epochs = np.arange(window_size, len(all_train_auc) + 1) | |
| plt.plot(ma_epochs, moving_average(all_train_auc, window_size), '-', | |
| color=train_auc_line.get_color(), label=f"Train AUC (MA-{window_size})") | |
| plt.plot(ma_epochs, moving_average(all_val_auc, window_size), '-', | |
| color=val_auc_line.get_color(), label=f"Val AUC (MA-{window_size})") | |
| plt.legend() | |
| plt.xlabel("Epoch") | |
| plt.ylabel("AUC") | |
| plt.tight_layout() | |
| plt.savefig(save_path) | |
| plt.close() | |
| return | |
| def plot_pred_summary_bc(preds, probs, gts, save_path): | |
| """ | |
| Plot predictions for binary classification with confusion matrix, performance metrics, and ROC curve. | |
| """ | |
| probs = np.array(probs) | |
| preds = np.array(preds) | |
| gts = np.array(gts) | |
| # Compute metrics | |
| balanced_accuracy = balanced_accuracy_score(gts, preds) | |
| auc = roc_auc_score(gts, probs) | |
| # Confusion matrix | |
| cm = confusion_matrix(gts, preds) | |
| tn, fp, fn, tp = cm.ravel() | |
| # Additional metrics | |
| sensitivity = tp / (tp + fn) if (tp + fn) > 0 else 0 | |
| specificity = tn / (tn + fp) if (tn + fp) > 0 else 0 | |
| ppv = tp / (tp + fp) if (tp + fp) > 0 else 0 | |
| npv = tn / (tn + fn) if (tn + fn) > 0 else 0 | |
| # ROC curve | |
| fpr, tpr, _ = roc_curve(gts, probs[:, 1] if probs.ndim == 2 else probs) | |
| # Create figure with 3 subplots | |
| fig, (ax1, ax2, ax3) = plt.subplots(1, 3, figsize=(15, 5)) | |
| # Confusion matrix | |
| im = ax1.imshow(cm, interpolation='nearest', cmap=plt.cm.Blues) | |
| ax1.set_title('Confusion Matrix') | |
| ax1.set_xlabel('Predicted') | |
| ax1.set_ylabel('Actual') | |
| ax1.set_xticks([0, 1]) | |
| ax1.set_yticks([0, 1]) | |
| ax1.set_xticklabels(['Negative', 'Positive']) | |
| ax1.set_yticklabels(['Negative', 'Positive']) | |
| # Add text annotations | |
| thresh = cm.max() / 2 | |
| for i in range(2): | |
| for j in range(2): | |
| ax1.text(j, i, format(cm[i, j], 'd'), | |
| ha="center", va="center", | |
| color="white" if cm[i, j] > thresh else "black") | |
| # Performance metrics bar chart | |
| metrics = ['Sensitivity', 'Specificity', 'PPV', 'NPV', 'Balanced Accuracy'] | |
| values = [sensitivity, specificity, ppv, npv, balanced_accuracy] | |
| bars = ax2.bar(metrics, values, color=['skyblue', 'lightcoral', 'lightgreen', 'gold', 'lightpink']) | |
| ax2.set_title('Performance Metrics') | |
| ax2.set_ylabel('Value') | |
| ax2.set_ylim(0, 1) | |
| # Add value labels on bars | |
| for bar, value in zip(bars, values): | |
| height = bar.get_height() | |
| ax2.text(bar.get_x() + bar.get_width()/2., height + 0.01, | |
| f'{value:.3f}', ha='center', va='bottom') | |
| # ROC curve | |
| ax3.plot(fpr, tpr, color='darkorange', lw=2, label=f'ROC curve (AUC = {auc:.3f})') | |
| ax3.plot([0, 1], [0, 1], color='navy', lw=2, linestyle='--') | |
| ax3.set_xlim([0.0, 1.0]) | |
| ax3.set_ylim([0.0, 1.05]) | |
| ax3.set_xlabel('False Positive Rate') | |
| ax3.set_ylabel('True Positive Rate') | |
| ax3.set_title('ROC Curve') | |
| ax3.legend(loc="lower right") | |
| ax3.grid(True) | |
| plt.tight_layout() | |
| plt.savefig(save_path, dpi=150, bbox_inches='tight') | |
| plt.close() | |
| return | |
| def plot_pred_summary_mc(preds, probs, gts, n_classes, save_path): | |
| """ | |
| Plot predictions for multiclass classification with confusion matrix and performance 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) | |
| preds = np.array(preds) | |
| gts = np.array(gts) | |
| # Compute metrics | |
| balanced_accuracy = balanced_accuracy_score(gts, preds) | |
| auc = roc_auc_score(gts, probs, multi_class='ovo', average='macro', labels=list(range(n_classes))) | |
| # Confusion matrix | |
| cm = confusion_matrix(gts, preds, labels=list(range(n_classes))) | |
| # Create figure with 2 subplots for multiclass | |
| fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(12, 5)) | |
| # Confusion matrix | |
| im = ax1.imshow(cm, interpolation='nearest', cmap=plt.cm.Blues) | |
| ax1.set_title('Confusion Matrix') | |
| ax1.set_xlabel('Predicted Class') | |
| ax1.set_ylabel('True Class') | |
| ax1.set_xticks(range(n_classes)) | |
| ax1.set_yticks(range(n_classes)) | |
| ax1.set_xticklabels([f'Class {i}' for i in range(n_classes)]) | |
| ax1.set_yticklabels([f'Class {i}' for i in range(n_classes)]) | |
| # Add text annotations | |
| thresh = cm.max() / 2 | |
| for i in range(n_classes): | |
| for j in range(n_classes): | |
| ax1.text(j, i, format(cm[i, j], 'd'), | |
| ha="center", va="center", | |
| color="white" if cm[i, j] > thresh else "black") | |
| # Performance metrics bar chart | |
| metrics = ['Balanced Accuracy', 'Macro AUC'] | |
| values = [balanced_accuracy, auc] | |
| bars = ax2.bar(metrics, values, color=['lightpink', 'lightblue']) | |
| ax2.set_title('Performance Metrics') | |
| ax2.set_ylabel('Value') | |
| ax2.set_ylim(0, 1) | |
| # Add value labels on bars | |
| for bar, value in zip(bars, values): | |
| height = bar.get_height() | |
| ax2.text(bar.get_x() + bar.get_width()/2., height + 0.01, | |
| f'{value:.3f}', ha='center', va='bottom') | |
| plt.tight_layout() | |
| plt.savefig(save_path, dpi=150, bbox_inches='tight') | |
| plt.close() | |
| return | |
| def plot_multiclass_roc_curve(roc_curve_data, save_path=None): | |
| """ | |
| Plot the ROC curve for multiclass classification. | |
| """ | |
| n_classes = len(roc_curve_data) | |
| plt.figure(figsize=(n_classes * 5, 5)) | |
| for c, (fpr, tpr, _) in enumerate(roc_curve_data): | |
| plt.subplot(1, n_classes, c + 1) | |
| plt.plot(fpr, tpr, label=f'ROC curve of class {c}') | |
| plt.plot([0, 1], [0, 1], color='navy', lw=2, linestyle='--') | |
| plt.xlim([0.0, 1.0]) | |
| plt.ylim([0.0, 1.0]) | |
| auc_val = auc(fpr, tpr) | |
| plt.legend([f"AUC = {auc_val:.2f}"]) | |
| plt.xlabel("False Positive Rate") | |
| plt.ylabel("True Positive Rate") | |
| plt.title(f"ROC Curve of Class {c}") | |
| plt.grid(True) | |
| plt.savefig(save_path, dpi=150, bbox_inches="tight") | |
| plt.close() | |
| return | |
| def plot_multiclass_confusion_matrix(cm, save_path=None): | |
| """ | |
| Plot the confusion matrix. | |
| """ | |
| n_classes = cm.shape[0] | |
| fig, ax1 = plt.subplots(1, 1, figsize=(n_classes * 3, n_classes * 3)) | |
| # Confusion matrix | |
| im = ax1.imshow(cm, interpolation='nearest', cmap=plt.cm.Blues) | |
| ax1.set_title('Confusion Matrix') | |
| ax1.set_xlabel('Predicted Class') | |
| ax1.set_ylabel('True Class') | |
| ax1.set_xticks(range(n_classes)) | |
| ax1.set_yticks(range(n_classes)) | |
| ax1.set_xticklabels([f'Class {i}' for i in range(n_classes)]) | |
| ax1.set_yticklabels([f'Class {i}' for i in range(n_classes)]) | |
| # Add text annotations | |
| thresh = cm.max() / 2 | |
| for i in range(n_classes): | |
| for j in range(n_classes): | |
| ax1.text(j, i, format(cm[i, j], 'd'), | |
| ha="center", va="center", | |
| color="white" if cm[i, j] > thresh else "black") | |
| plt.tight_layout() | |
| plt.savefig(save_path, dpi=150, bbox_inches='tight') | |
| plt.close() | |
| return | |