File size: 4,170 Bytes
76bbe95
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
"""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()