"""Stage 2: per cluster, push M member texts through Qwen3 and cache layer-READ_LAYER residuals. The ONE expensive Qwen3 pass. Output: resids.f32 memmap [n_kept_members, d] + index mapping cluster -> member rows + the member texts (for centroid targets). Cheap probe fitting reads this. torchrun --standalone --nproc_per_node=8 scripts/cache_resids.py """ import argparse import json import os import numpy as np import torch from transformers import AutoModelForCausalLM, AutoTokenizer from mxf.config import D_MODEL, MODEL, READ_LAYER, ProbeCacheConfig from mxf.inject import read_resid def main(): cfg = ProbeCacheConfig() ap = argparse.ArgumentParser() ap.add_argument("--clusters-dir", default="data/clusters") ap.add_argument("--members", type=int, default=cfg.members_per_cluster) ap.add_argument("--pool", default=cfg.pool) ap.add_argument("--out-dir", default=cfg.out_dir) a = ap.parse_args() os.makedirs(a.out_dir, exist_ok=True) rank = int(os.environ.get("RANK", 0)); world = int(os.environ.get("WORLD_SIZE", 1)) local = int(os.environ.get("LOCAL_RANK", 0)) if world > 1: torch.distributed.init_process_group("nccl"); torch.cuda.set_device(local) device = f"cuda:{local}" meta = json.load(open(f"{a.clusters_dir}/meta.json")) assign = np.load(f"{a.clusters_dir}/assign.npy") texts = [json.loads(l)["t"] for l in open(f"{a.clusters_dir}/texts.jsonl")] K = meta["clusters"] # pick up to M member row-ids per cluster, sharded across ranks by cluster id rng = np.random.default_rng(0) members = {} order = np.argsort(assign) sorted_a = assign[order] bounds = np.searchsorted(sorted_a, np.arange(K + 1)) for c in range(rank, K, world): rows = order[bounds[c] : bounds[c + 1]] if len(rows) == 0: continue members[c] = rows[rng.permutation(len(rows))[: a.members]].tolist() tok = AutoTokenizer.from_pretrained(MODEL) if tok.pad_token is None: tok.pad_token = tok.eos_token tok.padding_side = "right" model = AutoModelForCausalLM.from_pretrained(MODEL, torch_dtype=torch.bfloat16, attn_implementation="sdpa", # flash-attn has no sm_103 build device_map={"": device}).eval() flat = [(c, r) for c, rows in members.items() for r in rows] out = np.memmap(f"{a.out_dir}/resids_rank{rank}.f32", dtype=np.float32, mode="w+", shape=(len(flat), D_MODEL)) idx = [] for s in range(0, len(flat), cfg.batch): chunk = flat[s : s + cfg.batch] enc = tok([texts[r][:1000] for _, r in chunk], padding=True, truncation=True, max_length=64, return_tensors="pt", add_special_tokens=True).to(device) v = read_resid(model, READ_LAYER, dict(enc), pool=a.pool).cpu().numpy() out[s : s + len(chunk)] = v for j, (c, r) in enumerate(chunk): idx.append({"row": s + j, "cluster": int(c), "text_row": int(r)}) if rank == 0 and s % (cfg.batch * 40) == 0: print(f"cached {s}/{len(flat)}", flush=True) out.flush() json.dump(idx, open(f"{a.out_dir}/index_rank{rank}.json", "w")) if rank == 0: print(f"RESID_CACHE_DONE rank0 {len(flat)} members", flush=True) if world > 1: torch.distributed.barrier(); torch.distributed.destroy_process_group() if __name__ == "__main__": main()