diff options
Diffstat (limited to 'experiments/bci_td_confirmation.py')
| -rwxr-xr-x | experiments/bci_td_confirmation.py | 199 |
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() |
