summaryrefslogtreecommitdiff
path: root/experiments/analyze_traffic_timescale.py
blob: e266d8fadec293e31cc7c70e583934a84ced71f4 (plain)
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
"""Audit predictor-timescale validation development runs."""
import glob
import json
import math
import os
import statistics


ROOT = os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))), "results")


def finite_mean(values):
    values = [value for value in values if value is not None and math.isfinite(value)]
    return statistics.mean(values) if values else float("nan")


def main():
    rows = []
    for path in sorted(glob.glob(os.path.join(ROOT, "traffic_time_dev_v1_*.json"))):
        with open(path) as handle:
            row = json.load(handle)
        if row.get("final", {}).get("eval_split") != "validation":
            raise RuntimeError(f"non-validation timescale result: {path}")
        if row.get("provenance", {}).get("git_dirty") is not False:
            raise RuntimeError(f"dirty or unknown provenance: {path}")
        if not row.get("split", {}).get("validation_index_sha256"):
            raise RuntimeError(f"missing validation hash: {path}")
        rows.append(row)

    print("| scale | eta_P | warmup | n | last val (%) | best val (%) | initial traffic R2 | final traffic R2 |")
    print("|---:|---:|---:|---:|---:|---:|---:|---:|")
    groups = {}
    for row in rows:
        args = row["args"]
        key = (args["nuis_rho"], args["eta_P"], args["p_warmup_steps"])
        groups.setdefault(key, []).append(row)

    for key in sorted(groups):
        group = groups[key]
        last = [100 * row["final"]["val_acc"] for row in group]
        best = [100 * max(step["val_acc"] for step in row["steps"] if "val_acc" in step)
                for row in group]
        initial_r2 = [finite_mean(next(step["traffic_r2"] for step in row["steps"]
                                       if "traffic_r2" in step)) for row in group]
        final_r2 = [finite_mean(row["final"]["traffic_r2"]) for row in group]
        print(f"| {key[0]:g} | {key[1]:g} | {key[2]} | {len(group)} | "
              f"{statistics.mean(last):.3f} | {statistics.mean(best):.3f} | "
              f"{statistics.mean(initial_r2):.3f} | {statistics.mean(final_r2):.3f} |")


if __name__ == "__main__":
    main()