Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Comment thread
micahwoodard marked this conversation as resolved.
logger.debug("Maximum trial count exceeded.")
return True

Expand Down Expand Up @@ -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.")
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Comment thread
micahwoodard marked this conversation as resolved.
logger.info("Maximum trial count exceeded.")
return True

Expand Down
32 changes: 32 additions & 0 deletions tests/trial_generators/test_block_based_trial_generator.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
12 changes: 7 additions & 5 deletions tests/trial_generators/test_coupled_trial_generator.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
)
Expand Down Expand Up @@ -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
)
Expand Down
10 changes: 10 additions & 0 deletions tests/trial_generators/test_uncoupled_trial_generator.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
import logging
import unittest
from datetime import timedelta

import numpy as np

Expand Down Expand Up @@ -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()