From dae172b1425c646f0441609c28fc5156b50f820d Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Thu, 23 Jul 2026 08:07:50 -0500 Subject: protocol: freeze oral-B-v2 cold-start recovery --- .../analyze_bci_v2_recovery_confirmation.py | 498 +++++++++++++++++++++ 1 file changed, 498 insertions(+) create mode 100644 experiments/analyze_bci_v2_recovery_confirmation.py (limited to 'experiments/analyze_bci_v2_recovery_confirmation.py') 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() -- cgit v1.2.3