From bcd6b1a302d9591b9db0d13c8cd231c6575d160c Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Wed, 22 Jul 2026 05:42:26 -0500 Subject: bci: add continuous local-learning dynamics --- experiments/bci_smoke.py | 99 ++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 99 insertions(+) create mode 100644 experiments/bci_smoke.py (limited to 'experiments') 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") -- cgit v1.2.3