#!/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()