From dae172b1425c646f0441609c28fc5156b50f820d Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Thu, 23 Jul 2026 08:07:50 -0500 Subject: protocol: freeze oral-B-v2 cold-start recovery --- sdil/bci_v2_recovery_metrics.py | 44 +++++++++++++++++++++++++++++++++++++++++ 1 file changed, 44 insertions(+) create mode 100644 sdil/bci_v2_recovery_metrics.py (limited to 'sdil/bci_v2_recovery_metrics.py') diff --git a/sdil/bci_v2_recovery_metrics.py b/sdil/bci_v2_recovery_metrics.py new file mode 100644 index 0000000..78d3449 --- /dev/null +++ b/sdil/bci_v2_recovery_metrics.py @@ -0,0 +1,44 @@ +"""Target-ladder wrapper for the oral-B-v2 cold-start recovery.""" + +from sdil.bci_v2_metrics import ( + annotate_events, + signature_metrics_v2, +) + + +def annotate_target_events( + events, + day, + final_success, + episode_offset, + target_index, +): + """Reuse paired grouping while encoding the fixed target-ladder index.""" + return annotate_events( + events, + day, + final_success, + episode_offset, + target_index, + ) + + +def recovery_signature_metrics( + train_events, + challenge_mode_events, + cfg, + role, + targets, +): + values = signature_metrics_v2( + train_events, + challenge_mode_events, + cfg, + role, + ) + by_index = values.pop("horizon_success_fraction") + values["target_success_fraction"] = { + str(target): by_index[str(index)] + for index, target in enumerate(targets, start=1) + } + return values -- cgit v1.2.3