File size: 3,483 Bytes
8505f8e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
af8715e
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
"""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()