"""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])