File size: 12,822 Bytes
17fc9a9
8505f8e
17fc9a9
320b589
17fc9a9
 
 
 
 
8505f8e
320b589
 
 
 
 
 
8505f8e
320b589
 
8505f8e
 
 
 
 
 
320b589
8505f8e
 
 
 
320b589
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
17fc9a9
 
8505f8e
 
 
320b589
 
8505f8e
 
 
17fc9a9
 
 
 
 
444e414
 
 
320b589
 
 
 
 
 
8505f8e
 
 
 
320b589
 
 
48db85f
320b589
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
444e414
320b589
 
 
 
 
 
 
 
 
 
 
444e414
320b589
8505f8e
 
 
17fc9a9
 
8505f8e
320b589
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
8505f8e
 
 
17fc9a9
 
 
 
 
 
8505f8e
 
 
 
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
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
"""Stage 3: mint (direction, target_texts) pairs — with REAL, VALIDATED probes.

For each sampled cluster pair (A,B) we TRAIN a logistic-regression probe on a train split of the
READ_LAYER residuals and only keep the direction if it actually works:
  - held-out separation AUC(A vs B) >= --min-auc   (else the "concept direction" is meaningless)
  - cluster-A val residuals project HIGH on the unit probe and clear cluster-B by a margin
    (else the direction isn't something the corpus strongly exhibits → nothing to maximize).
The kept unit direction is the conditioning vector; targets = centroid-closest cluster-A texts.
We log the AUC / projection-margin distributions and the fraction of pairs dropped.

OUT-OF-CORE: residuals come straight from the clustering stage's emb.f32 ([n_docs, d] memmap,
row-aligned with assign.npy / texts.jsonl — no separate resid cache). Cluster member rows are
gathered on demand through a bounded LRU (--cache-gb), texts through a byte-offset index, so RAM
never scales with corpus size. Probes are GPU-batched logistic regressions: --probe-batch pairs
fit in parallel per exact Newton/IRLS run (sklearn's objective: sum-BCE + 0.5/C·||w||²) —
replaces sklearn-per-pair, which is far too slow at millions of pairs.

    python scripts/build_data.py --acts-dir data/actclusters --n-examples 8000000 \
        --min-auc 0.9 --min-margin 1.0 --hard-neg-k 16
"""
import argparse
import json
import os

import numpy as np
import torch

from mxf.config import D_MODEL, BuildDataConfig


class TextBank:
    """Random access into texts.jsonl by row via a byte-offset index (corpus never in RAM)."""
    def __init__(self, path):
        self.f = open(path, "rb")
        off = [0]
        for line in self.f:
            off.append(off[-1] + len(line))
        self.off = np.asarray(off[:-1], dtype=np.int64)

    def __getitem__(self, i):
        self.f.seek(self.off[i])
        return json.loads(self.f.readline())["t"]


class ResidLRU:
    """Bounded LRU over per-cluster member-row gathers from the emb memmap."""
    def __init__(self, emb, budget_gb):
        self.emb, self.d, self.used, self.max = emb, {}, 0, int(budget_gb * 1e9)

    def get(self, c, rows):
        v = self.d.pop(c, None)
        if v is None:
            v = np.asarray(self.emb[rows])
            self.used += v.nbytes
            while self.used > self.max and self.d:
                self.used -= self.d.pop(next(iter(self.d))).nbytes
        self.d[c] = v   # (re)insert = most recent
        return v


