"""Write data/STATS.md: counts per source x type x split, option-count histograms, noul label balance, gold-position balance, key-style mix, and (whitespace-word) length stats.""" from __future__ import annotations import re from collections import Counter from jevlike.types import QTYPES, Record SPLITS = ("train", "val", "test", "heldout") OPT_BINS = [(2, 2), (3, 3), (4, 4), (5, 5), (6, 10), (11, 20), (21, 40), (41, 100), (101, 255)] def _pct(vals: list[int], p: float) -> int: if not vals: return 0 s = sorted(vals) return s[min(len(s) - 1, int(p * (len(s) - 1) + 0.5))] def _words(s: str) -> int: return len(s.split()) def _qwords(r: Record) -> int: q = r.question return _words(q.instructions) + sum(_words(o.key) + _words(o.text) for o in q.options) def _key_style(keys: list[str]) -> str: if all(re.fullmatch(r"[A-Za-z]", k) for k in keys): return "letters" if all(re.fullmatch(r"\d+", k) for k in keys): return "numbers" if all(re.fullmatch(r"[A-Za-z_]+\d+", k) for k in keys): return "prefixed ids" if all(re.fullmatch(r"[A-Z]{2,4}", k) for k in keys): return "codes / caps" return "readable" def write_stats(path: str, splits: dict[str, list[Record]], info: dict) -> None: L: list[str] = [] w = L.append w("# jevlike dataset statistics\n") w(f"Built with `python -m jevlike.data.build_dataset {info['argv']}`".rstrip() + f" (seed {info['seed']}).\n") w(f"Records: generated {info['generated']}, after dedupe {info['deduped']}, after type balancing " f"{info['balanced']}; leak removal dropped {info['dropped']['cross_split']} (state shared across splits) + " f"{info['dropped']['heldout_shingle']} (8-word overlap with a heldout state). " f"Target type mix (in-domain): {', '.join(f'{t} {v:.0%}' for t, v in info['mix'].items())}.\n") # ---- totals w("## Totals\n") w("| split | records | states | choice | score | noul |") w("|---|---:|---:|---:|---:|---:|") for s in SPLITS: rs = splits.get(s, []) c = Counter(r.question.type for r in rs) n = len(rs) or 1 w(f"| {s} | {len(rs)} | {len({r.state for r in rs})} | " + " | ".join( f"{c[t]} ({c[t] / n:.0%})" for t in QTYPES) + " |") w("") # ---- per source w("## Records per source x type x split\n") w("Cells are `choice / score / noul`. Primary = the dataset's native task.\n") w("| source | primary | train | val | test | heldout |") w("|---|---|---|---|---|---|") srcs = sorted({r.source for rs in splits.values() for r in rs}) cnt = {s: Counter((r.source, r.question.type) for r in splits.get(s, [])) for s in SPLITS} for src in srcs: cells = [] for s in SPLITS: v = [cnt[s][(src, t)] for t in QTYPES] cells.append(" / ".join(str(x) for x in v) if any(v) else "") w(f"| {src} | {info['primary'].get(src, '?')} | " + " | ".join(cells) + " |") w("") failed = {k: v for k, v in info["status"].items() if not v.startswith("ok")} if failed: w("**Failed / skipped sources:**\n") for k, v in failed.items(): w(f"- {k}: {v}") w("") # ---- option counts w("## Option-count histogram\n") for t in ("choice", "score"): w(f"**{t}**\n") w("| split | " + " | ".join(f"{a}" if a == b else f"{a}-{b}" for a, b in OPT_BINS) + " |") w("|---|" + "---:|" * len(OPT_BINS)) for s in SPLITS: ns = [len(r.question.options) for r in splits.get(s, []) if r.question.type == t] row = [sum(1 for n in ns if a <= n <= b) for a, b in OPT_BINS] w(f"| {s} | " + " | ".join(str(x) for x in row) + " |") w("") # ---- choice gold position + key styles w("## Choice answer position and key style\n") w("Gold position (first / middle / last option) should be roughly uniform; key style shows how often the " "meaning must come from option descriptions (opaque keys).\n") w("| split | first | middle | last | readable keys | letters | numbers | prefixed ids | codes / caps | options with empty text |") w("|---|---:|---:|---:|---:|---:|---:|---:|---:|---:|") for s in SPLITS: rs = [r for r in splits.get(s, []) if r.question.type == "choice"] if not rs: continue pos = Counter("first" if r.label == 0 else "last" if r.label == len(r.question.options) - 1 else "middle" for r in rs) ks = Counter(_key_style([o.key for o in r.question.options]) for r in rs) n_opt = sum(len(r.question.options) for r in rs) empty = sum(1 for r in rs for o in r.question.options if not o.text) n = len(rs) w(f"| {s} | {pos['first'] / n:.0%} | {pos['middle'] / n:.0%} | {pos['last'] / n:.0%} | " + " | ".join(f"{ks[k] / n:.0%}" for k in ("readable", "letters", "numbers", "prefixed ids", "codes / caps")) + f" | {empty / max(1, n_opt):.0%} |") w("") # ---- score level balance w("## Score: gold level position\n") w("Share of score questions whose gold is the lowest level, the highest level, or in between.\n") w("| split | lowest | middle | highest | levels=2 | 3 | 4 | 5 | 6+ |") w("|---|---:|---:|---:|---:|---:|---:|---:|---:|") for s in SPLITS: rs = [r for r in splits.get(s, []) if r.question.type == "score"] if not rs: continue n = len(rs) lo = sum(1 for r in rs if r.label == 0) hi = sum(1 for r in rs if r.label == len(r.question.options) - 1) lv = Counter(min(6, len(r.question.options)) for r in rs) w(f"| {s} | {lo / n:.0%} | {(n - lo - hi) / n:.0%} | {hi / n:.0%} | " + " | ".join(f"{lv[k] / n:.0%}" for k in (2, 3, 4, 5, 6)) + " |") w("") # ---- noul balance w("## Noul label balance (share true)\n") w("| source | train | val | test | heldout |") w("|---|---:|---:|---:|---:|") noul_srcs = sorted({r.source for rs in splits.values() for r in rs if r.question.type == "noul"}) for src in ["ALL"] + noul_srcs: cells = [] for s in SPLITS: ls = [r.label for r in splits.get(s, []) if r.question.type == "noul" and (src == "ALL" or r.source == src)] cells.append(f"{sum(ls) / len(ls):.0%} (n={len(ls)})" if ls else "") w(f"| {'**all**' if src == 'ALL' else src} | " + " | ".join(cells) + " |") w("") # ---- lengths w("## Length (whitespace words)\n") w("State is truncated to ~2000 characters. Question = instructions + option keys + option texts.\n") w("| split | part | mean | p50 | p90 | p99 | max |") w("|---|---|---:|---:|---:|---:|---:|") for s in SPLITS: rs = splits.get(s, []) if not rs: continue for part, vals in (("state", [_words(r.state) for r in rs]), ("question", [_qwords(r) for r in rs]), ("total", [_words(r.state) + _qwords(r) for r in rs])): w(f"| {s} | {part} | {sum(vals) / len(vals):.1f} | {_pct(vals, .5)} | {_pct(vals, .9)} | " f"{_pct(vals, .99)} | {max(vals)} |") w("") w("## Questions per state (train)\n") per = Counter(Counter(r.state for r in splits.get("train", [])).values()) w(" / ".join(f"{k} question(s): {per[k]} states" for k in sorted(per)) + "\n") with open(path, "w", encoding="utf-8") as f: f.write("\n".join(L) + "\n")