| """ |
| run_eval.py -- unified benchmark harness reproducing the Escarda-family model-card |
| suite for BOTH model families in one script: |
| |
| * SpikeWhale V2 (config.SpikeWhaleConfig / model_v2.SpikeWhaleLM) -> Byrne-15M |
| * Quazimoto (model.QuazimotoLM via the repo's own generate.load_model) |
| -> Positronic-144M-Anti-DEG, Wheeler-DeWitt-62M (hybrid) |
| |
| Metrics (identical definitions / scoring to Mark2/eval_*.py, so numbers line up |
| with the published cards): |
| |
| WikiText-2 byte_ppl = exp( sum_token_NLL_nats / total_UTF8_bytes ) (down) |
| BLiMP acc = frac. of minimal pairs with logp(good) > logp(bad), |
| 12 paradigms x 150, full-sentence sum-logprob w/ <bos> (up) |
| MC suite = argmax continuation log-likelihood; acc (raw sum) and |
| acc_norm (sum / UTF-8 bytes). Winogrande/BoolQ report acc only. |
| ArithMark-3.0 = HellaSwag-format MC arithmetic; acc + acc_norm, n & chance. |
| |
| *_stderr = sqrt(p(1-p)/n) (binomial). |
| |
| Usage: |
| python run_eval.py --family quazimoto --code_dir <repo> --ckpt <ckpt.pt> --name Positronic-144M |
| python run_eval.py --family spikewhale --code_dir <Mark2> --ckpt checkpoints/best --name Byrne-15M-Looped |
| """ |
| import os, sys, math, json, ast, argparse |
| from collections import defaultdict |
|
|
| import datasets |
| import torch |
| import torch.nn.functional as F |
|
|
| ARITHMARK_REPO = "AxiomicLabs/ArithMark-3.0" |
|
|
|
|
| |
| |
| |
| class _LogitsOut: |
| __slots__ = ("logits",) |
| def __init__(self, logits): |
| self.logits = logits |
|
|
|
|
| class QuazWrap: |
| """Wrap a QuazimotoLM so it looks like an HF model: model(x).logits. |
| Quazimoto's forward is model(idx)->(logits, loss, aux); positional idx only.""" |
| def __init__(self, model): |
| self.m = model |
| def __call__(self, input_ids=None, use_cache=False, **kw): |
| out = self.m(input_ids) |
| return _LogitsOut(out[0]) |
|
|
|
|
| def load_tokenizer(tok_path): |
| cand = tok_path if tok_path.endswith(".json") else os.path.join(tok_path, "tokenizer.json") |
| from spike_tokenizer import SpikeTokenizer |
| return SpikeTokenizer(vocab_file=cand) |
|
|
|
|
| def load_family(family, code_dir, ckpt, device): |
| """Returns (model_callable, tokenizer, block_size, meta_str).""" |
| sys.path.insert(0, code_dir) |
| tok = load_tokenizer(os.path.join(code_dir, "tokenizer.json")) |
|
|
| if family == "spikewhale": |
| from config import SpikeWhaleConfig |
| from model_v2 import SpikeWhaleLM |
| if os.path.isdir(ckpt): |
| |
| |
| |
| |
| from safetensors.torch import load_file |
| cfg = SpikeWhaleConfig.from_pretrained(ckpt) |
| model = SpikeWhaleLM(cfg).to(device).eval() |
| state = load_file(os.path.join(ckpt, "model.safetensors"), device=str(device)) |
| missing, unexpected = model.load_state_dict(state, strict=False) |
| if missing: print(f" [warn] missing keys: {len(missing)} (e.g. {missing[:3]})") |
| if unexpected: print(f" [warn] unexpected keys: {len(unexpected)} (e.g. {unexpected[:3]})") |
| else: |
| c = torch.load(ckpt, map_location=device, weights_only=False) |
| cfg = SpikeWhaleConfig(**c["config"]) |
| model = SpikeWhaleLM(cfg).to(device).eval() |
| model.load_state_dict(c["model_state"], strict=False) |
| block = getattr(cfg, "max_position_embeddings", 1024) |
| return model, tok, block, f"spikewhale block={block}" |
|
|
| elif family == "quazimoto": |
| |
| from generate import load_model as q_load |
| model, cfg = q_load(ckpt, device) |
| |
| if hasattr(model, "reset_ring_memory"): |
| model.reset_ring_memory() |
| if hasattr(model, "set_ring_memory_writing"): |
| model.set_ring_memory_writing(False) |
| block = getattr(cfg, "block_size", 1024) |
| return QuazWrap(model), tok, block, f"quazimoto block={block}" |
|
|
| raise ValueError(family) |
|
|
|
|
| |
| |
| |
| @torch.no_grad() |
| def continuation_logprob(model, tok, context, continuation, device, max_len): |
| ctx_ids = tok.encode(context, add_special_tokens=False) |
| cont_ids = tok.encode(continuation, add_special_tokens=False) |
| if len(cont_ids) == 0: |
| return -1e9, 1, max(1, len(continuation)) |
| full = ctx_ids + cont_ids |
| if len(full) > max_len: |
| full = full[len(full) - max_len:] |
| |
| if len(cont_ids) > max_len: |
| cont_ids = cont_ids[len(cont_ids) - max_len:] |
| x = torch.tensor([full], dtype=torch.long, device=device) |
| logits = model(input_ids=x, use_cache=False).logits |
| logprobs = F.log_softmax(logits[:, :-1, :].float(), dim=-1) |
| targets = x[:, 1:] |
| tok_lp = logprobs.gather(-1, targets.unsqueeze(-1)).squeeze(-1)[0] |
| cont_lp = tok_lp[-len(cont_ids):] |
| return float(cont_lp.sum().item()), len(cont_ids), max(1, len(continuation)) |
|
|
|
|
| def stderr(p, n): |
| return math.sqrt(p * (1.0 - p) / n) if n > 0 else 0.0 |
|
|
|
|
| |
| |
| |
| def iter_piqa(limit): |
| from datasets import load_dataset |
| ds = load_dataset("parquet", |
| data_files="hf://datasets/ybisk/piqa@refs/convert/parquet/plain_text/validation/*.parquet", |
| split="train") |
| for n, ex in enumerate(ds): |
| if limit and n >= limit: break |
| yield ex["goal"] + " ", [ex["sol1"], ex["sol2"]], int(ex["label"]) |
|
|
| def iter_hellaswag(limit): |
| from datasets import load_dataset |
| ds = load_dataset("parquet", |
| data_files="hf://datasets/Rowan/hellaswag@refs/convert/parquet/default/validation/*.parquet", |
| split="train") |
| for n, ex in enumerate(ds): |
| if limit and n >= limit: break |
| ctx = (ex["activity_label"] + ": " + ex["ctx_a"] + " " + ex["ctx_b"]).strip() + " " |
| gold = int(ex["label"]) if ex["label"] != "" else 0 |
| yield ctx, ex["endings"], gold |
|
|
| def _label_index(labels, key): |
| if key in labels: return labels.index(key) |
| alt = {"1": "A", "2": "B", "3": "C", "4": "D", "5": "E"} |
| return labels.index(alt.get(key, key)) |
|
|
| def iter_arc(config, limit): |
| from datasets import load_dataset |
| ds = load_dataset("allenai/ai2_arc", config, split="test") |
| for n, ex in enumerate(ds): |
| if limit and n >= limit: break |
| labels = list(ex["choices"]["label"]) |
| ctx = "Question: " + ex["question"].strip() + "\nAnswer:" |
| choices = [" " + t for t in ex["choices"]["text"]] |
| yield ctx, choices, _label_index(labels, ex["answerKey"]) |
|
|
| def iter_winogrande(limit): |
| from datasets import load_dataset |
| ds = load_dataset("allenai/winogrande", "winogrande_xl", split="validation") |
| n = 0 |
| for ex in ds: |
| s = ex["sentence"] |
| if "_" not in s: continue |
| prefix, suffix = s.split("_", 1) |
| yield prefix, [ex["option1"] + suffix, ex["option2"] + suffix], int(ex["answer"]) - 1 |
| n += 1 |
| if limit and n >= limit: break |
|
|
| def iter_openbookqa(limit): |
| from datasets import load_dataset |
| ds = load_dataset("allenai/openbookqa", "main", split="test") |
| for n, ex in enumerate(ds): |
| if limit and n >= limit: break |
| labels = list(ex["choices"]["label"]) |
| ctx = "Question: " + ex["question_stem"].strip() + "\nAnswer:" |
| choices = [" " + t for t in ex["choices"]["text"]] |
| yield ctx, choices, _label_index(labels, ex["answerKey"]) |
|
|
| def iter_boolq(limit): |
| from datasets import load_dataset |
| ds = load_dataset("google/boolq", split="validation") |
| for n, ex in enumerate(ds): |
| if limit and n >= limit: break |
| ctx = ex["passage"].strip() + "\nQuestion: " + ex["question"].strip() + "?\nAnswer:" |
| yield ctx, [" no", " yes"], int(bool(ex["answer"])) |
|
|
|
|
| TASKS = { |
| "arc_easy": (lambda lim: iter_arc("ARC-Easy", lim), True), |
| "arc_challenge": (lambda lim: iter_arc("ARC-Challenge", lim), True), |
| "hellaswag": (lambda lim: iter_hellaswag(lim), True), |
| "winogrande": (lambda lim: iter_winogrande(lim), False), |
| "piqa": (lambda lim: iter_piqa(lim), True), |
| "openbookqa": (lambda lim: iter_openbookqa(lim), True), |
| "boolq": (lambda lim: iter_boolq(lim), False), |
| } |
| MC_ORDER = ["arc_easy", "arc_challenge", "hellaswag", "winogrande", "piqa", |
| "openbookqa", "boolq"] |
|
|
|
|
| @torch.no_grad() |
| def run_mc_task(name, iterator, model, tok, device, max_len): |
| correct = correct_norm = total = 0 |
| for ctx, choices, gold in iterator: |
| raw, norm = [], [] |
| for c in choices: |
| lp, ntok, nchar = continuation_logprob(model, tok, ctx, c, device, max_len) |
| raw.append(lp) |
| norm.append(lp / max(1, len(c.encode("utf-8")))) |
| correct += int(max(range(len(raw)), key=lambda i: raw[i]) == gold) |
| correct_norm += int(max(range(len(norm)), key=lambda i: norm[i]) == gold) |
| total += 1 |
| if total % 500 == 0: |
| print(f" [{name}] {total} acc={correct/total:.3f} acc_norm={correct_norm/total:.3f}", flush=True) |
| return correct / total, correct_norm / total, total |
|
|
|
|
| |
| |
| |
| def _parse_list(v): |
| return v if isinstance(v, list) else ast.literal_eval(v) |
|
|
| def _meta(ex): |
| m = ex.get("metadata") |
| if isinstance(m, dict): return m |
| try: return ast.literal_eval(m) |
| except Exception: return {} |
|
|
| def _difficulty(ex): |
| d = _meta(ex).get("difficulty") |
| if d: return d |
| parts = (ex.get("activity_label") or "").split("::") |
| return parts[2] if len(parts) >= 3 else "unknown" |
|
|
| @torch.no_grad() |
| def run_arithmark(model, tok, device, max_len, limit): |
| from datasets import load_dataset |
| ds = load_dataset(ARITHMARK_REPO, split="train") |
| correct = correct_norm = total = 0 |
| by_diff = defaultdict(lambda: [0, 0]) |
| for n, ex in enumerate(ds): |
| if limit and n >= limit: break |
| ctx = ex["ctx"]; choices = _parse_list(ex["endings"]); gold = int(ex["label"]) |
| raw, norm = [], [] |
| for c in choices: |
| lp, ntok, nchar = continuation_logprob(model, tok, ctx, c, device, max_len) |
| raw.append(lp); norm.append(lp / max(1, len(c.encode("utf-8")))) |
| ok = int(max(range(len(norm)), key=lambda i: norm[i]) == gold) |
| correct += int(max(range(len(raw)), key=lambda i: raw[i]) == gold) |
| correct_norm += ok |
| d = _difficulty(ex); by_diff[d][0] += ok; by_diff[d][1] += 1 |
| total += 1 |
| if total % 500 == 0: |
| print(f" [arithmark] {total} acc={correct/total:.3f} acc_norm={correct_norm/total:.3f}", flush=True) |
| return correct / total, correct_norm / total, total, dict(by_diff) |
|
|
|
|
| |
| |
| |
| @torch.no_grad() |
| def wikitext2_byte_ppl(model, tok, device, max_chars, ctx, stride): |
| from datasets import load_dataset |
| ds = load_dataset("wikitext", "wikitext-2-raw-v1", split="test") |
| text = "".join(r["text"] for r in ds)[:max_chars] |
| ids = tok.encode(text, add_special_tokens=False) |
| bos = getattr(tok, "bos_token_id", None) |
| if bos is not None: ids = [bos] + ids |
| total_nll, n_scored, prev_end = 0.0, 0, 0 |
| for begin in range(0, len(ids), stride): |
| end = min(begin + ctx, len(ids)) |
| trg = end - prev_end |
| x = torch.tensor([ids[begin:end]], device=device) |
| logits = model(input_ids=x, use_cache=False).logits |
| lp = F.log_softmax(logits[:, :-1].float(), -1) |
| tok_lp = lp.gather(-1, x[:, 1:].unsqueeze(-1)).squeeze(-1)[0] |
| sel = tok_lp[-trg:] if 0 < trg <= tok_lp.numel() else tok_lp |
| total_nll += float(-sel.sum().item()); n_scored += sel.numel() |
| prev_end = end |
| if end == len(ids): break |
| n_bytes = max(1, len(text.encode("utf-8"))) |
| return math.exp(total_nll / n_bytes), n_bytes, n_scored |
|
|
| @torch.no_grad() |
| def seq_logprob(model, tok, sentence, device): |
| ids = tok.encode(sentence, add_special_tokens=False) |
| bos = getattr(tok, "bos_token_id", None) |
| if bos is not None: ids = [bos] + ids |
| x = torch.tensor([ids], device=device) |
| logits = model(input_ids=x, use_cache=False).logits |
| lp = F.log_softmax(logits[:, :-1].float(), -1) |
| return float(lp.gather(-1, x[:, 1:].unsqueeze(-1)).squeeze(-1)[0].sum().item()) |
|
|
| @torch.no_grad() |
| def blimp(model, tok, device, n_paradigms, limit): |
| from datasets import load_dataset, get_dataset_config_names |
| cfgs = get_dataset_config_names("nyu-mll/blimp") |
| step = max(1, len(cfgs) // n_paradigms) |
| chosen = cfgs[::step][:n_paradigms] |
| per, tot_c, tot_n = {}, 0, 0 |
| for cfg in chosen: |
| ds = load_dataset("nyu-mll/blimp", cfg, split="train") |
| c = n = 0 |
| for ex in ds: |
| if limit and n >= limit: break |
| g = seq_logprob(model, tok, ex["sentence_good"], device) |
| b = seq_logprob(model, tok, ex["sentence_bad"], device) |
| c += int(g > b); n += 1 |
| per[cfg] = c / n; tot_c += c; tot_n += n |
| print(f" [blimp] {cfg:<42} {c/n:.3f} ({c}/{n})", flush=True) |
| return tot_c / tot_n, per |
|
|
|
|
| |
| def main(): |
| ap = argparse.ArgumentParser() |
| ap.add_argument("--family", required=True, choices=["quazimoto", "spikewhale"]) |
| ap.add_argument("--code_dir", required=True) |
| ap.add_argument("--ckpt", required=True) |
| ap.add_argument("--name", required=True) |
| ap.add_argument("--out", required=True, help="output JSON path") |
| ap.add_argument("--limit", type=int, default=None, help="cap per MC task (debug)") |
| ap.add_argument("--arith-limit", type=int, default=None) |
| ap.add_argument("--wt-chars", type=int, default=300_000) |
| ap.add_argument("--blimp-paradigms", type=int, default=12) |
| ap.add_argument("--blimp-limit", type=int, default=150) |
| ap.add_argument("--skip", nargs="*", default=[], help="stages to skip: mc lm arithmark") |
| args = ap.parse_args() |
|
|
| device = "cuda" if torch.cuda.is_available() else "cpu" |
| model, tok, block, meta = load_family(args.family, args.code_dir, args.ckpt, device) |
| max_len = min(2048, block) |
| wt_ctx = min(1024, block) |
| print(f"\n===== {args.name} ({meta}, max_len={max_len}) =====", flush=True) |
|
|
| R = {"name": args.name, "family": args.family, "ckpt": args.ckpt, |
| "block_size": block, "arithmark": ARITHMARK_REPO, "results": {}} |
|
|
| if "lm" not in args.skip: |
| print("\n--- WikiText-2 byte_ppl ---", flush=True) |
| bppl, nb, ns = wikitext2_byte_ppl(model, tok, device, args.wt_chars, wt_ctx, wt_ctx // 2) |
| print(f" byte_ppl = {bppl:.4f} ({ns} tok / {nb} bytes)", flush=True) |
| print("\n--- BLiMP ---", flush=True) |
| bacc, bper = blimp(model, tok, device, args.blimp_paradigms, args.blimp_limit) |
| print(f" blimp_acc = {bacc:.4f}", flush=True) |
| R["results"]["wikitext2_byte_ppl"] = bppl |
| R["results"]["blimp_acc"] = bacc |
| R["results"]["blimp_per_paradigm"] = bper |
|
|
| if "mc" not in args.skip: |
| mc = {} |
| for name in MC_ORDER: |
| itfac, want_norm = TASKS[name] |
| print(f"\n--- {name} ---", flush=True) |
| acc, acc_norm, n = run_mc_task(name, itfac(args.limit), model, tok, device, max_len) |
| print(f"[{name}] n={n} acc={acc:.4f} acc_norm={acc_norm:.4f}", flush=True) |
| entry = {"acc": acc, "acc_stderr": stderr(acc, n), "n": n} |
| if want_norm: |
| entry["acc_norm"] = acc_norm |
| entry["acc_norm_stderr"] = stderr(acc_norm, n) |
| mc[name] = entry |
| R["results"]["mc"] = mc |
|
|
| if "arithmark" not in args.skip: |
| print(f"\n--- ArithMark-3.0 ---", flush=True) |
| a, an, n, bd = run_arithmark(model, tok, device, max_len, args.arith_limit) |
| print(f"[arithmark] n={n} acc={a:.4f} acc_norm={an:.4f}", flush=True) |
| R["results"]["arithmark"] = {"acc": a, "acc_stderr": stderr(a, n), |
| "acc_norm": an, "acc_norm_stderr": stderr(an, n), |
| "n": n, "chance": 0.25, "by_difficulty": bd} |
|
|
| os.makedirs(os.path.dirname(os.path.abspath(args.out)), exist_ok=True) |
| with open(args.out, "w", encoding="utf-8") as f: |
| json.dump(R, f, indent=2) |
| print(f"\nwrote {args.out}", flush=True) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|