Harshtech1/CARZero-Replication
0
1"""2generate_visualizations.py — Publication-quality plots for CARZero evaluation.3Generates: ROC curves, confusion matrices, score distributions, architecture4diagram, and performance comparison chart.5"""6 7import os8import numpy as np9import pandas as pd10import matplotlib11matplotlib.use('Agg')12import matplotlib.pyplot as plt13import matplotlib.patches as mpatches14from matplotlib.patches import FancyBboxPatch, FancyArrowPatch15from sklearn.metrics import roc_curve, roc_auc_score, confusion_matrix, f1_score16import seaborn as sns17 18OUT_DIR = "visualizations"19os.makedirs(OUT_DIR, exist_ok=True)20 21# ── Style ────────────────────────────────────────────────────────────────────22plt.rcParams.update({23 'font.family': 'sans-serif', 'font.size': 12,24 'axes.spines.top': False, 'axes.spines.right': False,25 'figure.dpi': 150, 'savefig.bbox': 'tight',26})27COLORS = ['#2196F3', '#FF5722', '#4CAF50', '#9C27B0']28 29DISEASE_SYNONYMS = {30 "cardiomegaly": ["cardiomegaly", "cardiac enlargement", "enlarged heart",31 "cardiomegaly/borderline"],32 "pleural effusion": ["pleural effusion", "effusion", "pleural fluid"],33 "pneumonia": ["pneumonia", "pneumonitis", "consolidation",34 "pneumonia/bacterial", "pneumonia/viral"],35 "pneumothorax": ["pneumothorax"],36}37DISEASES = ["cardiomegaly", "pleural effusion", "pneumonia", "pneumothorax"]38 39def has_disease(text, disease):40 t = str(text).lower()41 return int(any(s in t for s in DISEASE_SYNONYMS[disease]))42 43def load_data():44 preds = pd.read_csv("zero_shot_results.csv")45 gt = pd.read_csv("data/carzero_cleaned_reports.csv")46 labels = {}47 for d in DISEASES:48 labels[d] = gt["Problems"].apply(lambda x: has_disease(x, d)).values49 return preds, labels50 51 52# ═══════════════════════════════════════════════════════════════════════════53# PLOT 1: ROC Curves (all diseases on one figure)54# ═══════════════════════════════════════════════════════════════════════════55def plot_roc_curves(preds, labels):56 fig, ax = plt.subplots(figsize=(8, 7))57 for i, d in enumerate(DISEASES):58 fpr, tpr, _ = roc_curve(labels[d], preds[d].values)59 auc = roc_auc_score(labels[d], preds[d].values)60 ax.plot(fpr, tpr, color=COLORS[i], lw=2.5,61 label=f'{d.title()} (AUROC = {auc:.4f})')62 ax.plot([0, 1], [0, 1], 'k--', lw=1, alpha=0.4, label='Random (0.5)')63 ax.set_xlabel('False Positive Rate', fontsize=13)64 ax.set_ylabel('True Positive Rate', fontsize=13)65 ax.set_title('CARZero Zero-Shot ROC Curves — Open-I Dataset', fontsize=14, fontweight='bold')66 ax.legend(loc='lower right', fontsize=11, framealpha=0.9)67 ax.set_xlim([-0.01, 1.01]); ax.set_ylim([-0.01, 1.01])68 ax.grid(True, alpha=0.2)69 fig.savefig(f"{OUT_DIR}/01_roc_curves.png"); plt.close(fig)70 print("✅ 01_roc_curves.png")71 72 73# ═══════════════════════════════════════════════════════════════════════════74# PLOT 2: Confusion Matrices (2x2 grid)75# ═══════════════════════════════════════════════════════════════════════════76def plot_confusion_matrices(preds, labels):77 fig, axes = plt.subplots(2, 2, figsize=(12, 10))78 for i, (d, ax) in enumerate(zip(DISEASES, axes.flat)):79 y_true = labels[d]80 y_scores = preds[d].values81 # Best threshold via Youden's J82 fpr, tpr, thresholds = roc_curve(y_true, y_scores)83 j = tpr - fpr84 best_t = thresholds[np.argmax(j)]85 y_pred = (y_scores >= best_t).astype(int)86 cm = confusion_matrix(y_true, y_pred)87 sns.heatmap(cm, annot=True, fmt='d', cmap='Blues', ax=ax,88 xticklabels=['Negative', 'Positive'],89 yticklabels=['Negative', 'Positive'],90 annot_kws={'size': 14})91 f1 = f1_score(y_true, y_pred, zero_division=0)92 ax.set_title(f'{d.title()}\n(F1={f1:.3f}, thresh={best_t:.2f})',93 fontsize=12, fontweight='bold')94 ax.set_ylabel('True Label'); ax.set_xlabel('Predicted Label')95 fig.suptitle('Confusion Matrices at Optimal Threshold (Youden\'s J)',96 fontsize=14, fontweight='bold', y=1.02)97 fig.tight_layout()98 fig.savefig(f"{OUT_DIR}/02_confusion_matrices.png"); plt.close(fig)99 print("✅ 02_confusion_matrices.png")100 101 102# ═══════════════════════════════════════════════════════════════════════════103# PLOT 3: Score Distributions (positive vs negative)104# ═══════════════════════════════════════════════════════════════════════════105def plot_score_distributions(preds, labels):106 fig, axes = plt.subplots(2, 2, figsize=(14, 10))107 for i, (d, ax) in enumerate(zip(DISEASES, axes.flat)):108 y_true = labels[d]109 scores = preds[d].values110 neg_scores = scores[y_true == 0]111 pos_scores = scores[y_true == 1]112 ax.hist(neg_scores, bins=50, alpha=0.6, color='#90CAF9', label=f'Negative (n={len(neg_scores)})', density=True)113 ax.hist(pos_scores, bins=50, alpha=0.7, color='#EF5350', label=f'Positive (n={len(pos_scores)})', density=True)114 ax.set_title(f'{d.title()}', fontsize=13, fontweight='bold')115 ax.set_xlabel('Similarity Score'); ax.set_ylabel('Density')116 ax.legend(fontsize=10); ax.grid(True, alpha=0.15)117 fig.suptitle('Score Distribution: Positive vs Negative Cases', fontsize=14, fontweight='bold')118 fig.tight_layout(); fig.savefig(f"{OUT_DIR}/03_score_distributions.png"); plt.close(fig)119 print("✅ 03_score_distributions.png")120 121 122# ═══════════════════════════════════════════════════════════════════════════123# PLOT 4: Performance Comparison (Ours vs Paper)124# ═══════════════════════════════════════════════════════════════════════════125def plot_performance_comparison(preds, labels):126 paper_auroc = [0.8730, 0.8560, 0.7210, 0.8920]127 our_auroc = [roc_auc_score(labels[d], preds[d].values) for d in DISEASES]128 x = np.arange(len(DISEASES))129 w = 0.35130 fig, ax = plt.subplots(figsize=(10, 6))131 bars1 = ax.bar(x - w/2, our_auroc, w, label='Our Implementation', color='#42A5F5', edgecolor='white')132 bars2 = ax.bar(x + w/2, paper_auroc, w, label='Paper (Reported)', color='#66BB6A', edgecolor='white')133 for bar, val in zip(bars1, our_auroc):134 ax.text(bar.get_x() + bar.get_width()/2, bar.get_height() + 0.008, f'{val:.3f}', ha='center', va='bottom', fontsize=10, fontweight='bold')135 for bar, val in zip(bars2, paper_auroc):136 ax.text(bar.get_x() + bar.get_width()/2, bar.get_height() + 0.008, f'{val:.3f}', ha='center', va='bottom', fontsize=10, fontweight='bold')137 ax.set_xticks(x)138 ax.set_xticklabels([d.title() for d in DISEASES], fontsize=11)139 ax.set_ylabel('AUROC', fontsize=13); ax.set_ylim(0, 1.05)140 ax.set_title('Zero-Shot AUROC: Our Implementation vs. Paper', fontsize=14, fontweight='bold')141 ax.legend(fontsize=12); ax.grid(True, axis='y', alpha=0.2)142 fig.tight_layout(); fig.savefig(f"{OUT_DIR}/04_performance_comparison.png"); plt.close(fig)143 print("✅ 04_performance_comparison.png")144 145 146# ═══════════════════════════════════════════════════════════════════════════147# PLOT 5: Architecture Diagram (Forward & Backward Path)148# ═══════════════════════════════════════════════════════════════════════════149def plot_architecture():150 fig, ax = plt.subplots(figsize=(16, 10))151 ax.set_xlim(0, 16); ax.set_ylim(0, 10)152 ax.axis('off')153 ax.set_title('CARZero Architecture — Forward & Backward Path', fontsize=16, fontweight='bold', pad=20)154 155 def draw_box(ax, x, y, w, h, text, color, fontsize=9):156 box = FancyBboxPatch((x, y), w, h, boxstyle="round,pad=0.15",157 facecolor=color, edgecolor='#333', lw=1.5, alpha=0.85)158 ax.add_patch(box)159 ax.text(x + w/2, y + h/2, text, ha='center', va='center',160 fontsize=fontsize, fontweight='bold', wrap=True)161 162 def draw_arrow(ax, x1, y1, x2, y2, color='#333', style='->', lw=1.5):163 ax.annotate('', xy=(x2, y2), xytext=(x1, y1),164 arrowprops=dict(arrowstyle=style, color=color, lw=lw))165 166 # ── Input Layer ──167 draw_box(ax, 0.5, 7.5, 2.5, 1.2, 'Chest X-Ray\n(224×224×3)', '#BBDEFB')168 draw_box(ax, 0.5, 1.5, 2.5, 1.2, 'Clinical Text\n"There is {d}."', '#C8E6C9')169 170 # ── Encoders ──171 draw_box(ax, 4, 7.5, 3, 1.2, 'Vision Encoder\n(ViT-B/16, timm)', '#90CAF9')172 draw_box(ax, 4, 1.5, 3, 1.2, 'Text Encoder\n(BioBERT)', '#A5D6A7')173 174 # ── Features ──175 draw_box(ax, 8, 8.2, 2.2, 0.6, 'img_global\n[B, 768]', '#E3F2FD', fontsize=8)176 draw_box(ax, 8, 7.2, 2.2, 0.6, 'img_local\n[B, 196, 768]', '#E3F2FD', fontsize=8)177 draw_box(ax, 8, 2.2, 2.2, 0.6, 'txt_global\n[B, 768]', '#E8F5E9', fontsize=8)178 draw_box(ax, 8, 1.2, 2.2, 0.6, 'txt_local\n[B, M, 768]', '#E8F5E9', fontsize=8)179 180 # ── SimR Module ──181 draw_box(ax, 11, 4, 3.5, 3.5, '', '#FFF3E0')182 ax.text(12.75, 7.2, 'SimR Alignment Engine', ha='center', fontsize=11, fontweight='bold', color='#E65100')183 draw_box(ax, 11.3, 5.8, 2.9, 0.8, 'TransformerDecoder\n(4 layers, 8 heads)', '#FFE0B2', fontsize=8)184 draw_box(ax, 11.3, 4.7, 2.9, 0.7, 'decoder_norm\n(LayerNorm)', '#FFE0B2', fontsize=8)185 draw_box(ax, 11.3, 4.0, 2.9, 0.5, 'MLP Head → logit', '#FFE0B2', fontsize=8)186 187 # ── Output ──188 draw_box(ax, 12, 1.5, 2.5, 1.2, 'Similarity\nScore', '#FFCDD2')189 190 # ── Forward arrows (blue) ──191 draw_arrow(ax, 3, 8.1, 4, 8.1, '#1565C0')192 draw_arrow(ax, 3, 2.1, 4, 2.1, '#2E7D32')193 draw_arrow(ax, 7, 8.1, 8, 8.5, '#1565C0')194 draw_arrow(ax, 7, 7.8, 8, 7.5, '#1565C0')195 draw_arrow(ax, 7, 2.4, 8, 2.5, '#2E7D32')196 draw_arrow(ax, 7, 1.8, 8, 1.5, '#2E7D32')197 # text_global → decoder (Q)198 draw_arrow(ax, 10.2, 2.5, 11.3, 5.8, '#2E7D32')199 # img_local → decoder (K/V)200 draw_arrow(ax, 10.2, 7.5, 11.3, 6.6, '#1565C0')201 # decoder → norm → mlp → output202 draw_arrow(ax, 12.75, 4.0, 13.25, 2.7, '#E65100')203 204 # ── Backward path (red dashed) ──205 ax.annotate('', xy=(4, 7.7), xytext=(11, 4.5),206 arrowprops=dict(arrowstyle='->', color='#C62828', lw=1.5, linestyle='dashed'))207 ax.annotate('', xy=(4, 1.7), xytext=(11, 4.2),208 arrowprops=dict(arrowstyle='->', color='#C62828', lw=1.5, linestyle='dashed'))209 210 # ── Legend ──211 fwd = mpatches.Patch(color='#1565C0', label='Forward Path')212 bwd = mpatches.Patch(color='#C62828', label='Backward Path (Gradients)')213 loss = mpatches.Patch(color='#E65100', label='SimR Alignment')214 ax.legend(handles=[fwd, bwd, loss], loc='lower left', fontsize=11, framealpha=0.9)215 216 # ── Loss annotation ──217 ax.text(13.25, 0.8, 'InfoNCE Loss\nL = L_t2i + L_i2t', ha='center',218 fontsize=10, fontstyle='italic', color='#C62828',219 bbox=dict(boxstyle='round', facecolor='#FFEBEE', alpha=0.8))220 221 fig.savefig(f"{OUT_DIR}/05_architecture_diagram.png"); plt.close(fig)222 print("✅ 05_architecture_diagram.png")223 224 225# ═══════════════════════════════════════════════════════════════════════════226# PLOT 6: Metrics Summary Table227# ═══════════════════════════════════════════════════════════════════════════228def plot_metrics_table(preds, labels):229 paper = [0.8730, 0.8560, 0.7210, 0.8920]230 rows = []231 for i, d in enumerate(DISEASES):232 y_true = labels[d]233 y_scores = preds[d].values234 auc = roc_auc_score(y_true, y_scores)235 fpr, tpr, thresh = roc_curve(y_true, y_scores)236 best_t = thresh[np.argmax(tpr - fpr)]237 y_pred = (y_scores >= best_t).astype(int)238 f1 = f1_score(y_true, y_pred, zero_division=0)239 n_pos = int(y_true.sum())240 rows.append([d.title(), f'{auc:.4f}', f'{paper[i]:.4f}',241 f'{auc - paper[i]:+.4f}', f'{f1:.4f}',242 f'{n_pos}', f'{len(y_true)}'])243 244 fig, ax = plt.subplots(figsize=(14, 4))245 ax.axis('off')246 cols = ['Pathology', 'Our AUROC', 'Paper AUROC', 'Gap', 'F1-Score', 'Positives', 'Total']247 table = ax.table(cellText=rows, colLabels=cols, loc='center', cellLoc='center')248 table.auto_set_font_size(False); table.set_fontsize(11)249 table.scale(1, 1.8)250 # Style header251 for j in range(len(cols)):252 table[0, j].set_facecolor('#1565C0')253 table[0, j].set_text_props(color='white', fontweight='bold')254 # Alternate row colors255 for i in range(len(rows)):256 color = '#E3F2FD' if i % 2 == 0 else 'white'257 for j in range(len(cols)):258 table[i+1, j].set_facecolor(color)259 fig.suptitle('CARZero Zero-Shot Classification — Full Metrics Summary',260 fontsize=14, fontweight='bold', y=0.95)261 fig.savefig(f"{OUT_DIR}/06_metrics_summary.png"); plt.close(fig)262 print("✅ 06_metrics_summary.png")263 264 265# ═══════════════════════════════════════════════════════════════════════════266# PLOT 7: Per-Disease AUROC Radar Chart267# ═══════════════════════════════════════════════════════════════════════════268def plot_radar(preds, labels):269 our = [roc_auc_score(labels[d], preds[d].values) for d in DISEASES]270 paper = [0.8730, 0.8560, 0.7210, 0.8920]271 cats = [d.title() for d in DISEASES]272 N = len(cats)273 angles = [n / float(N) * 2 * np.pi for n in range(N)]274 angles += angles[:1]275 our += our[:1]; paper += paper[:1]276 277 fig, ax = plt.subplots(figsize=(8, 8), subplot_kw=dict(polar=True))278 ax.plot(angles, our, 'o-', lw=2.5, color='#2196F3', label='Ours')279 ax.fill(angles, our, alpha=0.15, color='#2196F3')280 ax.plot(angles, paper, 's-', lw=2.5, color='#4CAF50', label='Paper')281 ax.fill(angles, paper, alpha=0.15, color='#4CAF50')282 ax.set_xticks(angles[:-1]); ax.set_xticklabels(cats, fontsize=11)283 ax.set_ylim(0, 1); ax.set_title('AUROC Radar — Ours vs Paper', fontsize=14, fontweight='bold', pad=20)284 ax.legend(loc='lower right', fontsize=12)285 fig.savefig(f"{OUT_DIR}/07_radar_chart.png"); plt.close(fig)286 print("✅ 07_radar_chart.png")287 288 289# ═══════════════════════════════════════════════════════════════════════════290# MAIN291# ═══════════════════════════════════════════════════════════════════════════292if __name__ == "__main__":293 print("🎨 Generating CARZero Visualization Suite...\n")294 preds, labels = load_data()295 296 plot_roc_curves(preds, labels)297 plot_confusion_matrices(preds, labels)298 plot_score_distributions(preds, labels)299 plot_performance_comparison(preds, labels)300 plot_architecture()301 plot_metrics_table(preds, labels)302 plot_radar(preds, labels)303 304 print(f"\n🏆 All 7 visualizations saved to '{OUT_DIR}/'")305 