summaryrefslogtreecommitdiff
path: root/experiments/analyze_bci_v2_recovery_confirmation.py
diff options
context:
space:
mode:
Diffstat (limited to 'experiments/analyze_bci_v2_recovery_confirmation.py')
-rw-r--r--experiments/analyze_bci_v2_recovery_confirmation.py498
1 files changed, 498 insertions, 0 deletions
diff --git a/experiments/analyze_bci_v2_recovery_confirmation.py b/experiments/analyze_bci_v2_recovery_confirmation.py
new file mode 100644
index 0000000..94143be
--- /dev/null
+++ b/experiments/analyze_bci_v2_recovery_confirmation.py
@@ -0,0 +1,498 @@
+#!/usr/bin/env python3
+"""Audit untouched oral-B-v2 cold-start recovery confirmation."""
+import argparse
+import glob
+import hashlib
+import json
+import math
+import os
+import statistics
+
+from experiments.bci_v2_recovery_confirmation import (
+ MODEL_SEEDS,
+ TASK_SEEDS,
+)
+from experiments.bci_v2_recovery_run import (
+ CHALLENGE_EPISODES,
+ FIXED_CONFIG,
+ TARGETS,
+ build_recovery_config,
+ source_paths,
+)
+from experiments.bci_v2_run import CONDITIONS
+
+
+ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
+RUNNER_PATH = os.path.join(
+ ROOT, "experiments", "bci_v2_recovery_confirmation.py"
+)
+T_CRITICAL_ONE_SIDED_95_DF5 = 2.015048373
+
+
+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_bound(values):
+ return (
+ statistics.mean(values)
+ - T_CRITICAL_ONE_SIDED_95_DF5
+ * statistics.stdev(values)
+ / math.sqrt(len(values))
+ )
+
+
+def upper_bound(values):
+ return (
+ statistics.mean(values)
+ + T_CRITICAL_ONE_SIDED_95_DF5
+ * statistics.stdev(values)
+ / math.sqrt(len(values))
+ )
+
+
+def summarize(values):
+ return {
+ "by_task_seed": values,
+ "mean": statistics.mean(values),
+ "one_sided_95pct_lower": lower_bound(values),
+ "one_sided_95pct_upper": upper_bound(values),
+ }
+
+
+def task_clusters(records, getter):
+ return [
+ statistics.mean(
+ getter(records[(task_seed, model_seed)])
+ for model_seed in MODEL_SEEDS
+ )
+ for task_seed in TASK_SEEDS
+ ]
+
+
+def confirmation_paths(r1_gate):
+ paths = source_paths()
+ paths["development_runner"] = paths["runner"]
+ paths["runner"] = RUNNER_PATH
+ paths["r1_gate"] = os.path.abspath(r1_gate)
+ return paths
+
+
+def validate(row, path, digests):
+ require(row.get("schema_version") == 3, 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 and model_seed in MODEL_SEEDS,
+ f"{path}: seeds",
+ )
+ protocol = row.get("protocol", {})
+ require(
+ protocol.get("name")
+ == "oral_b_v2_cold_start_recovery_confirmation_v1"
+ and protocol.get("split") == "untouched_confirmation"
+ and protocol.get("training_task_seed") == task_seed
+ and protocol.get("model_seed") == model_seed
+ and protocol.get("fixed_config") == FIXED_CONFIG
+ and protocol.get("no_further_selection") is True
+ and protocol.get("confirmation_grid_size") == 30
+ and protocol.get("protocol_sha256") == digests["protocol"]
+ and protocol.get("r1_gate_sha256") == digests["r1_gate"],
+ f"{path}: protocol",
+ )
+ provenance = row.get("provenance", {})
+ require(
+ provenance.get("git_tracked_dirty") is False
+ and provenance.get("tracked_inputs")
+ and all(provenance["tracked_inputs"].values())
+ and set(provenance["tracked_inputs"]) == set(digests)
+ and provenance.get("input_sha256") == digests,
+ f"{path}: provenance",
+ )
+ require(
+ row.get("config") == vars(build_recovery_config()),
+ f"{path}: config",
+ )
+ require(
+ row.get("split") == "untouched_confirmation"
+ and row.get("finite") is True
+ and finite_tree(row),
+ f"{path}: finite/split",
+ )
+ require(
+ set(row.get("conditions", {})) == set(CONDITIONS),
+ f"{path}: conditions",
+ )
+ for name in CONDITIONS:
+ warmup = row["warmup"][name]
+ condition = row["conditions"][name]
+ cost = condition["cost"]
+ require(
+ warmup["batches"] == 100
+ and warmup["examples"] == 6400
+ and warmup["instruction_present"] is False
+ and warmup["role_cursor_scalar_observations"] == 12800
+ and warmup["predictor_max_abs_error"] <= 1e-5
+ and len(condition["daily_success"]) == 14
+ and cost["maximum_state_episode_steps"] == 25088
+ and cost["cursor_scalar_observations"]
+ == 2 * cost["role_probe_examples"]
+ and cost["task_loss_queries"] == 0
+ and cost["reverse_mode_calls"] == 0,
+ f"{path}: condition {name}",
+ )
+ challenge = row["assays"]["challenge"]
+ require(
+ challenge["targets"] == list(TARGETS)
+ and challenge["episodes_per_target"] == CHALLENGE_EPISODES
+ and challenge["selection_over_targets"] is False
+ and challenge["maximum_steps_per_episode"] == 28
+ and row["signatures"]["challenge_episodes"]
+ == len(TARGETS) * CHALLENGE_EPISODES,
+ f"{path}: target ladder",
+ )
+
+
+def checks(metrics, clustered, all_values):
+ return {
+ "all_records_finite_paired_and_cost_audited": True,
+ "learning_and_plasticity": {
+ "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_role_gap_at_least_0p20":
+ metrics["fixed_role_gap"]["mean"] >= 0.20,
+ "fixed_role_gap_lower_bound_at_least_0p10":
+ metrics["fixed_role_gap"][
+ "one_sided_95pct_lower"
+ ] >= 0.10,
+ "mean_oracle_deficit_at_most_0p10":
+ metrics["oracle_deficit"]["mean"] <= 0.10,
+ "oracle_deficit_upper_bound_at_most_0p20":
+ metrics["oracle_deficit"][
+ "one_sided_95pct_upper"
+ ] <= 0.20,
+ "plasticity_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_values["role_cosine"]) >= 0.70,
+ },
+ "innovation_and_network_prediction": {
+ "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_accuracy_at_least_0p52":
+ metrics["surrounding_accuracy"]["mean"] >= 0.52,
+ "surrounding_accuracy_lower_bound_at_least_0p50":
+ metrics["surrounding_accuracy"][
+ "one_sided_95pct_lower"
+ ] >= 0.50,
+ "mean_decoder_corr_at_least_0p05":
+ metrics["decoder_corr"]["mean"] >= 0.05,
+ "decoder_corr_lower_bound_nonnegative":
+ metrics["decoder_corr"][
+ "one_sided_95pct_lower"
+ ] >= 0.0,
+ "positive_sign_in_at_least_25_of_30":
+ sum(
+ value > 0 for value in all_values["sign_inversion"]
+ ) >= 25,
+ "positive_sign_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,
+ },
+ "target_ladder_outcome_surprise": {
+ "every_challenge_fraction_between_0p10_and_0p90":
+ min(all_values["challenge_fraction"]) >= 0.10
+ and max(all_values["challenge_fraction"]) <= 0.90,
+ "mean_terminal_accuracy_at_least_0p80":
+ metrics["terminal_accuracy"]["mean"] >= 0.80,
+ "terminal_accuracy_lower_bound_at_least_0p75":
+ metrics["terminal_accuracy"][
+ "one_sided_95pct_lower"
+ ] >= 0.75,
+ "mean_terminal_separation_at_least_0p20":
+ metrics["terminal_separation"]["mean"] >= 0.20,
+ "terminal_separation_lower_bound_at_least_0p15":
+ metrics["terminal_separation"][
+ "one_sided_95pct_lower"
+ ] >= 0.15,
+ "mean_outcome_lesion_drop_at_least_0p20":
+ metrics["outcome_lesion_drop"]["mean"] >= 0.20,
+ "outcome_lesion_drop_lower_bound_at_least_0p15":
+ metrics["outcome_lesion_drop"][
+ "one_sided_95pct_lower"
+ ] >= 0.15,
+ "mean_critic_expectedness_at_least_0p05":
+ metrics["critic_expectedness"]["mean"] >= 0.05,
+ "critic_expectedness_lower_bound_at_least_0p02":
+ metrics["critic_expectedness"][
+ "one_sided_95pct_lower"
+ ] >= 0.02,
+ "every_critic_value_corr_at_least_0p95":
+ min(all_values["critic_value_corr"]) >= 0.95,
+ },
+ }
+
+
+def main():
+ parser = argparse.ArgumentParser()
+ parser.add_argument(
+ "--results",
+ default="results/bci_v2_recovery_confirmation",
+ )
+ parser.add_argument(
+ "--r1-gate",
+ default="results/bci_v2_recovery_dev_gate.json",
+ )
+ parser.add_argument(
+ "--out",
+ default="results/bci_v2_recovery_confirmation_gate.json",
+ )
+ args = parser.parse_args()
+ with open(args.r1_gate) as handle:
+ r1 = json.load(handle)
+ require(
+ r1.get("protocol")
+ == "oral_b_v2_cold_start_recovery_development_v1"
+ and r1.get("status") == "passed"
+ and r1.get("complete_grid") is True
+ and r1.get("fixed_config") == FIXED_CONFIG
+ and r1.get("confirmation_seeds_touched") is False
+ and r1.get("recovery_confirmation_opened") is True,
+ "recovery R1 gate",
+ )
+ development_digests = {
+ name: sha256(path)
+ for name, path in source_paths().items()
+ }
+ require(
+ r1.get("input_sha256") == development_digests,
+ "R1 source binding",
+ )
+ digests = {
+ name: sha256(path)
+ for name, path in confirmation_paths(args.r1_gate).items()
+ }
+ expected = {
+ f"bci_v2_recovery_confirm_t{task_seed}_m{model_seed}.json"
+ for task_seed in TASK_SEEDS
+ for model_seed in MODEL_SEEDS
+ }
+ observed = {
+ os.path.basename(path)
+ for path in glob.glob(os.path.join(args.results, "*.json"))
+ }
+ require(
+ observed == expected,
+ (
+ f"confirmation drift: missing={sorted(expected-observed)}, "
+ f"extra={sorted(observed-expected)}"
+ ),
+ )
+ 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_v2_recovery_confirm_t{task_seed}"
+ f"_m{model_seed}.json"
+ ),
+ )
+ with open(path) as handle:
+ row = json.load(handle)
+ validate(row, path, digests)
+ records[(task_seed, model_seed)] = row
+ commits.add(row["provenance"]["git_commit"])
+ source_sha256[path] = sha256(path)
+ require(len(commits) == 1, "confirmation source revision drift")
+
+ getters = {
+ "intact_final": lambda row:
+ row["conditions"]["intact"]["final_success"],
+ "intact_gain": lambda row:
+ row["conditions"]["intact"]["learning_gain"],
+ "fixed_role_gap": lambda row: (
+ row["conditions"]["intact"]["final_success"]
+ - row["conditions"]["fixed_role"]["final_success"]
+ ),
+ "oracle_deficit": lambda row: (
+ row["conditions"]["oracle_role"]["final_success"]
+ - row["conditions"]["intact"]["final_success"]
+ ),
+ "plasticity_half_margin": lambda row: (
+ 0.5 * row["conditions"]["intact"]["learning_gain"]
+ - row["conditions"]["plasticity_lesion"]["learning_gain"]
+ ),
+ "role_cosine": lambda row:
+ row["conditions"]["intact"]["role_cosine_after_training"],
+ "critic_training_gap": lambda row: (
+ row["conditions"]["intact"]["final_success"]
+ - row["conditions"]["critic_training_lesion"][
+ "final_success"
+ ]
+ ),
+ "outcome_training_gap": lambda row: (
+ row["conditions"]["intact"]["final_success"]
+ - row["conditions"]["outcome_training_lesion"][
+ "final_success"
+ ]
+ ),
+ "residual_soma_corr": lambda row:
+ row["signatures"]["mean_abs_residual_soma_corr"],
+ "raw_residual_corr_gap": lambda row:
+ row["signatures"][
+ "raw_minus_residual_abs_soma_corr"
+ ],
+ "surrounding_accuracy": lambda row:
+ row["signatures"][
+ "surrounding_event_decoder_balanced_acc"
+ ],
+ "decoder_corr": lambda row:
+ row["signatures"]["decoder_distance_residual_corr"],
+ "sign_inversion": lambda row:
+ row["signatures"][
+ "causal_role_sign_inversion_index"
+ ],
+ "velocity_advantage": lambda row:
+ row["signatures"][
+ "velocity_minus_error_abs_cv_corr"
+ ],
+ "challenge_fraction": lambda row:
+ row["signatures"]["challenge_success_fraction"],
+ "terminal_accuracy": lambda row:
+ row["signatures"][
+ "terminal_residual_outcome_balanced_acc"
+ ],
+ "terminal_separation": lambda row:
+ row["signatures"][
+ "terminal_role_aligned_outcome_separation"
+ ],
+ "outcome_lesion_drop": lambda row:
+ row["signatures"][
+ "terminal_outcome_separation_drop_under_acute_lesion"
+ ],
+ "critic_expectedness": lambda row:
+ row["signatures"][
+ "mean_critic_expectedness_contribution"
+ ],
+ "critic_value_corr": lambda row:
+ row["signatures"][
+ "critic_contribution_value_prediction_corr"
+ ],
+ }
+ clustered = {
+ name: task_clusters(records, getter)
+ for name, getter in getters.items()
+ }
+ metrics = {
+ name: summarize(values)
+ for name, values in clustered.items()
+ }
+ all_values = {
+ name: [getter(row) for row in records.values()]
+ for name, getter in getters.items()
+ }
+ gate_checks = checks(metrics, clustered, all_values)
+ passed = all(
+ value
+ for category in gate_checks.values()
+ for value in (
+ [category]
+ if isinstance(category, bool)
+ else category.values()
+ )
+ )
+ output = {
+ "protocol":
+ "oral_b_v2_cold_start_recovery_confirmation_v1",
+ "status": "passed" if passed else "failed",
+ "complete_grid": True,
+ "fixed_config": FIXED_CONFIG,
+ "checks": gate_checks,
+ "metrics": metrics,
+ "positive_sign_count": sum(
+ value > 0 for value in all_values["sign_inversion"]
+ ),
+ "source_commit": next(iter(commits)),
+ "source_sha256": source_sha256,
+ "input_sha256": digests,
+ "oral_b_v2_outcome_surprise_established": passed,
+ "old_r2_status_preserved": "failed",
+ "failed_v2_development_status_preserved": "failed",
+ "old_oral_a_gate_remains_closed": True,
+ "new_oral_a_v2_protocol_may_be_frozen": passed,
+ "review_score_before": 7,
+ "review_score_after": 8 if passed else 7,
+ "score_change_rule": (
+ "only a complete untouched recovery R2 pass establishes "
+ "role-vectorized TD outcome surprise; both prior failures "
+ "and the old oral-A gate remain unchanged"
+ ),
+ }
+ 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 recovery R2 gate differs")
+ else:
+ 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()