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