def kl_divergence(p, q): return np.sum(p * np.log(p / q)) def js_divergence(p, q): m = 0.5 * (p + q) return 0.5 * (kl_divergence(p, m) + kl_divergence(q, m)) def visualize_preds(y, y_pred, model_name): df = pd.DataFrame({'label': y, 'pred_proba': y_pred}) # Compute ROCAUC metrics rocauc = roc_auc_score(df['label'], df['pred_proba']) fpr, tpr, thresholds = roc_curve(df['label'], df['pred_proba']) baseline = np.sum(df['label']) / len(df) # Compute PRAUC metrics prauc = average_precision_score(df['label'], df['pred_proba']) prec, rec, thresholds = precision_recall_curve(df['label'], df['pred_proba']) # Split into consistent and inconsistent for prob distribution inconsistent = df[df['label'] == 1].reset_index(drop=True) consistent = df[df['label'] == 0].reset_index(drop=True) js_div = js_divergence(inconsistent['pred_proba'], consistent['pred_proba']) # Set up plots fig, (ax0, ax1, ax2, ax3) = plt.subplots(1, 4, figsize=(13, 3), tight_layout=True) title_font_size = 10 fig.suptitle(f'{model_name}', fontsize=title_font_size+2, y=1) # Plot ROC ax0.grid() ax0.plot(fpr, tpr, label='ROC') ax0.plot([0, 1], [0, 1], label='Random chance', linestyle='--', color='red') ax0.set_xlabel('False positive rate') ax0.set_ylabel('True positive rate') ax0.set_title(f'ROC AUC = {rocauc:.2f}', fontsize=title_font_size) ax0.legend() # Plot PRAUC ax1.grid() ax1.plot(rec, prec, label='PRAUC') ax1.axhline(y=baseline, label='Baseline', linestyle='--', color='red') ax1.set_xlabel('Recall') ax1.set_ylabel('Precision') ax1.set_xlim((-0.1, 1.1)) ax1.set_ylim((-0.1, 1.1)) ax1.set_title(f'PR AUC = {prauc:.2f}', fontsize=title_font_size) # Plot Precision & Recall ax2.grid() ax2.plot(thresholds, prec[1:], color='red', label='Precision') ax2.plot(thresholds, rec[1:], color='blue', label='Recall') ax2.invert_xaxis() ax2.set_xlabel('Thresholds (1.0 - 0.0)') ax2.set_ylabel('Precision / Recall') ax2.set_xlim((1.1, -0.1)) ax2.set_ylim((-0.1, 1.1)) ax2.legend() ax2.set_title(f'PR AUC = {prauc:.2f}', fontsize=title_font_size) # Plot prob distribution ax3.grid() ax3.hist(inconsistent['pred_proba'], color='red', alpha=0.5, density=True, label='Inconsistent', bins=max(int(inconsistent['pred_proba'].nunique()/20), 20)) ax3.hist(consistent['pred_proba'], color='green', alpha=0.5, density=True, label='Consistent', bins=max(int(inconsistent['pred_proba'].nunique()/20), 20)) ax3.set_xlabel('Prob of inconsistent') ax3.set_ylabel('Density') ax3.set_title(f'JS Divergence = {js_div:.3f}', fontsize=title_font_size) ax3.legend() plt.show()