"""Summarise probe.json files into the README's metrics, plus the ones it omits. forgetting mean over the subjects NOT read of (loss_end - loss_start) retained 1 - forgetting / (chance - mean start loss of those subjects) chance = ln(265); this reproduces upstream's "retained vs chance" gain loss_start - loss_end on the subjects that WERE read (upstream reports this only in passing: chess +0.013) python harness/summarize.py runs/probes/*/probe.json """ import json import sys import numpy as np def summarise(p): ev = p["evals"] a, b = ev[0]["loss"], ev[-1]["loss"] lanes = p["args"]["lanes"] read = set(a) if lanes == "all" else set(lanes.split(",")) unread = [k for k in a if k not in read] or list(a) f = float(np.mean([b[k] - a[k] for k in unread])) start = float(np.mean([a[k] for k in unread])) g = float(np.mean([a[k] - b[k] for k in read if k in a])) peak = max(float(np.mean([e["loss"][k] - a[k] for k in unread])) for e in ev) out = {"arm": p["args"]["arm"], "lanes": lanes, "trunk": p["args"]["trunk_mult"], "pool": p["args"]["pool_mult"], "lr_scale": round(p["lr_scale"], 3), "chars": ev[-1]["chars"], "forgetting": f, "peak_forgetting": peak, "retained": 1 - f / (p["chance"] - start), "gain": g, "trained_experts": p.get("trained_experts"), "n_experts": p["n_experts"]} if p.get("recover"): r = p["recover"][-1]["loss"] dmg = {k: ev[-1]["loss"][k] - a[k] for k in unread} back = [(ev[-1]["loss"][k] - r[k]) / dmg[k] for k in unread if dmg[k] > 0.05] out["recovered_frac"] = float(np.mean(back)) if back else None out["recover_chars"] = p["recover"][-1]["chars"] return out def main(paths): rows = [] for path in paths: with open(path) as fh: s = summarise(json.load(fh)) s["path"] = path rows.append(s) cols = ["arm", "lanes", "trunk", "pool", "chars", "forgetting", "peak_forgetting", "retained", "gain", "trained_experts", "recovered_frac"] print(" ".join(f"{c:>15}" for c in cols)) for s in rows: vals = [] for c in cols: v = s.get(c) vals.append(f"{v:>15.4f}" if isinstance(v, float) else f"{str(v):>15}") print(" ".join(vals)) return rows if __name__ == "__main__": main(sys.argv[1:])