| 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) |
| |
| |
| 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") |
| |
| |
| 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") |
| |
| |
| 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) |
| |
| |
| balanced_accuracy = balanced_accuracy_score(gts, preds) |
| auc = roc_auc_score(gts, probs) |
| |
| |
| cm = confusion_matrix(gts, preds) |
| tn, fp, fn, tp = cm.ravel() |
| |
| |
| 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 |
| |
| |
| fpr, tpr, _ = roc_curve(gts, probs[:, 1] if probs.ndim == 2 else probs) |
| |
| |
| fig, (ax1, ax2, ax3) = plt.subplots(1, 3, figsize=(15, 5)) |
| |
| |
| 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']) |
| |
| |
| 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") |
| |
| |
| 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) |
| |
| |
| 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') |
| |
| |
| 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. |
| """ |
| |
| |
| probs = np.array(probs) |
| probs = probs / probs.sum(axis=1, keepdims=True) |
| |
| preds = np.array(preds) |
| gts = np.array(gts) |
| |
| |
| balanced_accuracy = balanced_accuracy_score(gts, preds) |
| auc = roc_auc_score(gts, probs, multi_class='ovo', average='macro', labels=list(range(n_classes))) |
| |
| |
| cm = confusion_matrix(gts, preds, labels=list(range(n_classes))) |
| |
| |
| fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(12, 5)) |
| |
| |
| 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)]) |
| |
| |
| 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") |
| |
| |
| 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) |
| |
| |
| 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)) |
| |
| |
| 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)]) |
| |
| |
| 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 |
| |