#!/usr/bin/env python3 """G-BASE's comparator, fixed before training (PLAN-deem-T1-tune-v2.md section 5.2). Stdlib only. A multinomial naive Bayes over unigram tokens `[a-z']+` (lowercased), Laplace alpha = 1, the empirical class prior, no masking, plain argmax (an exact tie goes to the first guest in kit order, R-1's rule). Trained on the T-poc pool (or, for a fold, the pool minus the held-out table). The named-guest exclusion variant (argmax over guests the line does not name, S0's construction) is printed beside it, never gated. python3 -B baseline_nb.py --set s0 # rows/t1-baseline-nb.s0.P.rep1.jsonl (G-BASE's registered read) python3 -B baseline_nb.py --set h48 # rows/t1-baseline-nb.h48.P.rep1.jsonl python3 -B baseline_nb.py --set fold1 # trained on folds 2-4's tables, read on the ancients table Rows carry the runner's keys (uid, text_sha256, label, choice, correct, probs) so tables_t1.py prints the baseline like any arm. The file is rewritten whole each run (a pure function of the pool and the set). """ from __future__ import annotations import argparse import collections import json import math import os import re import sys sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) import t1common as C # noqa: E402 TOKEN = re.compile(r"[a-z']+") ALPHA = 1.0 ARM = "t1-baseline-nb" SPEC = ("multinomial naive Bayes, unigram tokens [a-z']+ (lowercased), Laplace alpha 1, empirical prior, " "no masking, plain argmax (exact tie: first in kit order)") def tokens(text: str) -> list[str]: return TOKEN.findall(text.lower()) class NB: def __init__(self, rows: list[dict]): self.prior = collections.Counter(r["label"] for r in rows) self.n = len(rows) self.counts = {g: collections.Counter() for g in C.GUESTS} for r in rows: self.counts[r["label"]].update(tokens(r["text"])) self.vocab = set().union(*self.counts.values()) self.totals = {g: sum(c.values()) for g, c in self.counts.items()} def log_joint(self, text: str) -> dict: v = len(self.vocab) out = {} for g in C.GUESTS: lp = math.log(self.prior[g] / self.n) if self.prior[g] else -math.inf denom = self.totals[g] + ALPHA * v for t in tokens(text): if t in self.vocab: # a token never seen in training carries no evidence lp += math.log((self.counts[g][t] + ALPHA) / denom) out[g] = lp return out @staticmethod def argmax(scores: dict, allowed) -> str: best = None for g in C.GUESTS: # kit order: the first maximum wins if g in allowed and (best is None or scores[g] > scores[best]): best = g return best def predict(self, text: str, exclude_named_for: str | None = None) -> tuple[str, str, dict]: lj = self.log_joint(text) m = max(lj.values()) z = sum(math.exp(v - m) for v in lj.values()) probs = {g: math.exp(lj[g] - m) / z for g in C.GUESTS} plain = self.argmax(lj, C.GUESTS) named = set(C.names_in(text)) allowed = [g for g in C.GUESTS if g not in named] or list(C.GUESTS) return plain, self.argmax(lj, allowed), probs def load_pool() -> tuple[list[dict], dict]: pool_path = os.path.join(C.HERE, "kit", "t1-pool.json") man = json.load(open(os.path.join(C.HERE, "kit", "t1-pool.manifest.json"))) raw = open(pool_path, "rb").read() if C.sha(raw) != man["fingerprints"]["t1_pool_json_sha256"]: raise SystemExit("kit/t1-pool.json does not match its manifest") return json.loads(raw), man def items_for(set_id: str, pool: list[dict]) -> tuple[list[dict], list[dict], str]: """(training rows, items to read, the set's sha256).""" if set_id == "s0": items = C.load_pinned(C.S0_PATH, C.S0_SHA) return pool, [dict(it, uid="c%03d" % i) for i, it in enumerate(items)], C.S0_SHA if set_id == "h48": return pool, C.load_pinned(C.H48_PATH, C.H48_SHA), C.H48_SHA m = re.fullmatch(r"fold([1-4])", set_id) if not m: raise SystemExit("--set takes s0, h48 or fold1..fold4") table = C.POC_TABLES[int(m.group(1)) - 1] path = os.path.join(C.HERE, "kit", "fold%s-%s.json" % (m.group(1), table)) return ([r for r in pool if r["source"]["table"] != table], json.load(open(path, encoding="utf-8")), C.file_sha(path)) def run(set_id: str) -> tuple[list[dict], dict]: pool, man = load_pool() train, items, set_sha = items_for(set_id, pool) nb = NB(train) rows = [] for it in items: plain, excl, probs = nb.predict(it["text"]) rows.append({"arm": ARM, "set": set_id, "render": "P", "rep": 1, "order_k": 0, "uid": it["uid"], "text_sha256": C.sha(it["text"]), "label": it["label"], "author_family": it.get("author_family"), "choice": plain, "correct": plain == it["label"], "choice_excl": excl, "correct_excl": excl == it["label"], "named_guests": C.names_in(it["text"], exclude=it["label"]), "probs": probs, "baseline_spec": SPEC, "pool_fingerprint": man["fingerprints"]["text_label_16"], "train_rows": len(train), "set_sha256": set_sha}) summary = {"set": set_id, "n": len(rows), "correct": sum(r["correct"] for r in rows), "correct_excl": sum(r["correct_excl"] for r in rows), "train_rows": len(train), "pool_fingerprint": man["fingerprints"]["text_label_16"]} return rows, summary def main() -> None: ap = argparse.ArgumentParser() ap.add_argument("--set", dest="set_id", required=True) ap.add_argument("--print-only", action="store_true") a = ap.parse_args() rows, s = run(a.set_id) body = "".join(json.dumps(r, ensure_ascii=False, sort_keys=True) + "\n" for r in rows) if not a.print_only: path = os.path.join(C.HERE, "rows", "%s.%s.P.rep1.jsonl" % (ARM, a.set_id)) with open(path, "w") as fh: fh.write(body) s["rows_file"] = os.path.relpath(path, C.HERE) s["rows_sha256"] = C.sha(body) print(json.dumps(s)) if __name__ == "__main__": main()