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 --- experiments/bci_v2_recovery_smoke.py | 104 +++++++++++++++++++++++++++++++++++ 1 file changed, 104 insertions(+) create mode 100644 experiments/bci_v2_recovery_smoke.py (limited to 'experiments/bci_v2_recovery_smoke.py') diff --git a/experiments/bci_v2_recovery_smoke.py b/experiments/bci_v2_recovery_smoke.py new file mode 100644 index 0000000..39055c1 --- /dev/null +++ b/experiments/bci_v2_recovery_smoke.py @@ -0,0 +1,104 @@ +#!/usr/bin/env python3 +"""Endpoint-free checks for the v2 cold-start recovery.""" +from dataclasses import replace +import os +import sys + +import torch + +sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) +from experiments.bci_v2_recovery_run import ( + TARGETS, + build_recovery_config, +) +from sdil.bci_v2 import BCIV2 + + +def copy_state(source, target): + for name in ("W", "A", "coupling", "P", "P_bias", "critic"): + setattr(target, name, getattr(source, name).clone()) + + +def check_dense_velocity_cold_start(): + recovered_cfg = replace( + build_recovery_config(), + days=2, + episodes_per_day=8, + steps_per_episode=2, + target=10.0, + inertia=0.0, + process_noise=0.0, + forward_eta=0.1, + perturb_every=99, + ) + failed_cfg = replace(recovered_cfg, velocity_reward_scale=0.25) + recovered = BCIV2(recovered_cfg, model_seed=1) + recovered.P.copy_(recovered.coupling) + recovered.A.copy_(recovered.role) + recovered.critic.zero_() + failed = BCIV2(failed_cfg, model_seed=1) + copy_state(recovered, failed) + context = torch.zeros(8, recovered_cfg.context_dim) + context[:, 0] = 1.0 + noise = torch.zeros(8, recovered_cfg.n_neurons) + noise[:, :recovered_cfg.n_plus] = torch.linspace( + 0.1, 0.8, 8 + ).unsqueeze(1) + perturbations = torch.ones_like(noise) + recovered_initial_w = recovered.W.clone() + failed_initial_w = failed.W.clone() + recovered_state = recovered.initial_episode_state(8) + failed_state = failed.initial_episode_state(8) + _, recovered_record = recovered.step( + context, + noise, + perturbations, + recovered_state, + global_step=1, + final_step=False, + learn_role=False, + probe_role=False, + learn_predictor=False, + learn_critic=False, + critic_enabled=False, + ) + _, failed_record = failed.step( + context, + noise, + perturbations, + failed_state, + global_step=1, + final_step=False, + learn_role=False, + probe_role=False, + learn_predictor=False, + learn_critic=False, + critic_enabled=False, + ) + assert torch.allclose( + recovered_record["td_innovation"], + 4.0 * failed_record["td_innovation"], + ) + assert torch.allclose( + recovered.W - recovered_initial_w, + 4.0 * (failed.W - failed_initial_w), + atol=1e-7, + ) + print("recovery dense cold-start drive is exactly 4x failed v2") + + +def check_fixed_target_ladder(): + assert TARGETS == (1.55, 1.60, 1.65, 1.70, 1.75, 1.80) + maxima = torch.tensor([1.58, 1.63, 1.68, 1.73, 1.78]) + fractions = torch.tensor([ + (maxima >= target).float().mean() for target in TARGETS + ]) + assert torch.all(fractions[:-1] >= fractions[1:]) + assert bool((fractions > 0).any()) and bool((fractions < 1).any()) + print("recovery target ladder is fixed, ordered, and monotone") + + +if __name__ == "__main__": + check_dense_velocity_cold_start() + check_fixed_target_ladder() + print("ALL V2 COLD-START RECOVERY MECHANICS CHECKS PASSED") -- cgit v1.2.3