287 lines
10 KiB
Python
287 lines
10 KiB
Python
#!/usr/bin/env python3
|
||
"""
|
||
Visualise aggregated GradCAM heatmaps produced by aggregate_gradcam.py.
|
||
|
||
Produces two figures:
|
||
|
||
Figure 1 — Mean heatmaps grid
|
||
Rows: classes (e.g. Normal, Glaucoma)
|
||
Cols: OD_all | OS_all | OD_correct | OD_incorrect | OS_correct | OS_incorrect
|
||
|
||
Figure 2 — Attention stats
|
||
Panel A: disc_frac distribution per class (violin/box), OD and OS side by side
|
||
Panel B: entropy distribution per class
|
||
Panel C: disc_frac correct vs incorrect per class (scatter means + error bars)
|
||
|
||
Figure 3 — Disc attention vs correct confidence
|
||
Scatter of disc_frac vs correct_conf (confidence if correct, 1-confidence if wrong)
|
||
One panel per class, OD and OS overlaid, Pearson r annotated
|
||
|
||
Usage
|
||
-----
|
||
python scripts/output_analysis/explainability/plot_gradcam_aggregate.py \
|
||
--agg-dir analysis_data/pipeline_nocrop/binary/single/gradcam_aggregate
|
||
"""
|
||
from __future__ import annotations
|
||
|
||
import argparse
|
||
from pathlib import Path
|
||
|
||
import matplotlib
|
||
matplotlib.use("Agg")
|
||
import matplotlib.pyplot as plt
|
||
import matplotlib.patches as mpatches
|
||
import numpy as np
|
||
import pandas as pd
|
||
|
||
|
||
def load_agg(agg_dir: Path):
|
||
npz = np.load(agg_dir / "mean_heatmaps.npz")
|
||
stats = pd.read_csv(agg_dir / "attention_stats.csv")
|
||
return npz, stats
|
||
|
||
|
||
def _classes_from_npz(npz) -> list[str]:
|
||
classes = []
|
||
for key in npz.files:
|
||
parts = key.split("_")
|
||
# key format: {EYE}_{ClassName}_{split}_{stat}
|
||
# ClassName may be multi-word (e.g. "Glaucoma", "Normal", "Suspect")
|
||
if parts[-1] == "mean" and parts[-2] == "all" and parts[0] == "OD":
|
||
classes.append(parts[1])
|
||
return sorted(set(classes))
|
||
|
||
|
||
def plot_mean_heatmaps(npz, classes: list[str], out_path: Path):
|
||
eyes = ["OD", "OS"]
|
||
splits = ["all", "correct", "incorrect"]
|
||
cols = [(e, s) for e in eyes for s in splits] # 6 columns
|
||
|
||
n_rows = len(classes)
|
||
n_cols = len(cols)
|
||
fig, axes = plt.subplots(n_rows, n_cols, figsize=(n_cols * 2.8, n_rows * 2.8))
|
||
if n_rows == 1:
|
||
axes = axes[np.newaxis, :]
|
||
|
||
for r, cls in enumerate(classes):
|
||
for c, (eye, split) in enumerate(cols):
|
||
ax = axes[r, c]
|
||
key = f"{eye}_{cls}_{split}_mean"
|
||
if key not in npz:
|
||
ax.axis("off")
|
||
ax.set_title(f"{eye} {split}\n(no data)", fontsize=7)
|
||
continue
|
||
cam = npz[key]
|
||
count = int(npz.get(f"{eye}_{cls}_{split}_count", np.array(0)))
|
||
ax.imshow(cam, cmap="jet", vmin=0, vmax=1)
|
||
ax.axis("off")
|
||
title = f"{cls} | {eye} {split}\n(N={count})"
|
||
ax.set_title(title, fontsize=7)
|
||
|
||
fig.suptitle("Mean GradCAM heatmaps by class / eye / outcome", fontsize=12)
|
||
fig.tight_layout()
|
||
fig.savefig(out_path, dpi=150, bbox_inches="tight")
|
||
plt.close(fig)
|
||
print(f"Saved → {out_path}")
|
||
|
||
|
||
def plot_attention_stats(stats: pd.DataFrame, classes: list[str], out_path: Path):
|
||
eyes = ["OD", "OS"]
|
||
cmap = plt.get_cmap("tab10")
|
||
class_colors = {cls: cmap(i) for i, cls in enumerate(classes)}
|
||
|
||
fig, axes = plt.subplots(1, 3, figsize=(16, 5))
|
||
|
||
# ---- Panel A: disc_frac per class × eye ----
|
||
ax = axes[0]
|
||
positions = []
|
||
labels = []
|
||
data_viol = []
|
||
tick_pos = []
|
||
pos = 0
|
||
for cls in classes:
|
||
for eye in eyes:
|
||
sub = stats[(stats["true_name"] == cls) & (stats["eye"] == eye)]["disc_frac"].dropna()
|
||
data_viol.append(sub.values)
|
||
positions.append(pos)
|
||
labels.append(f"{cls[:3]}\n{eye}")
|
||
tick_pos.append(pos)
|
||
pos += 1
|
||
pos += 0.5 # gap between classes
|
||
|
||
vp = ax.violinplot(data_viol, positions=positions, showmedians=True, widths=0.7)
|
||
for i, (pc, cls) in enumerate(zip(vp["bodies"], [c for c in classes for _ in eyes])):
|
||
pc.set_facecolor(class_colors[cls])
|
||
pc.set_alpha(0.65)
|
||
ax.set_xticks(tick_pos)
|
||
ax.set_xticklabels(labels, fontsize=8)
|
||
ax.set_ylabel("Disc fraction (attention mass within GT disc mask)")
|
||
ax.set_title("Disc attention by class")
|
||
ax.axhline(0.5, color="black", linewidth=0.8, linestyle="--", alpha=0.4)
|
||
ax.grid(axis="y", linestyle="--", alpha=0.3)
|
||
|
||
# ---- Panel B: entropy per class × eye (same layout) ----
|
||
ax = axes[1]
|
||
data_ent = []
|
||
for cls in classes:
|
||
for eye in eyes:
|
||
sub = stats[(stats["true_name"] == cls) & (stats["eye"] == eye)]["entropy"].dropna()
|
||
data_ent.append(sub.values)
|
||
|
||
vp2 = ax.violinplot(data_ent, positions=positions, showmedians=True, widths=0.7)
|
||
for pc, cls in zip(vp2["bodies"], [c for c in classes for _ in eyes]):
|
||
pc.set_facecolor(class_colors[cls])
|
||
pc.set_alpha(0.65)
|
||
ax.set_xticks(tick_pos)
|
||
ax.set_xticklabels(labels, fontsize=8)
|
||
ax.set_ylabel("Attention entropy (higher = more diffuse)")
|
||
ax.set_title("Attention entropy by class")
|
||
ax.grid(axis="y", linestyle="--", alpha=0.3)
|
||
|
||
# ---- Panel C: disc_frac correct vs incorrect, mean ± std ----
|
||
ax = axes[2]
|
||
x_ticks = []
|
||
x_labels = []
|
||
pos = 0
|
||
for cls in classes:
|
||
for eye in eyes:
|
||
for split, marker, ls in [("correct", "o", "-"), ("incorrect", "X", "--")]:
|
||
sub = stats[
|
||
(stats["true_name"] == cls) &
|
||
(stats["eye"] == eye) &
|
||
(stats["correct"] == (split == "correct"))
|
||
]["disc_frac"].dropna()
|
||
if len(sub) == 0:
|
||
continue
|
||
ax.errorbar(
|
||
pos, sub.mean(), yerr=sub.std(),
|
||
fmt=marker, color=class_colors[cls], linestyle=ls,
|
||
capsize=4, markersize=7, alpha=0.85,
|
||
label=f"{cls[:3]} {eye} {split}" if pos < 4 else "_",
|
||
)
|
||
pos += 1
|
||
x_ticks.append(pos - 1.5)
|
||
x_labels.append(f"{cls[:3]}\n{eye}")
|
||
pos += 0.5
|
||
|
||
ax.axhline(0.5, color="black", linewidth=0.8, linestyle="--", alpha=0.4)
|
||
ax.set_ylabel("Disc fraction")
|
||
ax.set_title("Disc fraction: correct vs incorrect\n(circle=correct, X=incorrect)")
|
||
ax.grid(axis="y", linestyle="--", alpha=0.3)
|
||
|
||
# legend: one patch per class
|
||
patches = [mpatches.Patch(color=class_colors[c], label=c) for c in classes]
|
||
patches += [
|
||
plt.Line2D([0], [0], marker="o", color="grey", label="correct", linestyle="none"),
|
||
plt.Line2D([0], [0], marker="X", color="grey", label="incorrect", linestyle="none"),
|
||
]
|
||
ax.legend(handles=patches, fontsize=7, loc="lower right")
|
||
|
||
fig.suptitle("GradCAM attention statistics", fontsize=12)
|
||
fig.tight_layout()
|
||
fig.savefig(out_path, dpi=150, bbox_inches="tight")
|
||
plt.close(fig)
|
||
print(f"Saved → {out_path}")
|
||
|
||
|
||
def plot_disc_attention_correlation(stats: pd.DataFrame, classes: list[str], out_path: Path):
|
||
"""
|
||
Scatter disc_frac vs correct_conf per class.
|
||
|
||
correct_conf = confidence if correct
|
||
= 1 - confidence if incorrect
|
||
|
||
This asks: does focusing attention on the disc region correlate with
|
||
the model being more confident about the right answer?
|
||
"""
|
||
import scipy.stats as scipy_stats
|
||
|
||
stats = stats.copy()
|
||
stats["correct_conf"] = np.where(
|
||
stats["correct"],
|
||
stats["confidence"],
|
||
1.0 - stats["confidence"],
|
||
)
|
||
|
||
corr_colors = {True: "steelblue", False: "tomato"}
|
||
corr_labels = {True: "Correct", False: "Incorrect"}
|
||
is_binary = len(classes) == 2
|
||
|
||
n_cls = len(classes)
|
||
fig, axes = plt.subplots(1, n_cls, figsize=(5 * n_cls, 5), sharey=True)
|
||
if n_cls == 1:
|
||
axes = [axes]
|
||
|
||
for ax, cls in zip(axes, classes):
|
||
sub = stats[stats["true_name"] == cls]
|
||
x_all, y_all = [], []
|
||
|
||
for correct_val, color in corr_colors.items():
|
||
csub = sub[sub["correct"] == correct_val]
|
||
x = csub["disc_frac"].values
|
||
y = csub["correct_conf"].values
|
||
ax.scatter(x, y, marker="o", color=color,
|
||
alpha=0.75, s=30,
|
||
label=corr_labels[correct_val],
|
||
edgecolors="none")
|
||
x_all.extend(x.tolist())
|
||
y_all.extend(y.tolist())
|
||
|
||
# pooled regression line
|
||
x_arr = np.array(x_all)
|
||
y_arr = np.array(y_all)
|
||
if len(x_arr) >= 3:
|
||
r, p = scipy_stats.pearsonr(x_arr, y_arr)
|
||
m, b = np.polyfit(x_arr, y_arr, 1)
|
||
xs = np.linspace(0, 1, 100)
|
||
ax.plot(xs, m * xs + b, color="black", linewidth=1.5, linestyle="--", alpha=0.7)
|
||
p_str = f"p={p:.3f}" if p >= 0.001 else "p<0.001"
|
||
ax.annotate(f"r={r:+.3f}\n{p_str}", xy=(0.05, 0.93), xycoords="axes fraction",
|
||
fontsize=9, va="top",
|
||
bbox=dict(boxstyle="round,pad=0.3", facecolor="white", alpha=0.7))
|
||
|
||
if is_binary:
|
||
ax.axhline(0.5, color="red", linewidth=1.0, linestyle=":",
|
||
alpha=0.7, label="Decision boundary (0.50)")
|
||
ax.set_xlim(0, 1)
|
||
ax.set_xlabel("Disc fraction\n(attention mass within GT disc mask)", fontsize=9)
|
||
ax.set_title(cls, fontsize=11)
|
||
ax.set_ylim(-0.02, 1.05)
|
||
ax.grid(linestyle="--", alpha=0.3)
|
||
ax.legend(fontsize=8, loc="lower right")
|
||
|
||
axes[0].set_ylabel("Correct-class confidence\n(conf if correct, 1−conf if wrong)", fontsize=9)
|
||
fig.suptitle("Disc attention vs correct-class confidence", fontsize=12)
|
||
fig.tight_layout()
|
||
fig.savefig(out_path, dpi=150, bbox_inches="tight")
|
||
plt.close(fig)
|
||
print(f"Saved → {out_path}")
|
||
|
||
|
||
def main():
|
||
ap = argparse.ArgumentParser()
|
||
ap.add_argument("--agg-dir", required=True,
|
||
help="Directory produced by aggregate_gradcam.py")
|
||
ap.add_argument("--out-heatmaps", default=None)
|
||
ap.add_argument("--out-stats", default=None)
|
||
ap.add_argument("--out-corr", default=None)
|
||
args = ap.parse_args()
|
||
|
||
agg_dir = Path(args.agg_dir)
|
||
out_hm = Path(args.out_heatmaps) if args.out_heatmaps else agg_dir / "mean_heatmaps_plot.png"
|
||
out_st = Path(args.out_stats) if args.out_stats else agg_dir / "attention_stats_plot.png"
|
||
out_corr = Path(args.out_corr) if args.out_corr else agg_dir / "disc_attention_correlation.png"
|
||
|
||
npz, stats = load_agg(agg_dir)
|
||
classes = _classes_from_npz(npz)
|
||
print(f"Classes found: {classes}")
|
||
print(f"Total eye records in stats: {len(stats)}")
|
||
|
||
plot_mean_heatmaps(npz, classes, out_hm)
|
||
plot_attention_stats(stats, classes, out_st)
|
||
plot_disc_attention_correlation(stats, classes, out_corr)
|
||
|
||
|
||
if __name__ == "__main__":
|
||
main()
|