Cem13's picture
Release Kodama Core with weights, runtime, attribution, and evaluation
c7893fa verified
Raw History Blame Contribute Delete
7.48 kB
"""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")