"""Long-context evals for a checkpoint: (1) per-position CE over 8192-token windows — if carried state helps, late positions beat early ones; (2) passkey retrieval scored by log-prob against distractor keys at several depths. python eval_longctx.py --run runs/c_gala_s0 --out runs/c_gala_s0/longctx.json """ import argparse, json, os import numpy as np import mlx.core as mx import mlx.nn as nn from model import HOARD, HoardConfig try: import tiktoken ENC = tiktoken.get_encoding("gpt2") except Exception: ENC = None def load(run_dir): meta = json.load(open(os.path.join(run_dir, "config.json"))) cfg = HoardConfig.from_dict(meta["config"]) model = HOARD(cfg) model.load_weights(os.path.join(run_dir, "model.safetensors")) model.eval() return model def position_ce(model, toks, seq=8192, n_windows=6): bins = [(0, 256), (256, 1024), (1024, 2048), (2048, 4096), (4096, 8192)] sums = np.zeros(len(bins)); cnts = np.zeros(len(bins)) for w in range(n_windows): st = w * (seq + 1) x = np.asarray(toks[st: st + seq + 1]).astype(np.int32)[None] if x.shape[1] < seq + 1: break x = mx.array(x) logits = model(x[:, :-1]) ce = nn.losses.cross_entropy(logits.astype(mx.float32), x[:, 1:], reduction="none")[0] mx.eval(ce) ce = np.array(ce) for i, (a, b) in enumerate(bins): sums[i] += ce[a:b].sum(); cnts[i] += b - a return {f"{a}-{b}": round(float(s / c), 4) for (a, b), s, c in zip(bins, sums, cnts)} def seq_logprob(model, ids, tail_ids): """log P(tail | prefix) under the model, ids = prefix+tail.""" x = mx.array(np.array(ids, dtype=np.int32))[None] logits = model(x[:, :-1]).astype(mx.float32) lp = nn.log_softmax(logits[0], axis=-1) n = len(tail_ids) tgt = x[0, 1:] sel = mx.take_along_axis(lp, tgt[:, None], axis=-1)[:, 0] mx.eval(sel) return float(np.array(sel)[-n:].sum()) def passkey(model, toks, seq=8192, depths=(0.1, 0.5, 0.9), n_trials=6, seed=0): if ENC is None: return {"error": "tiktoken unavailable"} rng = np.random.default_rng(seed) out = {} for depth in depths: correct = 0 for t in range(n_trials): key = int(rng.integers(10000, 99999)) distractors = [int(rng.integers(10000, 99999)) for _ in range(4)] needle = ENC.encode_ordinary(f" The pass key is {key}. Remember the pass key.") query = ENC.encode_ordinary(" The pass key is") filler_start = int(rng.integers(0, len(toks) - seq - 1)) filler = list(np.asarray(toks[filler_start: filler_start + seq]).astype(int)) pos = int(depth * (seq - len(needle) - len(query) - 24)) ctx = filler[:pos] + needle + filler[pos: seq - len(needle) - len(query) - 16] scores = [] for cand in [key] + distractors: tail = ENC.encode_ordinary(f" {cand}") scores.append(seq_logprob(model, ctx + query + tail, tail)) correct += int(np.argmax(scores) == 0) out[f"depth_{depth}"] = round(correct / n_trials, 3) return out def main(): ap = argparse.ArgumentParser() ap.add_argument("--run", required=True) ap.add_argument("--data", default="data/fineweb/val.bin") ap.add_argument("--out", default=None) ap.add_argument("--trials", type=int, default=6) ap.add_argument("--depths", type=float, nargs="+", default=[0.1, 0.5, 0.9]) ap.add_argument("--seqs", type=int, nargs="+", default=[8192, 32768], help="context lengths for the passkey test") a = ap.parse_args() model = load(a.run) toks = np.memmap(a.data, dtype=np.uint16, mode="r") res = {"position_ce": position_ce(model, toks), "passkey_acc": {str(sq): passkey(model, toks, seq=sq, depths=tuple(a.depths), n_trials=a.trials) for sq in a.seqs}} out = a.out or os.path.join(a.run, "longctx.json") json.dump(res, open(out, "w"), indent=2) print(json.dumps(res)) if __name__ == "__main__": main()