Ines-1 / training /code /reference_rows.py
Endikavi's picture
Ines-1 RC1 (private staging; release commit b7f5644)
61b6fb9 verified
Raw History Blame Contribute Delete
2.99 kB
"""Rebuild the typed-decisions card's two reference rows (Uniform, Prior) to pin metric definitions.
Card values (test, 2,000 decisions): Uniform acc 0.308 KL 0.444 Brier 0.238 ECE 0.169;
Prior (each question's train label frequencies) acc 0.470 KL 0.347 Brier 0.189 ECE 0.088.
Also writes the per-question teacher-agreement flag (label_agreement.argmax_agree) for the test.
"""
import collections
import json
import sys
from pathlib import Path
sys.path.insert(0, str(Path(__file__).resolve().parent))
from metrics import summarize # noqa: E402
REPO, REV = "LocalLLaMA/typed-decisions", "c76749ec58bd8c3d2ea706b31c333a9059c38f90"
def rows(split):
import pyarrow.parquet as pq
from huggingface_hub import hf_hub_download
return pq.read_table(hf_hub_download(REPO, "all/%s-00000-of-00001.parquet" % split, repo_type="dataset",
revision=REV)).to_pylist()
def keys_of(q):
if q["type"] == "choice":
return list(q["criteria"])
if q["type"] == "score":
return [str(i) for i in range(len(q["criteria"]))]
return ["true", "false"]
def items(test, dist):
out = []
for r in test:
qs, gold = json.loads(r["questions"]), json.loads(r["gold"])
for qid, q in qs.items():
k = keys_of(q)
g = gold[qid]
lab = str(g["label"]).lower() if q["type"] == "noul" else str(g["label"])
soft = {(str(a).lower() if q["type"] == "noul" else str(a)): float(b) for a, b in g["probabilities"].items()}
out.append({"id": r["id"], "qid": qid, "type": q["type"], "keys": k, "probs": dist(r, qid, k),
"gold": [lab], "soft": soft})
return out
def main(out_path):
test, train = rows("test"), rows("train")
freq = collections.defaultdict(collections.Counter)
for r in train:
gold = json.loads(r["gold"])
for qid, g in gold.items():
lab = str(g["label"]).lower() if g.get("type") == "noul" else str(g["label"])
freq[(r["workflow"], qid)][lab] += 1
def uniform(r, qid, k):
return [1 / len(k)] * len(k)
def prior(r, qid, k):
c = freq[(r["workflow"], qid)]
t = sum(c[x] for x in k)
return [c[x] / t for x in k]
res = {}
for name, f in (("uniform", uniform), ("prior", prior)):
it = items(test, f)
s = summarize(it)
if name == "uniform":
s["accuracy_expected_1_over_k"] = sum(1 / len(i["keys"]) for i in it) / len(it)
res[name] = s
agree = {}
for r in test:
la = json.loads(r["label_agreement"])
for qid, v in la.items():
agree["%s|%s" % (r["id"], qid)] = bool(v["argmax_agree"])
res["argmax_agree_counts"] = dict(collections.Counter(agree.values()))
Path(out_path).write_text(json.dumps({"reference_rows": res, "argmax_agree": agree}, indent=1) + "\n")
print(json.dumps(res, indent=1))
if __name__ == "__main__":
main(sys.argv[1])