""" 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(""): 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())