Skip to content

Commit 03ee8dc

Browse files
committed
Merge branch 'main' into feat-metadata-mapping
2 parents ce0d701 + a9ed4d9 commit 03ee8dc

1 file changed

Lines changed: 14 additions & 5 deletions

File tree

  • workspace/aind_behavior_dynamic_foraging_curricula/src/aind_behavior_dynamic_foraging_curricula

workspace/aind_behavior_dynamic_foraging_curricula/src/aind_behavior_dynamic_foraging_curricula/metrics.py

Lines changed: 14 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -1,21 +1,30 @@
11
import logging
22
import os
3-
from typing import List, Literal
3+
from typing import Annotated, List, Literal, Optional
44

55
import numpy as np
66
from aind_behavior_curriculum import Metrics
77
from aind_behavior_dynamic_foraging.data_contract import dataset as df_foraging_dataset
8-
from pydantic import Field
8+
from pydantic import BeforeValidator, Field
99

1010
STAGE_NAMES = Literal["stage_1_warmup", "stage_1", "stage_2", "stage_3", "final", "graduated"]
1111

1212
logger = logging.getLogger(__name__)
1313

1414

15+
def coerce_none_to_nan(v: Optional[float]) -> float:
16+
if v is None:
17+
return float("nan")
18+
return v
19+
20+
21+
NoneToNan = Annotated[float, BeforeValidator(coerce_none_to_nan)]
22+
23+
1524
class DynamicForagingMetrics(Metrics):
1625
"""Metrics for dynamic foraging"""
1726

18-
foraging_efficiency_per_session: List[float] = Field(
27+
foraging_efficiency_per_session: List[NoneToNan] = Field(
1928
min_length=1, description="Full history of foraging efficiency per session"
2029
)
2130
unignored_trials_per_session: List[int] = Field(
@@ -87,7 +96,7 @@ def metrics_from_dataset(
8796
)
8897

8998
return DynamicForagingMetrics(
90-
foraging_efficiency_per_session=foraging_efficiency_per_session + [foraging_efficiency],
99+
foraging_efficiency_per_session=foraging_efficiency_per_session + [coerce_none_to_nan(foraging_efficiency)],
91100
unignored_trials_per_session=unignored_trials_per_session + [sum(x is not None for x in is_right_choice)],
92101
total_sessions=total_sessions + 1,
93102
consecutive_sessions_at_current_stage=consecutive_sessions_at_current_stage + 1,
@@ -97,7 +106,7 @@ def metrics_from_dataset(
97106

98107
def compute_foraging_efficiency(
99108
is_baiting: bool, is_rewarded: list[bool], p_right_reward: list[float], p_left_reward: list[float]
100-
) -> float:
109+
) -> Optional[float]:
101110
"""
102111
Compute foraging efficiency for a two-arm bandit task.
103112

0 commit comments

Comments
 (0)