summaryrefslogtreecommitdiff
path: root/experiments/bci_v2_smoke.py
diff options
context:
space:
mode:
Diffstat (limited to 'experiments/bci_v2_smoke.py')
-rw-r--r--experiments/bci_v2_smoke.py68
1 files changed, 68 insertions, 0 deletions
diff --git a/experiments/bci_v2_smoke.py b/experiments/bci_v2_smoke.py
index fc6509b..6b7a900 100644
--- a/experiments/bci_v2_smoke.py
+++ b/experiments/bci_v2_smoke.py
@@ -8,6 +8,7 @@ import torch
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
from sdil.bci import generate_trajectories
from sdil.bci_v2 import BCIV2, BCIV2Config, run_day_v2
+from sdil.bci_v2_metrics import annotate_events, signature_metrics_v2
def check_role_estimator():
@@ -284,10 +285,77 @@ def check_horizon_and_pairing():
print("v2 fixed-horizon terminal accounting and paired replay: exact")
+def check_signature_metric_mechanics():
+ cfg = BCIV2Config(
+ days=2, episodes_per_day=128, steps_per_episode=5,
+ target=0.15, forward_eta=0.03,
+ feedback="performance_velocity",
+ )
+ trajectories = generate_trajectories(cfg, 9)
+ model = BCIV2(cfg, model_seed=10)
+ model.P.copy_(model.coupling)
+ model.A.copy_(model.role)
+ training = run_day_v2(
+ model,
+ trajectories,
+ 0,
+ collect=True,
+ learn_role=False,
+ probe_role=False,
+ learn_predictor=False,
+ )
+ annotate_events(
+ training["events"], 0, training["success"], 0,
+ cfg.steps_per_episode,
+ )
+ modes = {}
+ settings = {
+ "intact": {},
+ "acute_critic_lesion": {"critic_enabled": False},
+ "acute_outcome_lesion": {
+ "terminal_outcome_enabled": False
+ },
+ }
+ for name, setting in settings.items():
+ report = run_day_v2(
+ model.clone(),
+ trajectories,
+ 1,
+ collect=True,
+ plasticity_gain=0.0,
+ learn_role=False,
+ probe_role=False,
+ learn_predictor=False,
+ learn_critic=False,
+ **setting,
+ )
+ annotate_events(
+ report["events"], 0, report["success"], 10_000,
+ cfg.steps_per_episode,
+ )
+ modes[name] = report["events"]
+ metrics = signature_metrics_v2(
+ training["events"], modes, cfg, model.role
+ )
+ assert metrics["challenge_episodes"] == cfg.episodes_per_day
+ assert (
+ metrics["challenge_success_count"]
+ + metrics["challenge_failure_count"]
+ == cfg.episodes_per_day
+ )
+ assert all(
+ torch.isfinite(torch.tensor(value))
+ for value in metrics.values()
+ if isinstance(value, (int, float))
+ )
+ print("v2 paired terminal-outcome signature metrics: finite")
+
+
if __name__ == "__main__":
check_role_estimator()
check_predictor()
check_td_and_eligibility_updates()
check_terminal_outcomes_and_lesions()
check_horizon_and_pairing()
+ check_signature_metric_mechanics()
print("ALL ORAL-B-V2 MECHANICS CHECKS PASSED")