Team Ai
Modelpublic

Harshtech1/CARZero-Replication

sourceHugging Faceupdated 4mo agoView on Hugging Face
0likes
generate_visualizations.py305 linesDownload Raw Back to root
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