diff options
Diffstat (limited to 'experiments/analyze_traffic.py')
| -rw-r--r-- | experiments/analyze_traffic.py | 73 |
1 files changed, 73 insertions, 0 deletions
diff --git a/experiments/analyze_traffic.py b/experiments/analyze_traffic.py new file mode 100644 index 0000000..92faa13 --- /dev/null +++ b/experiments/analyze_traffic.py @@ -0,0 +1,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() |
