File size: 4,386 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
"""Build an SAE-feature post-train dataset in the exact vec-bank format pretrain.py consumes.

Per feature: direction = unit(W_enc[:,f]) (encoder column, probe side — the same side eval_sae.py
rewards); targets = the feature's top --targets corpus windows (by per-window peak act) decoded to
text. Features are split (seeded, --eval-frac) into train/eval and ONLY train features enter
vecs.f32/records.jsonl — the eval half is reserved for cross-uplift (eval_sae.py --split
<out>/split.json). build_stats.json carries n_examples so scripts/rl.py --data-dir <out> also
works on this bank unchanged.

    python scripts/build_sae_data.py --out-dir data/sae --targets 3
"""
import argparse
import json
import os
import random

import numpy as np
import torch
from transformers import AutoTokenizer

from mxf.config import D_MODEL, MODEL
from mxf.sae import load_max_acts, load_sae


def main():
    ap = argparse.ArgumentParser()
    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("--out-dir", default="data/sae")
    ap.add_argument("--n-features", type=int, default=0, help="0 = ALL alive features; else seeded sample")
    ap.add_argument("--targets", type=int, default=3, help="top corpus windows per feature")
    ap.add_argument("--min-act", type=float, default=0.0, help="drop features with corpus peak <= this")
    ap.add_argument("--eval-frac", type=float, default=0.5)
    ap.add_argument("--seed", type=int, default=0)
    a = ap.parse_args()
    os.makedirs(a.out_dir, exist_ok=True)
    rng = random.Random(a.seed)

    tok = AutoTokenizer.from_pretrained(MODEL)
    sae = load_sae(a.sae_path)                                # cpu fp32
    data = load_max_acts(a.maxacts_path)
    tokens, acts = data["max_tokens"], data["max_acts"]       # [F, N, L]
    assert acts.shape[0] == sae.d_sae, f"max-acts F={acts.shape[0]} != SAE F={sae.d_sae}"
    dataset_max = acts.amax(dim=(1, 2))                       # [F] corpus peak per feature

    alive = (dataset_max > a.min_act).nonzero(as_tuple=True)[0].tolist()
    if a.n_features and a.n_features < len(alive):
        alive = rng.sample(alive, a.n_features)
    rng.shuffle(alive)
    n_eval = int(len(alive) * a.eval_frac)
    ev, train = sorted(alive[:n_eval]), sorted(alive[n_eval:])
    json.dump({"seed": a.seed, "eval_frac": a.eval_frac, "min_act": a.min_act,
               "train": train, "eval": ev}, open(f"{a.out_dir}/split.json", "w"))
    print(f"{len(alive)} alive features (min_act {a.min_act}) -> {len(train)} train / {len(ev)} eval",
          flush=True)

    dirs = torch.nn.functional.normalize(sae.W_enc, dim=0)    # [d, F] unit encoder columns
    rows, skipped = [], 0
    for j, f in enumerate(train):
        peak = acts[f].amax(dim=-1)                           # [N] per-window peaks
        for w in peak.argsort(descending=True)[: a.targets].tolist():
            if peak[w] <= 0:
                break                                         # desc order: rest are dead too
            text = tok.decode(tokens[f, w].tolist(), skip_special_tokens=True).strip()
            if len(text) < 3:
                skipped += 1
                continue
            rows.append((f, text, peak[w].item()))
        if (j + 1) % 5000 == 0:
            print(f"{j + 1}/{len(train)} features, {len(rows)} records", flush=True)

    assert rows, "no records minted — check --min-act / max-acts file"
    vecs = np.memmap(f"{a.out_dir}/vecs.f32", dtype=np.float32, mode="w+",
                     shape=(len(rows), D_MODEL))
    with open(f"{a.out_dir}/records.jsonl", "w") as recs:
        for n, (f, text, act) in enumerate(rows):             # one row per (feature, target)
            vecs[n] = dirs[:, f].numpy()
            recs.write(json.dumps({"vec_idx": n, "target_text": text, "feature": f,
                                   "act": round(act, 3)}) + "\n")
    vecs.flush()
    stats = {"n_examples": len(rows), "n_train_features": len(train), "n_eval_features": len(ev),
             "targets_per_feature": a.targets, "skipped_short": skipped, "seed": a.seed}
    json.dump(stats, open(f"{a.out_dir}/build_stats.json", "w"), indent=1)
    print(f"BUILD_SAE_DATA_DONE {stats}", flush=True)


if __name__ == "__main__":
    main()