File size: 8,611 Bytes
64c6b48 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 | """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()
|