| """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)) |
| 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") |
| |
| 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() |
|
|