Download eval/bpb_bin.py from SlayerLab/gollem-v5-ckpts: direct link, hf CLI and curl.
- Browser
- Download file 3.6 kB
-
https://huggingface.co/SlayerLab/gollem-v5-ckpts/resolve/437a1d431ea63562ee5c4bbf065b22b14616aaaf/eval/bpb_bin.py
- Command line
-
hf download hf://SlayerLab/gollem-v5-ckpts@437a1d431ea63562ee5c4bbf065b22b14616aaaf/eval/bpb_bin.py
-
curl -L -o bpb_bin.py https://huggingface.co/SlayerLab/gollem-v5-ckpts/resolve/437a1d431ea63562ee5c4bbf065b22b14616aaaf/eval/bpb_bin.py
3.6 kB
| # 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() | |