refactor: update evaluation metrics to include pmass_allowed and nll_json

This commit is contained in:
wassname
2026-05-21 06:05:16 +00:00
parent ef20a83504
commit 49751fbeab
5 changed files with 160 additions and 63 deletions
+8 -5
View File
@@ -84,7 +84,8 @@ def main() -> None:
else {f: float(r["label"][i]) for i, f in enumerate(_DEFAULT_FORCED_FOUNDATIONS)}),
"top1": r["top1"],
"margin": float(r["margin"]),
"nll_prompt": float(r["nll_prompt"]),
"pmass_allowed": float(r["pmass_allowed"]),
"nll_json": float(r["nll_json"]),
}
f.write(json.dumps(rec) + "\n")
logger.info(f"wrote {len(out['per_row'])} rows to {out_path}")
@@ -103,6 +104,8 @@ def main() -> None:
print(f" median_nll_T = {out['median_nll_T']} (temperature-scaled, nats)")
print(f" T = {out['T']}")
print(f" mean_js = {out['mean_js']} (max possible = ln 2 = 0.693)")
print(f" mean_pmass_allowed = {out['mean_pmass_allowed']} (valid-token mass)")
print(f" mean_nll_json = {out['mean_nll_json']} (assistant prefill, nats/tok)")
if out["profile"] is not None:
print("\n=== mean profile (human vs model) ===")
@@ -114,13 +117,13 @@ def main() -> None:
f"{np.median(p_top1):.3f} / {p_top1.mean():.3f} / {p_top1.max():.3f}")
print(" SHOULD: median > 0.4 (clear winner per row); <0.2 -> probe broken")
# Prompt-NLL degradation probe (free; teacher-forced on rendered chat).
nll = np.array([float(r["nll_prompt"]) for r in out["per_row"]])
# JSON-prefill NLL degradation probe (teacher-forced on assistant prefill).
nll = np.array([float(r["nll_json"]) for r in out["per_row"]])
nll = nll[np.isfinite(nll)]
if len(nll):
print(f"\n nll_prompt (nats/tok) min/median/mean/max: "
print(f"\n nll_json (nats/tok) min/median/mean/max: "
f"{nll.min():.3f} / {np.median(nll):.3f} / {nll.mean():.3f} / {nll.max():.3f}")
print(" SHOULD: stable across runs at fixed model; rises under steering/ablation -> degradation")
print(" SHOULD: stable across runs at fixed model; rises under steering/ablation -> JSON-prefill degradation")
if __name__ == "__main__":