MBG-1.0 / source /evaluate_quality.py
deeprcurs-staff's picture
Upload source/evaluate_quality.py with huggingface_hub
8191d39 verified
Raw History Blame Contribute Delete
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())