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