"""Benchmark: factual generalization — the full 1b+semantic path. Measures factual correctness on UNSEEN slots, comparing: 1. char-level (baseline) — 0% expected 2. structural only (format generalization) — 0% factual expected 3. structural + embedding expert with a learned relation — factual gains Requires a fastText .vec file (download once): python -c "from ensemble import EmbeddingExpert; EmbeddingExpert.download_fasttext('data')" Usage: python bench/run_factual.py --vec data/wiki-news-300d-1M.vec --D 5000 """ from __future__ import annotations import argparse import sys from pathlib import Path sys.path.insert(0, str(Path(__file__).resolve().parent.parent)) from ensemble import Expert, Brain, EmbeddingExpert TRAIN_PAIRS = [ ("france", "paris"), ("germany", "berlin"), ("italy", "rome"), ("japan", "tokyo"), ("egypt", "cairo"), ("russia", "moscow"), ] HOLDOUT = [ ("spain", "madrid"), ("portugal", "lisbon"), ("greece", "athens"), ("norway", "oslo"), ("china", "beijing"), ("india", "delhi"), ("brazil", "brasilia"), ("turkey", "ankara"), ] PATTERN = "what is the capital of {country}" def make_qa(pairs): return [(f"what is the capital of {c}", f"the capital of {c} is {C}") for c, C in pairs] * 2 def run(mode, D, emb=None): qa = make_qa(TRAIN_PAIRS) if mode == "char": lex = Expert.from_qa_pairs(qa, domain="geo", D=D) brain = Brain(); brain.add_expert(lex) elif mode == "structural": lex = Expert.from_qa_pairs(qa, domain="geo", D=D, patterns=[PATTERN]) brain = Brain(); brain.add_expert(lex) else: # "factual" assert emb is not None emb.learn_relation("capital_of", {c: C for c, C in TRAIN_PAIRS}) lex = Expert.from_qa_pairs(qa, domain="geo", D=D, patterns=[PATTERN], embedding=emb) brain = Brain(); brain.add_expert(lex); brain.add_expert(emb) nonempty = 0 correct = 0 print(f"\n=== {mode.upper()} ===") for slot, truth in HOLDOUT: q = f"what is the capital of {slot}" a = brain.query(q, max_new_tokens=15).answer.strip().lower() if a: nonempty += 1 hit = a == truth or truth in a if hit: correct += 1 print(f" [{('Y' if hit else 'n')}] {slot:10} -> {a!r:16} (truth {truth})") n = len(HOLDOUT) print(f" non-empty: {nonempty}/{n} factual-correct: {correct}/{n}") return nonempty / n, correct / n def main() -> int: ap = argparse.ArgumentParser() ap.add_argument("--vec", required=True, help="fastText .vec path") ap.add_argument("--D", type=int, default=5000) ap.add_argument("--vocab", type=int, default=30000) args = ap.parse_args() print(f"loading fastText ({args.vec}, cap {args.vocab} words)...", flush=True) emb = EmbeddingExpert.from_fasttext(args.vec, D=args.D, max_vocab=args.vocab) print(f" {emb}", flush=True) results = {} for mode in ("char", "structural", "factual"): ne, fc = run(mode, args.D, emb if mode == "factual" else None) results[mode] = (ne, fc) print(f"\n{'='*60}") print(f"FACTUAL GENERALIZATION (capitals holdout, unseen countries)") print(f"{'mode':14} {'non-empty':>12} {'factual-correct':>16}") for mode, (ne, fc) in results.items(): print(f" {mode:12} {ne:>12.0%} {fc:>16.0%}") print(f"{'='*60}") return 0 if __name__ == "__main__": sys.exit(main())