sae-null-step / tables /scripts /real_arm_replicate.py
anonymous
Add tables
ba7ef34 verified
Raw History Blame Contribute Delete
6.85 kB
"""Replicate the REAL arm of the common-split chain experiment over seeds.
The paper's null step is measured over twelve seeds; the real transition it is
compared against was drawn once. A reviewer's objection follows directly: with
n=1 on the treatment arm there is no real-arm distribution, so "the null step
reproduces 86% of the real drift" is a statement about one draw.
This runs extra REAL chains from the *same* base SAE saved by
`jumprelu_sft_chain.py`, so the real arm gets a seed distribution of its own,
and (optionally) re-runs the NULL chains with decoder weights saved so that the
two clouds can be compared directionally rather than only in magnitude.
Seeding follows `jumprelu_sft_chain.py` exactly -- torch.manual_seed(1000*seed+i)
-- so real seed 0 reproduces the chain already on disk. Real seed j and null
seed j therefore share an RNG stream while training on different activations:
that is common random numbers, a paired design that removes shuffle-order
variance from the real-minus-null difference.
python3 scripts/real_arm_replicate.py --layer 18 --real-seeds 0 1 2 3 \
--null-seeds 0 1 2 3 --save-weights --out real_arm_L18
"""
import argparse
import json
import pathlib
import sys
import time
import pandas as pd
import torch
HERE = pathlib.Path(__file__).resolve().parent
sys.path.insert(0, str(HERE))
sys.path.insert(0, ROOT_DEFAULT)
from jumprelu_sft_chain import (DRIFT, SFT_STAGES, SAE, evaluate, fit, loaders,
normalize_decoder)
from sae_tuned_lens.readout import (compose_lens, load_readout, matched_and_null,
token_overlap, top_tokens)
from sae_tuned_lens.tuned_lens import load_lens
# Paths below were absolute in the authors' environment. Set SAE_RL_ROOT to the
# directory holding drift_run/ and sae_rl/, or edit ROOT_DEFAULT.
import os as _os
ROOT_DEFAULT = _os.environ.get("SAE_RL_ROOT", ".")
ARMS = {"TopK_k64": dict(arch="topk", k=64, lam=0.0),
"JumpReLU": dict(arch="jumprelu", k=None, lam=3e-4)}
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--layer", type=int, default=18)
ap.add_argument("--arm", default="TopK_k64")
ap.add_argument("--base", type=str, default=None,
help="base SAE state_dict; defaults to the run that made it")
ap.add_argument("--real-seeds", type=int, nargs="*", default=[0, 1, 2, 3])
ap.add_argument("--null-seeds", type=int, nargs="*", default=[])
ap.add_argument("--epochs", type=int, default=20)
ap.add_argument("--lr", type=float, default=1e-4)
ap.add_argument("--batch-size", type=int, default=512)
ap.add_argument("--save-weights", action="store_true",
help="save decoder matrices at every step, for the "
"cloud-vs-cloud direction test")
ap.add_argument("--out", type=str, required=True)
ap.add_argument("--device", default="cuda")
args = ap.parse_args()
default_base = {6: "jr_L6_topk", 12: "jr_L12_topk",
18: "jumprelu_sft_topk", 23: "jr_L23_topk"}
if args.base is None:
args.base = str(DRIFT / default_base[args.layer] /
f"{args.arm}_L{args.layer}_base.pt")
out = DRIFT / args.out
out.mkdir(parents=True, exist_ok=True)
wdir = out / "weights"
if args.save_weights:
wdir.mkdir(exist_ok=True)
logf = open(out / "run.log", "a", buffering=1)
def log(m):
print(m, flush=True)
logf.write(m + "\n")
dev, cfg = args.device, ARMS[args.arm]
torch.backends.cuda.matmul.allow_tf32 = True
t0 = time.time()
log(f"\n{'='*70}\nreal-arm replication | layer {args.layer} | {args.arm}")
log(f"base {args.base}")
log(f"real seeds {args.real_seeds} | null seeds {args.null_seeds} | "
f"{args.epochs} epochs/step")
base_state = torch.load(args.base, map_location=dev)
base_tr, base_va = loaders("instruct_base", args.layer, args.batch_size, dev)
U = load_readout(str(DRIFT / "qwen05b/model.safetensors"), fold_final_norm=True,
device=dev)
P = compose_lens(U, load_lens(
str(DRIFT / f"drift_out/tuned_lens_layer{args.layer}.pt")).A.to(dev))
sae0 = SAE(cfg["arch"], cfg["k"]).to(dev)
sae0.load_state_dict(base_state)
Wd0 = sae0.decoder.weight.data.T.contiguous()
top0 = top_tokens(Wd0, P, k=10)
bmse, bl0, bdead = evaluate(sae0, base_va, dev)
log(f"base: val_mse={bmse:.6f} L0={bl0:.1f} dead={100*bdead:.1f}%")
if args.save_weights:
torch.save(Wd0.cpu(), wdir / "base.pt")
rows = []
def record(**kw):
rows.append(kw)
pd.DataFrame(rows).to_csv(out / "real_arm_replicate.csv", index=False)
def run_chain(chain, seed):
stages = ["instruct_base"] * len(SFT_STAGES) if chain == "null" else SFT_STAGES
state = base_state
for i, stage in enumerate(stages, start=1):
tr, va = (base_tr, base_va) if chain == "null" else \
loaders(stage, args.layer, args.batch_size, dev)
torch.manual_seed(1000 * seed + i)
s = SAE(cfg["arch"], cfg["k"]).to(dev)
s.load_state_dict(state)
log(f" -- {chain} seed{seed} step {i}/{len(stages)} ({stage})")
state, vmse, kept = fit(s, tr, va, args.epochs, args.lr, cfg["lam"],
dev, lambda m: None)
m = SAE(cfg["arch"], cfg["k"]).to(dev)
m.load_state_dict(state)
Wd = m.decoder.weight.data.T.contiguous()
r = matched_and_null(Wd0, Wd, None)
tok = token_overlap(top0, top_tokens(Wd, P, k=10))
_, l0, dead = evaluate(m, base_va, dev)
log(f" dec_cos={r['cos_matched_mean']:.4f} "
f"tok_J={tok['tok_jaccard_matched']:.4f} L0={l0:.1f} "
f"dead={100*dead:.1f}% kept_ep={kept}")
if args.save_weights:
torch.save(Wd.cpu(), wdir / f"{chain}_s{seed}_step{i}.pt")
record(layer=args.layer, arm=args.arm, chain=chain, seed=seed,
step=i, stage=stage, dec_cos=r["cos_matched_mean"],
drift=1.0 - r["cos_matched_mean"],
cos_null_pairing=r["cos_null_mean"],
drift_norm=r["drift_norm"],
tok_jaccard=tok["tok_jaccard_matched"],
frac_identical_top1=tok["frac_identical_top1"],
val_mse=vmse, mean_l0=l0, dead_frac=dead, kept_epoch=kept)
for sd in args.real_seeds:
run_chain("real", sd)
for sd in args.null_seeds:
run_chain("null", sd)
(out / "config.json").write_text(json.dumps(vars(args), indent=2))
log(f"\ndone in {(time.time()-t0)/60:.1f} min -> {out/'real_arm_replicate.csv'}")
if __name__ == "__main__":
main()