Pimed / training_code /utils /plots.py
deboraJ23's picture
upload training_code
64fd08f verified
Raw
History Blame Contribute Delete
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