summaryrefslogtreecommitdiff
path: root/experiments/analyze_traffic.py
blob: 92faa13d7268b4cc4b25f714f3e0a2eabe0905fd (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
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
"""Audit development validation runs for endogenous apical traffic.

This script deliberately refuses test-evaluated files: its output is for
protocol selection, not a paper-facing test table.
"""
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")
PREFIX = "traffic_dev_v1_"


def mean_sd(values):
    mean = statistics.mean(values)
    sd = statistics.stdev(values) if len(values) > 1 else float("nan")
    return mean, sd


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


def main():
    rows = load_rows()
    print("| dataset | traffic | scale | signal | n | validation acc (%) | traffic R2 |")
    print("|:---|:---|---:|:---|---:|---:|---:|")
    groups = {}
    for row in rows:
        args = row["args"]
        signal = "residual"
        if not args["use_residual"]:
            signal = "matched" if args["raw_scale_control"] == "match_innovation_norm" else "raw"
        elif not args["p_neutral"]:
            signal = "taskfit"
        key = (args["dataset"], args["traffic_mode"], args["nuis_rho"], signal)
        groups.setdefault(key, []).append(row)

    for key in sorted(groups):
        group = groups[key]
        acc = [100 * row["final"]["val_acc"] for row in group]
        r2 = [statistics.mean(v for v in row["final"]["traffic_r2"]
                              if v is not None and math.isfinite(v))
              for row in group
              if any(v is not None and math.isfinite(v)
                     for v in row["final"]["traffic_r2"])]
        am, asd = mean_sd(acc)
        r2_text = "—"
        if r2:
            rm, rsd = mean_sd(r2)
            r2_text = f"{rm:.3f}" if len(r2) == 1 else f"{rm:.3f} ± {rsd:.3f}"
        acc_text = f"{am:.3f}" if len(acc) == 1 else f"{am:.3f} ± {asd:.3f}"
        print(f"| {key[0]} | {key[1]} | {key[2]:g} | {key[3]} | {len(group)} | "
              f"{acc_text} | {r2_text} |")


if __name__ == "__main__":
    main()