# experiment/probing/plot_exemplar_heatmaps.py """Render the heatmap / trajectory / per-feature figures per image_id. Pure parquet reader — no model loading, no GPU required. """ from __future__ import annotations import argparse import json import os import sys import matplotlib.pyplot as plt import numpy as np import pandas as pd sys.path.insert(0, os.path.join(os.path.dirname(__file__), "../..")) from experiment.probing._helpers import heatmap_matrix_from_parquet METHODS = ["base", "adv", "efuf", "nullu", "lora_finetune"] def _load_tau(tau_path: str, alpha: float) -> dict[str, float]: with open(tau_path) as f: d = json.load(f) return {m: d[m][f"{alpha:.2f}"] for m in METHODS if m in d} def _flatten_feature_ids(features_path: str, set_key: str) -> list[int]: with open(features_path) as f: feats = json.load(f) ids: set[int] = set() for _, v in feats[set_key].items(): ids.update(v) return sorted(ids) def render_heatmap(parquet_path: str, output_dir: str, feature_ids: list[int], tau: dict[str, float]) -> None: fig, axes = plt.subplots(1, len(METHODS), figsize=(4 * len(METHODS), 6), sharey=True) vmin, vmax = 0.0, max(tau.values()) * 2.0 im = None for ax, m in zip(axes, METHODS): try: mat, tokens, is_tt = heatmap_matrix_from_parquet( parquet_path, method=m, image_id=_iid_from_parquet(parquet_path), pass_="teacher", feature_ids=feature_ids, ) except ValueError: ax.set_title(f"{m}\n(no data)") ax.axis("off") continue im = ax.imshow(mat, aspect="auto", origin="lower", cmap="hot", vmin=vmin, vmax=vmax) ax.set_title(f"{m} τ_c={tau.get(m, float('nan')):.2f}") ax.set_xlabel("token") ax.set_xticks(range(len(tokens))) ax.set_xticklabels( [t.strip() for t in tokens], rotation=90, fontsize=6, ) for i, tt in enumerate(is_tt): if tt: ax.get_xticklabels()[i].set_color("red") axes[0].set_ylabel("layer") if im is not None: fig.colorbar(im, ax=axes.tolist(), shrink=0.6) fig.suptitle(f"image_id={_iid_from_parquet(parquet_path)} (teacher-forced)") fig.savefig(os.path.join(output_dir, "heatmap.png"), dpi=140, bbox_inches="tight") plt.close(fig) def render_trajectory(parquet_path: str, output_dir: str, feature_ids: list[int], tau: dict[str, float]) -> None: fig, ax = plt.subplots(figsize=(8, 4)) for m in METHODS: try: mat, _, is_tt = heatmap_matrix_from_parquet( parquet_path, method=m, image_id=_iid_from_parquet(parquet_path), pass_="teacher", feature_ids=feature_ids, ) except ValueError: continue toilet_cols = [i for i, tt in enumerate(is_tt) if tt] if not toilet_cols: traj = mat.max(axis=1) else: traj = mat[:, toilet_cols].max(axis=1) ax.plot(traj, label=f"{m} (τ_c={tau.get(m, float('nan')):.2f})") ax.axhline(tau.get(m, 0.0), linestyle=":", alpha=0.4) ax.set_xlabel("layer") ax.set_ylabel("max_{f ∈ Φ_toilet, t ∈ toilet-tokens} z") ax.set_title(f"image_id={_iid_from_parquet(parquet_path)}") ax.legend(fontsize=8) fig.savefig(os.path.join(output_dir, "trajectory.png"), dpi=140, bbox_inches="tight") plt.close(fig) def _iid_from_parquet(parquet_path: str) -> str: return os.path.basename(os.path.dirname(parquet_path)) def main(): p = argparse.ArgumentParser() p.add_argument("--toilet_features", required=True) p.add_argument("--tau_c", required=True) p.add_argument("--feature_set", choices=["A", "B", "C"], default="A") p.add_argument("--alpha", type=float, default=0.05) p.add_argument("--image_ids", default="") p.add_argument("--image_ids_file", default="") p.add_argument("--output_root", default="outputs/feature_ks") args = p.parse_args() feature_ids = _flatten_feature_ids(args.toilet_features, args.feature_set) tau = _load_tau(args.tau_c, args.alpha) if args.image_ids: ids = [s.strip() for s in args.image_ids.split(",") if s.strip()] elif args.image_ids_file: with open(args.image_ids_file) as f: ids = [line.strip() for line in f if line.strip()] else: ids = [d for d in os.listdir(args.output_root) if os.path.isdir(os.path.join(args.output_root, d))] for iid in ids: d = os.path.join(args.output_root, iid) parquet = os.path.join(d, "activations.parquet") if not os.path.exists(parquet): print(f" SKIP {iid}: no parquet") continue render_heatmap(parquet, d, feature_ids, tau) render_trajectory(parquet, d, feature_ids, tau) print(f" rendered {iid}") if __name__ == "__main__": main()