Gala-598M-MLX / eval_longctx.py
junafinity's picture
Gala-598M: weights, code, logs, report
76bbe95 verified
Raw History Blame Contribute Delete
4.17 kB
"""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()