Download tables/scripts/real_arm_replicate.py from sae-anon/sae-null-step: direct link, hf CLI and curl.
- Browser
- Download file 6.85 kB
-
https://huggingface.co/sae-anon/sae-null-step/resolve/main/tables/scripts/real_arm_replicate.py
- Command line
-
hf download hf://sae-anon/sae-null-step/tables/scripts/real_arm_replicate.py
-
curl -L -o real_arm_replicate.py https://huggingface.co/sae-anon/sae-null-step/resolve/main/tables/scripts/real_arm_replicate.py
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() | |