"""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 # drop positions with resid norm > mult * batch median (sink guard) @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} # <|im_end|> or <|endoftext|> (also pad) 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" # position 0 must be the first real token 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") # fp32 [b,T,d],[b,T] 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] # row i, feature i: [b,T] best = per.masked_fill(~keep, 0.0).max(1).values # relu>=0: 0-fill == old no-kept guard 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)) # [F] corpus peaks 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() # scoring must see the clean base; with adapter='none' there is nothing to disable 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()