summaryrefslogtreecommitdiff
path: root/experiments/oral_a_dynamic_scaling_v2.py
diff options
context:
space:
mode:
Diffstat (limited to 'experiments/oral_a_dynamic_scaling_v2.py')
-rw-r--r--experiments/oral_a_dynamic_scaling_v2.py149
1 files changed, 149 insertions, 0 deletions
diff --git a/experiments/oral_a_dynamic_scaling_v2.py b/experiments/oral_a_dynamic_scaling_v2.py
new file mode 100644
index 0000000..2e58720
--- /dev/null
+++ b/experiments/oral_a_dynamic_scaling_v2.py
@@ -0,0 +1,149 @@
+#!/usr/bin/env python3
+"""Run shards of the calibrated-BCI-v2-gated standard-depth panel."""
+import argparse
+import hashlib
+import json
+import os
+import subprocess
+
+from oral_a_dynamic_scaling import (
+ DEPTHS,
+ METHODS,
+ SEEDS,
+ existing_matches,
+ jobs,
+)
+
+
+ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
+PROTOCOL_PATH = os.path.join(ROOT, "ORAL_A_RECOVERY_V2.md")
+LEGACY_RUNNER_PATH = os.path.join(
+ ROOT, "experiments", "oral_a_dynamic_scaling.py")
+
+
+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 git_output(*args):
+ return subprocess.run(
+ ["git", *args], cwd=ROOT, check=True, capture_output=True,
+ text=True).stdout.strip()
+
+
+def require_prerequisites(d4_path, bci_v2_path):
+ with open(d4_path) as handle:
+ d4 = json.load(handle)
+ with open(bci_v2_path) as handle:
+ bci_v2 = json.load(handle)
+ 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-A recovery v2 requires the complete D4 pass")
+ if not (
+ bci_v2.get("protocol")
+ == "oral_b_v2_calibrated_recovery_confirmation_v1"
+ and bci_v2.get("status") == "passed"
+ and bci_v2.get("complete_grid") is True
+ and bci_v2.get("oral_b_v2_outcome_surprise_established") is True
+ and bci_v2.get("review_score_after") == 8
+ and bci_v2.get("new_oral_a_v2_protocol_may_be_frozen") is True
+ and bci_v2.get("old_oral_a_gate_remains_closed") is True
+ and bci_v2.get("all_prior_failures_preserved") is True
+ ):
+ raise ValueError(
+ "oral-A recovery v2 requires the complete calibrated BCI-v2 pass")
+
+ paths = [
+ os.path.abspath(__file__),
+ PROTOCOL_PATH,
+ LEGACY_RUNNER_PATH,
+ os.path.abspath(d4_path),
+ os.path.abspath(bci_v2_path),
+ os.path.join(ROOT, "experiments", "conv_run.py"),
+ ]
+ relative = [os.path.relpath(path, ROOT) for path in paths]
+ tracked = all(
+ subprocess.run(
+ ["git", "ls-files", "--error-unmatch", path],
+ cwd=ROOT,
+ capture_output=True,
+ ).returncode == 0
+ for path in relative
+ )
+ dirty = bool(git_output(
+ "status", "--porcelain", "--untracked-files=no"))
+ if dirty or not tracked:
+ raise RuntimeError(
+ "oral-A recovery v2 requires a clean tracked source")
+ return {
+ "git_commit": git_output("rev-parse", "HEAD"),
+ "protocol_sha256": sha256(PROTOCOL_PATH),
+ "d4_gate_sha256": sha256(d4_path),
+ "bci_v2_gate_sha256": sha256(bci_v2_path),
+ "legacy_runner_sha256": sha256(LEGACY_RUNNER_PATH),
+ }
+
+
+def main():
+ parser = argparse.ArgumentParser()
+ parser.add_argument(
+ "--d4_gate",
+ default="results/kp_dynamic_projection_confirmation_gate.json",
+ )
+ parser.add_argument(
+ "--bci_v2_gate",
+ default="results/bci_v2_calibrated_confirmation_gate.json",
+ )
+ parser.add_argument("--device", default="cuda")
+ parser.add_argument("--method", choices=("all",) + METHODS, default="all")
+ parser.add_argument("--depth", type=int, choices=DEPTHS)
+ parser.add_argument("--seed", type=int, choices=SEEDS)
+ parser.add_argument("--shard-index", type=int, default=0)
+ parser.add_argument("--num-shards", type=int, default=1)
+ parser.add_argument(
+ "--outdir", default="results/oral_a_dynamic_scaling_v2")
+ parser.add_argument("--dry-run", action="store_true")
+ args = parser.parse_args()
+ if not 0 <= args.shard_index < args.num_shards:
+ raise ValueError("invalid shard index")
+ source = require_prerequisites(args.d4_gate, args.bci_v2_gate)
+ selected = jobs(args.device, args.outdir)
+ if args.method != "all":
+ selected = [job for job in selected if job[0] == args.method]
+ if args.depth is not None:
+ selected = [job for job in selected if job[1] == args.depth]
+ if args.seed is not None:
+ selected = [job for job in selected if job[2] == args.seed]
+ selected = [
+ job
+ for index, job in enumerate(selected)
+ if index % args.num_shards == args.shard_index
+ ]
+ if not selected:
+ raise ValueError("no oral-A recovery-v2 jobs match the requested shard")
+
+ os.makedirs(args.outdir, exist_ok=True)
+ for method, depth, seed, path, command in selected:
+ tag = f"{method}_d{depth}_s{seed}"
+ if os.path.exists(path):
+ if not existing_matches(
+ path, method, depth, seed, source["git_commit"]
+ ):
+ raise RuntimeError(
+ f"refusing mismatched existing record: {path}")
+ print(f"preserving {tag}", flush=True)
+ continue
+ print(tag, " ".join(command), flush=True)
+ if not args.dry_run:
+ subprocess.run(command, cwd=ROOT, check=True)
+
+
+if __name__ == "__main__":
+ main()