File size: 2,425 Bytes
0662d8e | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 | """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:])
|