summaryrefslogtreecommitdiff
path: root/experiments/bci_td_confirmation.py
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 20:32:19 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 20:32:19 -0500
commit32122d026176ed199271aea243ccb0eef78eb30f (patch)
treeaaa78f5439625d7f39663a38275eacc447d85eb4 /experiments/bci_td_confirmation.py
parent9ca51a97119efc90bc2a0b6b65624dc87689e088 (diff)
protocol: freeze oral B recovery confirmation
Diffstat (limited to 'experiments/bci_td_confirmation.py')
-rwxr-xr-xexperiments/bci_td_confirmation.py199
1 files changed, 199 insertions, 0 deletions
diff --git a/experiments/bci_td_confirmation.py b/experiments/bci_td_confirmation.py
new file mode 100755
index 0000000..2a5590d
--- /dev/null
+++ b/experiments/bci_td_confirmation.py
@@ -0,0 +1,199 @@
+#!/usr/bin/env python3
+"""Run one cell of the frozen oral-B temporal-difference confirmation."""
+import argparse
+import hashlib
+import json
+import math
+import os
+import platform
+import resource
+import subprocess
+import sys
+import time
+
+import torch
+
+sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
+from experiments.bci_td_run import (CONDITIONS, build_config, evaluate,
+ neutral_warmup, train)
+from sdil.bci import BCISDIL, generate_trajectories
+from sdil.bci_metrics import signature_metrics
+
+
+ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
+PROTOCOL_PATH = os.path.join(ROOT, "ORAL_B_RECOVERY.md")
+TASK_SEEDS = tuple(range(10, 16))
+MODEL_SEEDS = tuple(range(5))
+
+
+def sha256(path):
+ digest = hashlib.sha256()
+ with open(path, "rb") as handle:
+ for chunk in iter(lambda: handle.read(1024 * 1024), b""):
+ digest.update(chunk)
+ return digest.hexdigest()
+
+
+def finite_tree(value):
+ if isinstance(value, dict):
+ return all(finite_tree(item) for item in value.values())
+ if isinstance(value, (list, tuple)):
+ return all(finite_tree(item) for item in value)
+ if isinstance(value, (int, float)):
+ return math.isfinite(value)
+ return True
+
+
+def provenance(d4_gate, r1_gate):
+ def run(command):
+ return subprocess.run(
+ command, cwd=ROOT, check=True, capture_output=True,
+ text=True).stdout.strip()
+
+ paths = {
+ "runner": os.path.abspath(__file__),
+ "development_runner": os.path.join(ROOT, "experiments", "bci_td_run.py"),
+ "protocol": PROTOCOL_PATH,
+ "d4_gate": os.path.abspath(d4_gate),
+ "r1_gate": os.path.abspath(r1_gate),
+ }
+ relative = {name: os.path.relpath(path, ROOT)
+ for name, path in paths.items()}
+ tracked = {
+ name: subprocess.run(
+ ["git", "ls-files", "--error-unmatch", path], cwd=ROOT,
+ capture_output=True).returncode == 0
+ for name, path in relative.items()
+ }
+ return {
+ "git_commit": run(["git", "rev-parse", "HEAD"]),
+ "git_tracked_dirty": bool(run(
+ ["git", "status", "--porcelain", "--untracked-files=no"])),
+ "tracked_inputs": tracked,
+ "input_sha256": {name: sha256(path) for name, path in paths.items()},
+ }
+
+
+def require_gate(d4, r1, d4_digest):
+ if not (d4.get("protocol") ==
+ "kp_dynamic_neutral_projection_confirmation_v1"
+ and d4.get("status") == "passed"
+ and d4.get("review_score_after") == 7):
+ raise ValueError("oral-B R2 requires the complete audited D4 pass")
+ if not (r1.get("protocol") == "oral_b_td_development_v1"
+ and r1.get("status") == "passed"
+ and r1.get("complete_grid") is True
+ and r1.get("confirmation_seeds_touched") is False
+ and r1.get("oral_b_confirmation_opened") is True
+ and r1.get("review_score_after") == 7
+ and r1.get("d4_gate_sha256") == d4_digest
+ and r1.get("selected", {}).get("eta") in (0.03, 0.1)):
+ raise ValueError("oral-B R2 requires the complete eligible R1 gate")
+
+
+def main():
+ parser = argparse.ArgumentParser()
+ parser.add_argument(
+ "--d4_gate",
+ default="results/kp_dynamic_projection_confirmation_gate.json")
+ parser.add_argument(
+ "--r1_gate", default="results/bci_td_dev_gate.json")
+ parser.add_argument("--task-seed", type=int, choices=TASK_SEEDS, required=True)
+ parser.add_argument("--model-seed", type=int, choices=MODEL_SEEDS, required=True)
+ parser.add_argument("--outdir", default="results/bci_td_confirmation")
+ args = parser.parse_args()
+
+ with open(args.d4_gate) as handle:
+ d4 = json.load(handle)
+ with open(args.r1_gate) as handle:
+ r1 = json.load(handle)
+ d4_digest = sha256(args.d4_gate)
+ require_gate(d4, r1, d4_digest)
+ eta = float(r1["selected"]["eta"])
+ source = provenance(args.d4_gate, args.r1_gate)
+ if (source["git_tracked_dirty"]
+ or not all(source["tracked_inputs"].values())):
+ raise RuntimeError(
+ "R2 requires clean tracked code, protocol, D4 gate, and R1 gate")
+
+ cfg = build_config(eta)
+ trajectories = generate_trajectories(cfg, args.task_seed)
+ evaluation_seed = args.task_seed + 300_000
+ evaluation = generate_trajectories(
+ cfg, evaluation_seed, days=1, episodes=256)
+ initial = BCISDIL(cfg, args.model_seed)
+ conditions = {}
+ trained = {}
+ warmups = {}
+ training_events = None
+ started = time.perf_counter()
+ for name in CONDITIONS:
+ model, warmup = neutral_warmup(
+ initial, args.task_seed, args.model_seed, name)
+ events, report = train(
+ model, trajectories, name, collect=name == "intact")
+ warmups[name] = warmup
+ trained[name] = model
+ conditions[name] = report
+ if name == "intact":
+ training_events = events
+
+ evaluation_events = None
+ for name in CONDITIONS:
+ report = evaluate(
+ trained[name], evaluation, collect=name == "intact")
+ conditions[name]["final_success"] = report["success_rate"]
+ if name == "intact":
+ evaluation_events = report["events"]
+ signatures = signature_metrics(
+ training_events, evaluation_events, cfg, trained["intact"].role)
+
+ result = {
+ "schema_version": 1,
+ "protocol": {
+ "name": "oral_b_td_confirmation_v1",
+ "split": "untouched_confirmation",
+ "training_task_seed": args.task_seed,
+ "evaluation_task_seed": evaluation_seed,
+ "selected_eta": eta,
+ "no_further_selection": True,
+ "confirmation_grid_size": len(TASK_SEEDS) * len(MODEL_SEEDS),
+ "d4_gate_sha256": source["input_sha256"]["d4_gate"],
+ "r1_gate_sha256": source["input_sha256"]["r1_gate"],
+ "protocol_sha256": source["input_sha256"]["protocol"],
+ },
+ "args": vars(args),
+ "config": vars(cfg),
+ "provenance": source,
+ "warmup": warmups,
+ "conditions": conditions,
+ "signatures": signatures,
+ "wall_s": time.perf_counter() - started,
+ "peak_rss_mib": resource.getrusage(resource.RUSAGE_SELF).ru_maxrss / 1024,
+ "hardware": {
+ "device": "cpu", "platform": platform.platform(),
+ "torch_version": torch.__version__, "threads": torch.get_num_threads(),
+ },
+ }
+ result["finite"] = finite_tree(result)
+ if not result["finite"]:
+ raise RuntimeError("non-finite oral-B R2 record")
+ os.makedirs(args.outdir, exist_ok=True)
+ path = os.path.join(
+ args.outdir,
+ f"bci_td_confirm_t{args.task_seed}_m{args.model_seed}.json")
+ if os.path.exists(path):
+ raise FileExistsError(f"refusing to overwrite {path}")
+ with open(path, "w") as handle:
+ json.dump(result, handle, indent=2, sort_keys=True)
+ handle.write("\n")
+ print(json.dumps({
+ "path": path,
+ "finite": result["finite"],
+ "intact_final": conditions["intact"]["final_success"],
+ "sign_inversion": signatures["causal_role_sign_inversion_index"],
+ }, indent=2))
+
+
+if __name__ == "__main__":
+ main()