gollem-v5-ckpts / eval /bpb_bin.py
Maggio33's picture
Upload eval/bpb_bin.py with huggingface_hub
94b0512 verified
Raw History Blame
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()