From 98cda96c8d8495876c0b9c9d2c6704f4635238cd Mon Sep 17 00:00:00 2001 From: Micah Woodard Date: Wed, 9 Sep 2026 13:11:18 -0700 Subject: [PATCH 1/6] fixes one off errors and auto water --- .../trial_generators/block_based_trial_generator.py | 5 ++++- .../coupled_trial_generators/coupled_trial_generator.py | 7 ++----- .../trial_generators/uncoupled_trial_gnerator.py | 7 ++----- 3 files changed, 8 insertions(+), 11 deletions(-) diff --git a/src/aind_behavior_dynamic_foraging/task_logic/trial_generators/block_based_trial_generator.py b/src/aind_behavior_dynamic_foraging/task_logic/trial_generators/block_based_trial_generator.py index 1672ce8..0637116 100644 --- a/src/aind_behavior_dynamic_foraging/task_logic/trial_generators/block_based_trial_generator.py +++ b/src/aind_behavior_dynamic_foraging/task_logic/trial_generators/block_based_trial_generator.py @@ -229,7 +229,10 @@ def next(self) -> Trial | None: # determine autowater if is_autowater := self._are_autowater_conditions_met(): - is_auto_reward_right = True if self.block.p_right_reward > self.block.p_left_reward else False + if self.block.p_right_reward != self.block.p_left_reward: + is_auto_reward_right = True if self.block.p_right_reward > self.block.p_left_reward else False + else: + is_auto_reward_right = np.random.choice([True, False]) reward_fraction = self.spec.autowater_parameters.reward_fraction logger.debug("Delivering autowater: is_auto_reward_right = %s" % is_auto_reward_right) diff --git a/src/aind_behavior_dynamic_foraging/task_logic/trial_generators/coupled_trial_generators/coupled_trial_generator.py b/src/aind_behavior_dynamic_foraging/task_logic/trial_generators/coupled_trial_generators/coupled_trial_generator.py index 11d75d3..ab4edaa 100644 --- a/src/aind_behavior_dynamic_foraging/task_logic/trial_generators/coupled_trial_generators/coupled_trial_generator.py +++ b/src/aind_behavior_dynamic_foraging/task_logic/trial_generators/coupled_trial_generators/coupled_trial_generator.py @@ -130,10 +130,7 @@ def _are_end_conditions_met(self) -> bool: frac = end_conditions.ignore_ratio_threshold win = end_conditions.ignore_window_length - if ( - time_elapsed > timedelta(seconds=end_conditions.min_time) - and choice_history[-win:].count(None) >= frac * win - ): + if time_elapsed > timedelta(seconds=end_conditions.min_time) and choice_history[-win:].count(None) > frac * win: logger.debug("Minimum time and ignored trial count exceeded.") return True @@ -141,7 +138,7 @@ def _are_end_conditions_met(self) -> bool: logger.debug("Maximum session time exceeded.") return True - if end_conditions.max_trial < len(choice_history): + if end_conditions.max_trial <= len(choice_history): logger.debug("Maximum trial count exceeded.") return True diff --git a/src/aind_behavior_dynamic_foraging/task_logic/trial_generators/uncoupled_trial_gnerator.py b/src/aind_behavior_dynamic_foraging/task_logic/trial_generators/uncoupled_trial_gnerator.py index a8f1732..df3e6ed 100644 --- a/src/aind_behavior_dynamic_foraging/task_logic/trial_generators/uncoupled_trial_gnerator.py +++ b/src/aind_behavior_dynamic_foraging/task_logic/trial_generators/uncoupled_trial_gnerator.py @@ -184,10 +184,7 @@ def _are_end_conditions_met(self) -> bool: frac = end_conditions.ignore_ratio_threshold win = end_conditions.ignore_window_length - if ( - time_elapsed > timedelta(seconds=end_conditions.min_time) - and choice_history[-win:].count(None) >= frac * win - ): + if time_elapsed > timedelta(seconds=end_conditions.min_time) and choice_history[-win:].count(None) > frac * win: logger.info("Minimum time and ignored trial count exceeded.") return True @@ -195,7 +192,7 @@ def _are_end_conditions_met(self) -> bool: logger.info("Maximum session time exceeded.") return True - if end_conditions.max_trial < len(choice_history): + if end_conditions.max_trial <= len(choice_history): logger.info("Maximum trial count exceeded.") return True From 7d9634209fdf5d54a4927d5feb130462fe9692ba Mon Sep 17 00:00:00 2001 From: Micah Woodard Date: Wed, 9 Sep 2026 13:20:17 -0700 Subject: [PATCH 2/6] min_consecutive_stable_trials fix --- .../coupled_trial_generators/coupled_trial_generator.py | 7 +++++-- .../trial_generators/uncoupled_trial_gnerator.py | 5 ++++- 2 files changed, 9 insertions(+), 3 deletions(-) diff --git a/src/aind_behavior_dynamic_foraging/task_logic/trial_generators/coupled_trial_generators/coupled_trial_generator.py b/src/aind_behavior_dynamic_foraging/task_logic/trial_generators/coupled_trial_generators/coupled_trial_generator.py index ab4edaa..420a9c2 100644 --- a/src/aind_behavior_dynamic_foraging/task_logic/trial_generators/coupled_trial_generators/coupled_trial_generator.py +++ b/src/aind_behavior_dynamic_foraging/task_logic/trial_generators/coupled_trial_generators/coupled_trial_generator.py @@ -130,7 +130,10 @@ def _are_end_conditions_met(self) -> bool: frac = end_conditions.ignore_ratio_threshold win = end_conditions.ignore_window_length - if time_elapsed > timedelta(seconds=end_conditions.min_time) and choice_history[-win:].count(None) > frac * win: + if ( + time_elapsed > timedelta(seconds=end_conditions.min_time) + and choice_history[-win:].count(None) >= frac * win + ): logger.debug("Minimum time and ignored trial count exceeded.") return True @@ -229,7 +232,7 @@ def _is_behavior_stable( run_len += 1 else: run_len = 0 - if run_len >= min_stable: + if run_len > min_stable: logger.info("Behavior stable at trial index %s." % i) return True logger.info("Behavior not stable in block anytime evaluation.") diff --git a/src/aind_behavior_dynamic_foraging/task_logic/trial_generators/uncoupled_trial_gnerator.py b/src/aind_behavior_dynamic_foraging/task_logic/trial_generators/uncoupled_trial_gnerator.py index df3e6ed..5322400 100644 --- a/src/aind_behavior_dynamic_foraging/task_logic/trial_generators/uncoupled_trial_gnerator.py +++ b/src/aind_behavior_dynamic_foraging/task_logic/trial_generators/uncoupled_trial_gnerator.py @@ -184,7 +184,10 @@ def _are_end_conditions_met(self) -> bool: frac = end_conditions.ignore_ratio_threshold win = end_conditions.ignore_window_length - if time_elapsed > timedelta(seconds=end_conditions.min_time) and choice_history[-win:].count(None) > frac * win: + if ( + time_elapsed > timedelta(seconds=end_conditions.min_time) + and choice_history[-win:].count(None) >= frac * win + ): logger.info("Minimum time and ignored trial count exceeded.") return True From 9e754ce0cd5c1450b2a834ad52d049475ad3ed57 Mon Sep 17 00:00:00 2001 From: Micah Woodard Date: Wed, 9 Sep 2026 13:30:57 -0700 Subject: [PATCH 3/6] fixes p_reward_left eval --- .../task_logic/trial_generators/block_based_trial_generator.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/aind_behavior_dynamic_foraging/task_logic/trial_generators/block_based_trial_generator.py b/src/aind_behavior_dynamic_foraging/task_logic/trial_generators/block_based_trial_generator.py index 0637116..4d7df96 100644 --- a/src/aind_behavior_dynamic_foraging/task_logic/trial_generators/block_based_trial_generator.py +++ b/src/aind_behavior_dynamic_foraging/task_logic/trial_generators/block_based_trial_generator.py @@ -232,7 +232,7 @@ def next(self) -> Trial | None: if self.block.p_right_reward != self.block.p_left_reward: is_auto_reward_right = True if self.block.p_right_reward > self.block.p_left_reward else False else: - is_auto_reward_right = np.random.choice([True, False]) + is_auto_reward_right = np.random.choice([True, False]).item() reward_fraction = self.spec.autowater_parameters.reward_fraction logger.debug("Delivering autowater: is_auto_reward_right = %s" % is_auto_reward_right) From 5c2230551f6b95b1257c19dc3d714817a3b3281b Mon Sep 17 00:00:00 2001 From: Micah Woodard Date: Wed, 9 Sep 2026 13:41:22 -0700 Subject: [PATCH 4/6] adds unit tests --- .../block_based_trial_generator.py | 2 +- .../test_block_based_trial_generator.py | 32 +++++++++++++++++++ 2 files changed, 33 insertions(+), 1 deletion(-) diff --git a/src/aind_behavior_dynamic_foraging/task_logic/trial_generators/block_based_trial_generator.py b/src/aind_behavior_dynamic_foraging/task_logic/trial_generators/block_based_trial_generator.py index 4d7df96..5d659bc 100644 --- a/src/aind_behavior_dynamic_foraging/task_logic/trial_generators/block_based_trial_generator.py +++ b/src/aind_behavior_dynamic_foraging/task_logic/trial_generators/block_based_trial_generator.py @@ -232,7 +232,7 @@ def next(self) -> Trial | None: if self.block.p_right_reward != self.block.p_left_reward: is_auto_reward_right = True if self.block.p_right_reward > self.block.p_left_reward else False else: - is_auto_reward_right = np.random.choice([True, False]).item() + is_auto_reward_right = bool(np.random.choice([True, False])) reward_fraction = self.spec.autowater_parameters.reward_fraction logger.debug("Delivering autowater: is_auto_reward_right = %s" % is_auto_reward_right) diff --git a/tests/trial_generators/test_block_based_trial_generator.py b/tests/trial_generators/test_block_based_trial_generator.py index 44cdad9..e171167 100644 --- a/tests/trial_generators/test_block_based_trial_generator.py +++ b/tests/trial_generators/test_block_based_trial_generator.py @@ -55,6 +55,38 @@ def test_next_returns_correct_reward_probs(self): self.assertEqual(trial.p_reward_left, self.generator.block.p_left_reward) self.assertEqual(trial.p_reward_right, self.generator.block.p_right_reward) + def test_next_autowater_equal_probs_choice_right_sets_expected_p_rewards(self): + spec = ConcreteBlockBasedTrialGeneratorSpec( + autowater_parameters=AutoWaterParameters(min_ignored_trials=0, min_unrewarded_trials=0), + bias_intervention_parameters=None, + ) + generator = spec.create_generator() + generator.block = Block(p_left_reward=0.5, p_right_reward=0.5, left_length=10, right_length=10) + + with patch("numpy.random.choice", return_value=np.bool_(True)): + trial = generator.next() + + assert trial is not None + self.assertTrue(trial.is_auto_reward_right) + self.assertEqual(trial.p_reward_left, generator.block.p_left_reward) + self.assertEqual(trial.p_reward_right, 1.0) + + def test_next_autowater_equal_probs_choice_left_sets_expected_p_rewards(self): + spec = ConcreteBlockBasedTrialGeneratorSpec( + autowater_parameters=AutoWaterParameters(min_ignored_trials=0, min_unrewarded_trials=0), + bias_intervention_parameters=None, + ) + generator = spec.create_generator() + generator.block = Block(p_left_reward=0.5, p_right_reward=0.5, left_length=10, right_length=10) + + with patch("numpy.random.choice", return_value=np.bool_(False)): + trial = generator.next() + + assert trial is not None + self.assertFalse(trial.is_auto_reward_right) + self.assertEqual(trial.p_reward_left, 1.0) + self.assertEqual(trial.p_reward_right, generator.block.p_right_reward) + class TestAntiBiasBlockBasedTrialGenerator(unittest.TestCase): def _patch_bias(self, bias_value: float) -> Any: From fa4cd9df3120b1d47f5aaa4450ad3cbc9f02e280 Mon Sep 17 00:00:00 2001 From: Micah Woodard Date: Wed, 9 Sep 2026 13:52:47 -0700 Subject: [PATCH 5/6] fixes and adds tests --- .../trial_generators/test_coupled_trial_generator.py | 12 +++++++----- .../test_uncoupled_trial_generator.py | 10 ++++++++++ 2 files changed, 17 insertions(+), 5 deletions(-) diff --git a/tests/trial_generators/test_coupled_trial_generator.py b/tests/trial_generators/test_coupled_trial_generator.py index 2d259e7..f68f411 100644 --- a/tests/trial_generators/test_coupled_trial_generator.py +++ b/tests/trial_generators/test_coupled_trial_generator.py @@ -81,14 +81,16 @@ def test_behavior_stable_anytime(self): high_reward_is_right = right_prob > left_prob beh_params.behavior_evaluation_mode = "anytime" - # stable run early, then drifts off — should still pass - choices = [high_reward_is_right] * (min_stable + kernel_size - 1) + [not high_reward_is_right] * 10 + # stable run early, then drifts off — should still pass. + # For "anytime" mode, implementation requires run_len > min_stable. + choices = [high_reward_is_right] * (min_stable + kernel_size) + [not high_reward_is_right] * 10 self.assertTrue( self.generator._is_behavior_stable(choices, right_prob, left_prob, beh_params, len(choices), kernel_size) ) - # stable at end: wrong side early, correct side at end - choices = [not high_reward_is_right] * 10 + [high_reward_is_right] * (min_stable + kernel_size - 1) + # stable at end: wrong side early, correct side at end. + # Use one additional trial so stable windows are strictly greater than min_stable. + choices = [not high_reward_is_right] * 10 + [high_reward_is_right] * (min_stable + kernel_size) self.assertTrue( self.generator._is_behavior_stable(choices, right_prob, left_prob, beh_params, len(choices), kernel_size) ) @@ -257,7 +259,7 @@ def test_update_block_does_not_switch_before_right_length(self): #### Test next #### def test_next_returns_none_after_max_trials(self): - self.generator.is_right_choice_history = [True] * (self.spec.trial_generation_end_parameters.max_trial + 1) + self.generator.is_right_choice_history = [True] * (self.spec.trial_generation_end_parameters.max_trial) self.generator.start_time = self.generator.start_time - timedelta( self.spec.trial_generation_end_parameters.min_time ) diff --git a/tests/trial_generators/test_uncoupled_trial_generator.py b/tests/trial_generators/test_uncoupled_trial_generator.py index b590588..e51dbcd 100644 --- a/tests/trial_generators/test_uncoupled_trial_generator.py +++ b/tests/trial_generators/test_uncoupled_trial_generator.py @@ -1,3 +1,4 @@ +from datetime import timedelta import logging import unittest @@ -170,6 +171,15 @@ def test_left_counter_resets_on_left_switch(self): self.generator.update(TrialOutcome(trial=Trial(), is_right_choice=True, is_rewarded=True)) self.assertEqual(self.generator.trials_in_left_block, 0) + def test_next_returns_none_after_max_trials(self): + self.generator.is_right_choice_history = [True] * (self.spec.trial_generation_end_parameters.max_trial) + self.generator.start_time = self.generator.start_time - timedelta( + self.spec.trial_generation_end_parameters.min_time + ) + + trial = self.generator.next() + self.assertIsNone(trial) + if __name__ == "__main__": unittest.main() From fa39b1b3bcde1b8ff7377633b1b2218552282b6b Mon Sep 17 00:00:00 2001 From: Micah Woodard Date: Wed, 9 Sep 2026 14:02:09 -0700 Subject: [PATCH 6/6] lints --- tests/trial_generators/test_uncoupled_trial_generator.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/trial_generators/test_uncoupled_trial_generator.py b/tests/trial_generators/test_uncoupled_trial_generator.py index e51dbcd..bdc43b2 100644 --- a/tests/trial_generators/test_uncoupled_trial_generator.py +++ b/tests/trial_generators/test_uncoupled_trial_generator.py @@ -1,6 +1,6 @@ -from datetime import timedelta import logging import unittest +from datetime import timedelta import numpy as np