summaryrefslogtreecommitdiff
path: root/experiments/analyze_bci_td_confirmation.py
diff options
context:
space:
mode:
Diffstat (limited to 'experiments/analyze_bci_td_confirmation.py')
-rwxr-xr-xexperiments/analyze_bci_td_confirmation.py374
1 files changed, 374 insertions, 0 deletions
diff --git a/experiments/analyze_bci_td_confirmation.py b/experiments/analyze_bci_td_confirmation.py
new file mode 100755
index 0000000..8697e38
--- /dev/null
+++ b/experiments/analyze_bci_td_confirmation.py
@@ -0,0 +1,374 @@
+#!/usr/bin/env python3
+"""Audit the frozen oral-B temporal-difference confirmation panel."""
+import argparse
+import glob
+import hashlib
+import json
+import math
+import os
+import statistics
+
+
+TASK_SEEDS = tuple(range(10, 16))
+MODEL_SEEDS = tuple(range(5))
+CONDITIONS = ("intact", "fixed_vectorizer", "plasticity_lesion", "oracle_role")
+T_CRITICAL_ONE_SIDED_95_DF5 = 2.015048373
+ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
+PROTOCOL_PATH = os.path.join(ROOT, "ORAL_B_RECOVERY.md")
+RUNNER_PATH = os.path.join(ROOT, "experiments", "bci_td_confirmation.py")
+DEVELOPMENT_RUNNER_PATH = os.path.join(ROOT, "experiments", "bci_td_run.py")
+
+
+def require(condition, message):
+ if not condition:
+ raise ValueError(message)
+
+
+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 lower_confidence_bound(values):
+ return (statistics.mean(values)
+ - T_CRITICAL_ONE_SIDED_95_DF5
+ * statistics.stdev(values) / math.sqrt(len(values)))
+
+
+def upper_confidence_bound(values):
+ return (statistics.mean(values)
+ + T_CRITICAL_ONE_SIDED_95_DF5
+ * statistics.stdev(values) / math.sqrt(len(values)))
+
+
+def task_cluster_values(records, getter):
+ return [statistics.mean(
+ getter(records[(task_seed, model_seed)])
+ for model_seed in MODEL_SEEDS)
+ for task_seed in TASK_SEEDS]
+
+
+def summarize(values):
+ return {
+ "by_task_seed": values,
+ "mean": statistics.mean(values),
+ "one_sided_95pct_lower": lower_confidence_bound(values),
+ "one_sided_95pct_upper": upper_confidence_bound(values),
+ }
+
+
+def validate(row, path, eta, digests):
+ require(row.get("schema_version") == 1, f"{path}: schema")
+ args = row.get("args", {})
+ task_seed = args.get("task_seed")
+ model_seed = args.get("model_seed")
+ require(task_seed in TASK_SEEDS, f"{path}: task seed")
+ require(model_seed in MODEL_SEEDS, f"{path}: model seed")
+ protocol = row.get("protocol", {})
+ require(protocol.get("name") == "oral_b_td_confirmation_v1",
+ f"{path}: protocol")
+ require(protocol.get("split") == "untouched_confirmation",
+ f"{path}: split")
+ require(protocol.get("training_task_seed") == task_seed,
+ f"{path}: training seed")
+ require(protocol.get("evaluation_task_seed") == task_seed + 300_000,
+ f"{path}: evaluation seed")
+ require(protocol.get("selected_eta") == eta, f"{path}: selected eta")
+ require(protocol.get("no_further_selection") is True,
+ f"{path}: post-selection tuning")
+ require(protocol.get("confirmation_grid_size") == 30,
+ f"{path}: grid size")
+ for key in ("d4_gate", "r1_gate", "protocol"):
+ require(protocol.get(f"{key}_sha256") == digests[key],
+ f"{path}: {key} digest")
+
+ source = row.get("provenance", {})
+ require(source.get("git_tracked_dirty") is False, f"{path}: dirty")
+ require(all(source.get("tracked_inputs", {}).values())
+ and set(source.get("tracked_inputs", {})) == {
+ "runner", "development_runner", "protocol", "d4_gate", "r1_gate"},
+ f"{path}: untracked input")
+ expected_source_digests = {
+ "runner": digests["runner"],
+ "development_runner": digests["development_runner"],
+ "protocol": digests["protocol"],
+ "d4_gate": digests["d4_gate"],
+ "r1_gate": digests["r1_gate"],
+ }
+ require(source.get("input_sha256") == expected_source_digests,
+ f"{path}: source digest drift")
+ require(isinstance(source.get("git_commit"), str)
+ and len(source["git_commit"]) == 40, f"{path}: commit")
+
+ config = row.get("config", {})
+ expected = {
+ "n_plus": 5, "n_minus": 5, "n_background": 30,
+ "context_dim": 16, "steps_per_episode": 28,
+ "episodes_per_day": 64, "days": 14, "target": 0.8,
+ "inertia": 0.65, "process_noise": 0.12, "context_ar": 0.8,
+ "coupling_scale": 1.0, "predictor_eta": 0.2,
+ "vectorizer_eta": 0.03, "forward_eta": eta,
+ "perturb_sigma": 0.03, "perturb_every": 4,
+ "kappa": 0.0, "feedback": "performance_velocity",
+ }
+ for key, value in expected.items():
+ require(config.get(key) == value, f"{path}: config {key}")
+
+ require(row.get("finite") is True and finite_tree(row), f"{path}: finite")
+ for name in CONDITIONS:
+ warmup = row["warmup"][name]
+ require(warmup["batches"] == 100 and warmup["examples"] == 6400,
+ f"{path}: warmup count {name}")
+ require(warmup["instruction_present"] is False,
+ f"{path}: warmup instruction {name}")
+ require(warmup["role_cursor_scalar_observations"] == 12800,
+ f"{path}: warmup observations {name}")
+ require(warmup["predictor_max_abs_error"] <= 1e-5,
+ f"{path}: predictor {name}")
+ condition = row["conditions"][name]
+ require(len(condition["daily_success"]) == 14,
+ f"{path}: trajectory {name}")
+ cost = condition["cost"]
+ require(cost["ordinary_state_episode_steps"] == 25088,
+ f"{path}: ordinary cost {name}")
+ require(cost["online_role_perturbation_events"] == 98,
+ f"{path}: event cost {name}")
+ expected_online = 12544 if name in (
+ "intact", "plasticity_lesion") else 0
+ require(cost["conservative_online_cursor_scalar_observations"]
+ == expected_online, f"{path}: cursor cost {name}")
+ require(row["signatures"]["evaluation_episodes"] == 256,
+ f"{path}: evaluation episodes")
+
+
+def main():
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--results", default="results/bci_td_confirmation")
+ 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("--out", default="results/bci_td_confirmation_gate.json")
+ 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)
+ require(d4.get("protocol") ==
+ "kp_dynamic_neutral_projection_confirmation_v1"
+ and d4.get("status") == "passed"
+ and d4.get("review_score_after") == 7, "D4 gate")
+ require(r1.get("protocol") == "oral_b_td_development_v1"
+ and r1.get("status") == "passed"
+ and r1.get("complete_grid") is True
+ and r1.get("oral_b_confirmation_opened") is True
+ and r1.get("confirmation_seeds_touched") is False,
+ "R1 gate")
+ eta = float(r1["selected"]["eta"])
+ require(eta in (0.03, 0.1), "R1 selected eta")
+ digests = {
+ "runner": sha256(RUNNER_PATH),
+ "development_runner": sha256(DEVELOPMENT_RUNNER_PATH),
+ "protocol": sha256(PROTOCOL_PATH),
+ "d4_gate": sha256(args.d4_gate),
+ "r1_gate": sha256(args.r1_gate),
+ }
+ require(r1.get("d4_gate_sha256") == digests["d4_gate"],
+ "R1/D4 digest binding")
+
+ expected_names = {
+ f"bci_td_confirm_t{task_seed}_m{model_seed}.json"
+ for task_seed in TASK_SEEDS for model_seed in MODEL_SEEDS
+ }
+ observed_names = {
+ os.path.basename(path)
+ for path in glob.glob(os.path.join(args.results, "*.json"))
+ }
+ require(observed_names == expected_names,
+ f"confirmation grid drift: missing={sorted(expected_names-observed_names)}, "
+ f"extra={sorted(observed_names-expected_names)}")
+ records = {}
+ commits = set()
+ source_sha256 = {}
+ for task_seed in TASK_SEEDS:
+ for model_seed in MODEL_SEEDS:
+ path = os.path.join(
+ args.results,
+ f"bci_td_confirm_t{task_seed}_m{model_seed}.json")
+ with open(path) as handle:
+ row = json.load(handle)
+ validate(row, path, eta, digests)
+ records[(task_seed, model_seed)] = row
+ commits.add(row["provenance"]["git_commit"])
+ source_sha256[path] = sha256(path)
+ require(len(commits) == 1, "R2 source revision drift")
+
+ metrics = {}
+ getters = {
+ "intact_final": lambda r: r["conditions"]["intact"]["final_success"],
+ "intact_gain": lambda r: r["conditions"]["intact"]["learning_gain"],
+ "fixed_final_gap": lambda r: (
+ r["conditions"]["intact"]["final_success"]
+ - r["conditions"]["fixed_vectorizer"]["final_success"]),
+ "oracle_final_deficit": lambda r: (
+ r["conditions"]["oracle_role"]["final_success"]
+ - r["conditions"]["intact"]["final_success"]),
+ "plasticity_half_margin": lambda r: (
+ 0.5 * r["conditions"]["intact"]["learning_gain"]
+ - r["conditions"]["plasticity_lesion"]["learning_gain"]),
+ "role_cosine": lambda r: r["conditions"]["intact"][
+ "role_cosine_after_training"],
+ "residual_soma_corr": lambda r: r["signatures"][
+ "mean_abs_residual_soma_corr"],
+ "raw_residual_corr_gap": lambda r: r["signatures"][
+ "raw_minus_residual_abs_soma_corr"],
+ "surrounding_event_accuracy": lambda r: r["signatures"][
+ "surrounding_event_decoder_balanced_acc"],
+ "decoder_distance_corr": lambda r: r["signatures"][
+ "decoder_distance_residual_corr"],
+ "residual_outcome_accuracy": lambda r: r["signatures"][
+ "residual_outcome_decoder_balanced_acc"],
+ "residual_soma_outcome_gap": lambda r: r["signatures"][
+ "residual_minus_soma_outcome_acc"],
+ "sign_inversion": lambda r: r["signatures"][
+ "causal_role_sign_inversion_index"],
+ "velocity_advantage": lambda r: r["signatures"][
+ "velocity_minus_error_abs_cv_corr"],
+ "longitudinal_prediction": lambda r: r["signatures"][
+ "early_residual_late_activity_change_corr"],
+ }
+ clustered = {}
+ for name, getter in getters.items():
+ values = task_cluster_values(records, getter)
+ clustered[name] = values
+ metrics[name] = summarize(values)
+
+ all_role_cosines = [getters["role_cosine"](row)
+ for row in records.values()]
+ all_signs = [getters["sign_inversion"](row)
+ for row in records.values()]
+ learning_checks = {
+ "mean_intact_final_at_least_0p70":
+ metrics["intact_final"]["mean"] >= 0.70,
+ "every_task_mean_final_at_least_0p60":
+ min(clustered["intact_final"]) >= 0.60,
+ "intact_final_lower_bound_at_least_0p60":
+ metrics["intact_final"]["one_sided_95pct_lower"] >= 0.60,
+ "mean_learning_gain_at_least_0p10":
+ metrics["intact_gain"]["mean"] >= 0.10,
+ "learning_gain_lower_bound_at_least_0p05":
+ metrics["intact_gain"]["one_sided_95pct_lower"] >= 0.05,
+ "mean_fixed_gap_at_least_0p20":
+ metrics["fixed_final_gap"]["mean"] >= 0.20,
+ "fixed_gap_lower_bound_at_least_0p10":
+ metrics["fixed_final_gap"]["one_sided_95pct_lower"] >= 0.10,
+ "mean_oracle_deficit_at_most_0p10":
+ metrics["oracle_final_deficit"]["mean"] <= 0.10,
+ "oracle_deficit_upper_bound_at_most_0p20":
+ metrics["oracle_final_deficit"]["one_sided_95pct_upper"] <= 0.20,
+ "plasticity_lesion_half_margin_lower_bound_nonnegative":
+ metrics["plasticity_half_margin"]["one_sided_95pct_lower"] >= 0.0,
+ "mean_role_cosine_at_least_0p80":
+ metrics["role_cosine"]["mean"] >= 0.80,
+ "every_role_cosine_at_least_0p70": min(all_role_cosines) >= 0.70,
+ }
+ innovation_checks = {
+ "mean_residual_soma_corr_at_most_0p10":
+ metrics["residual_soma_corr"]["mean"] <= 0.10,
+ "residual_soma_corr_upper_bound_at_most_0p12":
+ metrics["residual_soma_corr"]["one_sided_95pct_upper"] <= 0.12,
+ "mean_raw_residual_corr_gap_at_least_0p20":
+ metrics["raw_residual_corr_gap"]["mean"] >= 0.20,
+ "raw_residual_corr_gap_lower_bound_at_least_0p15":
+ metrics["raw_residual_corr_gap"]["one_sided_95pct_lower"] >= 0.15,
+ "mean_surrounding_event_accuracy_at_least_0p55":
+ metrics["surrounding_event_accuracy"]["mean"] >= 0.55,
+ "surrounding_event_accuracy_lower_bound_at_least_0p52":
+ metrics["surrounding_event_accuracy"]["one_sided_95pct_lower"] >= 0.52,
+ "mean_decoder_distance_corr_at_least_0p10":
+ metrics["decoder_distance_corr"]["mean"] >= 0.10,
+ "decoder_distance_corr_lower_bound_at_least_0p02":
+ metrics["decoder_distance_corr"]["one_sided_95pct_lower"] >= 0.02,
+ }
+ vectorization_checks = {
+ "mean_residual_outcome_accuracy_at_least_0p57":
+ metrics["residual_outcome_accuracy"]["mean"] >= 0.57,
+ "residual_outcome_accuracy_lower_bound_at_least_0p53":
+ metrics["residual_outcome_accuracy"]["one_sided_95pct_lower"] >= 0.53,
+ "mean_residual_soma_outcome_gap_at_least_0p03":
+ metrics["residual_soma_outcome_gap"]["mean"] >= 0.03,
+ "residual_soma_outcome_gap_lower_bound_nonnegative":
+ metrics["residual_soma_outcome_gap"]["one_sided_95pct_lower"] >= 0.0,
+ "positive_sign_inversion_in_at_least_25_of_30":
+ sum(value > 0 for value in all_signs) >= 25,
+ "positive_sign_inversion_in_every_task_cluster":
+ min(clustered["sign_inversion"]) > 0.0,
+ "mean_velocity_advantage_at_least_0p05":
+ metrics["velocity_advantage"]["mean"] >= 0.05,
+ "velocity_advantage_lower_bound_nonnegative":
+ metrics["velocity_advantage"]["one_sided_95pct_lower"] >= 0.0,
+ "mean_longitudinal_prediction_at_least_0p30":
+ metrics["longitudinal_prediction"]["mean"] >= 0.30,
+ "longitudinal_prediction_lower_bound_at_least_0p10":
+ metrics["longitudinal_prediction"]["one_sided_95pct_lower"] >= 0.10,
+ }
+ checks = {
+ "all_records_finite_paired_and_cost_audited": True,
+ "learning_and_plasticity": learning_checks,
+ "innovation_identification_and_network_prediction": innovation_checks,
+ "outcome_and_causal_role_vectorization": vectorization_checks,
+ }
+ passed = all(
+ value
+ for category, values in checks.items()
+ for value in ([values] if isinstance(values, bool) else values.values()))
+ output = {
+ "protocol": "oral_b_td_confirmation_v1",
+ "status": "passed" if passed else "failed",
+ "complete_grid": True,
+ "selected_eta": eta,
+ "checks": checks,
+ "metrics": metrics,
+ "positive_sign_count": sum(value > 0 for value in all_signs),
+ "source_commit": next(iter(commits)),
+ "source_sha256": source_sha256,
+ "d4_gate_sha256": digests["d4_gate"],
+ "r1_gate_sha256": digests["r1_gate"],
+ "protocol_sha256": digests["protocol"],
+ "oral_b_plasticity_innovation_established": passed,
+ "online_control_or_desired_velocity_established": False,
+ "review_score_before": 7,
+ "review_score_after": 8 if passed else 7,
+ "score_change_rule": (
+ "only a complete untouched R2 pass establishes oral-B "
+ "innovation-guided plasticity; kappa=0 cannot establish online control"),
+ }
+ os.makedirs(os.path.dirname(os.path.abspath(args.out)), exist_ok=True)
+ if os.path.exists(args.out):
+ with open(args.out) as handle:
+ existing = json.load(handle)
+ require(existing == output,
+ "existing R2 gate differs from deterministic re-audit")
+ print(json.dumps(output, indent=2))
+ return
+ with open(args.out, "w") as handle:
+ json.dump(output, handle, indent=2, sort_keys=True)
+ handle.write("\n")
+ print(json.dumps(output, indent=2))
+
+
+if __name__ == "__main__":
+ main()