summaryrefslogtreecommitdiff
path: root/experiments
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-07-23 07:53:33 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-07-23 07:53:33 -0500
commitce10335dcc3aa12f4f1e6dd53352d7e71d0cccf7 (patch)
tree1389b61c3e6a1dbdabe7df557128c16e5a3f2122 /experiments
parentc34675664b949ac7c3aa7a61e714d69183d0ea6f (diff)
bci: audit v2 outcome-surprise signatures
Diffstat (limited to 'experiments')
-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")