summaryrefslogtreecommitdiff
path: root/experiments/bci_v2_recovery_smoke.py
diff options
context:
space:
mode:
Diffstat (limited to 'experiments/bci_v2_recovery_smoke.py')
-rw-r--r--experiments/bci_v2_recovery_smoke.py104
1 files changed, 104 insertions, 0 deletions
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")