summaryrefslogtreecommitdiff
path: root/experiments/bci_smoke.py
diff options
context:
space:
mode:
Diffstat (limited to 'experiments/bci_smoke.py')
-rw-r--r--experiments/bci_smoke.py99
1 files changed, 99 insertions, 0 deletions
diff --git a/experiments/bci_smoke.py b/experiments/bci_smoke.py
new file mode 100644
index 0000000..fe3df78
--- /dev/null
+++ b/experiments/bci_smoke.py
@@ -0,0 +1,99 @@
+"""CPU-only mechanics checks for the continuous Harnett-signature BCI."""
+import os
+import sys
+
+import torch
+
+sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
+from sdil.bci import (BCIConfig, BCISDIL, causal_roles, generate_trajectories,
+ run_day)
+
+
+def check_roles_and_pairing():
+ cfg = BCIConfig(days=2, episodes_per_day=8, steps_per_episode=4)
+ role = causal_roles(cfg)
+ assert torch.allclose(role[:5], torch.full((5,), 0.2))
+ assert torch.allclose(role[5:10], torch.full((5,), -0.2))
+ assert torch.equal(role[10:], torch.zeros(30))
+ first = generate_trajectories(cfg, 11)
+ second = generate_trajectories(cfg, 11)
+ third = generate_trajectories(cfg, 12)
+ assert torch.equal(first.context, second.context)
+ assert torch.equal(first.process_noise, second.process_noise)
+ assert torch.equal(first.perturbations, second.perturbations)
+ assert not torch.equal(first.context, third.context)
+ assert torch.equal(first.context[..., 0], torch.ones_like(first.context[..., 0]))
+ print("BCI causal roles and task-seed pairing: exact")
+
+
+def check_causal_estimator():
+ cfg = BCIConfig(days=2, episodes_per_day=4096, steps_per_episode=2)
+ model = BCISDIL(cfg, model_seed=3)
+ generator = torch.Generator().manual_seed(4)
+ soma = torch.randn(4096, cfg.n_neurons, generator=generator) * 0.2
+ xi = torch.empty_like(soma).bernoulli_(0.5, generator=generator).mul_(2).sub_(1)
+ estimate = model.causal_targets(soma, cfg.target, xi)
+ exact = model.exact_causal_direction(soma, cfg.target)
+ mean_estimate = estimate.mean(0)
+ mean_exact = exact.mean(0)
+ cosine = torch.nn.functional.cosine_similarity(
+ mean_estimate, mean_exact, dim=0).item()
+ assert cosine > 0.99
+ assert estimate[:, 10:].abs().mean() > 0 # simultaneous cross-cell variance
+ assert mean_estimate[10:].abs().max() < 0.03
+ print(f"BCI simultaneous causal estimator mean cosine: {cosine:.6f}")
+
+
+def check_predictor_identification():
+ cfg = BCIConfig(days=2, episodes_per_day=256, steps_per_episode=2,
+ predictor_eta=0.2)
+ model = BCISDIL(cfg, model_seed=5)
+ generator = torch.Generator().manual_seed(6)
+ active = torch.ones(256, dtype=torch.bool)
+ for _ in range(100):
+ soma = torch.randn(256, cfg.n_neurons, generator=generator).tanh()
+ ordinary = model.coupling * soma
+ model.update_predictor(soma, ordinary, active)
+ assert (model.P - model.coupling).abs().max() < 1e-4
+ assert model.P_bias.abs().max() < 1e-5
+ print("BCI neutral per-cell predictor identifies ordinary coupling: exact")
+
+
+def check_phase_masks_and_velocity_reset():
+ cfg = BCIConfig(days=2, episodes_per_day=8, steps_per_episode=4,
+ kappa=0.3, feedback="error_velocity")
+ trajectories = generate_trajectories(cfg, 7)
+ initial = BCISDIL(cfg, model_seed=8)
+ intact = initial.clone()
+ online_lesion = initial.clone()
+ plasticity_lesion = initial.clone()
+
+ active = torch.ones(8, dtype=torch.bool)
+ args = (
+ trajectories.context[0, :, 0], trajectories.process_noise[0, :, 0],
+ trajectories.perturbations[0, :, 0], torch.zeros(8, cfg.n_neurons),
+ None, active, 1)
+ out_intact = intact.step(*args, control_gain=1.0, plasticity_gain=1.0,
+ learn_vectorizer=False)
+ out_online = online_lesion.step(*args, control_gain=0.0, plasticity_gain=1.0,
+ learn_vectorizer=False)
+ out_plastic = plasticity_lesion.step(
+ *args, control_gain=1.0, plasticity_gain=0.0, learn_vectorizer=False)
+ assert not torch.equal(out_intact["soma"], out_online["soma"])
+ assert torch.equal(intact.W, online_lesion.W)
+ assert torch.equal(out_intact["soma"], out_plastic["soma"])
+ assert torch.equal(plasticity_lesion.W, initial.W)
+ assert torch.equal(out_intact["feedback"][:, 1], torch.zeros(8))
+
+ day = run_day(initial.clone(), trajectories, 0, collect=True)
+ assert len(day["events"]) == cfg.steps_per_episode
+ assert day["global_step"] == cfg.steps_per_episode
+ print("BCI online/plasticity phase masks and episode velocity reset: exact")
+
+
+if __name__ == "__main__":
+ check_roles_and_pairing()
+ check_causal_estimator()
+ check_predictor_identification()
+ check_phase_masks_and_velocity_reset()
+ print("ALL BCI MECHANICS CHECKS PASSED")