@torch.no_grad()
def fit_probes_gpu(feats, C, iters, device):
    """One L2 logistic regression per (Atr, Btr, Ava, Bva) tuple, ALL fitted in parallel on GPU.
    Exact Newton/IRLS in the n-dim dual: with n<=~128 samples << d=4096, w* = Xᵀα (representer),
    and Woodbury keeps every solve at [P, n, n] fp64 — matches sklearn's optimum of
    sum-BCE + 0.5/C·||w||² (+ ~free intercept, δ=1e-4) in ~10 iterations. Padded batch + masks;
    val AUC / medians batched too. Returns per pair (unit_w, val_auc, projA_med, projB_med) or
    None (degenerate)."""
    P, d = len(feats), feats[0][0].shape[1]
    nt = max(len(A) + len(B) for A, B, _, _ in feats)
    va = max(1, max(len(v) for _, _, v, _ in feats)); vb = max(1, max(len(v) for _, _, _, v in feats))
    X = np.zeros((P, nt, d), np.float32); y = np.zeros((P, nt), np.float64); m = np.zeros((P, nt), np.float64)
    Xa = np.zeros((P, va, d), np.float32); ma = np.zeros((P, va), bool)
    Xb = np.zeros((P, vb, d), np.float32); mb = np.zeros((P, vb), bool)
    for i, (A, B, Av, Bv) in enumerate(feats):   # pad on CPU, one H2D per tensor (not per pair)
        X[i, : len(A)] = A; X[i, len(A) : len(A) + len(B)] = B
        y[i, : len(A)] = 1.0; m[i, : len(A) + len(B)] = 1.0
        Xa[i, : len(Av)] = Av; ma[i, : len(Av)] = True
        Xb[i, : len(Bv)] = Bv; mb[i, : len(Bv)] = True
    X, y, m, Xa, ma, Xb, mb = (torch.from_numpy(t).to(device) for t in (X, y, m, Xa, ma, Xb, mb))
    G = (C * (X @ X.transpose(1, 2))).double() + 1e4      # X̃D⁻¹X̃ᵀ with X̃=[X,1], D=diag(1/C..,δ)
    eye = torch.eye(nt, device=device, dtype=torch.float64)
    z = torch.zeros(P, nt, device=device, dtype=torch.float64)
    c = torch.zeros_like(z)
    for _ in range(iters):                                # damped-free Newton: θ' = -D⁻¹X̃ᵀc
        p = torch.sigmoid(z)
        r = (p - y) * m
        sh = ((p * (1 - p)).clamp(min=1e-12) * m).sqrt()
        q = (G @ r.unsqueeze(-1)).squeeze(-1) + z         # X̃D⁻¹∇: pads contribute 0 (r,c=0 there)
        A_ = eye + sh.unsqueeze(-1) * G * sh.unsqueeze(-2)
        v = torch.linalg.solve(A_, (sh * q).unsqueeze(-1)).squeeze(-1)
        c = r - sh * v
        z = -(G @ c.unsqueeze(-1)).squeeze(-1)
    w = -C * torch.einsum("pn,pnd->pd", c.float(), X)
    nrm = w.norm(dim=1)
    wu = w / nrm[:, None].clamp(min=1e-12)
    pa = torch.einsum("pvd,pd->pv", Xa, wu).masked_fill(~ma, torch.nan)
    pb = torch.einsum("pvd,pd->pv", Xb, wu).masked_fill(~mb, torch.nan)
    both = ma[:, :, None] & mb[:, None, :]
    gt = (pa[:, :, None] > pb[:, None, :]).double() + 0.5 * (pa[:, :, None] == pb[:, None, :]).double()
    auc = (gt.where(both, 0.0)).sum((1, 2)) / both.sum((1, 2)).clamp(min=1)
    medA, medB = pa.nanquantile(0.5, dim=1), pb.nanquantile(0.5, dim=1)
    nrm, WU, auc, medA, medB = (t.cpu().numpy() for t in (nrm, wu, auc, medA, medB))
    return [None if (nrm[i] < 1e-8 or not len(f[2]) or not len(f[3]))
            else (WU[i], float(auc[i]), float(medA[i]), float(medB[i])) for i, f in enumerate(feats)]


