"""Re-render a gradient-ascent heatmap JSON with a SHARED color scale across alphas. The original per-panel scaling normalizes each alpha to its own max, so a large-alpha panel (bigger real changes) gets a stretched scale and looks PALER than a small-alpha panel — visually backwards. This re-plots all panels on ONE shared, robust (percentile-clipped) diverging scale so color magnitude is comparable across alpha: bigger change → more saturated. Usage: python -m mechanistic_interp.scripts.replot_gradient_ascent \ --json mechanistic_interp/graph/gradient_ascent_bath2toilet.json \ --out mechanistic_interp/graph/gradient_ascent_bath2toilet_shared.png \ [--pct 99] [--drop_early 0] """ import argparse import json import numpy as np import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt def main(): ap = argparse.ArgumentParser() ap.add_argument("--json", required=True) ap.add_argument("--out", required=True) ap.add_argument("--pct", type=float, default=99.0, help="Percentile of |Δ| (over ALL panels) used as the shared vmax — " "robust to the early-layer outlier. 100 = plain max.") ap.add_argument("--normalize", action="store_true", help="Alias for --norm_by max (divide by largest |Δ|, scale ±1).") ap.add_argument("--norm_by", choices=["none", "max", "mean"], default="none", help="Divide every cell by max|Δ| or mean|Δ| (over ALL panels). " "'mean' = color is in multiples of the typical effect (crisp); " "'max' = biggest cell = ±1 (artifact-dominated, blurry).") ap.add_argument("--drop_early", type=int, default=0, help="Blank the first N intervention layers (rows) — the off-manifold " "artifact — from both the scale and the display.") ap.add_argument("--title", default=None) args = ap.parse_args() d = json.load(open(args.json)) alphas = d["alphas"] mats = {} for a in alphas: M = np.array([[np.nan if v is None else v for v in row] for row in d["delta"][str(a)]]) if args.drop_early > 0: M[: args.drop_early, :] = np.nan mats[a] = M allvals = np.concatenate([m[np.isfinite(m)].ravel() for m in mats.values()]) absvals = np.abs(allvals) norm_by = "max" if args.normalize else args.norm_by if norm_by == "max": norm = float(absvals.max()) for a in mats: mats[a] = mats[a] / norm vmax = 1.0 cbar_label = "Δ toilet σ / max|Δ| (±1)" print(f"normalized by max|Δ| = {norm:.5f} → scale [-1, 1]") elif norm_by == "mean": norm = float(absvals.mean()) for a in mats: mats[a] = mats[a] / norm znorm = np.abs(np.concatenate([m[np.isfinite(m)].ravel() for m in mats.values()])) vmax = float(np.percentile(znorm, args.pct)) # robust crisp cap, in mean-units cbar_label = "Δ toilet σ / mean|Δ| (× typical effect)" print(f"normalized by mean|Δ| = {norm:.5f}; crisp cap p{args.pct} = {vmax:.2f}× mean " f"(max = {absvals.max()/norm:.1f}× mean)") else: vmax = float(np.percentile(absvals, args.pct)) cbar_label = "Δ toilet score σ (shared scale)" print(f"shared vmax (p{args.pct}) = {vmax:.4f} (raw max |Δ| = {absvals.max():.4f})") n = len(alphas) ncol = min(3, n) nrow = -(-n // ncol) fig, axes = plt.subplots(nrow, ncol, figsize=(4.8 * ncol, 4.3 * nrow), squeeze=False) im = None for k, a in enumerate(alphas): ax = axes[k // ncol][k % ncol] im = ax.imshow(mats[a], origin="upper", cmap="RdBu_r", vmin=-vmax, vmax=vmax, aspect="auto") ax.set_title(f"α = {a}") ax.set_xlabel("readout layer l′ (toilet)") ax.set_ylabel("intervention layer l (bathroom)") for k in range(n, nrow * ncol): axes[k // ncol][k % ncol].axis("off") # One shared colorbar for all panels. cbar = fig.colorbar(im, ax=axes, fraction=0.025, pad=0.02) cbar.set_label(cbar_label) title = args.title or (f"{d.get('n_images','?')} images — shared color scale " f"(p{args.pct} clip, drop_early={args.drop_early})") fig.suptitle(title, fontsize=12) fig.savefig(args.out, dpi=150, bbox_inches="tight") print(f"saved {args.out}") if __name__ == "__main__": main()