# 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()