11import logging
22import os
3- from typing import List , Literal
3+ from typing import Annotated , List , Literal , Optional
44
55import numpy as np
66from aind_behavior_curriculum import Metrics
77from aind_behavior_dynamic_foraging .data_contract import dataset as df_foraging_dataset
8- from pydantic import Field
8+ from pydantic import BeforeValidator , Field
99
1010STAGE_NAMES = Literal ["stage_1_warmup" , "stage_1" , "stage_2" , "stage_3" , "final" , "graduated" ]
1111
1212logger = 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+
1524class 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
98107def 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