botanic1-report / scripts /sync_leaderboard.py
jean-livingmodels's picture
Sync S_bal^test scores and intervals with the technical report
43ee4eb verified
Raw History Blame
6.11 kB
#!/usr/bin/env python3
"""Sync data/leaderboard.json with the technical report's S_bal^test data.
Sources, both in the technical report repository (dotomics/BOTANIC1-technical-report):
figures/data/botanic1_scorecard_28.csv the 22 skill cells per model (skill == 1)
figures/data/fig3_rerun/rerun_ci.json per-example paired bootstrap of the same
evaluation (20,000 replicates, draws shared
across models within each cell)
Convention (figures/scripts/fig3_rerun_ci.py, CENTRE = "rerun"; and
figures/scripts/rebuild_sbal_aggregates.py in that repository):
- point estimates are the scorecard cells; a family score is the mean of its cells
grouped by the scorecard's `task` column, and S_bal is the unweighted mean of the
nine family scores;
- intervals are the rerun's own 95% percentile intervals, not translated onto the
point estimate;
- vs_plantcad2l.delta is the scorecard S_bal difference, and its interval is the
rerun's paired example-level interval (as in sections/tables/sbal_uncertainty.tex).
Run: python3 scripts/sync_leaderboard.py path/to/BOTANIC1-technical-report
"""
import collections
import csv
import json
import subprocess
import sys
from pathlib import Path
from statistics import fmean
ROOT = Path(__file__).resolve().parents[1]
LB = ROOT / "data" / "leaderboard.json"
# leaderboard family key -> scorecard `task` column / rerun family key
TASK = {"chromatin": "chromatin_access", "grc": "genomic_region_classification_v2",
"conservation": "plantcad_conservation", "splicing": "splicing", "tis": "plantcad_tis",
"tts": "plantcad_tts", "proseq": "pro_seq", "llr": "llr", "causal": "gwas"}
TASK_ALIASES = {"gwas": {"gwas", "gwas_eval_benchmark"}, "llr": {"llr", "llr_eval"}}
def norm(model_id: str) -> str:
"""Scorecard / leaderboard spelling -> rerun_ci.json name."""
if model_id.startswith("botanic1-"):
return "Botanic1-" + model_id.split("-", 1)[1]
if model_id.startswith("CARBON-"):
return "Carbon-" + model_id.split("-", 1)[1]
return model_id
def r4(x: float) -> float:
return round(x, 4)
def interval(entry: dict) -> list[float]:
return [r4(entry["lo"]), r4(entry["hi"])]
def fam_key(task: str) -> str:
for k, t in TASK.items():
if task == t or task in TASK_ALIASES.get(t, ()):
return k
raise SystemExit(f"unmapped scorecard task {task!r}")
def main(report: str) -> None:
rep = Path(report)
data = rep / "figures" / "data"
commit = subprocess.run(["git", "-C", str(rep), "rev-parse", "--short=8", "HEAD"],
capture_output=True, text=True).stdout.strip() or "unknown"
cells: dict[str, dict[str, tuple[str, str, float]]] = collections.defaultdict(dict)
with open(data / "botanic1_scorecard_28.csv") as f:
for r in csv.DictReader(f):
if int(r["skill"]):
cells[r["model"]][r["label"]] = (fam_key(r["task"]), r["metric_key"], float(r["value"]))
rerun = json.loads((data / "fig3_rerun" / "rerun_ci.json").read_text())
lb = json.loads(LB.read_text())
ref = "PlantCAD2-L"
sbal: dict[str, float] = {}
for m in lb["models"]:
c = cells.get(m["id"])
if c is None or len(c) != 22:
raise SystemExit(f"{m['id']}: expected 22 skill cells, found {0 if c is None else len(c)}")
fam: dict[str, list[float]] = collections.defaultdict(list)
for k, _, v in c.values():
fam[k].append(v)
if set(fam) != set(TASK):
raise SystemExit(f"{m['id']}: families {sorted(fam)}")
m["families"] = {k: r4(fmean(fam[k])) for k in m["families"]}
sbal[m["id"]] = fmean(fmean(v) for v in fam.values())
m["s_bal"] = r4(sbal[m["id"]])
for mt in m["metrics"]:
if mt["label"] not in c:
raise SystemExit(f"{m['id']}: no scorecard cell {mt['label']!r}")
mt["value"] = r4(c[mt["label"]][2])
for m in lb["models"]:
r = rerun["models"].get(norm(m["id"]))
for k in ("s_bal_ci", "families_ci", "vs_plantcad2l"):
m.pop(k, None)
for mt in m["metrics"]:
mt.pop("ci", None)
if r is None:
print(f" no rerun entry for {m['id']}")
continue
m["s_bal_ci"] = interval(r["sbal"])
m["families_ci"] = {k: interval(r["families"][TASK[k]])
for k in m["families"] if TASK[k] in r["families"]}
c = cells[m["id"]]
for mt in m["metrics"]:
e = r["cells"].get(c[mt["label"]][1])
if e is not None:
mt["ci"] = interval(e)
if m["id"] != ref:
p = rerun["pairs"].get(f"{norm(m['id'])} vs {ref}")
if p is not None:
lo, hi = interval(p)
m["vs_plantcad2l"] = {"delta": r4(sbal[m["id"]] - sbal[ref]), "lo": lo, "hi": hi,
"significant": bool(lo > 0 or hi < 0)}
gap = abs(r["sbal"]["mean"] - sbal[m["id"]])
flag = " <-- check" if gap > 0.005 else ""
print(f"{m['id']:14s} s_bal={m['s_bal']:.4f} rerun={r['sbal']['mean']:.4f} "
f"ci={m['s_bal_ci']}{flag}")
lb["models"].sort(key=lambda m: -m["s_bal"])
lb["ci_note"] = ("95% confidence intervals from the per-sample paired bootstrap of the technical "
"report (20,000 replicates, draws shared across models within each cell). "
"Scores and intervals come from the same evaluation; intervals are not "
"translated onto the point estimate.")
lb["ci_meta"] = {"replicates": rerun["meta"]["replicates"], "level": 95,
"source": "BOTANIC1-technical-report figures/data/botanic1_scorecard_28.csv + "
f"figures/data/fig3_rerun/rerun_ci.json @ {commit}"}
LB.write_text(json.dumps(lb, indent=1, ensure_ascii=False) + "\n")
print(f"wrote {LB}")
if __name__ == "__main__":
main(sys.argv[1])