Download source/evaluate_quality.py from deeprcurs/MBG-1.0: direct link, hf CLI and curl.
- Browser
- Download file 10.8 kB
-
https://huggingface.co/deeprcurs/MBG-1.0/resolve/main/source/evaluate_quality.py
- Command line
-
hf download hf://deeprcurs/MBG-1.0/source/evaluate_quality.py
-
curl -L -o evaluate_quality.py https://huggingface.co/deeprcurs/MBG-1.0/resolve/main/source/evaluate_quality.py
10.8 kB
| """ | |
| MBG 1.0 — src/evaluate_quality.py | |
| Reusable OUTPUT-QUALITY evaluation harness for MBG checkpoints. | |
| Beyond plain loss, this measures how good the model's *outputs* actually are: | |
| * Held-out perplexity, overall + per Trinity-Mirror domain (CAUSAL/SPATIAL/ | |
| TEMPORAL/GENERAL). | |
| * Trinity-Mirror probe disambiguation accuracy (does the domain probe steer the | |
| correct reading?). | |
| * Generation quality on domain prefixes: lexical diversity (distinct-1/2), | |
| mean length, repetition (fluency proxy), mean top-1 confidence. | |
| * JSON + Markdown report per run, and a side-by-side table when multiple | |
| checkpoints are compared (e.g. quality-vs-scale). | |
| The tokenizer is re-trained on the SAME corpus + vocab_size as training, which is | |
| deterministic, so encodings match training. Checkpoints can be bf16 goldens. | |
| English-only (project rule). | |
| Usage: | |
| /home/user/.venv/bin/python src/evaluate_quality.py \ | |
| --checkpoint checkpoints/golden/mbg_l1-17m_20260901-173235.pt \ | |
| --data data/english_L1.txt --samples 8 --gen_len 20 --seed 0 | |
| # multiple checkpoints for a comparison table: | |
| --checkpoint A.pt --checkpoint B.pt ... | |
| """ | |
| from __future__ import annotations | |
| import argparse, json, math, os, random, sys, time | |
| import torch | |
| import torch.nn.functional as F | |
| sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) | |
| from src.mbg_mini_gpt import ( | |
| MbGPT, build_data_from_file, N_PROBES, PROBE2IDX, DOMAIN2PROBE, | |
| ) | |
| from src.tokenizer_bpe import BpeTokenizer | |
| from src.gen_probe_challenge import AMBIG | |
| DOMAINS = ["CAUSAL", "SPATIAL", "TEMPORAL", "GENERAL"] | |
| # --------------------------------------------------------------------------- | |
| # model / tokenizer loading | |
| # --------------------------------------------------------------------------- | |
| def load_model(ckpt_path: str) -> tuple: | |
| ck = torch.load(ckpt_path, map_location="cpu") | |
| cfg = ck["config"] | |
| model = MbGPT( | |
| vocab_size=cfg["vocab_size"], n_embd=cfg["n_embd"], | |
| n_head=cfg.get("n_head", 4), n_layer=cfg["n_layer"], | |
| ffn_dim=cfg["ffn_dim"], n_expert=cfg["n_expert"], | |
| n_probe=cfg.get("n_probe", N_PROBES), block_size=cfg["block_size"], | |
| ) | |
| model.load_state_dict(ck["model_state"], strict=False) | |
| model.eval() | |
| return model, cfg | |
| def make_encoder(cfg, texts, data_path) -> BpeTokenizer: | |
| if cfg.get("tokenizer") == "char": | |
| from src.mbg_mini_gpt import CharTokenizer | |
| return CharTokenizer("".join(texts)) | |
| enc = BpeTokenizer(vocab_size=max(cfg["vocab_size"], 16)) | |
| enc.train(texts) | |
| return enc | |
| # --------------------------------------------------------------------------- | |
| # held-out split (stratified by domain, seeded) — matches dataset build | |
| # --------------------------------------------------------------------------- | |
| def val_split(lines, seed=0, val_frac=0.10): | |
| rng = random.Random(seed) | |
| by = {} | |
| for t, d in lines: | |
| by.setdefault(d, []).append((t, d)) | |
| val = [] | |
| for d, items in by.items(): | |
| idx = list(range(len(items))) | |
| rng.shuffle(idx) | |
| n = max(1, int(round(len(items) * val_frac))) | |
| for i in idx[:n]: | |
| val.append(items[i]) | |
| return val | |
| def seq_loss(model, enc, text, domain): | |
| ids = enc.encode(text, add_bos_eos=True) | |
| if len(ids) < 3: | |
| return None | |
| x = torch.tensor(ids[:-1], dtype=torch.long).unsqueeze(0) | |
| y = torch.tensor(ids[1:], dtype=torch.long).unsqueeze(0) | |
| with torch.no_grad(): | |
| lg, _ = model(x, probe_idx=PROBE2IDX[DOMAIN2PROBE[domain]]) | |
| return F.cross_entropy(lg.transpose(1, 2), y).item() | |
| def per_domain_ppl(model, enc, val, max_examples=600): | |
| by = {d: [] for d in DOMAINS} | |
| for t, d in val: | |
| by.setdefault(d, []).append((t, d)) | |
| out = {} | |
| for d, items in by.items(): | |
| losses, n = 0.0, 0 | |
| for t, _ in items[:max_examples]: | |
| l = seq_loss(model, enc, t, d) | |
| if l is not None: | |
| losses += l; n += 1 | |
| mean = losses / max(n, 1) | |
| out[d] = {"n": n, "loss": round(mean, 4), | |
| "ppl": round(math.exp(min(mean, 20)), 3)} | |
| return out | |
| # --------------------------------------------------------------------------- | |
| # probe disambiguation (reuse AMBIG readings) | |
| # --------------------------------------------------------------------------- | |
| def disambiguation(model, enc): | |
| total = correct = 0 | |
| for head, readings in AMBIG.items(): | |
| doms = list(readings.keys()) | |
| for domA in doms: | |
| for domB in doms: | |
| if domA == domB: | |
| continue | |
| fullA = f"{head}{readings[domA]}" | |
| lA = seq_loss(model, enc, fullA, domA) | |
| lB = seq_loss(model, enc, fullA, domB) | |
| if lA is not None and lB is not None: | |
| total += 1 | |
| if lA < lB: | |
| correct += 1 | |
| return correct, total | |
| # --------------------------------------------------------------------------- | |
| # generation quality | |
| # --------------------------------------------------------------------------- | |
| def generate(model, enc, prefix, domain, length, temperature=1.0): | |
| ids = enc.encode(prefix, add_bos_eos=True) | |
| x = torch.tensor(ids, dtype=torch.long).unsqueeze(0) | |
| out = list(ids) | |
| with torch.no_grad(): | |
| for _ in range(length): | |
| if len(out) >= model.block_size: | |
| break | |
| xx = torch.tensor(out, dtype=torch.long).unsqueeze(0)[:, -model.block_size:] | |
| lg, _ = model(xx, probe_idx=PROBE2IDX[DOMAIN2PROBE[domain]]) | |
| logits = lg[0, -1] / max(temperature, 1e-6) | |
| if temperature < 0.9: # greedy-ish / low temp | |
| nxt = int(logits.argmax().item()) | |
| else: | |
| p = F.softmax(logits, dim=-1) | |
| nxt = int(torch.multinomial(p, 1).item()) | |
| out.append(nxt) | |
| if nxt == enc._special_id("<eos>"): | |
| break | |
| return out | |
| def distinct(ids, n): | |
| grams = set() | |
| for i in range(len(ids) - n + 1): | |
| grams.add(tuple(ids[i:i + n])) | |
| return len(grams) / max(1, len(ids) - n + 1) | |
| def generation_metrics(model, enc, lines, samples, gen_len, seed=0): | |
| rng = random.Random(seed) | |
| by = {d: [] for d in DOMAINS} | |
| for t, d in lines: | |
| by.setdefault(d, []).append((t, d)) | |
| metrics = {} | |
| for d in DOMAINS: | |
| cands = [t for t, _ in by[d]] | |
| texts, confs = [], [] | |
| for _ in range(samples): | |
| prefix = rng.choice(cands).split()[:3] | |
| prefix = " ".join(prefix) | |
| ids = generate(model, enc, prefix, d, gen_len, temperature=0.8) | |
| texts.append(enc.decode(ids)) | |
| # metrics on the generated token streams | |
| all_ids = [enc.encode(t) for t in texts] | |
| lens = [len(i) for i in all_ids] | |
| d1 = sum(len(set(i)) for i in all_ids) / max(1, sum(len(i) for i in all_ids)) | |
| d2 = sum(len(set(tuple(i[j:j+2]) for j in range(len(i)-1))) for i in all_ids) \ | |
| / max(1, sum(max(0, len(i)-1) for i in all_ids)) | |
| metrics[d] = { | |
| "n": len(texts), | |
| "distinct_1": round(d1, 4), | |
| "distinct_2": round(d2, 4), | |
| "mean_len": round(sum(lens) / max(1, len(lens)), 2), | |
| } | |
| return metrics | |
| # --------------------------------------------------------------------------- | |
| # driver | |
| # --------------------------------------------------------------------------- | |
| def evaluate(ckpt_path, data_path, samples, gen_len, seed): | |
| lines = build_data_from_file(data_path) | |
| texts = [t for t, _ in lines] | |
| model, cfg = load_model(ckpt_path) | |
| enc = make_encoder(cfg, texts, data_path) | |
| val = val_split(lines, seed) | |
| r = {"checkpoint": ckpt_path, "params": sum(p.numel() for p in model.parameters())} | |
| t0 = time.time() | |
| r["per_domain_ppl"] = per_domain_ppl(model, enc, val) | |
| c, t = disambiguation(model, enc) | |
| r["disambiguation"] = {"correct": c, "total": t, | |
| "acc": round(c / max(t, 1), 4)} | |
| r["generation"] = generation_metrics(model, enc, lines, samples, gen_len, seed) | |
| r["wall_s"] = round(time.time() - t0, 1) | |
| return r, cfg | |
| def main(argv=None) -> int: | |
| ap = argparse.ArgumentParser() | |
| ap.add_argument("--checkpoint", action="append", required=True) | |
| ap.add_argument("--data", default="data/english_L1.txt") | |
| ap.add_argument("--samples", type=int, default=8) | |
| ap.add_argument("--gen_len", type=int, default=20) | |
| ap.add_argument("--seed", type=int, default=0) | |
| ap.add_argument("--tag", default="quality") | |
| args = ap.parse_args(argv) | |
| torch.set_num_threads(2) | |
| results = [] | |
| for ck in args.checkpoint: | |
| r, cfg = evaluate(ck, args.data, args.samples, args.gen_len, args.seed) | |
| results.append(r) | |
| print(f"[eval] {ck} params={r['params']:,} wall={r['wall_s']}s") | |
| # markdown + json report | |
| os.makedirs("checkpoints", exist_ok=True) | |
| ts = time.strftime("%Y%m%d-%H%M%S") | |
| js = os.path.join("checkpoints", f"QUALITY-{args.tag}-{ts}.json") | |
| with open(js, "w") as fh: | |
| json.dump(results, fh, indent=2) | |
| md = [f"# MBG 1.0 — Output Quality Evaluation ({args.tag})\n", | |
| f"*Timestamp: {ts} · data={args.data} · samples={args.samples} " | |
| f"gen_len={args.gen_len} seed={args.seed}*\n"] | |
| md.append("| Checkpoint | Params | PPL(CAUSAL) | PPL(SPATIAL) | PPL(TEMP) | " | |
| "PPL(GEN) | Disambig acc | d1 | d2 |") | |
| md.append("|---|---|---|---|---|---|---|---|---|") | |
| for r in results: | |
| p = r["per_domain_ppl"] | |
| g = r["generation"] | |
| d1 = sum(v["distinct_1"] for v in g.values()) / 4 | |
| d2 = sum(v["distinct_2"] for v in g.values()) / 4 | |
| name = os.path.basename(r["checkpoint"]) | |
| md.append(f"| {name} | {r['params']/1e6:.2f}M | " | |
| f"{p['CAUSAL']['ppl']} | {p['SPATIAL']['ppl']} | " | |
| f"{p['TEMPORAL']['ppl']} | {p['GENERAL']['ppl']} | " | |
| f"{r['disambiguation']['acc']*100:.1f}% | {d1:.3f} | {d2:.3f} |") | |
| md.append("") | |
| md.append("### Per-domain perplexity detail") | |
| for r in results: | |
| md.append(f"\n**{os.path.basename(r['checkpoint'])}**") | |
| for d, v in r["per_domain_ppl"].items(): | |
| md.append(f"- {d}: loss={v['loss']} ppl={v['ppl']} (n={v['n']})") | |
| md.append(f"- disambiguation: {r['disambiguation']['correct']}/" | |
| f"{r['disambiguation']['total']} = " | |
| f"{r['disambiguation']['acc']*100:.1f}%") | |
| rpt = os.path.join("checkpoints", f"REPORT-QUALITY-{args.tag}-{ts}.md") | |
| with open(rpt, "w") as fh: | |
| fh.write("\n".join(md) + "\n") | |
| print(f"[report] {rpt}") | |
| print(f"[json] {js}") | |
| return 0 | |
| if __name__ == "__main__": | |
| sys.exit(main()) | |