| """Stage 7: CROSS-UPLIFT eval — does cluster-direction training transfer to unseen SAE features? |
| |
| Per held-out feature: condition on unit(W_enc[:,f]) (inject@INJECT_LAYER at the marker), generate |
| greedy + best-of-N at T via HF generate, then score every text STANDALONE through the CLEAN base |
| (adapter disabled, no injection) at READ_LAYER: SAE-encode feature f, max over kept positions |
| (pos-0 attention-sink + norm-outlier guards — the old SAE line's exact protocol, so numbers are |
| directly comparable). dataset_max = the feature's corpus peak from the max-acts dump. |
| |
| python scripts/eval_sae.py --adapter checkpoints/pretrain/final --split data/sae/split.json \ |
| --best-of 16 --out results/xuplift_pretrain.json |
| python scripts/eval_sae.py --adapter none --n-features 256 --out results/xuplift_base.json |
| """ |
| import argparse |
| import contextlib |
| import json |
| import os |
| import random |
|
|
| import torch |
| from peft import PeftModel |
| from transformers import AutoModelForCausalLM, AutoTokenizer |
|
|
| from mxf.config import INJECT_LAYER, MODEL, READ_LAYER, STEER_COEFF |
| from mxf.inject import get_layer, hooked, make_inject_hook, read_resid |
| from mxf.prompts import build_prompt_ids |
| from mxf.sae import load_max_acts, load_sae |
|
|
| NORM_FILTER_MULT = 10.0 |
|
|
|
|
| @torch.no_grad() |
| def generate(model, tok, prompt_ids, marker, dirs, a, device, greedy): |
| """One batched HF generate() with per-row direction injection. All prompts are the identical |
| token sequence, so `marker` is prefill-absolute; decode steps (len-1 forwards under KV cache) |
| skip the hook — see make_inject_hook.""" |
| B = dirs.shape[0] |
| ids = torch.tensor([prompt_ids] * B, dtype=torch.long, device=device) |
| attn = torch.ones_like(ids, dtype=torch.bool) |
| hook = make_inject_hook([dirs[i : i + 1] for i in range(B)], [[marker]] * B, |
| STEER_COEFF, device, torch.bfloat16) |
| kw = dict(max_new_tokens=a.max_new_tokens, min_new_tokens=a.min_new_tokens, |
| pad_token_id=tok.pad_token_id, do_sample=not greedy, use_cache=True) |
| if not greedy: |
| kw.update(temperature=a.temperature, top_p=1.0, top_k=0) |
| with hooked(get_layer(model, INJECT_LAYER), hook): |
| out = model.generate(input_ids=ids, attention_mask=attn, **kw) |
| stop = {tok.eos_token_id, tok.pad_token_id} |
| texts = [] |
| for row in out: |
| comp = row[len(prompt_ids):].tolist() |
| stops = [j for j, t in enumerate(comp) if t in stop] |
| texts.append(tok.decode(comp[: stops[0]] if stops else comp, |
| skip_special_tokens=True).strip()) |
| return texts |
|
|
|
|
| @torch.no_grad() |
| def score(texts, feats, model, clean, tok, sae, device, a): |
| """enc_act[i] = max_t relu((x_t-b_dec)·W_enc[:,f_i]+b_enc[f_i]) at READ_LAYER — standalone |
| re-tokenization (no chat template), clean base. Empty/all-filtered rows score 0.""" |
| r = torch.zeros(len(texts)) |
| valid = [i for i, t in enumerate(texts) if t.strip()] |
| prev = tok.padding_side |
| tok.padding_side = "right" |
| try: |
| for s in range(0, len(valid), a.score_batch): |
| idxs = valid[s : s + a.score_batch] |
| enc = tok([texts[i] for i in idxs], return_tensors="pt", padding=True, truncation=True, |
| max_length=a.max_new_tokens + 32, add_special_tokens=False).to(device) |
| if enc["input_ids"].shape[1] == 0: |
| continue |
| with clean(): |
| h, mask = read_resid(model, READ_LAYER, dict(enc), pool="all") |
| norms = h.norm(dim=-1) |
| keep = mask & (norms <= NORM_FILTER_MULT * norms[mask].median()) |
| keep[:, 0] = False |
| b = torch.arange(len(idxs), device=device) |
| per = sae.encode_features(h, [feats[i] for i in idxs])[b, :, b] |
| best = per.masked_fill(~keep, 0.0).max(1).values |
| for row, i in enumerate(idxs): |
| r[i] = best[row].item() |
| finally: |
| tok.padding_side = prev |
| return r |
|
|
|
|
| def main(): |
| ap = argparse.ArgumentParser() |
| ap.add_argument("--adapter", required=True, help="maxact-fast checkpoint dir, or 'none' for base") |
| ap.add_argument("--sae-path", default=None, help="ae.pt path; default: hf_hub_download") |
| ap.add_argument("--maxacts-path", default=None, help="max-acts .pt path; default: hf_hub_download") |
| ap.add_argument("--split", default=None, help="split.json from build_sae_data; uses its 'eval' half") |
| ap.add_argument("--n-features", type=int, default=0, |
| help="cap the split's eval list; without --split, seeded sample of alive features") |
| ap.add_argument("--best-of", type=int, default=16) |
| ap.add_argument("--temperature", type=float, default=1.0) |
| ap.add_argument("--batch-features", type=int, default=8, help="sampled gen batch = this * best-of") |
| ap.add_argument("--max-new-tokens", type=int, default=96) |
| ap.add_argument("--min-new-tokens", type=int, default=16) |
| ap.add_argument("--score-batch", type=int, default=128) |
| ap.add_argument("--seed", type=int, default=0) |
| ap.add_argument("--out", required=True) |
| a = ap.parse_args() |
| assert a.split or a.n_features > 0, "need --split or --n-features" |
| torch.manual_seed(a.seed) |
| device = "cuda:0" |
|
|
| tok = AutoTokenizer.from_pretrained(MODEL) |
| if tok.pad_token is None: |
| tok.pad_token = tok.eos_token |
| prompt_ids, mpos = build_prompt_ids(tok) |
| marker = mpos[0] |
|
|
| dataset_max = load_max_acts(a.maxacts_path)["max_acts"].amax(dim=(1, 2)) |
| if a.split: |
| feats = json.load(open(a.split))["eval"] |
| if a.n_features: |
| feats = feats[: a.n_features] |
| else: |
| alive = (dataset_max > 0).nonzero(as_tuple=True)[0].tolist() |
| feats = sorted(random.Random(a.seed).sample(alive, min(a.n_features, len(alive)))) |
|
|
| sae = load_sae(a.sae_path, device) |
| model = AutoModelForCausalLM.from_pretrained(MODEL, torch_dtype=torch.bfloat16, |
| attn_implementation="sdpa", device_map={"": device}) |
| if a.adapter != "none": |
| model = PeftModel.from_pretrained(model, a.adapter) |
| model.eval() |
| |
| clean = model.disable_adapter if a.adapter != "none" else contextlib.nullcontext |
| print(f"{len(feats)} eval features | adapter {a.adapter} | marker @{marker}", flush=True) |
|
|
| results = [] |
| for s in range(0, len(feats), a.batch_features): |
| fb = feats[s : s + a.batch_features] |
| g_texts = generate(model, tok, prompt_ids, marker, sae.enc_dirs(fb), a, device, greedy=True) |
| flat = [f for f in fb for _ in range(a.best_of)] |
| s_texts = generate(model, tok, prompt_ids, marker, sae.enc_dirs(flat), a, device, greedy=False) |
| g_act = score(g_texts, fb, model, clean, tok, sae, device, a) |
| s_act = score(s_texts, flat, model, clean, tok, sae, device, a).view(len(fb), a.best_of) |
| for i, f in enumerate(fb): |
| bi = int(s_act[i].argmax()) |
| results.append({"feature": int(f), "dataset_max": dataset_max[f].item(), |
| "greedy_text": g_texts[i], "greedy_act": g_act[i].item(), |
| "best_text": s_texts[i * a.best_of + bi], "best_act": s_act[i, bi].item()}) |
| print(f"{len(results)}/{len(feats)} features", flush=True) |
|
|
| dmax = torch.tensor([max(r["dataset_max"], 1e-6) for r in results]) |
| g = torch.tensor([r["greedy_act"] for r in results]) |
| b = torch.tensor([r["best_act"] for r in results]) |
| summary = { |
| "adapter": a.adapter, "n_features": len(results), "best_of": a.best_of, |
| "temperature": a.temperature, |
| "greedy/normalized_act": (g / dmax).mean().item(), |
| "greedy/normalized_act_median": (g / dmax).median().item(), |
| "greedy/beat_frac": (g > dmax).float().mean().item(), |
| f"best_of_{a.best_of}/normalized_act": (b / dmax).mean().item(), |
| f"best_of_{a.best_of}/normalized_act_median": ((b / dmax).median().item()), |
| f"best_of_{a.best_of}/beat_frac": (b > dmax).float().mean().item(), |
| } |
| print(json.dumps(summary, indent=2)) |
| os.makedirs(os.path.dirname(a.out) or ".", exist_ok=True) |
| json.dump({"summary": summary, "results": results}, open(a.out, "w")) |
| print("EVAL_SAE_DONE", flush=True) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|