maxact-fast / scripts /eval_sae.py
ceselder's picture
sae: loader + build_sae_data + eval_sae (cross-uplift metric); weights_only=False for torch 2.13
64c6b48
Raw
History Blame Contribute Delete
8.61 kB
"""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()