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()