gollem-v5-ckpts / eval /board_eval_fabryka.py
Maggio33's picture
Upload eval/board_eval_fabryka.py with huggingface_hub
5792063 verified
Raw History Blame
9.18 kB
# Board-native eval (protokol fabryka tiny_ml) na GPT (Qwen3-arch v1/v2). Hart N-02.
# BLiMP/ARC: no-BOS, raw-acc argmax, sliding-window @block_size, lm_eval pinned-rev.
# WikiText: byte_perplexity/BPB (parquet doc-level, fabryka loglikelihood_rolling).
# eff: tiny_ml_suite formula (BLiMP/ARC/normWiki x sizeMult).
import sys, os, json, argparse, math, re
from collections import namedtuple
import torch
from types import SimpleNamespace
from huggingface_hub import hf_hub_download
from tokenizers import Tokenizer
import lm_eval
from lm_eval.api.model import LM
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) if not os.path.exists(ckpt_rel) else 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)
n = sum(p.numel() for p in m.parameters())
return m, Tokenizer.from_file(tokp), c, n
class FabrykaLM(LM):
def __init__(self, model, tok, block, device, batch=256):
super().__init__()
self.model = model
self.tok = tok
self.context_length = block
self.dev = device
self.batch = batch
self.truncated = 0
self.total = 0
def _score_many(self, pairs):
totals = [0.0] * len(pairs)
greedy = [True] * len(pairs)
buckets = {}
for req, (ctx, cont) in enumerate(pairs):
prefix = self.tok.encode(ctx).ids or [self.tok.encode(" ").ids[0]]
target = self.tok.encode(cont).ids
self.total += 1
if len(prefix) + len(target) - 1 > self.context_length:
self.truncated += 1
tokens = prefix + target
for index, token in enumerate(target, len(prefix)):
window = tokens[max(0, index - self.context_length):index]
buckets.setdefault(len(window), []).append((req, window, token))
with torch.inference_mode():
for rows in buckets.values():
for s in range(0, len(rows), self.batch):
chunk = rows[s:s + self.batch]
x = torch.tensor([it[1] for it in chunk], dtype=torch.long, device=self.dev)
y = torch.tensor([it[2] for it in chunk], dtype=torch.long, device=self.dev)
logits = self.model(x)[0][:, -1, :].log_softmax(-1)
vals = logits.gather(1, y[:, None]).flatten().tolist()
gs = logits.argmax(-1).eq(y).tolist()
for (req, _, _), v, g in zip(chunk, vals, gs):
totals[req] += v
greedy[req] = greedy[req] and g
return list(zip(totals, greedy))
def loglikelihood(self, requests):
return self._score_many([tuple(r.args) for r in requests])
def loglikelihood_rolling(self, requests):
return [p[0] for p in self._score_many([("", r.args[0]) for r in requests])]
def generate_until(self, requests):
raise NotImplementedError
def wikitext_detokenizer(string):
string = string.replace("s '", "s'")
string = re.sub(r"/' [0-9]/", r"/'[0-9]/", string)
string = string.replace(" @-@ ", "-").replace(" @,@ ", ",").replace(" @.@ ", ".")
string = string.replace(" : ", ": ").replace(" ; ", "; ").replace(" . ", ". ")
string = string.replace(" ! ", "! ").replace(" ? ", "? ").replace(" , ", ", ")
string = re.sub(r"\(\s*([^\)]*?)\s*\)", r"(\1)", string)
string = re.sub(r"\[\s*([^\]]*?)\s*\]", r"[\1]", string)
string = re.sub(r"{\s*([^}]*?)\s*}", r"{\1}", string)
string = re.sub(r'"\s*([^"]*?)\s*"', r'"\1"', string)
string = re.sub(r"'\s*([^']*?)\s*'", r"'\1'", string)
string = string.replace("= = = =", "====").replace("= = =", "===").replace("= =", "==")
string = string.replace(" " + chr(176) + " ", chr(176))
string = string.replace(" \n", "\n").replace("\n ", "\n")
string = string.replace(" N ", " 1 ").replace(" 's", "'s")
return string
def _wikitext_pages(parquet_path):
import pyarrow.parquet as pq
texts = pq.read_table(parquet_path).column("text").to_pylist()
pages, cur = [], []
for s in texts:
if re.match(r"^ = [^=]", s):
if cur:
pages.append("".join(cur)); cur = []
cur.append(s)
if cur:
pages.append("".join(cur))
return [p for p in pages if p.strip()]
def compute_wikitext(lm, parquet_path, limit=None):
# Efficient teacher-forced reset-chunk BPB (block-boundary reset, Glint-1.3 style;
# ~block-faster than per-token sliding, byte_ppl diff <1% for block=1024).
pages = _wikitext_pages(parquet_path)
if limit:
pages = pages[:int(limit)]
block = lm.context_length
sum_nll = 0.0
sum_bytes = 0
sum_words = 0
with torch.inference_mode():
for p in pages:
det = wikitext_detokenizer(p)
ids = lm.tok.encode(det).ids
sum_bytes += len(p.encode("utf-8"))
sum_words += len(re.split(r"\s+", p))
for i in range(0, len(ids) - 1, block):
chunk = ids[i:i + block + 1]
if len(chunk) < 2:
continue
x = torch.tensor([chunk[:-1]], dtype=torch.long, device=lm.dev)
y = torch.tensor([chunk[1:]], dtype=torch.long, device=lm.dev)
logits = lm.model(x)[0]
nll = torch.nn.functional.cross_entropy(
logits.view(-1, logits.size(-1)), y.view(-1), reduction="sum")
sum_nll += nll.item()
return {"byte_perplexity": math.exp(sum_nll / sum_bytes),
"bits_per_byte": sum_nll / (sum_bytes * math.log(2)),
"word_perplexity": math.exp(sum_nll / sum_words),
"n_pages": len(pages), "n_bytes": sum_bytes,
"method": f"teacher-forced reset-chunk block={block}"}
def eff_score(blimp, arc, wiki_byte_ppl, params):
wiki_score = 100 * max(0, min(1, 1 - math.log(min(wiki_byte_ppl, 500) / 1.86) / math.log(500 / 1.86)))
overall = (100 * blimp + 100 * arc + wiki_score) / 3
size_pos = math.log(150_000_000 / params) / math.log(150_000_000 / 1000)
mult = 1 + 0.5 * max(0, min(1, size_pos))
return {"wiki_score": wiki_score, "overall": overall, "size_multiplier": mult,
"efficiency": overall * mult}
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--ckpt", required=True)
ap.add_argument("--blimp-limit", type=int, default=None)
ap.add_argument("--wikitext-parquet", default="/root/Salesforce_wt2_test.parquet")
ap.add_argument("--batch", type=int, default=256)
ap.add_argument("--wiki-only", action="store_true")
ap.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu")
ap.add_argument("--out", default=None)
a = ap.parse_args()
m, tok, c, nparams = load_model(a.ckpt, a.device)
print(f"[load] {a.ckpt} params={nparams} config={c} device={a.device}", flush=True)
lm = FabrykaLM(m, tok, c["block"], a.device, batch=a.batch)
if a.wiki_only:
wiki = compute_wikitext(lm, a.wikitext_parquet)
out = {"ckpt": a.ckpt, "params": nparams, "wikitext": wiki}
print("[RESULT]", json.dumps(out), flush=True)
if a.out:
json.dump(out, open(a.out, "w"), indent=2)
return
g = lm._score_many([("", "The cat is sleeping"), ("", "The cat are sleeping")])
ok = g[0][0] > g[1][0]
print(f"[sanity] prefers_grammatical={ok}", flush=True)
seeds = dict(num_fewshot=0, bootstrap_iters=0, random_seed=42, numpy_random_seed=42,
torch_random_seed=42, fewshot_random_seed=42)
res_b = lm_eval.simple_evaluate(model=lm, tasks=["blimp"], limit=a.blimp_limit, **seeds)
res_a = lm_eval.simple_evaluate(model=lm, tasks=["arc_easy"], limit=None, **seeds)
rb = res_b["results"]
r = res_a["results"]
bl = [k for k in rb if k.startswith("blimp_")]
blimp = sum(rb[k]["acc,none"] for k in bl) / len(bl) if bl else rb.get("blimp", {}).get("acc,none")
arc = r["arc_easy"]["acc,none"]
wiki = None
if a.wikitext_parquet and os.path.exists(a.wikitext_parquet):
wiki = compute_wikitext(lm, a.wikitext_parquet)
out = {"ckpt": a.ckpt, "params": nparams, "blimp": blimp, "blimp_leaves": len(bl),
"arc_easy_raw_acc": arc, "arc_easy_acc_norm": r["arc_easy"].get("acc_norm,none"),
"wikitext": wiki, "sanity_minpair": ok, "truncated": lm.truncated,
"blimp_limit": a.blimp_limit}
if wiki:
out["eff"] = eff_score(blimp, arc, wiki["byte_perplexity"], nparams)
print("[RESULT]", json.dumps(out), flush=True)
if a.out:
json.dump(out, open(a.out, "w"), indent=2)
if __name__ == "__main__":
main()