Maggio33 commited on
Commit
94b0512
·
verified ·
1 Parent(s): 5792063

Upload eval/bpb_bin.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. eval/bpb_bin.py +86 -0
eval/bpb_bin.py ADDED
@@ -0,0 +1,86 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Held-out byte-normalized BPB on a uint16 token .bin (forward-hold collapse-test). Hart N-02.
2
+ # Teacher-forced NLL over non-overlapping block windows; byte-denom = decoded UTF-8 bytes.
3
+ import sys, os, json, argparse, math
4
+ import numpy as np
5
+ import torch
6
+ from types import SimpleNamespace
7
+ from huggingface_hub import hf_hub_download
8
+ from tokenizers import Tokenizer
9
+
10
+ R = "SlayerLab/gollem-v5-ckpts"
11
+
12
+
13
+ def load_model(ckpt_rel, device):
14
+ mdef = hf_hub_download(R, "train_gpt_ref.py")
15
+ sys.path.insert(0, os.path.dirname(mdef))
16
+ from train_gpt_ref import GPT
17
+ tokp = hf_hub_download(R, "tokenizer.json")
18
+ ck = hf_hub_download(R, ckpt_rel)
19
+ d = torch.load(ck, map_location="cpu")
20
+ c = d["config"]
21
+ cfg = SimpleNamespace(**c)
22
+ m = GPT(c["vocab"], c["n_layer"], c["n_embd"], c["n_head"], c["block"], cfg)
23
+ miss, unexp = m.load_state_dict(d["model"], strict=False)
24
+ assert not unexp and miss in ([], ["head.weight"]), (miss, unexp)
25
+ m.eval().to(device)
26
+ return m, Tokenizer.from_file(tokp), c
27
+
28
+
29
+ def bpb_on_bin(m, tok, binpath, block, device, batch=16, max_tokens=None):
30
+ data = np.fromfile(binpath, dtype=np.uint16)
31
+ if max_tokens:
32
+ data = data[:int(max_tokens)]
33
+ n = len(data)
34
+ # byte-denominator: decode whole stream once (tokenizer-agnostic normalization)
35
+ total_bytes = 0
36
+ for j in range(0, n, 200000):
37
+ total_bytes += len(tok.decode(data[j:j + 200000].tolist()).encode("utf-8"))
38
+ print(f" [decode-done] bytes={total_bytes} tokens={n}", flush=True)
39
+ windows = []
40
+ for i in range(0, n - 1, block):
41
+ chunk = data[i:i + block + 1]
42
+ if len(chunk) < 2:
43
+ break
44
+ windows.append(chunk)
45
+ total_nll = 0.0
46
+ total_tok = 0
47
+ with torch.inference_mode():
48
+ for s in range(0, len(windows), batch):
49
+ grp = windows[s:s + batch]
50
+ L = min(len(w) for w in grp)
51
+ xs = torch.tensor(np.stack([w[:L][:-1].astype(np.int64) for w in grp]), device=device)
52
+ ys = torch.tensor(np.stack([w[:L][1:].astype(np.int64) for w in grp]), device=device)
53
+ logits = m(xs)[0]
54
+ nll = torch.nn.functional.cross_entropy(
55
+ logits.reshape(-1, logits.size(-1)), ys.reshape(-1), reduction="sum")
56
+ total_nll += nll.item()
57
+ total_tok += ys.numel()
58
+ return {"bin": os.path.basename(binpath), "tokens": total_tok, "bytes": total_bytes,
59
+ "nll_per_tok": total_nll / total_tok,
60
+ "byte_ppl": math.exp(total_nll / total_bytes),
61
+ "bpb": total_nll / (total_bytes * math.log(2))}
62
+
63
+
64
+ def main():
65
+ ap = argparse.ArgumentParser()
66
+ ap.add_argument("--ckpt", required=True)
67
+ ap.add_argument("--bins", required=True, help="comma-separated HF-rel paths under repo, e.g. vals/fwe_only_val.bin")
68
+ ap.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu")
69
+ ap.add_argument("--max-tokens", type=int, default=None)
70
+ ap.add_argument("--out", default=None)
71
+ a = ap.parse_args()
72
+ m, tok, c = load_model(a.ckpt, a.device)
73
+ print(f"[load] {a.ckpt} params={sum(p.numel() for p in m.parameters())} device={a.device}", flush=True)
74
+ res = {"ckpt": a.ckpt, "bpb": {}}
75
+ for rel in a.bins.split(","):
76
+ rel = rel.strip()
77
+ local = hf_hub_download(R, rel)
78
+ r = bpb_on_bin(m, tok, local, c["block"], a.device, max_tokens=a.max_tokens)
79
+ res["bpb"][os.path.basename(rel)] = r
80
+ print("[BPB]", rel, json.dumps(r), flush=True)
81
+ if a.out:
82
+ json.dump(res, open(a.out, "w"), indent=2)
83
+
84
+
85
+ if __name__ == "__main__":
86
+ main()