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