hallucination / mechanistic_interp /scripts /diff_gradient_ascent.py
ToiTenBao's picture
Upload hallucination folder
a2ffd07 verified
Raw
History Blame Contribute Delete
1.93 kB
"""
Per-cell difference of two gradient-ascent heatmaps: Δ = A − B (per alpha, per (l,l')).
A and B are the JSON outputs of gradient_ascent.py (each {alphas, delta:{alpha:matrix}}).
Use to compare two models/conditions, e.g. base induction (A) minus LoRA induction (B):
positive cell ⇒ A's bathroom→toilet influence is larger there (LoRA reduced it).
Plots with the shared plot_heatmaps (normal | meannorm), same as gradient_ascent.py.
"""
import argparse
import json
import numpy as np
from mechanistic_interp.gradient_ascent import plot_heatmaps
def load(path):
d = json.load(open(path))
mats = {float(a): np.array([[np.nan if v is None else v for v in row] for row in M])
for a, M in d["delta"].items()}
return mats, d
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--json_a", required=True, help="minuend JSON (A)")
ap.add_argument("--json_b", required=True, help="subtrahend JSON (B)")
ap.add_argument("--out", required=True)
ap.add_argument("--plot_mode", default="meannorm", choices=["normal", "meannorm"])
ap.add_argument("--title", default=None)
args = ap.parse_args()
A, da = load(args.json_a)
B, db = load(args.json_b)
alphas = sorted(set(A) & set(B))
if not alphas:
raise SystemExit(f"No common alphas: A={sorted(A)} B={sorted(B)}")
diff = {a: A[a] - B[a] for a in alphas} # NaN where either is NaN (below diagonal)
nA, nB = da.get("n_images", "?"), db.get("n_images", "?")
title = args.title or f"A − B (A: n={nA}, B: n={nB}) per cell — {args.plot_mode}"
# integer-looking alpha labels read better
alphas_disp = [int(a) if float(a).is_integer() else a for a in alphas]
plot_heatmaps({ad: diff[a] for a, ad in zip(alphas, alphas_disp)}, alphas_disp,
args.out, title=title, mode=args.plot_mode, cbar_label="Δ_A − Δ_B")
if __name__ == "__main__":
main()