maxact-fast / scripts /build_data.py
ceselder's picture
out-of-core rework: embed_cluster_acts streams resids to disk memmap + streaming multi-GPU k-means (--kmeans-sample); build_data reads unified emb.f32 + GPU-batched Newton/IRLS probes
320b589
Raw
History Blame Contribute Delete
12.8 kB
"""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()