summaryrefslogtreecommitdiff
path: root/experiments/analyze_verified.py
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-07-21 08:26:54 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-07-21 08:26:54 -0500
commit5b9de3f36c91ea8ad57fc4dae5aa0c04f08a7a26 (patch)
treef70cc3489cfea3510d51975007d444901aae2639 /experiments/analyze_verified.py
parentd6a00aab979a50394cd4d5423f7304bb7819ce25 (diff)
feat: aggregate audited multi-seed results
Diffstat (limited to 'experiments/analyze_verified.py')
-rwxr-xr-xexperiments/analyze_verified.py153
1 files changed, 153 insertions, 0 deletions
diff --git a/experiments/analyze_verified.py b/experiments/analyze_verified.py
new file mode 100755
index 0000000..8a1a841
--- /dev/null
+++ b/experiments/analyze_verified.py
@@ -0,0 +1,153 @@
+#!/usr/bin/env python3
+"""Aggregate the audited multi-seed experiments into compact Markdown tables."""
+import argparse
+import glob
+import json
+import math
+import os
+import statistics
+
+
+def read_many(patterns):
+ paths = sorted({path for pattern in patterns for path in glob.glob(pattern)})
+ records = []
+ for path in paths:
+ with open(path) as f:
+ record = json.load(f)
+ record["_path"] = path
+ records.append(record)
+ return records
+
+
+def mean_sd(values):
+ if not values:
+ return float("nan"), float("nan")
+ return statistics.mean(values), statistics.stdev(values) if len(values) > 1 else 0.0
+
+
+def fmt(values, percent=False, digits=3):
+ mean, sd = mean_sd(values)
+ if percent:
+ mean, sd = 100 * mean, 100 * sd
+ return f"{mean:.{digits}f} ± {sd:.{digits}f}"
+
+
+def early_layer_alignment(record):
+ values = (record.get("final", {}).get("cos_r_negg")
+ or record.get("final", {}).get("cos_fa_negg"))
+ if not values:
+ return None
+ n = max(1, len(values) // 3)
+ finite = [v for v in values[:n] if math.isfinite(v)]
+ return statistics.mean(finite) if finite else None
+
+
+def method_name(record):
+ args = record["args"]
+ return (args.get("mode") or args.get("method")).upper()
+
+
+def print_scaling(root):
+ records = read_many([
+ os.path.join(root, "scale_v3_cifar10_*.json"),
+ os.path.join(root, "fa_scale_v2_cifar10_*.json"),
+ ])
+ groups = {}
+ for record in records:
+ key = (record["args"]["depth"], method_name(record))
+ groups.setdefault(key, []).append(record)
+ print("## CIFAR-10 depth scaling\n")
+ print("| depth | method | n | test acc (%) | early-third alignment | wall (s) |")
+ print("|---:|:---|---:|---:|---:|---:|")
+ order = {"BP": 0, "FA": 1, "DFA": 2, "SDIL": 3}
+ for (depth, method), rows in sorted(groups.items(), key=lambda x: (x[0][0], order[x[0][1]])):
+ aligns = [v for row in rows if (v := early_layer_alignment(row)) is not None]
+ align = fmt(aligns) if aligns else "—"
+ print(f"| {depth} | {method} | {len(rows)} | "
+ f"{fmt([r['final']['test_acc'] for r in rows], percent=True)} | {align} | "
+ f"{fmt([r['final']['wall_s'] for r in rows], digits=1)} |")
+
+ indexed = {(r["args"]["depth"], r["args"]["seed"], method_name(r)): r for r in records}
+ for depth in sorted({r["args"]["depth"] for r in records}):
+ gaps = []
+ for seed in range(100):
+ sdil = indexed.get((depth, seed, "SDIL"))
+ dfa = indexed.get((depth, seed, "DFA"))
+ if sdil and dfa:
+ gaps.append(sdil["final"]["test_acc"] - dfa["final"]["test_acc"])
+ if gaps:
+ print(f"\nPaired SDIL−DFA gap at d{depth}: {fmt(gaps, percent=True)} points (n={len(gaps)}).")
+
+
+def nuisance_label(record):
+ args = record["args"]
+ if args["use_residual"]:
+ return "residual"
+ if args["raw_scale_control"] == "match_innovation_norm":
+ return "matched raw"
+ return "raw"
+
+
+def print_nuisance(root):
+ records = read_many([os.path.join(root, "nuis_ctrl_v1_*.json")])
+ groups = {}
+ for record in records:
+ key = (record["args"]["nuis_rho"], nuisance_label(record))
+ groups.setdefault(key, []).append(record)
+ print("\n\n## Somato-dendritic nuisance control\n")
+ print("| rho | signal | n | test acc (%) |")
+ print("|---:|:---|---:|---:|")
+ order = {"raw": 0, "matched raw": 1, "residual": 2}
+ for (rho, label), rows in sorted(groups.items(), key=lambda x: (x[0][0], order[x[0][1]])):
+ print(f"| {rho:g} | {label} | {len(rows)} | "
+ f"{fmt([r['final']['test_acc'] for r in rows], percent=True)} |")
+
+
+def print_local_baselines(root):
+ records = read_many([
+ os.path.join(root, "pepita_tuned_v1_*.json"),
+ os.path.join(root, "baseline_budget_v1_mnist_ff_*.json"),
+ os.path.join(root, "ep_original_v2_*.json"),
+ ])
+ groups = {}
+ for record in records:
+ args = record["args"]
+ key = (method_name(record), args["dataset"], args["depth"], args["width"], args["epochs"])
+ groups.setdefault(key, []).append(record)
+ print("\n\n## Native-protocol local baselines\n")
+ print("| method | dataset | depth × width | budget | n | test acc (%) | wall (s) |")
+ print("|:---|:---|:---|:---|---:|---:|---:|")
+ for (method, dataset, depth, width, epochs), rows in sorted(groups.items()):
+ budget = f"{epochs}/layer" if method == "FF" else str(epochs)
+ print(f"| {method} | {dataset} | {depth} × {width} | {budget} | {len(rows)} | "
+ f"{fmt([r['final']['test_acc'] for r in rows], percent=True)} | "
+ f"{fmt([r['final']['wall_s'] for r in rows], digits=1)} |")
+
+
+def provenance_audit(root):
+ records = read_many([
+ os.path.join(root, "scale_v3_*.json"),
+ os.path.join(root, "nuis_ctrl_v1_*.json"),
+ os.path.join(root, "fa_scale_v2_*.json"),
+ os.path.join(root, "pepita_tuned_v1_*.json"),
+ os.path.join(root, "baseline_budget_v1_mnist_ff_*.json"),
+ os.path.join(root, "ep_original_v2_*.json"),
+ ])
+ dirty = [r["_path"] for r in records if r.get("provenance", {}).get("git_dirty") is not False]
+ print(f"\n\nProvenance: {len(records)} files audited; {len(dirty)} dirty or unknown.")
+ for path in dirty:
+ print(f"- {path}")
+
+
+def main():
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--results", default="results")
+ args = parser.parse_args()
+ print_scaling(args.results)
+ print_nuisance(args.results)
+ print_local_baselines(args.results)
+ provenance_audit(args.results)
+
+
+if __name__ == "__main__":
+ main()