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.py41
1 files changed, 41 insertions, 0 deletions
diff --git a/experiments/bci_smoke.py b/experiments/bci_smoke.py
index fe3df78..5cd5b65 100644
--- a/experiments/bci_smoke.py
+++ b/experiments/bci_smoke.py
@@ -7,6 +7,8 @@ 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)
+from sdil.bci_metrics import (annotate_day_events, grouped_classification_accuracy,
+ grouped_regression_correlation, signature_metrics)
def check_roles_and_pairing():
@@ -91,9 +93,48 @@ def check_phase_masks_and_velocity_reset():
print("BCI online/plasticity phase masks and episode velocity reset: exact")
+def check_grouped_decoders_and_metrics():
+ generator = torch.Generator().manual_seed(9)
+ groups = torch.arange(200)
+ x = torch.randn(200, 3, generator=generator)
+ labels = x[:, 0] + 0.2 * x[:, 1] > 0
+ accuracy, _ = grouped_classification_accuracy(x, labels, groups)
+ correlation = grouped_regression_correlation(x, 2 * x[:, 0] - x[:, 2], groups)
+ assert accuracy > 0.9
+ assert correlation > 0.99
+
+ cfg = BCIConfig(days=4, episodes_per_day=24, steps_per_episode=6,
+ kappa=0.1)
+ trajectories = generate_trajectories(cfg, 10)
+ model = BCISDIL(cfg, model_seed=11)
+ training_events = []
+ global_step = 0
+ for day in range(cfg.days):
+ report = run_day(
+ model, trajectories, day, collect=True, global_step=global_step)
+ global_step = report["global_step"]
+ annotate_day_events(
+ report["events"], day, report["success"],
+ episode_offset=day * cfg.episodes_per_day)
+ training_events.extend(report["events"])
+ evaluation = run_day(
+ model, trajectories, cfg.days - 1, collect=True,
+ global_step=global_step, plasticity_gain=0.0,
+ learn_vectorizer=False, learn_predictor=False)
+ annotate_day_events(evaluation["events"], 0, evaluation["success"], 10000)
+ metrics = signature_metrics(
+ training_events, evaluation["events"], cfg, model.role)
+ assert metrics["active_training_events"] > 0
+ assert metrics["evaluation_episodes"] == cfg.episodes_per_day
+ assert all(torch.isfinite(torch.tensor(value)) for value in metrics.values())
+ print(f"BCI grouped decoder mechanics: acc={accuracy:.3f}, corr={correlation:.3f}")
+ print("BCI preregistered signature metrics: finite")
+
+
if __name__ == "__main__":
check_roles_and_pairing()
check_causal_estimator()
check_predictor_identification()
check_phase_masks_and_velocity_reset()
+ check_grouped_decoders_and_metrics()
print("ALL BCI MECHANICS CHECKS PASSED")