Ines-1 / training /code /audit_report.py
Endikavi's picture
Ines-1 RC1 (private staging; release commit b7f5644)
61b6fb9 verified
Raw History Blame Contribute Delete
7.72 kB
"""Audit report of the decision models from per-item outputs: lineage, accuracy + calibration,
paired comparisons, tasksource breakdown, Engram-off. Writes report.json and report.md.
python scripts/decisions/audit_report.py --audit AUDIT_DIR --models v2=RUN_DIR v3=RUN_DIR v3pre=RUN_DIR \
[--chain pretrain=CKPT_DIR]
AUDIT_DIR holds items/<model>[-engram_off]__<test>.jsonl (eval_items.py), reference_rows.json
(reference_rows.py), audit_data.json and tasksource_test_items.jsonl (audit_data.py).
"""
from __future__ import annotations
import argparse
import collections
import hashlib
import json
import sys
from pathlib import Path
sys.path.insert(0, str(Path(__file__).resolve().parent))
from metrics import bootstrap_ci, item_scores, load, paired, summarize # noqa: E402
TESTS = ("typed_en", "typed_es", "telepatia_es", "tasksource")
def sha256(p):
h = hashlib.sha256()
with open(p, "rb") as f:
for chunk in iter(lambda: f.read(1 << 24), b""):
h.update(chunk)
return h.hexdigest()
def lineage(chain):
rows = []
for name, d in chain:
ck = d / "checkpoint" if (d / "checkpoint").exists() else d
man = json.loads((ck / "manifest.json").read_text())
ident = man["identity"]
row = {"stage": name, "run_id": ident["run_id"], "parent_run_id": ident.get("parent_run_id") or None,
"parent_weights_sha256": ident.get("parent_checkpoint_sha256") or None,
"weights_sha256": sha256(ck / "model" / "model.pt" if (ck / "model" / "model.pt").exists() else ck / "model.pt"),
"global_step_at_save": man["global_step"], "tokens_seen": man.get("tokens_seen"),
"training_recipe_hash": ident.get("training_recipe_hash"), "dataset_subset_id": ident.get("dataset_subset_id"),
"code_commit": ident.get("code_commit"), "config_seed": ident.get("seed"), "saved_at": man["saved_at"],
"selected_epoch": man["extra"].get("selected_epoch"), "val_score": man["extra"].get("val_score")}
res_p = d / "results.json"
if res_p.exists():
r = json.loads(res_p.read_text())
row.update({"train_files": [Path(x).name for x in r["train"]], "val_files": [Path(x).name for x in r["val"]],
"epochs_scheduled": r["epochs"], "lr": r["lr"], "accum": r["accum"], "aux": r["aux"],
"sft_seed": r["seed"], "val_curve": [c.get("val_score") for c in r.get("curve", [])]})
log = d.parent / (d.name + ".log")
if log.exists():
for line in log.read_text().splitlines():
if line.startswith("train questions"):
q, s = line.split(",")
row["train_questions_per_epoch"] = int(q.split()[-1])
row["optimizer_steps_scheduled"] = int(s.split()[-1])
break
if "train_questions_per_epoch" in row:
per = row["optimizer_steps_scheduled"] / row["epochs_scheduled"]
row["optimizer_steps_to_selected_weights"] = round(per * row["selected_epoch"])
rows.append(row)
return rows
def fmt(x, d=3):
return "–" if x is None else ("%.*f" % (d, x))
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--audit", type=Path, required=True)
ap.add_argument("--models", nargs="+", required=True, help="name=run_dir, in lineage order")
ap.add_argument("--chain", nargs="*", default=[], help="name=checkpoint_dir before the fine-tuning runs")
a = ap.parse_args()
A = a.audit
models = [tuple(x.split("=", 1)) for x in a.models]
chain = [(n, Path(p)) for n, p in (x.split("=", 1) for x in a.chain)] + [(n, Path(p)) for n, p in models]
ref = json.loads((A / "reference_rows.json").read_text())
agree = ref["argmax_agree"]
ts_meta = {x["id"]: x for x in map(json.loads, open(A / "tasksource_test_items.jsonl"))}
rep = {"lineage": lineage(chain), "reference_rows": ref["reference_rows"], "results": {}, "paired": {},
"typed_en_by_teacher_agreement": {}, "tasksource": {}}
items = {}
for name, _ in models:
for off in ("", "-engram_off"):
for t in TESTS:
p = A / "items" / ("%s%s__%s.jsonl" % (name, off, t))
if p.exists():
items[(name + off, t)] = load(p)
for (m, t), it in sorted(items.items()):
s = summarize(it)
s["ci95_cases"] = bootstrap_ci(it)
s["mean_prompt_tokens"] = sum(i["prompt_tokens"] for i in it) / len(it)
rep["results"]["%s|%s" % (m, t)] = s
if t == "typed_en":
parts = collections.defaultdict(list)
for i in it:
parts["agreed" if agree["%s|%s" % (i["id"], i["qid"])] else "split"].append(i)
rep["typed_en_by_teacher_agreement"][m] = {k: {"n": len(v), "accuracy": summarize(v)["accuracy"]}
for k, v in parts.items()}
names = [n for n, _ in models]
for t in TESTS:
for x, y in [(names[i], names[j]) for i in range(len(names)) for j in range(i + 1, len(names))]:
if (x, t) in items and (y, t) in items:
rep["paired"]["%s vs %s|%s" % (x, y, t)] = paired(items[(x, t)], items[(y, t)])
for n in names:
if (n, t) in items and (n + "-engram_off", t) in items:
rep["paired"]["%s vs %s-engram_off|%s" % (n, n, t)] = paired(items[(n, t)], items[(n + "-engram_off", t)])
# tasksource: by family and by whether the state text also appears in train
for n in names:
it = items.get((n, "tasksource"))
if not it:
continue
fam = collections.defaultdict(list)
st = collections.defaultdict(list)
var = collections.defaultdict(list)
for i in it:
meta = ts_meta[i["id"]]
ok = item_scores(i)["ok"]
fam[meta["family"]].append(ok)
st["state_text_in_train" if meta["state_text_in_train"] else "state_text_not_in_train"].append(ok)
var[meta["variant"]].append(ok)
rep["tasksource"][n] = {
"by_family": {k: [sum(v), len(v)] for k, v in sorted(fam.items())},
"by_state_overlap": {k: [sum(v), len(v)] for k, v in st.items()},
"by_variant": {k: [sum(v), len(v)] for k, v in var.items()}}
(A / "report.json").write_text(json.dumps(rep, indent=1, ensure_ascii=False) + "\n")
# markdown
L = ["# Audit report", "", "## Accuracy and calibration", "",
"| model | test | correct/n | acc | 95% CI (cases) | choice | score | noul | NLL | KL(gold‖p) | Brier (soft) | ECE10 (ours) |",
"|---|---|---:|---:|---|---:|---:|---:|---:|---:|---:|---:|"]
for k, s in rep["results"].items():
m, t = k.split("|")
bt = s["by_type"]
L.append("| %s | %s | %d/%d | %.3f | %.3f–%.3f | %s | %s | %s | %s | %s | %s | %s |" % (
m, t, s["correct"], s["n"], s["accuracy"], *s["ci95_cases"],
*[fmt(bt[x]["accuracy"]) if x in bt else "–" for x in ("choice", "score", "noul")],
fmt(s["nll"]), fmt(s["kl"]), fmt(s["brier"]), fmt(s["ece10"])))
L += ["", "## Paired comparisons (same questions, exact McNemar)", "",
"| comparison | n | only first right | only second right | p |", "|---|---:|---:|---:|---:|"]
for k, v in rep["paired"].items():
L.append("| %s | %d | %d | %d | %.2g |" % (k, v["n_common"], v["only_a_right"], v["only_b_right"], v["p_mcnemar_exact"]))
(A / "report.md").write_text("\n".join(L) + "\n")
print("\n".join(L))
if __name__ == "__main__":
main()