diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-23 07:53:33 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-23 07:53:33 -0500 |
| commit | ce10335dcc3aa12f4f1e6dd53352d7e71d0cccf7 (patch) | |
| tree | 1389b61c3e6a1dbdabe7df557128c16e5a3f2122 /experiments | |
| parent | c34675664b949ac7c3aa7a61e714d69183d0ea6f (diff) | |
bci: audit v2 outcome-surprise signatures
Diffstat (limited to 'experiments')
| -rw-r--r-- | experiments/bci_v2_smoke.py | 68 |
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") |
