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..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 @@ -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 = 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/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..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 @@ -141,7 +141,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 @@ -232,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 a8f1732..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 @@ -195,7 +195,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 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: 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..bdc43b2 100644 --- a/tests/trial_generators/test_uncoupled_trial_generator.py +++ b/tests/trial_generators/test_uncoupled_trial_generator.py @@ -1,5 +1,6 @@ import logging import unittest +from datetime import timedelta import numpy as np @@ -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()