File size: 3,595 Bytes
94b0512
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
# Held-out byte-normalized BPB on a uint16 token .bin (forward-hold collapse-test). Hart N-02.
# Teacher-forced NLL over non-overlapping block windows; byte-denom = decoded UTF-8 bytes.
import sys, os, json, argparse, math
import numpy as np
import torch
from types import SimpleNamespace
from huggingface_hub import hf_hub_download
from tokenizers import Tokenizer

R = "SlayerLab/gollem-v5-ckpts"


def load_model(ckpt_rel, device):
    mdef = hf_hub_download(R, "train_gpt_ref.py")
    sys.path.insert(0, os.path.dirname(mdef))
    from train_gpt_ref import GPT
    tokp = hf_hub_download(R, "tokenizer.json")
    ck = hf_hub_download(R, ckpt_rel)
    d = torch.load(ck, map_location="cpu")
    c = d["config"]
    cfg = SimpleNamespace(**c)
    m = GPT(c["vocab"], c["n_layer"], c["n_embd"], c["n_head"], c["block"], cfg)
    miss, unexp = m.load_state_dict(d["model"], strict=False)
    assert not unexp and miss in ([], ["head.weight"]), (miss, unexp)
    m.eval().to(device)
    return m, Tokenizer.from_file(tokp), c


def bpb_on_bin(m, tok, binpath, block, device, batch=16, max_tokens=None):
    data = np.fromfile(binpath, dtype=np.uint16)
    if max_tokens:
        data = data[:int(max_tokens)]
    n = len(data)
    # byte-denominator: decode whole stream once (tokenizer-agnostic normalization)
    total_bytes = 0
    for j in range(0, n, 200000):
        total_bytes += len(tok.decode(data[j:j + 200000].tolist()).encode("utf-8"))
    print(f"  [decode-done] bytes={total_bytes} tokens={n}", flush=True)
    windows = []
    for i in range(0, n - 1, block):
        chunk = data[i:i + block + 1]
        if len(chunk) < 2:
            break
        windows.append(chunk)
    total_nll = 0.0
    total_tok = 0
    with torch.inference_mode():
        for s in range(0, len(windows), batch):
            grp = windows[s:s + batch]
            L = min(len(w) for w in grp)
            xs = torch.tensor(np.stack([w[:L][:-1].astype(np.int64) for w in grp]), device=device)
            ys = torch.tensor(np.stack([w[:L][1:].astype(np.int64) for w in grp]), device=device)
            logits = m(xs)[0]
            nll = torch.nn.functional.cross_entropy(
                logits.reshape(-1, logits.size(-1)), ys.reshape(-1), reduction="sum")
            total_nll += nll.item()
            total_tok += ys.numel()
    return {"bin": os.path.basename(binpath), "tokens": total_tok, "bytes": total_bytes,
            "nll_per_tok": total_nll / total_tok,
            "byte_ppl": math.exp(total_nll / total_bytes),
            "bpb": total_nll / (total_bytes * math.log(2))}


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--ckpt", required=True)
    ap.add_argument("--bins", required=True, help="comma-separated HF-rel paths under repo, e.g. vals/fwe_only_val.bin")
    ap.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu")
    ap.add_argument("--max-tokens", type=int, default=None)
    ap.add_argument("--out", default=None)
    a = ap.parse_args()
    m, tok, c = load_model(a.ckpt, a.device)
    print(f"[load] {a.ckpt} params={sum(p.numel() for p in m.parameters())} device={a.device}", flush=True)
    res = {"ckpt": a.ckpt, "bpb": {}}
    for rel in a.bins.split(","):
        rel = rel.strip()
        local = hf_hub_download(R, rel)
        r = bpb_on_bin(m, tok, local, c["block"], a.device, max_tokens=a.max_tokens)
        res["bpb"][os.path.basename(rel)] = r
        print("[BPB]", rel, json.dumps(r), flush=True)
    if a.out:
        json.dump(res, open(a.out, "w"), indent=2)


if __name__ == "__main__":
    main()