maxact-fast / scripts /cache_resids.py
ceselder's picture
scripts: attn sdpa (no flash-attn sm_103 build on Blackwell)
af8715e
Raw
History Blame Contribute Delete
3.48 kB
"""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()