summaryrefslogtreecommitdiff
path: root/sdil/bci_v2_recovery_metrics.py
blob: 78d3449f74ae56bda571239f402b95034e6e17b5 (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
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