dreddnafious's picture
Revision 2: step density matched; headline replicates
0662d8e verified
Raw History Blame Contribute Delete
2.43 kB
"""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:])