def main():
    cfg = BuildDataConfig()
    ap = argparse.ArgumentParser()
    ap.add_argument("--acts-dir", default="data/actclusters",
                    help="embed_cluster_acts out-dir (emb.f32 / assign.npy / centroids.npy / texts.jsonl)")
    ap.add_argument("--out-dir", default=cfg.out_dir)
    ap.add_argument("--n-examples", type=int, default=cfg.n_examples)
    ap.add_argument("--targets", type=int, default=cfg.targets_per_example)
    ap.add_argument("--min-auc", type=float, default=0.9, help="drop pairs the probe can't separate")
    ap.add_argument("--min-margin", type=float, default=1.0,
                    help="min (projA_med - projB_med) in resid units: cluster must activate the probe")
    ap.add_argument("--val-frac", type=float, default=0.3)
    ap.add_argument("--probe-c", type=float, default=cfg.probe_c)
    ap.add_argument("--hard-neg-k", type=int, default=0,
                    help="0 = random B (trivially-separable, cartoon directions). >0 = draw B from "
                         "A's k nearest clusters by centroid → subtle, information-rich probe directions.")
    ap.add_argument("--members-cap", type=int, default=64,
                    help="residual rows per cluster used to fit/val the probe")
    ap.add_argument("--probe-batch", type=int, default=512, help="cluster pairs fitted per GPU LR batch")
    ap.add_argument("--probe-iters", type=int, default=10, help="Newton/IRLS iterations per LR fit")
    ap.add_argument("--cache-gb", type=float, default=32.0, help="LRU budget for member-residual gathers")
    ap.add_argument("--device", default="cuda")
    a = ap.parse_args()
    os.makedirs(a.out_dir, exist_ok=True)
    rng = np.random.default_rng(cfg.seed)

    meta = json.load(open(f"{a.acts_dir}/meta.json"))
    K = meta["clusters"]
    emb = np.memmap(f"{a.acts_dir}/emb.f32", dtype=np.float32, mode="r",
                    shape=(meta["n_docs"], meta["d"]))   # -1 is illegal in memmap shapes
    assign = np.load(f"{a.acts_dir}/assign.npy")
    cent = np.load(f"{a.acts_dir}/centroids.npy")
    texts = TextBank(f"{a.acts_dir}/texts.jsonl")

    # group doc rows by cluster CSR-style: rows(c) = order[startx[c]:startx[c+1]] (no 1M py lists)
    order = np.argsort(assign, kind="stable")
    startx = np.zeros(K + 1, dtype=np.int64)
    np.cumsum(np.bincount(assign, minlength=K), out=startx[1:])

    # per usable cluster (>=8 members): shuffled member rows (probe fit/val, <= members-cap),
    # train/val split, target rows = centroid-closest among docs NOT in the probe-fit split —
    # targets come only from docs the probe never saw, so direction↔text can't be memorized
    mrows, split, tgt_rows, mus = {}, {}, {}, {}
    for c in range(K):
        rows = order[startx[c] : startx[c + 1]]
        if len(rows) < 8:
            continue
        perm = rows[rng.permutation(len(rows))]
        mem = perm[: a.members_cap]
        s = max(2, int(len(mem) * (1 - a.val_frac)))
        cand = np.sort(perm[s:])                          # all cluster docs minus probe-fit rows
        if not len(cand):
            continue
        dctr = np.linalg.norm(np.asarray(emb[cand]) - cent[c], axis=1)
        mrows[c], split[c], tgt_rows[c] = mem, s, cand[np.argsort(dctr)[:32]].tolist()
        if a.hard_neg_k > 0:
            mus[c] = np.asarray(emb[np.sort(mem[:s])]).mean(0)
        if len(mrows) % 50_000 == 0:
            print(f"  prepped {len(mrows)} clusters", flush=True)
    pool = np.array(sorted(mrows))
    nid = {int(c): i for i, c in enumerate(pool)}
    print(f"{len(pool)} clusters with >=8 members", flush=True)

    # hard-negative pairing: nearest clusters in QWEN3 layer-27 residual space (train-split mean
    # resid) — the probe/injection/reward live in Qwen3 residuals, so the genuinely hard negatives
    # are pairs the model itself struggles to linearly separate. GPU brute-force top-k.
    nearest = None
    if a.hard_neg_k > 0:
        mu = torch.from_numpy(np.stack([mus[c] for c in pool])).to(a.device)
        mn, mub = (mu * mu).sum(1), mu.to(torch.bfloat16)
        nearest = np.empty((len(pool), min(a.hard_neg_k, len(pool) - 1)), dtype=np.int64)
        cb = max(64, int(4e9 / (len(pool) * 4)))
        for s0 in range(0, len(pool), cb):
            d2 = mn[None] - 2 * (mub[s0 : s0 + cb] @ mub.T).float()
            d2[torch.arange(len(d2)), torch.arange(s0, s0 + len(d2))] = torch.inf   # exclude self
            nearest[s0 : s0 + len(d2)] = d2.topk(nearest.shape[1], dim=1, largest=False).indices.cpu().numpy()
        nearest = pool[nearest]   # pool index -> cluster id
        del mu, mub, mn, d2
        torch.cuda.empty_cache()

    lru = ResidLRU(emb, a.cache_gb)
    vec_bank = np.memmap(f"{a.out_dir}/vecs.f32", dtype=np.float32, mode="w+",
                         shape=(a.n_examples, D_MODEL))
    recs = open(f"{a.out_dir}/records.jsonl", "w")
    aucs, margins = [], []
    n, tried, dropped = 0, 0, 0
    while n < a.n_examples:
        pairs, feats = [], []
        while len(pairs) < a.probe_batch:                 # sample a GPU batch of cluster pairs
            A = int(rng.choice(pool))
            B = int(rng.choice(nearest[nid[A]])) if a.hard_neg_k > 0 else int(rng.choice(pool))
            if B == A:
                continue
            RA, RB = lru.get(A, mrows[A]), lru.get(B, mrows[B])
            sA, sB = split[A], split[B]
            pairs.append((A, B))
            feats.append((RA[:sA], RB[:sB], RA[sA:], RB[sB:]))
        for (A, B), pr in zip(pairs, fit_probes_gpu(feats, a.probe_c, a.probe_iters, a.device)):
            tried += 1
            if tried % 25_000 == 0:
                med = (f"auc med {np.median(aucs):.3f} margin med {np.median(margins):.2f}"
                       if aucs else "none kept yet")
                print(f"minted {n}/{a.n_examples} | kept {tried-dropped}/{tried} pairs "
                      f"(drop {dropped/tried:.0%}) | {med}", flush=True)
            if pr is None:
                dropped += 1; continue
            wu, auc, pA, pB = pr
            if auc < a.min_auc or (pA - pB) < a.min_margin:   # weak probe OR cluster doesn't activate it
                dropped += 1; continue
            aucs.append(auc); margins.append(pA - pB)
            for tr in rng.permutation(tgt_rows[A])[: a.targets]:
                if n >= a.n_examples:
                    break
                vec_bank[n] = wu
                recs.write(json.dumps({"vec_idx": n, "target_text": texts[int(tr)][:1200],
                                       "cluster": A, "val_auc": round(auc, 3),
                                       "proj_margin": round(pA - pB, 2)}) + "\n")
                n += 1
            if n >= a.n_examples:
                break
    recs.close(); vec_bank.flush()
    stats = {"n_examples": n, "pairs_tried": tried, "pairs_dropped": dropped,
             "drop_frac": dropped / max(tried, 1), "auc_median": float(np.median(aucs)),
             "auc_p10": float(np.percentile(aucs, 10)), "margin_median": float(np.median(margins)),
             "min_auc": a.min_auc, "min_margin": a.min_margin}
    json.dump(stats, open(f"{a.out_dir}/build_stats.json", "w"), indent=1)
    print(f"BUILD_DATA_DONE {stats}", flush=True)


if __name__ == "__main__":
    main()