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