Files
hypertower/v4/figures/F8_explainability.py
T
rpotter6298 3d7777f010 Add new experiments and analysis scripts for dropzero features
- Introduced `fused_importance_with_axial.py` to evaluate the importance of Axial_Length in the fused-head model.
- Created JSON configurations for various experiments excluding zero-importance clinical features:
  - `cd_solo_bilateral_dropzero.json`: Bilateral clinical-only evaluation.
  - `cd_solo_single_dropzero.json`: Single-eye clinical-only evaluation.
  - `ensemble_refugelike_ckpt_dropzero.json`: Ensemble model with dropped zero-importance features.
  - `ensemble_single_refugelike_dropzero.json`: Single-eye ensemble model with dropped features.
2026-08-24 12:26:16 +02:00

1424 lines
59 KiB
Python

"""F8 - Explainability figures from the R50 checkpointed v4 run.
Sources predictions and Grad-CAM panels exclusively from the
``experiments/explainability/ensemble_refugelike_ckpt`` run (img+cd ensemble
with the refugelike R50 backbone, save_checkpoints=true). rep00 (seed=1234)
is the single rep used for the figure; matches the headline configuration in
section 3.
GradCAM machinery lives in ``v4.classes.accessory.explainability``; PAPILA
specific knowledge (disc contour rasterisation, OS→OD orientation flip) is
attached to ``ImageDataView`` in ``v4.classes.profiles.v4papila`` and consumed
here via getattr so future non-PAPILA profiles can opt in without changing
this file.
Usage
python -m v4.figures.F8_explainability --only-fusion
python -m v4.figures.F8_explainability --only-gradcam --run-gradcam
python -m v4.figures.F8_explainability --run-gradcam # both
"""
from __future__ import annotations
from pathlib import Path
import matplotlib
matplotlib.use("Agg")
import matplotlib.patches as mpatches
import matplotlib.pyplot as plt
import numpy as np
import pandas as pd
from PIL import Image
from sklearn.metrics import roc_auc_score
from v4.figures.util.loaders import REPO_ROOT
OUT_DIR = Path(__file__).parent / "output"
GRADCAM_DIR = OUT_DIR / "F8_gradcam"
V4_CKPT_RUN = REPO_ROOT / "v4" / "results" / "experiments" / "explainability" / "ensemble_refugelike_ckpt" / "rep00" / "binary"
LABEL_NAMES = {0: "Normal", 1: "Glaucoma"}
EVENT_ORDER = [
"full_correction", "img_assist", "md_assist",
"full_error", "img_drag", "md_drag",
"concordant_correct", "concordant_wrong",
]
EVENT_LABELS = {
"full_correction": "Both wrong -> fused right",
"img_assist": "Image right, clinical wrong",
"md_assist": "Clinical right, image wrong",
"full_error": "Both right -> fused wrong",
"img_drag": "Clinical right, image wrong -> fused wrong",
"md_drag": "Image right, clinical wrong -> fused wrong",
"concordant_correct": "All correct",
"concordant_wrong": "All wrong",
}
EVENT_COLORS = {
"full_correction": "#2f8f5b",
"img_assist": "#74b66b",
"md_assist": "#b7c85a",
"full_error": "#b23a48",
"img_drag": "#df7f5f",
"md_drag": "#d2a24c",
"concordant_correct": "#7da5c9",
"concordant_wrong": "#9b8fc2",
"net_positive": "#1f7a4d",
"net_negative": "#9f2735",
}
def _prediction_frame_from_probs(
*,
y_true: np.ndarray,
probs_fused: np.ndarray,
probs_img: np.ndarray,
probs_md: np.ndarray,
rep: str,
fold: str,
source: str,
) -> pd.DataFrame:
pred_fused = probs_fused.argmax(axis=1)
pred_img = probs_img.argmax(axis=1)
pred_md = probs_md.argmax(axis=1)
df = pd.DataFrame(
{
"idx": np.arange(len(y_true)),
"y_true": y_true.astype(int),
"pred_fused": pred_fused.astype(int),
"prob_fused_c0": probs_fused[:, 0],
"prob_fused_c1": probs_fused[:, 1],
"pred_img": pred_img.astype(int),
"prob_img_c0": probs_img[:, 0],
"prob_img_c1": probs_img[:, 1],
"pred_md": pred_md.astype(int),
"prob_md_c0": probs_md[:, 0],
"prob_md_c1": probs_md[:, 1],
"rep": rep,
"fold": fold,
"source": source,
}
)
for head in ("fused", "img", "md"):
pred = df[f"pred_{head}"].to_numpy()
y = df["y_true"].to_numpy()
df[f"tp_{head}"] = ((pred == 1) & (y == 1)).astype(int)
df[f"fp_{head}"] = ((pred == 1) & (y == 0)).astype(int)
df[f"tn_{head}"] = ((pred == 0) & (y == 0)).astype(int)
df[f"fn_{head}"] = ((pred == 0) & (y == 1)).astype(int)
return df
def _load_ckpt_predictions(split: str = "test") -> pd.DataFrame:
"""Reconstruct v4 component predictions from checkpointed fold modules.
predictions.h5 only stores the final hb head, so the image/clinical
component predictions needed for F8a are rebuilt from tower/head checkpoints.
"""
if split not in {"test", "val", "both"}:
raise ValueError(f"split must be 'test', 'val', or 'both', got {split!r}")
import importlib
import json
import torch
import torch.nn.functional as F
from v4.classes.split_manager import SplitManager
from v4.classes.v4_hypertower import load_data, build_towers, _make_loader
import v4.classes.v4_hypertower as orch
from v4.classes.heads.classifier import ClassificationHead
summary = json.loads((V4_CKPT_RUN / "summary.json").read_text())
cfg = summary["config"]
orch.cfg_ref = cfg
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
data = load_data(cfg)
label_filter = cfg.get("label_filter", None)
df_mode = data.df.copy()
if label_filter is not None:
df_mode = df_mode[df_mode[data.label_col].isin(label_filter)].reset_index(drop=True)
identity_cols = getattr(data, "identity_cols", [])
identity_level = cfg.get("split_identity_level", 1)
group_col = identity_cols[identity_level - 1] if identity_level and identity_cols else None
splits = SplitManager(group_col=group_col).build_plans(
df_mode,
label_col=data.label_col,
n_splits=cfg.get("folds", 5),
seed=cfg.get("fold_seed", 100),
)
stage_by_name = {s["name"]: s for s in cfg["stages"]}
nt_cfg = stage_by_name["nt"]
hb_cfg = stage_by_name["hb"]
rows: list[pd.DataFrame] = []
for fold_idx, split_obj in enumerate(splits):
ckpt_dir = V4_CKPT_RUN / "checkpoints" / f"fold{fold_idx}"
if not ckpt_dir.exists():
continue
towers = build_towers(cfg["towers"], data)
for name, tower in towers.items():
tower.load_state_dict(torch.load(ckpt_dir / f"tower_{name}.pt", map_location="cpu"))
tower.to(device).eval()
nt_mod = importlib.import_module(nt_cfg["module"])
nt = getattr(nt_mod, nt_cfg["class"])(
[towers[n].out_dim for n in nt_cfg["inputs"]],
**nt_cfg.get("args", {}),
).to(device)
nt.load_state_dict(torch.load(ckpt_dir / "stage_nt.pt", map_location="cpu"))
nt.eval()
hb_mod = importlib.import_module(hb_cfg["module"])
hb = getattr(hb_mod, hb_cfg["class"])(
{"a": nt.out_dim, "b": nt.out_dim},
**hb_cfg.get("args", {}),
).to(device)
hb.load_state_dict(torch.load(ckpt_dir / "stage_hb.pt", map_location="cpu"))
hb.eval()
img_aux = ClassificationHead(towers["img"].out_dim, cfg["num_classes"]).to(device)
cd_aux = ClassificationHead(towers["cd"].out_dim, cfg["num_classes"]).to(device)
hb_head = ClassificationHead(hb.out_dim, cfg["num_classes"], dropout=0.3).to(device)
img_aux.load_state_dict(torch.load(ckpt_dir / "stage_img_aux.pt", map_location="cpu"))
cd_aux.load_state_dict(torch.load(ckpt_dir / "stage_cd_aux.pt", map_location="cpu"))
hb_head.load_state_dict(torch.load(ckpt_dir / "stage_hb_head.pt", map_location="cpu"))
img_aux.eval(); cd_aux.eval(); hb_head.eval()
split_frames = []
if split in {"val", "both"}:
split_frames.append(("val", split_obj.val))
if split in {"test", "both"} and split_obj.test is not None:
split_frames.append(("test", split_obj.test))
for split_name, split_df in split_frames:
shell = data.build_shells(split_df, level="patient", label_filter=label_filter)
loader = _make_loader(
shell,
towers,
batch_size=cfg["training"].get("batch_size", 8),
shuffle=False,
)
y_all, pf_all, pi_all, pc_all, ids_all = [], [], [], [], []
with torch.no_grad():
for batch in loader:
y = batch["label"].detach().cpu().numpy()
img_a = towers["img"](batch["img"]["a"].to(device))
img_b = towers["img"](batch["img"]["b"].to(device))
cd_a = towers["cd"](batch["cd"]["a"].to(device))
cd_b = towers["cd"](batch["cd"]["b"].to(device))
z_a = nt([img_a, cd_a])
z_b = nt([img_b, cd_b])
z_hb = hb({"a": z_a, "b": z_b})
p_fused = F.softmax(hb_head(z_hb), dim=1).cpu().numpy()
p_img = 0.5 * (
F.softmax(img_aux(img_a), dim=1).cpu().numpy()
+ F.softmax(img_aux(img_b), dim=1).cpu().numpy()
)
p_cd = 0.5 * (
F.softmax(cd_aux(cd_a), dim=1).cpu().numpy()
+ F.softmax(cd_aux(cd_b), dim=1).cpu().numpy()
)
y_all.append(y)
pf_all.append(p_fused)
pi_all.append(p_img)
pc_all.append(p_cd)
ids_all.extend(batch.get("entity_id", []))
if y_all:
part = _prediction_frame_from_probs(
y_true=np.concatenate(y_all),
probs_fused=np.concatenate(pf_all, axis=0),
probs_img=np.concatenate(pi_all, axis=0),
probs_md=np.concatenate(pc_all, axis=0),
rep="v4",
fold=f"fold{fold_idx}",
source=split_name,
)
part["entity_id"] = [str(e) for e in ids_all]
rows.append(part)
del towers, nt, hb, img_aux, cd_aux, hb_head
if torch.cuda.is_available():
torch.cuda.empty_cache()
if not rows:
raise FileNotFoundError(f"No reconstructed predictions under {V4_CKPT_RUN}")
return pd.concat(rows, ignore_index=True)
def _classify_events(df: pd.DataFrame) -> pd.DataFrame:
df = df.copy()
y = df["y_true"].to_numpy()
fused_ok = df["pred_fused"].to_numpy() == y
img_ok = df["pred_img"].to_numpy() == y
md_ok = df["pred_md"].to_numpy() == y
def classify(fo: bool, io: bool, mo: bool) -> str:
if fo and io and mo:
return "concordant_correct"
if not fo and not io and not mo:
return "concordant_wrong"
if fo and not io and not mo:
return "full_correction"
if fo and io and not mo:
return "img_assist"
if fo and not io and mo:
return "md_assist"
if not fo and io and mo:
return "full_error"
if not fo and not io and mo:
return "img_drag"
if not fo and io and not mo:
return "md_drag"
return "other"
df["event_type"] = [classify(fo, io, mo) for fo, io, mo in zip(fused_ok, img_ok, md_ok)]
df["fused_ok"] = fused_ok
df["img_ok"] = img_ok
df["md_ok"] = md_ok
df["tower_state"] = np.select(
[
(~img_ok) & (~md_ok),
img_ok & (~md_ok),
(~img_ok) & md_ok,
img_ok & md_ok,
],
[
"Both networks wrong",
"Image only correct",
"Clinical only correct",
"Both networks correct",
],
default="Other",
)
df["tower_mean_c1"] = 0.5 * (df["prob_img_c1"] + df["prob_md_c1"])
df["fusion_delta_c1"] = df["prob_fused_c1"] - df["tower_mean_c1"]
df["fusion_margin"] = np.where(
df["y_true"] == 1,
df["prob_fused_c1"] - df["tower_mean_c1"],
(1.0 - df["prob_fused_c1"]) - (1.0 - df["tower_mean_c1"]),
)
df["tower_gap_abs"] = (df["prob_img_c1"] - df["prob_md_c1"]).abs()
return df
def make_fusion_event_panel(split: str = "test") -> None:
df = _classify_events(_load_ckpt_predictions(split=split))
auc = roc_auc_score(df["y_true"], df["prob_fused_c1"])
fold_groups = list(df.groupby(["rep", "fold"]))
counts = df["event_type"].value_counts().reindex(EVENT_ORDER, fill_value=0)
positive_keys = EVENT_ORDER[:3]
negative_keys = EVENT_ORDER[3:6]
total_positive = int(counts[positive_keys].sum())
total_negative = int(counts[negative_keys].sum())
fig = plt.figure(figsize=(15, 8.4))
gs = fig.add_gridspec(2, 2, height_ratios=[0.9, 1.45], width_ratios=[1.05, 1.0],
wspace=0.32, hspace=0.38)
ax = fig.add_subplot(gs[0, 0])
display_counts = counts.to_dict()
display_counts["net_positive"] = total_positive
display_counts["net_negative"] = total_negative
display_labels = {
**EVENT_LABELS,
"net_positive": "Net positive",
"net_negative": "Net negative",
}
display_order = [
"full_correction", "img_assist", "md_assist",
"full_error", "img_drag", "md_drag",
"net_positive", "net_negative",
]
bars = [k for k in display_order if display_counts.get(k, 0) > 0]
y = np.arange(len(bars))
ax.barh(y, [display_counts[k] for k in bars],
color=[EVENT_COLORS[k] for k in bars], height=0.68)
ax.set_yticks(y)
ax.set_yticklabels([display_labels[k] for k in bars], fontsize=8)
ax.invert_yaxis()
ax.set_xlabel("Count")
ax.set_title("Fusion event taxonomy", fontsize=10, fontweight="bold")
max_count = max(display_counts[k] for k in bars) if bars else 0
ax.set_xlim(0, max_count * 1.10 + 1)
for yi, k in enumerate(bars):
ax.text(display_counts[k] + 0.8, yi, str(int(display_counts[k])),
va="center", fontsize=8)
ax = fig.add_subplot(gs[0, 1])
per_fold = pd.DataFrame(
{
"fold": [str(f) for (_, f), _ in fold_groups],
"positive": [sum((g["event_type"] == k).sum() for k in positive_keys) for _, g in fold_groups],
"negative": [sum((g["event_type"] == k).sum() for k in negative_keys) for _, g in fold_groups],
}
)
x = np.arange(len(per_fold))
ax.bar(x, per_fold["positive"], color="#2f8f5b", width=0.72, label="positive correction")
ax.bar(x, -per_fold["negative"], color="#b23a48", width=0.72, label="negative correction")
ax.plot(x, per_fold["positive"] - per_fold["negative"], color="#222", lw=1.2,
marker="o", ms=3, label="net")
ax.axhline(0, color="#222", lw=0.8)
ax.set_xticks(x)
ax.set_xticklabels(per_fold["fold"], rotation=45, ha="right", fontsize=7)
ax.set_ylabel("Events per fold")
ax.set_title("Per-fold correction balance", fontsize=10, fontweight="bold")
ax.legend(fontsize=7, loc="upper right", frameon=False)
ax = fig.add_subplot(gs[1, :])
disagree = df[df["event_type"].isin(positive_keys + negative_keys)].copy()
concordant = df[df["event_type"].isin(["concordant_correct", "concordant_wrong"])].copy()
n_bins = 28
agreement_grid = np.zeros((n_bins, n_bins), dtype=float)
pos_grid = np.zeros((n_bins, n_bins), dtype=float)
neg_grid = np.zeros((n_bins, n_bins), dtype=float)
for _, row in concordant.iterrows():
xi = min(n_bins - 1, max(0, int(row["prob_img_c1"] * n_bins)))
yi = min(n_bins - 1, max(0, int(row["prob_md_c1"] * n_bins)))
agreement_grid[yi, xi] += 1
for _, row in disagree.iterrows():
xi = min(n_bins - 1, max(0, int(row["prob_img_c1"] * n_bins)))
yi = min(n_bins - 1, max(0, int(row["prob_md_c1"] * n_bins)))
if row["event_type"] in positive_keys:
pos_grid[yi, xi] += 1
else:
neg_grid[yi, xi] += 1
grey_rgba = np.zeros((n_bins, n_bins, 4), dtype=float)
grey_rgba[:, :, :3] = 0.52
if agreement_grid.max() > 0:
grey_rgba[:, :, 3] = 0.06 + 0.22 * np.sqrt(agreement_grid / agreement_grid.max())
grey_rgba[agreement_grid == 0, 3] = 0.0
ax.imshow(grey_rgba, extent=[0, 1, 0, 1], origin="lower", aspect="auto",
interpolation="nearest")
support = pos_grid + neg_grid
dominance = np.divide(pos_grid - neg_grid, support,
out=np.zeros_like(pos_grid), where=support > 0)
cmap = plt.get_cmap("RdYlGn")
rgba = cmap((dominance + 1.0) / 2.0)
if support.max() > 0:
rgba[:, :, 3] = 0.08 + 0.46 * np.sqrt(support / support.max())
rgba[support == 0, 3] = 0.0
ax.imshow(rgba, extent=[0, 1, 0, 1], origin="lower", aspect="auto",
interpolation="nearest")
point_handles = []
for correct, label, color in [
(True, "fused correct", "#2f8f5b"),
(False, "fused wrong", "#b23a48"),
]:
sub = df[df["fused_ok"] == correct]
h = ax.scatter(
sub["prob_img_c1"], sub["prob_md_c1"], s=18, alpha=0.45,
color=color, edgecolors="white", linewidths=0.18,
label=label,
)
point_handles.append(h)
ax.axvline(0.5, color="#222", lw=0.9, ls="--", alpha=0.75)
ax.axhline(0.5, color="#222", lw=0.9, ls="--", alpha=0.75)
ax.plot([0, 1], [0, 1], color="#222", lw=0.7, ls=":", alpha=0.55)
ax.set_xlim(0, 1)
ax.set_ylim(0, 1)
ax.set_xlabel("Image network P(glaucoma)")
ax.set_ylabel("Clinical network P(glaucoma)")
ax.set_title(
"Network confidence space - shaded by dominant empirical fusion outcome "
"(grey=agreement, green=positive correction, red=negative correction)",
fontsize=10, fontweight="bold",
)
ax.text(0.02, 0.96, "clinical says glaucoma", transform=ax.transAxes,
ha="left", va="top", fontsize=8, color="#333")
ax.text(0.98, 0.04, "image says glaucoma", transform=ax.transAxes,
ha="right", va="bottom", fontsize=8, color="#333")
shade_handles = [
mpatches.Patch(facecolor="#858585", alpha=0.28, label="agreement density"),
mpatches.Patch(facecolor="#2f8f5b", alpha=0.34, label="positive correction shade"),
mpatches.Patch(facecolor="#b23a48", alpha=0.34, label="negative correction shade"),
]
ax.legend(handles=point_handles + shade_handles, ncol=5, fontsize=7,
loc="upper center", bbox_to_anchor=(0.5, -0.14), frameon=False)
out = OUT_DIR / "S8a_comparison_panel.png"
fig.savefig(out, dpi=180, bbox_inches="tight")
plt.close(fig)
print(f"Saved: {out}")
def _drop_structural_features(
names: list[str],
groups: list[list[int]],
drop: set[str],
) -> tuple[list[str], list[list[int]]]:
"""Return (names, groups) with entries in ``drop`` removed.
Encoded-vector indices are left intact - the model still sees them - but
we simply do not permute/report those dims. Used to hide index-level
columns like eyeID that appear in the clinical feature list only because
they double as an index level.
"""
kept = [(n, g) for n, g in zip(names, groups) if n not in drop]
if not kept:
return list(names), list(groups)
kept_names, kept_groups = zip(*kept)
return list(kept_names), [list(g) for g in kept_groups]
def make_clinical_importance(n_permutations: int = 30, seed: int = 0) -> None:
"""S8e clinical permutation importance via the cd-tower → cd_aux head.
Isolates the clinical-only prediction path at the R50 ckpt run, then
column-shuffles the encoded clinical vector to measure per-feature AUC
drop. Per-original-column grouping comes from
``ClinicalDataView.feature_groups`` (one-hot encoded dims for a
categorical column shuffle together).
Aggregates per-fold mean drops, then averages across folds.
"""
import json
import torch
import torch.nn.functional as TF_
from tqdm import tqdm
from v4.classes.accessory.explainability import permutation_importance
from v4.classes.split_manager import SplitManager
from v4.classes.v4_hypertower import load_data, _make_loader
import v4.classes.v4_hypertower as orch
summary_path = V4_CKPT_RUN / "summary.json"
if not summary_path.exists():
print(f"[F8] make_clinical_importance: no summary.json under {V4_CKPT_RUN}")
return
cfg = json.loads(summary_path.read_text())["config"]
orch.cfg_ref = cfg
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
data = load_data(cfg)
clinical_view = data.matrix
feature_names = getattr(clinical_view, "feature_names", None)
feature_groups = getattr(clinical_view, "feature_groups", None)
if feature_names is None or feature_groups is None:
print(
"[F8] make_clinical_importance: profile lacks feature_names / "
"feature_groups; skipping S8e."
)
return
# eyeID is structural (used as the OD/OS index level in DataViews), not a
# learned variable. Strip it from the reported permutation analysis.
feature_names, feature_groups = _drop_structural_features(
feature_names, feature_groups, drop={"eyeID"},
)
label_filter = cfg.get("label_filter", None)
df_mode = data.df.copy()
if label_filter is not None:
df_mode = df_mode[df_mode[data.label_col].isin(label_filter)].reset_index(drop=True)
identity_cols = getattr(data, "identity_cols", [])
identity_level = cfg.get("split_identity_level", 1)
group_col = identity_cols[identity_level - 1] if identity_level and identity_cols else None
splits = SplitManager(group_col=group_col).build_plans(
df_mode,
label_col=data.label_col,
n_splits=cfg.get("folds", 5),
seed=cfg.get("fold_seed", 100),
)
all_results = []
for fold_idx in tqdm(range(cfg.get("folds", 5)), desc="Importance folds", unit="fold"):
ckpt_dir = V4_CKPT_RUN / "checkpoints" / f"fold{fold_idx}"
if not ckpt_dir.exists():
continue
towers, nt, hb, img_aux, cd_aux, hb_head = _build_v4_fold_modules(
fold_idx, cfg, data, device,
)
split_obj = splits[fold_idx]
if split_obj.test is None:
del towers, nt, hb, img_aux, cd_aux, hb_head
continue
shell = data.build_shells(split_obj.test, level="patient", label_filter=label_filter)
loader = _make_loader(
shell,
towers,
batch_size=cfg["training"].get("batch_size", 8),
shuffle=False,
)
X_list, y_list = [], []
with torch.no_grad():
for batch in loader:
X_list.append(batch["cd"]["a"].cpu().numpy()) # OD-side clinical
y_list.append(batch["label"].cpu().numpy())
if not X_list:
del towers, nt, hb, img_aux, cd_aux, hb_head
continue
X = np.concatenate(X_list).astype(np.float32)
y = np.concatenate(y_list).astype(int)
if len(np.unique(y)) < 2:
del towers, nt, hb, img_aux, cd_aux, hb_head
continue
def score_fn(X_in: np.ndarray) -> float:
x_t = torch.from_numpy(X_in).to(device)
with torch.no_grad():
z = towers["cd"](x_t)
logits = cd_aux(z)
probs = TF_.softmax(logits, dim=1).cpu().numpy()
try:
return roc_auc_score(y, probs[:, 1])
except Exception:
return float("nan")
result = permutation_importance(
score_fn=score_fn,
X=X,
groups=feature_groups,
n_permutations=n_permutations,
seed=seed + fold_idx,
feature_names=feature_names,
)
all_results.append(result)
del towers, nt, hb, img_aux, cd_aux, hb_head
if torch.cuda.is_available():
torch.cuda.empty_cache()
if not all_results:
print("[F8] make_clinical_importance: no folds with usable data; skipping S8e.")
return
fold_drops = np.stack([r["mean_drop"] for r in all_results])
mean_across_folds = fold_drops.mean(axis=0)
std_across_folds = fold_drops.std(axis=0)
baseline_mean = float(np.mean([r["baseline"] for r in all_results]))
names = all_results[0]["feature_names"]
order = np.argsort(mean_across_folds)[::-1]
sorted_names = [names[i] for i in order]
sorted_means = mean_across_folds[order]
sorted_stds = std_across_folds[order]
OUT_DIR.mkdir(parents=True, exist_ok=True)
fig, ax = plt.subplots(figsize=(7.5, max(3.5, 0.32 * len(sorted_names))))
y_pos = np.arange(len(sorted_names))
ax.barh(
y_pos, sorted_means, xerr=sorted_stds,
color="#3B6FB5", edgecolor="black", height=0.7,
error_kw=dict(ecolor="#444", lw=0.8, capsize=2),
)
ax.set_yticks(y_pos)
ax.set_yticklabels(sorted_names, fontsize=9)
ax.invert_yaxis()
ax.set_xlabel("Mean AUC drop on shuffling (averaged across folds)", fontsize=10)
ax.axvline(0, color="black", linewidth=0.7)
ax.set_title(
f"S8e — Clinical permutation importance (cd-only head, R50 ckpt run)\n"
f"baseline AUC = {baseline_mean:.3f}; n_permutations = {n_permutations}",
fontsize=10, fontweight="bold",
)
ax.grid(axis="x", alpha=0.3, linestyle="--")
fig.tight_layout()
out_png = OUT_DIR / "S8e_clinical_importance.png"
out_csv = OUT_DIR / "S8e_clinical_importance.csv"
fig.savefig(out_png, dpi=180, bbox_inches="tight")
plt.close(fig)
pd.DataFrame({
"feature": sorted_names,
"mean_drop": sorted_means,
"std_drop": sorted_stds,
"baseline_auc_mean": baseline_mean,
}).to_csv(out_csv, index=False)
print(f"Saved: {out_png}")
print(f"Saved: {out_csv}")
def make_fused_clinical_importance(n_permutations: int = 30, seed: int = 0) -> None:
"""S8e-b clinical permutation importance measured at the fused L2 head.
Answers a different question from the L1 (cd-only) importance: given the
image tower is already contributing, which clinical features still change
the patient-level fused prediction? For each permutation, the clinical
vector is re-embedded through the cd tower, fused with cached image
embeddings through the L1 (nt) bridge per eye, aggregated through the L2
(hb) bridge, and scored against the patient-level label at hb_head.
OD and OS clinical vectors are permuted together (same source-patient
replaces both eyes) so patient-level pairing is preserved.
Aggregation: all 10 reps of ``ensemble_refugelike_ckpt`` are looped; each
rep produces a per-fold mean AUC drop, averaged within-rep to a rep-level
drop; final bars show mean ± SD **across reps** (n=10). Baseline is the
rep-mean of within-rep fold-mean baseline AUCs. This matches the
rep-mean-of-fold-means aggregation used elsewhere in the manuscript.
"""
import json
import torch
import torch.nn.functional as TF_
from tqdm import tqdm
from v4.classes.accessory.explainability import permutation_importance
from v4.classes.split_manager import SplitManager
from v4.classes.v4_hypertower import load_data, _make_loader
import v4.classes.v4_hypertower as orch
rep_base = V4_CKPT_RUN.parent.parent # .../ensemble_refugelike_ckpt
rep_dirs = sorted(
p for p in rep_base.glob("rep*/binary")
if (p / "summary.json").exists() and (p / "checkpoints").exists()
)
if not rep_dirs:
print(f"[F8] make_fused_clinical_importance: no reps found under {rep_base}")
return
print(f"[F8] Fused clinical importance across {len(rep_dirs)} reps")
cfg = json.loads((rep_dirs[0] / "summary.json").read_text())["config"]
orch.cfg_ref = cfg
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
data = load_data(cfg)
clinical_view = data.matrix
feature_names = getattr(clinical_view, "feature_names", None)
feature_groups = getattr(clinical_view, "feature_groups", None)
if feature_names is None or feature_groups is None:
print("[F8] make_fused_clinical_importance: profile lacks feature_names / "
"feature_groups; skipping.")
return
# eyeID is structural (index level for OD/OS in DataViews), not a learned
# variable. Skip it in the reported analysis.
feature_names, feature_groups = _drop_structural_features(
feature_names, feature_groups, drop={"eyeID"},
)
label_filter = cfg.get("label_filter", None)
df_mode = data.df.copy()
if label_filter is not None:
df_mode = df_mode[df_mode[data.label_col].isin(label_filter)].reset_index(drop=True)
identity_cols = getattr(data, "identity_cols", [])
identity_level = cfg.get("split_identity_level", 1)
group_col = identity_cols[identity_level - 1] if identity_level and identity_cols else None
splits = SplitManager(group_col=group_col).build_plans(
df_mode,
label_col=data.label_col,
n_splits=cfg.get("folds", 5),
seed=cfg.get("fold_seed", 100),
)
rep_mean_drops = [] # list of (F,) arrays: within-rep fold-mean drops
rep_mean_baselines = [] # list of floats: within-rep fold-mean baseline AUCs
reference_names = None
for rep_idx, run_dir in enumerate(rep_dirs):
rep_cfg = json.loads((run_dir / "summary.json").read_text())["config"]
# matched-seed reps share the same fold-splitting seed / label_filter;
# verify config alignment before reusing splits computed from cfg above.
if rep_cfg.get("fold_seed") != cfg.get("fold_seed"):
print(f"[F8] {run_dir.parent.name}: fold_seed mismatch, rebuilding splits.")
splits_here = SplitManager(group_col=group_col).build_plans(
df_mode, label_col=data.label_col,
n_splits=rep_cfg.get("folds", 5), seed=rep_cfg.get("fold_seed", 100),
)
else:
splits_here = splits
fold_results = []
for fold_idx in tqdm(
range(rep_cfg.get("folds", 5)),
desc=f"rep {rep_idx:02d}/{len(rep_dirs)-1}",
unit="fold",
leave=False,
):
ckpt_dir = run_dir / "checkpoints" / f"fold{fold_idx}"
if not ckpt_dir.exists():
continue
towers, nt, hb, img_aux, cd_aux, hb_head = _build_v4_fold_modules(
fold_idx, rep_cfg, data, device, run_dir=run_dir,
)
split_obj = splits_here[fold_idx]
if split_obj.test is None:
del towers, nt, hb, img_aux, cd_aux, hb_head
continue
shell = data.build_shells(split_obj.test, level="patient", label_filter=label_filter)
loader = _make_loader(
shell,
towers,
batch_size=rep_cfg["training"].get("batch_size", 8),
shuffle=False,
)
cd_a_list, cd_b_list = [], []
z_img_a_list, z_img_b_list = [], []
y_list = []
with torch.no_grad():
for batch in loader:
cd_a_list.append(batch["cd"]["a"].cpu().numpy())
cd_b_list.append(batch["cd"]["b"].cpu().numpy())
z_img_a = towers["img"](batch["img"]["a"].to(device))
z_img_b = towers["img"](batch["img"]["b"].to(device))
z_img_a_list.append(z_img_a.cpu().numpy())
z_img_b_list.append(z_img_b.cpu().numpy())
y_list.append(batch["label"].cpu().numpy())
if not cd_a_list:
del towers, nt, hb, img_aux, cd_aux, hb_head
continue
cd_a = np.concatenate(cd_a_list).astype(np.float32)
cd_b = np.concatenate(cd_b_list).astype(np.float32)
z_img_a_cached = torch.from_numpy(np.concatenate(z_img_a_list)).to(device)
z_img_b_cached = torch.from_numpy(np.concatenate(z_img_b_list)).to(device)
y = np.concatenate(y_list).astype(int)
if len(np.unique(y)) < 2:
del towers, nt, hb, img_aux, cd_aux, hb_head
continue
F = cd_a.shape[1]
X_wide = np.concatenate([cd_a, cd_b], axis=1) # (N, 2F)
wide_groups = [[c for c in g] + [c + F for c in g] for g in feature_groups]
def score_fn(X_in: np.ndarray) -> float:
cd_a_in = torch.from_numpy(X_in[:, :F]).to(device)
cd_b_in = torch.from_numpy(X_in[:, F:]).to(device)
with torch.no_grad():
z_cd_a = towers["cd"](cd_a_in)
z_cd_b = towers["cd"](cd_b_in)
z_nt_a = nt([z_img_a_cached, z_cd_a])
z_nt_b = nt([z_img_b_cached, z_cd_b])
z_hb = hb({"a": z_nt_a, "b": z_nt_b})
logits = hb_head(z_hb)
probs = TF_.softmax(logits, dim=1).cpu().numpy()
try:
return roc_auc_score(y, probs[:, 1])
except Exception:
return float("nan")
result = permutation_importance(
score_fn=score_fn,
X=X_wide,
groups=wide_groups,
n_permutations=n_permutations,
seed=seed + rep_idx * 100 + fold_idx,
feature_names=feature_names,
)
fold_results.append(result)
del towers, nt, hb, img_aux, cd_aux, hb_head
del z_img_a_cached, z_img_b_cached
if torch.cuda.is_available():
torch.cuda.empty_cache()
if not fold_results:
print(f"[F8] {run_dir.parent.name}: no folds produced usable data; skipping rep.")
continue
rep_drops_matrix = np.stack([r["mean_drop"] for r in fold_results])
rep_mean_drops.append(rep_drops_matrix.mean(axis=0))
rep_mean_baselines.append(float(np.mean([r["baseline"] for r in fold_results])))
if reference_names is None:
reference_names = fold_results[0]["feature_names"]
print(f"[F8] {run_dir.parent.name}: baseline={rep_mean_baselines[-1]:.4f} "
f"(n_folds={len(fold_results)})")
if not rep_mean_drops:
print("[F8] make_fused_clinical_importance: no reps produced usable data; skipping.")
return
rep_stack = np.stack(rep_mean_drops) # (n_reps, F)
mean_across_reps = rep_stack.mean(axis=0)
std_across_reps = rep_stack.std(axis=0)
baseline_mean = float(np.mean(rep_mean_baselines))
baseline_std = float(np.std(rep_mean_baselines))
n_reps = rep_stack.shape[0]
names = reference_names
order = np.argsort(mean_across_reps)[::-1]
sorted_names = [names[i] for i in order]
sorted_means = mean_across_reps[order]
sorted_stds = std_across_reps[order]
OUT_DIR.mkdir(parents=True, exist_ok=True)
fig, ax = plt.subplots(figsize=(7.5, max(3.5, 0.32 * len(sorted_names))))
y_pos = np.arange(len(sorted_names))
ax.barh(
y_pos, sorted_means, xerr=sorted_stds,
color="#c44e52", edgecolor="black", height=0.7,
error_kw=dict(ecolor="#444", lw=0.8, capsize=2),
)
ax.set_yticks(y_pos)
ax.set_yticklabels(sorted_names, fontsize=9)
ax.invert_yaxis()
ax.set_xlabel(f"Mean fused-head AUC drop on clinical shuffling (mean +/- SD across {n_reps} reps)", fontsize=10)
ax.axvline(0, color="black", linewidth=0.7)
ax.set_title(
"Fused-head clinical permutation importance",
fontsize=10, fontweight="bold",
)
ax.grid(axis="x", alpha=0.3, linestyle="--")
fig.tight_layout()
out_png = OUT_DIR / "S8e_fused_clinical_importance.png"
out_csv = OUT_DIR / "S8e_fused_clinical_importance.csv"
fig.savefig(out_png, dpi=180, bbox_inches="tight")
plt.close(fig)
pd.DataFrame({
"feature": sorted_names,
"mean_drop_across_reps": sorted_means,
"std_drop_across_reps": sorted_stds,
"baseline_auc_mean": baseline_mean,
"baseline_auc_std": baseline_std,
"n_reps": n_reps,
}).to_csv(out_csv, index=False)
print(f"Saved: {out_png}")
print(f"Saved: {out_csv}")
def _build_v4_fold_modules(fold_idx: int, cfg: dict, data, device, run_dir: Path | None = None):
import importlib
import torch
from v4.classes.v4_hypertower import build_towers
from v4.classes.heads.classifier import ClassificationHead
stage_by_name = {s["name"]: s for s in cfg["stages"]}
nt_cfg = stage_by_name["nt"]
hb_cfg = stage_by_name["hb"]
base = run_dir if run_dir is not None else V4_CKPT_RUN
ckpt_dir = base / "checkpoints" / f"fold{fold_idx}"
if not ckpt_dir.exists():
raise FileNotFoundError(f"No v4 checkpoint directory: {ckpt_dir}")
towers = build_towers(cfg["towers"], data)
for name, tower in towers.items():
tower.load_state_dict(torch.load(ckpt_dir / f"tower_{name}.pt", map_location="cpu"))
tower.to(device).eval()
nt_mod = importlib.import_module(nt_cfg["module"])
nt = getattr(nt_mod, nt_cfg["class"])(
[towers[n].out_dim for n in nt_cfg["inputs"]],
**nt_cfg.get("args", {}),
).to(device)
nt.load_state_dict(torch.load(ckpt_dir / "stage_nt.pt", map_location="cpu"))
nt.eval()
hb_mod = importlib.import_module(hb_cfg["module"])
hb = getattr(hb_mod, hb_cfg["class"])(
{"a": nt.out_dim, "b": nt.out_dim},
**hb_cfg.get("args", {}),
).to(device)
hb.load_state_dict(torch.load(ckpt_dir / "stage_hb.pt", map_location="cpu"))
hb.eval()
img_aux = ClassificationHead(towers["img"].out_dim, cfg["num_classes"]).to(device)
cd_aux = ClassificationHead(towers["cd"].out_dim, cfg["num_classes"]).to(device)
hb_head = ClassificationHead(hb.out_dim, cfg["num_classes"], dropout=0.3).to(device)
img_aux.load_state_dict(torch.load(ckpt_dir / "stage_img_aux.pt", map_location="cpu"))
cd_aux.load_state_dict(torch.load(ckpt_dir / "stage_cd_aux.pt", map_location="cpu"))
hb_head.load_state_dict(torch.load(ckpt_dir / "stage_hb_head.pt", map_location="cpu"))
img_aux.eval(); cd_aux.eval(); hb_head.eval()
return towers, nt, hb, img_aux, cd_aux, hb_head
def _orient_eye_array(arr: np.ndarray, eye: str) -> np.ndarray:
"""Mirror OS so temporal/nasal anatomy is aligned to OD-style orientation."""
return np.fliplr(arr) if eye.upper() == "OS" else arr
def _disc_centred_patch_array(
arr: np.ndarray,
disc_mask: np.ndarray,
*,
span: int = 5,
out: int = 96,
) -> tuple[np.ndarray | None, float | None]:
if disc_mask is None or disc_mask.sum() == 0:
return None, None
ys, xs = np.where(disc_mask)
cy, cx = ys.mean(), xs.mean()
disc_r = float(np.sqrt(disc_mask.sum() / np.pi))
half = max(1, int(round(span * disc_r / 2)))
h, w = disc_mask.shape
y0, y1 = int(round(cy)) - half, int(round(cy)) + half
x0, x1 = int(round(cx)) - half, int(round(cx)) + half
pad_spec = ((max(0, -y0), max(0, y1 - h)), (max(0, -x0), max(0, x1 - w)))
if arr.ndim == 3:
pad_spec = (*pad_spec, (0, 0))
arr_pad = np.pad(arr, pad_spec, constant_values=0)
patch = arr_pad[y0 + pad_spec[0][0]: y1 + pad_spec[0][0],
x0 + pad_spec[1][0]: x1 + pad_spec[1][0]]
if arr.ndim == 2:
pil = Image.fromarray((np.clip(patch, 0, 1) * 255).astype(np.uint8))
patch_out = np.array(pil.resize((out, out), Image.BILINEAR)) / 255.0
else:
patch_out = np.array(Image.fromarray(patch.astype(np.uint8)).resize((out, out), Image.BILINEAR))
return patch_out.astype(np.float32), out * disc_r / (2 * half)
QUAD_ORDER = ("ST", "SN", "IT", "IN") # superotemporal, superonasal, inferotemporal, inferonasal
def _quadrant_fractions(
cam: np.ndarray,
disc_mask: np.ndarray,
*,
peri_inner: float = 1.0,
peri_outer: float = 2.0,
) -> tuple[dict[str, float], dict[str, float], dict[str, float]] | tuple[None, None, None]:
"""Per-quadrant Grad-CAM fractions in OD-oriented coordinates, for three
region scopes.
Quadrant boundaries are the disc-mask centroid (cx, cy). In the OD-oriented
frame nasal is left (x < cx) and temporal is right (x > cx); superior is
top (y < cy) and inferior is bottom (y > cy):
ST = x > cx, y < cy
SN = x < cx, y < cy
IT = x > cx, y > cy
IN = x < cx, y > cy
Three region scopes are returned:
disc_q : fractions of CAM intensity that fall inside the GT disc mask
peri_q : fractions inside a peri-disc annulus of disc-radius units
(peri_inner to peri_outer, default 1x-2x), excluding the disc
full_q : fractions over the entire image
Each dict sums to 1 (within floating-point error). Returns (None, None, None)
if the disc mask is empty.
"""
if disc_mask is None or disc_mask.sum() == 0:
return None, None, None
ys, xs = np.where(disc_mask)
cy = float(ys.mean()); cx = float(xs.mean())
disc_r = float(np.sqrt(disc_mask.sum() / np.pi))
h, w = cam.shape
yy, xx = np.mgrid[0:h, 0:w]
dist = np.sqrt((xx - cx) ** 2 + (yy - cy) ** 2)
peri_mask = (dist >= peri_inner * disc_r) & (dist <= peri_outer * disc_r) & ~disc_mask
quads = {
"ST": (xx > cx) & (yy < cy),
"SN": (xx < cx) & (yy < cy),
"IT": (xx > cx) & (yy > cy),
"IN": (xx < cx) & (yy > cy),
}
disc_total = float(cam[disc_mask].sum()) + 1e-8
peri_total = float(cam[peri_mask].sum()) + 1e-8
full_total = float(cam.sum()) + 1e-8
disc_q = {k: float(cam[disc_mask & q].sum()) / disc_total for k, q in quads.items()}
peri_q = {k: float(cam[peri_mask & q].sum()) / peri_total for k, q in quads.items()}
full_q = {k: float(cam[q].sum()) / full_total for k, q in quads.items()}
return disc_q, peri_q, full_q
def _annotate_nasal_temporal(ax, *, fontsize: int = 9, color: str = "white",
pad: float = 2.5) -> None:
"""Label the disc-side (nasal) and macula-side (temporal) edges of an
OD-oriented fundus axis. After OS is mirrored to OD orientation, the disc
sits on the LEFT (nasal) side and the macula on the RIGHT (temporal) side.
"""
ax.text(0.015, 0.5, "N", transform=ax.transAxes,
ha="left", va="center", fontsize=fontsize, fontweight="bold",
color=color,
bbox=dict(facecolor="black", alpha=0.55, edgecolor="none", pad=pad))
ax.text(0.985, 0.5, "T", transform=ax.transAxes,
ha="right", va="center", fontsize=fontsize, fontweight="bold",
color=color,
bbox=dict(facecolor="black", alpha=0.55, edgecolor="none", pad=pad))
def _make_oriented_disc_detail(mean_patches: dict, examples: dict, out_path: Path) -> None:
"""2x2 grid of mean Grad-CAM patches centred on the optic disc.
Rows = ground-truth class (Normal / Glaucoma)
Cols = prediction outcome relative to truth (Correct / Incorrect)
Each cell shows the per-pixel mean Grad-CAM intensity in a disc-centred,
OD-oriented patch (OS mirrored). A dashed circle traces the mean disc
boundary. Nasal (N) and Temporal (T) edges are labelled.
Sample overlays are intentionally NOT shown here; see overlay_grid_*.png
for browseable per-image overlays and hand-pick from those for any
figure that wants concrete examples.
"""
from matplotlib.patches import Circle
classes = ["Normal", "Glaucoma"]
splits = ["correct", "incorrect"]
fig, axes = plt.subplots(2, 2, figsize=(8, 7.8), constrained_layout=True)
# Column headers
for ci, split in enumerate(splits):
axes[0, ci].set_title(
f"{split.capitalize()} predictions",
fontsize=11, fontweight="bold", pad=10, color="#222",
)
for ri, cls in enumerate(classes):
# Row label drawn outside the leftmost cell
axes[ri, 0].text(
-0.13, 0.5, cls, transform=axes[ri, 0].transAxes,
ha="right", va="center", fontsize=12, fontweight="bold",
color="#c44e52" if cls == "Glaucoma" else "#3B6FB5", rotation=90,
)
for ci, split in enumerate(splits):
ax = axes[ri, ci]
ax.set_facecolor("#202020")
key = (cls, split)
if key in mean_patches:
patch, radius, count, *rest = mean_patches[key]
disc_frac = rest[0] if rest else None
ax.imshow(patch, cmap="jet", vmin=0, vmax=1)
ax.add_patch(Circle((48, 48), radius, fill=False,
edgecolor="white", linewidth=1.8, linestyle="--"))
_annotate_nasal_temporal(ax)
cap = f"n = {count}"
if disc_frac is not None and not np.isnan(disc_frac):
cap += f" · disc-frac = {disc_frac:.2f}"
ax.text(0.5, -0.06, cap, transform=ax.transAxes,
ha="center", va="top", fontsize=9, color="#222")
else:
ax.text(0.5, 0.5, "no data", transform=ax.transAxes,
ha="center", va="center", color="white", fontsize=9)
ax.set_xticks([]); ax.set_yticks([])
fig.savefig(out_path, dpi=180, bbox_inches="tight")
plt.close(fig)
print(f"Saved: {out_path}")
def _patient_id_from_entity_id(entity_id) -> int:
if isinstance(entity_id, (tuple, list)):
return int(entity_id[0])
return int(entity_id)
def _v4_gradcam_target_layer(img_tower):
blocks = getattr(img_tower, "_blocks", None)
if blocks:
return blocks[-1]
backbone = getattr(img_tower, "backbone", None)
if hasattr(backbone, "features"):
return backbone.features[-1]
children = list(backbone.children()) if backbone is not None else []
if children:
return children[-1]
raise RuntimeError(f"Could not infer GradCAM target layer for {type(img_tower).__name__}")
def _maybe(profile_view, method: str, *args, **kwargs):
"""Call profile_view.method(*args, **kwargs) if it exists, else return None."""
fn = getattr(profile_view, method, None)
return fn(*args, **kwargs) if fn is not None else None
def _orient_via_profile(profile_view, image_or_array, side: str):
"""Use the profile's orient_for_display if it provides one; else passthrough."""
return _maybe(profile_view, "orient_for_display", image_or_array, side) \
if hasattr(profile_view, "orient_for_display") else image_or_array
def _compute_image_aux_cam(
gcam,
batch: dict,
eye: str,
towers: dict,
img_aux,
device,
target_class: int | None,
) -> tuple[np.ndarray, int]:
"""Run image-tower GradCAM against the img_aux head for one (batch, eye) pair."""
for mod in [towers["img"], img_aux]:
mod.zero_grad(set_to_none=True)
side = "a" if eye.upper() == "OD" else "b"
img_t = batch["img"][side].to(device)
return gcam.compute(
forward_fn=lambda: img_aux(towers["img"](img_t)),
output_shape=img_t.shape[-2:],
target_class=target_class,
)
def _make_oriented_gradcam(n_grid: int = 16, alpha: float = 0.45,
target_class: int | None = None) -> None:
import json
import torch
from tqdm import tqdm
from v4.classes.accessory.explainability import GradCAM, overlay_gradcam
from v4.classes.split_manager import SplitManager
from v4.classes.v4_hypertower import load_data, _make_loader
import v4.classes.v4_hypertower as orch
summary = json.loads((V4_CKPT_RUN / "summary.json").read_text())
cfg = summary["config"]
orch.cfg_ref = cfg
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print(f"GradCAM device: {device}")
data = load_data(cfg)
image_view = data.image # profile data view — owns dataset-specific knowledge
label_filter = cfg.get("label_filter", None)
df_mode = data.df.copy()
if label_filter is not None:
df_mode = df_mode[df_mode[data.label_col].isin(label_filter)].reset_index(drop=True)
identity_cols = getattr(data, "identity_cols", [])
identity_level = cfg.get("split_identity_level", 1)
group_col = identity_cols[identity_level - 1] if identity_level and identity_cols else None
splits = SplitManager(group_col=group_col).build_plans(
df_mode,
label_col=data.label_col,
n_splits=cfg.get("folds", 5),
seed=cfg.get("fold_seed", 100),
)
cam_sum = {0: None, 1: None}
cam_count = {0: 0, 1: 0}
overlay_items = {0: [], 1: []}
disc_patch_sum: dict[tuple[str, str], np.ndarray] = {}
disc_patch_count: dict[tuple[str, str], int] = {}
disc_radius_sum: dict[tuple[str, str], float] = {}
disc_frac_sum: dict[tuple[str, str], float] = {}
# Per-eye quadrant fractions for three region scopes; aggregated per cell.
quad_disc_list: dict[tuple[str, str], list[dict[str, float]]] = {}
quad_peri_list: dict[tuple[str, str], list[dict[str, float]]] = {}
quad_full_list: dict[tuple[str, str], list[dict[str, float]]] = {}
examples: dict[tuple[str, str], tuple[np.ndarray, int, str, int]] = {}
fold_range = range(cfg.get("folds", 5))
for fold_idx in tqdm(fold_range, desc="GradCAM folds", unit="fold"):
ckpt_dir = V4_CKPT_RUN / "checkpoints" / f"fold{fold_idx}"
if not ckpt_dir.exists():
continue
towers, nt, hb, img_aux, _cd_aux, hb_head = _build_v4_fold_modules(
fold_idx, cfg, data, device,
)
target_layer = _v4_gradcam_target_layer(towers["img"])
gcam = GradCAM(target_layer)
split_obj = splits[fold_idx]
if split_obj.test is None:
continue
shell = data.build_shells(split_obj.test, level="patient", label_filter=label_filter)
loader = _make_loader(shell, towers, batch_size=1, shuffle=False)
for batch in tqdm(loader, desc=f"fold{fold_idx}", leave=False, unit="pt"):
pid = _patient_id_from_entity_id(batch["entity_id"][0])
label = int(batch["label"][0].item())
for eye in ("OD", "OS"):
if not image_view.get_image_path(pid, eye).exists():
continue
pil_eval = image_view.eval_image_pil(pid, eye)
roi_mask = _maybe(image_view, "build_roi_mask",
pid, eye, target_h=pil_eval.size[1], target_w=pil_eval.size[0])
cam_np, pred = _compute_image_aux_cam(
gcam, batch, eye, towers, img_aux, device, target_class,
)
# Orient for OD-style display (profile decides; passthrough if it doesn't define)
cam_np = _orient_via_profile(image_view, cam_np, eye)
pil_oriented = _orient_via_profile(image_view, pil_eval, eye)
if roi_mask is not None:
roi_mask = _orient_via_profile(image_view, roi_mask, eye)
if cam_sum[label] is None:
cam_sum[label] = cam_np.copy()
else:
cam_sum[label] += cam_np
cam_count[label] += 1
if len(overlay_items[label]) < n_grid:
overlay_items[label].append(
(overlay_gradcam(pil_oriented, cam_np, alpha), pid, eye, pred)
)
patch, roi_r_out = _disc_centred_patch_array(cam_np, roi_mask)
if patch is not None:
key = (LABEL_NAMES[label], "correct" if pred == label else "incorrect")
disc_patch_sum[key] = disc_patch_sum.get(key, 0) + patch
disc_patch_count[key] = disc_patch_count.get(key, 0) + 1
disc_radius_sum[key] = disc_radius_sum.get(key, 0.0) + float(roi_r_out)
disc_frac_sum[key] = disc_frac_sum.get(key, 0.0) + float(
cam_np[roi_mask].sum() / (cam_np.sum() + 1e-8)
)
disc_q, peri_q, full_q = _quadrant_fractions(cam_np, roi_mask)
if disc_q is not None:
quad_disc_list.setdefault(key, []).append(disc_q)
quad_peri_list.setdefault(key, []).append(peri_q)
quad_full_list.setdefault(key, []).append(full_q)
if key not in examples:
ov = overlay_gradcam(pil_oriented, cam_np, alpha)
ov_small = np.array(ov.resize(cam_np.shape[::-1], Image.BILINEAR))
ex_patch, _ = _disc_centred_patch_array(ov_small, roi_mask)
if ex_patch is not None:
examples[key] = (ex_patch, pid, eye, pred)
del cam_np
gcam.remove()
del towers, nt, hb, img_aux, _cd_aux, hb_head
if torch.cuda.is_available():
torch.cuda.empty_cache()
GRADCAM_DIR.mkdir(parents=True, exist_ok=True)
for cls in (0, 1):
if cam_count[cls] == 0:
continue
mean_cam = cam_sum[cls] / cam_count[cls]
mean_cam = (mean_cam - mean_cam.min()) / (mean_cam.max() - mean_cam.min() + 1e-8)
fig, ax = plt.subplots(figsize=(5, 5))
im = ax.imshow(mean_cam, cmap="jet", vmin=0, vmax=1)
ax.axis("off")
ax.set_title(f"Mean v4 image-aux GradCAM - {LABEL_NAMES[cls]} (oriented, n={cam_count[cls]})",
fontsize=10, fontweight="bold")
fig.colorbar(im, ax=ax, fraction=0.046, pad=0.04)
fig.savefig(GRADCAM_DIR / f"mean_cam_{LABEL_NAMES[cls].lower()}.png",
dpi=180, bbox_inches="tight")
plt.close(fig)
if cam_count[0] and cam_count[1]:
fig, axes = plt.subplots(1, 2, figsize=(11, 5.5), constrained_layout=True)
fig.patch.set_facecolor("#f6f6f6")
fig.suptitle(
"Mean image-tower Grad-CAM by ground-truth class (OD-oriented)",
fontsize=12, fontweight="bold",
)
cls_colors = {0: "#3B6FB5", 1: "#c44e52"}
for ax, cls in zip(axes, (0, 1)):
mean_cam = cam_sum[cls] / cam_count[cls]
mean_cam = (mean_cam - mean_cam.min()) / (mean_cam.max() - mean_cam.min() + 1e-8)
im = ax.imshow(mean_cam, cmap="jet", vmin=0, vmax=1)
_annotate_nasal_temporal(ax)
ax.set_title(
f"{LABEL_NAMES[cls]} n = {cam_count[cls]}",
fontsize=12, fontweight="bold", color=cls_colors[cls],
)
ax.set_xticks([]); ax.set_yticks([])
fig.colorbar(im, ax=ax, fraction=0.046, pad=0.04)
fig.savefig(GRADCAM_DIR / "mean_cam_comparison.png", dpi=180, bbox_inches="tight")
plt.close(fig)
for cls in (0, 1):
items = overlay_items[cls]
if not items:
continue
items.sort(key=lambda x: x[3] == cls)
ncols = 4
nrows = int(np.ceil(len(items) / ncols))
fig, axes = plt.subplots(nrows, ncols, figsize=(ncols * 3.1, nrows * 3.1))
axes = np.array(axes).reshape(-1)
fig.suptitle(
f"Per-eye Grad-CAM samples for browsing — ground truth = {LABEL_NAMES[cls]}\n"
"(OD-oriented; green title = model agreed, red = disagreed)",
fontsize=11, fontweight="bold",
)
for i, ax in enumerate(axes):
if i < len(items):
ov, pid, eye, pred = items[i]
ax.imshow(ov)
_annotate_nasal_temporal(ax, fontsize=7, pad=1.5)
color = "#2f7d46" if pred == cls else "#b23a48"
ax.set_title(f"RET{pid:03d}{eye} -> {LABEL_NAMES[pred]}",
fontsize=7.5, color=color)
ax.axis("off")
fig.tight_layout()
fig.savefig(GRADCAM_DIR / f"overlay_grid_{LABEL_NAMES[cls].lower()}.png",
dpi=160, bbox_inches="tight")
plt.close(fig)
mean_patches = {
key: (
disc_patch_sum[key] / disc_patch_count[key],
disc_radius_sum[key] / disc_patch_count[key],
disc_patch_count[key],
disc_frac_sum.get(key, float("nan")) / disc_patch_count[key],
)
for key in disc_patch_sum
}
if mean_patches:
_make_oriented_disc_detail(mean_patches, examples, GRADCAM_DIR / "disc_attention_detail.png")
# Per-quadrant CAM fractions (within-disc and full-image scopes).
# One CSV row per (class, outcome, scope, quadrant) cell, with mean and SD
# computed over the per-eye fractions in that cell.
rows = []
for key in sorted(quad_disc_list.keys()):
cls_name, outcome = key
n_eyes = len(quad_disc_list[key])
for scope_label, scope_list in (("disc", quad_disc_list[key]),
("peri", quad_peri_list[key]),
("full", quad_full_list[key])):
for q in QUAD_ORDER:
vals = np.array([d[q] for d in scope_list], dtype=np.float64)
rows.append({
"class": cls_name,
"outcome": outcome,
"scope": scope_label,
"quadrant": q,
"n_eyes": n_eyes,
"mean": float(vals.mean()),
"sd": float(vals.std(ddof=1)) if n_eyes > 1 else float("nan"),
})
if rows:
out_csv = OUT_DIR / "F8_quadrant_fractions.csv"
pd.DataFrame(rows).to_csv(out_csv, index=False)
print(f"saved quadrant fractions: {out_csv}")
def make_oriented_gradcam(n_grid: int = 16, alpha: float = 0.45,
target_class: int | None = None) -> None:
_make_oriented_gradcam(n_grid=n_grid, alpha=alpha, target_class=target_class)
def main(*, run_gradcam: bool = False, n_permutations: int = 30,
fusion_split: str = "test") -> None:
OUT_DIR.mkdir(parents=True, exist_ok=True)
make_fusion_event_panel(split=fusion_split)
make_clinical_importance(n_permutations=n_permutations)
make_fused_clinical_importance(n_permutations=n_permutations)
if run_gradcam:
make_oriented_gradcam()
else:
print("Skipped GradCAM. Re-run with --run-gradcam to generate oriented heatmaps.")
if __name__ == "__main__":
import argparse
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--only-fusion", action="store_true",
help="Generate only the fusion event panel; skip clinical and GradCAM figures.")
parser.add_argument("--only-gradcam", action="store_true",
help="Generate only oriented GradCAM outputs; skip fusion and clinical figures.")
parser.add_argument("--run-gradcam", action="store_true",
help="Generate GPU-heavy oriented GradCAM outputs.")
parser.add_argument("--n-grid", type=int, default=16,
help="Max overlays per class for GradCAM grids.")
parser.add_argument("--alpha", type=float, default=0.45,
help="GradCAM overlay opacity.")
parser.add_argument("--target-class", type=int, default=None,
help="GradCAM target class; default uses predicted class.")
parser.add_argument("--n-permutations", type=int, default=30,
help="Clinical permutation repeats per feature.")
parser.add_argument("--fusion-split", choices=["test", "val", "both"], default="test",
help="Prediction split for the fusion event panel.")
args = parser.parse_args()
if args.only_fusion:
OUT_DIR.mkdir(parents=True, exist_ok=True)
make_fusion_event_panel(split=args.fusion_split)
elif args.only_gradcam:
make_oriented_gradcam(
n_grid=args.n_grid,
alpha=args.alpha,
target_class=args.target_class,
)
else:
main(
run_gradcam=args.run_gradcam,
n_permutations=args.n_permutations,
fusion_split=args.fusion_split,
)