From c3709dc0dead3234a7423ae705da9703e4d52f68 Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Wed, 22 Jul 2026 05:44:14 -0500 Subject: bci: add preregistered signature metrics --- experiments/bci_smoke.py | 41 +++++++++++++++++++++++++++++++++++++++++ 1 file changed, 41 insertions(+) (limited to 'experiments/bci_smoke.py') 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") -- cgit v1.2.3