| """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"] |
|
|
| |
| 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", |
| 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() |
|
|