summaryrefslogtreecommitdiff
path: root/experiments/analyze_query_confirm.py
diff options
context:
space:
mode:
Diffstat (limited to 'experiments/analyze_query_confirm.py')
-rw-r--r--experiments/analyze_query_confirm.py96
1 files changed, 96 insertions, 0 deletions
diff --git a/experiments/analyze_query_confirm.py b/experiments/analyze_query_confirm.py
new file mode 100644
index 0000000..3a3bfd3
--- /dev/null
+++ b/experiments/analyze_query_confirm.py
@@ -0,0 +1,96 @@
+"""Audit frozen multi-seed confirmation of the low-query protocol."""
+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 mean_sd(values):
+ return statistics.mean(values), statistics.stdev(values) if len(values) > 1 else 0.0
+
+
+def key(row):
+ args = row["args"]
+ if args["mode"] == "dfa":
+ return "DFA"
+ return f"K{args['pert_ndirs']}/e{args['pert_every']}"
+
+
+def main():
+ paths = sorted(glob.glob(os.path.join(ROOT, "query_confirm_v1_*.json")))
+ rows = []
+ commits = set()
+ for path in paths:
+ with open(path) as handle:
+ row = json.load(handle)
+ if row.get("final", {}).get("eval_split") != "test":
+ raise RuntimeError(f"non-test confirmation result: {path}")
+ if row.get("provenance", {}).get("git_dirty") is not False:
+ raise RuntimeError(f"dirty or unknown provenance: {path}")
+ if "hardware" not in row or "training_forward_equivalent_examples" not in row.get("cost", {}):
+ raise RuntimeError(f"missing cost/memory audit: {path}")
+ commits.add(row["provenance"]["git_commit"])
+ rows.append(row)
+ if len(commits) != 1:
+ raise RuntimeError(f"mixed commits: {commits}")
+
+ groups = {}
+ for row in rows:
+ groups.setdefault(key(row), []).append(row)
+ expected = {"DFA", "K16/e4", "K1/e4"}
+ if set(groups) != expected or any(len(group) != 5 for group in groups.values()):
+ raise RuntimeError(f"need five rows for {expected}; got "
+ f"{ {name: len(group) for name, group in groups.items()} }")
+ for group in groups.values():
+ group.sort(key=lambda row: row["args"]["seed"])
+ seeds = [[row["args"]["seed"] for row in groups[name]] for name in sorted(groups)]
+ if len({tuple(value) for value in seeds}) != 1:
+ raise RuntimeError(f"unpaired seeds: {seeds}")
+
+ print(f"commit={next(iter(commits))} seeds={seeds[0]}")
+ print("| method | n | test (%) | batch queries | train forward-eq examples | "
+ "peak allocated (MiB) | wall (s) |")
+ print("|:---|---:|---:|---:|---:|---:|---:|")
+ for name in ("DFA", "K16/e4", "K1/e4"):
+ group = groups[name]
+ acc = mean_sd([100 * row["final"]["test_acc"] for row in group])
+ queries = mean_sd([row["cost"]["calibration_batch_loss_evaluations"] for row in group])
+ work = mean_sd([row["cost"]["training_forward_equivalent_examples"] for row in group])
+ memory = mean_sd([row["hardware"]["peak_memory_allocated_bytes"] / 2**20 for row in group])
+ wall = mean_sd([row["final"]["wall_s"] for row in group])
+ print(f"| {name} | {len(group)} | {acc[0]:.3f} +/- {acc[1]:.3f} | "
+ f"{queries[0]:.0f} | {work[0]:.0f} | {memory[0]:.1f} +/- {memory[1]:.1f} | "
+ f"{wall[0]:.1f} +/- {wall[1]:.1f} |")
+
+ dfa = [row["final"]["test_acc"] for row in groups["DFA"]]
+ ref = [row["final"]["test_acc"] for row in groups["K16/e4"]]
+ low = [row["final"]["test_acc"] for row in groups["K1/e4"]]
+ ref_gains = [a - b for a, b in zip(ref, dfa)]
+ low_gains = [a - b for a, b in zip(low, dfa)]
+ retention = statistics.mean(low_gains) / statistics.mean(ref_gains)
+ query_reduction = (groups["K16/e4"][0]["cost"]["calibration_batch_loss_evaluations"]
+ / groups["K1/e4"][0]["cost"]["calibration_batch_loss_evaluations"])
+ calibration_work_reduction = (
+ groups["K16/e4"][0]["cost"]["calibration_forward_equivalent_examples"]
+ / groups["K1/e4"][0]["cost"]["calibration_forward_equivalent_examples"])
+ total_work_reduction = (
+ groups["K16/e4"][0]["cost"]["training_forward_equivalent_examples"]
+ / groups["K1/e4"][0]["cost"]["training_forward_equivalent_examples"])
+ print(f"paired reference gain: {100 * statistics.mean(ref_gains):.3f} points")
+ print(f"paired low-query gain: {100 * statistics.mean(low_gains):.3f} points")
+ print(f"gain retention: {retention:.3f}")
+ print(f"query reduction: {query_reduction:.1f}x")
+ print(f"calibration/total forward-eq reduction: "
+ f"{calibration_work_reduction:.1f}x / {total_work_reduction:.1f}x")
+ passed = query_reduction >= 10 and retention >= 0.9
+ print(f"C3 confirmation gate: {'PASS' if passed else 'FAIL'}")
+ if not passed:
+ raise SystemExit(1)
+
+
+if __name__ == "__main__":
+ main()