From c44d221e39962859e6d6e5481e046115377a4be0 Mon Sep 17 00:00:00 2001 From: Micah Woodard Date: Wed, 9 Sep 2026 12:54:11 -0700 Subject: [PATCH 1/3] fixes v1 vs v2 curriculum differences --- .github/workflows/dynamic-foraging-cicd.yml | 15 +++- schema/coupled_baiting.json | 54 ++++++++---- schema/uncoupled.json | 27 ++++-- schema/uncoupled_baiting.json | 27 ++++-- .../coupled_baiting/stages.py | 18 ++-- .../uncoupled/curriculum.py | 6 +- .../uncoupled/stages.py | 9 +- .../uncoupled_baiting/curriculum.py | 6 +- .../uncoupled_baiting/stages.py | 9 +- .../tests/test_coupled_baiting.py | 83 ++++++++++--------- .../tests/test_metrics.py | 4 +- .../tests/test_uncoupled.py | 81 ++++++++++-------- .../tests/test_uncoupled_baiting.py | 81 ++++++++++-------- 13 files changed, 258 insertions(+), 162 deletions(-) diff --git a/.github/workflows/dynamic-foraging-cicd.yml b/.github/workflows/dynamic-foraging-cicd.yml index 242ef0ac..270998c2 100644 --- a/.github/workflows/dynamic-foraging-cicd.yml +++ b/.github/workflows/dynamic-foraging-cicd.yml @@ -54,8 +54,19 @@ jobs: working-directory: ./.bonsai run: ./setup.ps1 - - name: Run python unit tests - run: uv run python -m unittest + - name: Run root python unit tests + run: | + uv run python -m unittest discover tests + uv run --directory .\workspace\aind_behavior_dynamic_foraging_curricula\ python -m unittest discover tests + + - name: Run workspace python unit tests + shell: bash + run: | + for workspace in workspace/*; do + if [ -d "$workspace/tests" ]; then + uv run --directory "$workspace" python -m unittest discover tests + fi + done - name: Regenerate all schemas run: | diff --git a/schema/coupled_baiting.json b/schema/coupled_baiting.json index 56b66adc..a890441f 100644 --- a/schema/coupled_baiting.json +++ b/schema/coupled_baiting.json @@ -132,11 +132,14 @@ "rate": 0.2 }, "truncation_parameters": { - "truncation_mode": "exclude", - "min": 10.0, + "truncation_mode": "clamp", + "min": 0.0, "max": 20.0 }, - "scaling_parameters": null + "scaling_parameters": { + "scale": 1.0, + "offset": 10.0 + } }, "autowater_parameters": { "min_ignored_trials": 3, @@ -241,11 +244,14 @@ "rate": 0.2 }, "truncation_parameters": { - "truncation_mode": "exclude", - "min": 10.0, + "truncation_mode": "clamp", + "min": 0.0, "max": 20.0 }, - "scaling_parameters": null + "scaling_parameters": { + "scale": 1.0, + "offset": 10.0 + } }, "autowater_parameters": { "min_ignored_trials": 5, @@ -348,11 +354,14 @@ "rate": 0.1 }, "truncation_parameters": { - "truncation_mode": "exclude", - "min": 10.0, + "truncation_mode": "clamp", + "min": 0.0, "max": 40.0 }, - "scaling_parameters": null + "scaling_parameters": { + "scale": 1.0, + "offset": 10.0 + } }, "autowater_parameters": { "min_ignored_trials": 7, @@ -455,11 +464,14 @@ "rate": 0.05 }, "truncation_parameters": { - "truncation_mode": "exclude", - "min": 20.0, + "truncation_mode": "clamp", + "min": 0.0, "max": 60.0 }, - "scaling_parameters": null + "scaling_parameters": { + "scale": 1.0, + "offset": 20.0 + } }, "autowater_parameters": { "min_ignored_trials": 10, @@ -562,11 +574,14 @@ "rate": 0.05 }, "truncation_parameters": { - "truncation_mode": "exclude", - "min": 20.0, + "truncation_mode": "clamp", + "min": 0.0, "max": 60.0 }, - "scaling_parameters": null + "scaling_parameters": { + "scale": 1.0, + "offset": 20.0 + } }, "autowater_parameters": null, "bias_intervention_parameters": { @@ -673,11 +688,14 @@ "rate": 0.05 }, "truncation_parameters": { - "truncation_mode": "exclude", - "min": 20.0, + "truncation_mode": "clamp", + "min": 0.0, "max": 60.0 }, - "scaling_parameters": null + "scaling_parameters": { + "scale": 1.0, + "offset": 20.0 + } }, "autowater_parameters": null, "bias_intervention_parameters": { diff --git a/schema/uncoupled.json b/schema/uncoupled.json index 3e2a101b..f649fc24 100644 --- a/schema/uncoupled.json +++ b/schema/uncoupled.json @@ -132,11 +132,14 @@ "rate": 0.1 }, "truncation_parameters": { - "truncation_mode": "exclude", - "min": 10.0, + "truncation_mode": "clamp", + "min": 0.0, "max": 30.0 }, - "scaling_parameters": null + "scaling_parameters": { + "scale": 1.0, + "offset": 10.0 + } }, "autowater_parameters": { "min_ignored_trials": 3, @@ -241,11 +244,14 @@ "rate": 0.1 }, "truncation_parameters": { - "truncation_mode": "exclude", - "min": 10.0, + "truncation_mode": "clamp", + "min": 0.0, "max": 30.0 }, - "scaling_parameters": null + "scaling_parameters": { + "scale": 1.0, + "offset": 10.0 + } }, "autowater_parameters": { "min_ignored_trials": 5, @@ -348,11 +354,14 @@ "rate": 0.1 }, "truncation_parameters": { - "truncation_mode": "exclude", - "min": 10.0, + "truncation_mode": "clamp", + "min": 0.0, "max": 40.0 }, - "scaling_parameters": null + "scaling_parameters": { + "scale": 1.0, + "offset": 10.0 + } }, "autowater_parameters": { "min_ignored_trials": 7, diff --git a/schema/uncoupled_baiting.json b/schema/uncoupled_baiting.json index e792102c..fcd0ac7f 100644 --- a/schema/uncoupled_baiting.json +++ b/schema/uncoupled_baiting.json @@ -132,11 +132,14 @@ "rate": 0.1 }, "truncation_parameters": { - "truncation_mode": "exclude", - "min": 10.0, + "truncation_mode": "clamp", + "min": 0.0, "max": 30.0 }, - "scaling_parameters": null + "scaling_parameters": { + "scale": 1.0, + "offset": 10.0 + } }, "autowater_parameters": { "min_ignored_trials": 3, @@ -241,11 +244,14 @@ "rate": 0.1 }, "truncation_parameters": { - "truncation_mode": "exclude", - "min": 10.0, + "truncation_mode": "clamp", + "min": 0.0, "max": 30.0 }, - "scaling_parameters": null + "scaling_parameters": { + "scale": 1.0, + "offset": 10.0 + } }, "autowater_parameters": { "min_ignored_trials": 5, @@ -348,11 +354,14 @@ "rate": 0.1 }, "truncation_parameters": { - "truncation_mode": "exclude", - "min": 10.0, + "truncation_mode": "clamp", + "min": 0.0, "max": 40.0 }, - "scaling_parameters": null + "scaling_parameters": { + "scale": 1.0, + "offset": 10.0 + } }, "autowater_parameters": { "min_ignored_trials": 7, diff --git a/workspace/aind_behavior_dynamic_foraging_curricula/src/aind_behavior_dynamic_foraging_curricula/coupled_baiting/stages.py b/workspace/aind_behavior_dynamic_foraging_curricula/src/aind_behavior_dynamic_foraging_curricula/coupled_baiting/stages.py index 2b3b4094..a1525fdf 100644 --- a/workspace/aind_behavior_dynamic_foraging_curricula/src/aind_behavior_dynamic_foraging_curricula/coupled_baiting/stages.py +++ b/workspace/aind_behavior_dynamic_foraging_curricula/src/aind_behavior_dynamic_foraging_curricula/coupled_baiting/stages.py @@ -93,7 +93,8 @@ def make_s_stage_1_warmup(): ), block_length=ExponentialDistribution( distribution_parameters=ExponentialDistributionParameters(rate=0.2), - truncation_parameters=TruncationParameters(min=10, max=20), + truncation_parameters=TruncationParameters(truncation_mode="clamp", min=0, max=20), + scaling_parameters=ScalingParameters(offset=10), ), inter_trial_interval_duration=ExponentialDistribution( distribution_parameters=ExponentialDistributionParameters(rate=1.0 / 3), @@ -151,7 +152,8 @@ def make_s_stage_1(): ), block_length=ExponentialDistribution( distribution_parameters=ExponentialDistributionParameters(rate=0.2), - truncation_parameters=TruncationParameters(min=10, max=20), + truncation_parameters=TruncationParameters(truncation_mode="clamp", min=0, max=20), + scaling_parameters=ScalingParameters(offset=10), ), inter_trial_interval_duration=ExponentialDistribution( distribution_parameters=ExponentialDistributionParameters(rate=1.0 / 3), @@ -207,7 +209,8 @@ def make_s_stage_2(): ), block_length=ExponentialDistribution( distribution_parameters=ExponentialDistributionParameters(rate=0.1), - truncation_parameters=TruncationParameters(min=10, max=40), + truncation_parameters=TruncationParameters(truncation_mode="clamp", min=0, max=40), + scaling_parameters=ScalingParameters(offset=10), ), inter_trial_interval_duration=ExponentialDistribution( distribution_parameters=ExponentialDistributionParameters(rate=1.0 / 3), @@ -263,7 +266,8 @@ def make_s_stage_3(): ), block_length=ExponentialDistribution( distribution_parameters=ExponentialDistributionParameters(rate=0.05), - truncation_parameters=TruncationParameters(min=20, max=60), + truncation_parameters=TruncationParameters(truncation_mode="clamp", min=0, max=60), + scaling_parameters=ScalingParameters(offset=20), ), inter_trial_interval_duration=ExponentialDistribution( distribution_parameters=ExponentialDistributionParameters(rate=1.0 / 3), @@ -315,7 +319,8 @@ def make_s_stage_final(): ), block_length=ExponentialDistribution( distribution_parameters=ExponentialDistributionParameters(rate=0.05), - truncation_parameters=TruncationParameters(min=20, max=60), + truncation_parameters=TruncationParameters(truncation_mode="clamp", min=0, max=60), + scaling_parameters=ScalingParameters(offset=20), ), inter_trial_interval_duration=ExponentialDistribution( distribution_parameters=ExponentialDistributionParameters(rate=1.0 / 3), @@ -365,7 +370,8 @@ def make_s_stage_graduated(): ), block_length=ExponentialDistribution( distribution_parameters=ExponentialDistributionParameters(rate=0.05), - truncation_parameters=TruncationParameters(min=20, max=60), + truncation_parameters=TruncationParameters(truncation_mode="clamp", min=0, max=60), + scaling_parameters=ScalingParameters(offset=20), ), inter_trial_interval_duration=ExponentialDistribution( distribution_parameters=ExponentialDistributionParameters(rate=1.0 / 3), diff --git a/workspace/aind_behavior_dynamic_foraging_curricula/src/aind_behavior_dynamic_foraging_curricula/uncoupled/curriculum.py b/workspace/aind_behavior_dynamic_foraging_curricula/src/aind_behavior_dynamic_foraging_curricula/uncoupled/curriculum.py index 2567329d..2ac4bad5 100644 --- a/workspace/aind_behavior_dynamic_foraging_curricula/src/aind_behavior_dynamic_foraging_curricula/uncoupled/curriculum.py +++ b/workspace/aind_behavior_dynamic_foraging_curricula/src/aind_behavior_dynamic_foraging_curricula/uncoupled/curriculum.py @@ -49,7 +49,11 @@ def st_stage_1_to_stage_2(metrics: DynamicForagingMetrics) -> bool: # stage 2 @StageTransition def st_stage_2_to_stage_3(metrics: DynamicForagingMetrics) -> bool: - return bool(metrics.foraging_efficiency_per_session[-1] >= 0.65 and metrics.unignored_trials_per_session[-1] >= 300) + return bool( + metrics.foraging_efficiency_per_session[-1] >= 0.65 + and metrics.unignored_trials_per_session[-1] >= 300 + and metrics.consecutive_sessions_at_current_stage >= 2 + ) @StageTransition diff --git a/workspace/aind_behavior_dynamic_foraging_curricula/src/aind_behavior_dynamic_foraging_curricula/uncoupled/stages.py b/workspace/aind_behavior_dynamic_foraging_curricula/src/aind_behavior_dynamic_foraging_curricula/uncoupled/stages.py index 5ed6da45..98fe9b99 100644 --- a/workspace/aind_behavior_dynamic_foraging_curricula/src/aind_behavior_dynamic_foraging_curricula/uncoupled/stages.py +++ b/workspace/aind_behavior_dynamic_foraging_curricula/src/aind_behavior_dynamic_foraging_curricula/uncoupled/stages.py @@ -111,7 +111,8 @@ def make_s_stage_1_warmup(): ), block_length=ExponentialDistribution( distribution_parameters=ExponentialDistributionParameters(rate=0.1), - truncation_parameters=TruncationParameters(min=10, max=30), + truncation_parameters=TruncationParameters(truncation_mode="clamp", min=0, max=30), + scaling_parameters=ScalingParameters(offset=10), ), inter_trial_interval_duration=ExponentialDistribution( distribution_parameters=ExponentialDistributionParameters(rate=1.0 / 3), @@ -177,7 +178,8 @@ def make_s_stage_1(): ), block_length=ExponentialDistribution( distribution_parameters=ExponentialDistributionParameters(rate=0.1), - truncation_parameters=TruncationParameters(min=10, max=30), + truncation_parameters=TruncationParameters(truncation_mode="clamp", min=0, max=30), + scaling_parameters=ScalingParameters(offset=10), ), inter_trial_interval_duration=ExponentialDistribution( distribution_parameters=ExponentialDistributionParameters(rate=1.0 / 3), @@ -233,7 +235,8 @@ def make_s_stage_2(): ), block_length=ExponentialDistribution( distribution_parameters=ExponentialDistributionParameters(rate=0.1), - truncation_parameters=TruncationParameters(min=10, max=40), + truncation_parameters=TruncationParameters(truncation_mode="clamp", min=0, max=40), + scaling_parameters=ScalingParameters(offset=10), ), inter_trial_interval_duration=ExponentialDistribution( distribution_parameters=ExponentialDistributionParameters(rate=1.0 / 3), diff --git a/workspace/aind_behavior_dynamic_foraging_curricula/src/aind_behavior_dynamic_foraging_curricula/uncoupled_baiting/curriculum.py b/workspace/aind_behavior_dynamic_foraging_curricula/src/aind_behavior_dynamic_foraging_curricula/uncoupled_baiting/curriculum.py index 90a4c5f9..c5ecccf1 100644 --- a/workspace/aind_behavior_dynamic_foraging_curricula/src/aind_behavior_dynamic_foraging_curricula/uncoupled_baiting/curriculum.py +++ b/workspace/aind_behavior_dynamic_foraging_curricula/src/aind_behavior_dynamic_foraging_curricula/uncoupled_baiting/curriculum.py @@ -49,7 +49,11 @@ def st_stage_1_to_stage_2(metrics: DynamicForagingMetrics) -> bool: # stage 2 @StageTransition def st_stage_2_to_stage_3(metrics: DynamicForagingMetrics) -> bool: - return bool(metrics.foraging_efficiency_per_session[-1] >= 0.65 and metrics.unignored_trials_per_session[-1] >= 300) + return bool( + metrics.foraging_efficiency_per_session[-1] >= 0.65 + and metrics.unignored_trials_per_session[-1] >= 300 + and metrics.consecutive_sessions_at_current_stage >= 2 + ) @StageTransition diff --git a/workspace/aind_behavior_dynamic_foraging_curricula/src/aind_behavior_dynamic_foraging_curricula/uncoupled_baiting/stages.py b/workspace/aind_behavior_dynamic_foraging_curricula/src/aind_behavior_dynamic_foraging_curricula/uncoupled_baiting/stages.py index 3702879d..9d832b9b 100644 --- a/workspace/aind_behavior_dynamic_foraging_curricula/src/aind_behavior_dynamic_foraging_curricula/uncoupled_baiting/stages.py +++ b/workspace/aind_behavior_dynamic_foraging_curricula/src/aind_behavior_dynamic_foraging_curricula/uncoupled_baiting/stages.py @@ -108,7 +108,8 @@ def make_s_stage_1_warmup(): ), block_length=ExponentialDistribution( distribution_parameters=ExponentialDistributionParameters(rate=0.1), - truncation_parameters=TruncationParameters(min=10, max=30), + truncation_parameters=TruncationParameters(truncation_mode="clamp", min=0, max=30), + scaling_parameters=ScalingParameters(offset=10), ), inter_trial_interval_duration=ExponentialDistribution( distribution_parameters=ExponentialDistributionParameters(rate=1.0 / 3), @@ -166,7 +167,8 @@ def make_s_stage_1(): ), block_length=ExponentialDistribution( distribution_parameters=ExponentialDistributionParameters(rate=0.1), - truncation_parameters=TruncationParameters(min=10, max=30), + truncation_parameters=TruncationParameters(truncation_mode="clamp", min=0, max=30), + scaling_parameters=ScalingParameters(offset=10), ), inter_trial_interval_duration=ExponentialDistribution( distribution_parameters=ExponentialDistributionParameters(rate=1.0 / 3), @@ -222,7 +224,8 @@ def make_s_stage_2(): ), block_length=ExponentialDistribution( distribution_parameters=ExponentialDistributionParameters(rate=0.1), - truncation_parameters=TruncationParameters(min=10, max=40), + truncation_parameters=TruncationParameters(truncation_mode="clamp", min=0, max=40), + scaling_parameters=ScalingParameters(offset=10), ), inter_trial_interval_duration=ExponentialDistribution( distribution_parameters=ExponentialDistributionParameters(rate=1.0 / 3), diff --git a/workspace/aind_behavior_dynamic_foraging_curricula/tests/test_coupled_baiting.py b/workspace/aind_behavior_dynamic_foraging_curricula/tests/test_coupled_baiting.py index 2a809080..888a406a 100644 --- a/workspace/aind_behavior_dynamic_foraging_curricula/tests/test_coupled_baiting.py +++ b/workspace/aind_behavior_dynamic_foraging_curricula/tests/test_coupled_baiting.py @@ -17,7 +17,7 @@ def make_metrics( unignored_trials_per_session: list[int] = None, total_sessions: int = 1, consecutive_sessions_at_current_stage: int = 1, - stage_name: str = "stage_1_warmup", + stage_name: str = "STAGE_1_WARMUP", ) -> DynamicForagingMetrics: return DynamicForagingMetrics( foraging_efficiency_per_session=foraging_efficiency_per_session or [0.0], @@ -32,16 +32,16 @@ class TestCurriculumStructure(unittest.TestCase): def test_all_stages_in_curriculum(self): stages = CURRICULUM.see_stages() stage_names = [s.name for s in stages] - self.assertIn("stage_1_warmup", stage_names) - self.assertIn("stage_1", stage_names) - self.assertIn("stage_2", stage_names) - self.assertIn("stage_3", stage_names) - self.assertIn("final", stage_names) - self.assertIn("graduated", stage_names) + self.assertIn("STAGE_1_WARMUP", stage_names) + self.assertIn("STAGE_1", stage_names) + self.assertIn("STAGE_2", stage_names) + self.assertIn("STAGE_3", stage_names) + self.assertIn("STAGE_FINAL", stage_names) + self.assertIn("GRADUATED", stage_names) def test_enrollment_starts_at_stage_1_warmup(self): trainer_state = TRAINER.create_enrollment() - self.assertEqual(trainer_state.stage.name, "stage_1_warmup") + self.assertEqual(trainer_state.stage.name, "STAGE_1_WARMUP") class TestWarmupTransitions(unittest.TestCase): @@ -50,20 +50,20 @@ def setUp(self): def test_warmup_to_stage_2_on_good_performance(self): metrics = make_metrics( - unignored_trials_per_session=[250], foraging_efficiency_per_session=[0.65], stage_name="stage_1_warmup" + unignored_trials_per_session=[250], foraging_efficiency_per_session=[0.65], stage_name="STAGE_1_WARMUP" ) updated = TRAINER.evaluate(self.trainer_state, metrics) - self.assertEqual(updated.stage.name, "stage_2") + self.assertEqual(updated.stage.name, "STAGE_2") def test_warmup_to_stage_1_after_first_session(self): metrics = make_metrics( unignored_trials_per_session=[100], foraging_efficiency_per_session=[0.4], consecutive_sessions_at_current_stage=1, - stage_name="stage_1_warmup", + stage_name="STAGE_1_WARMUP", ) updated = TRAINER.evaluate(self.trainer_state, metrics) - self.assertEqual(updated.stage.name, "stage_1") + self.assertEqual(updated.stage.name, "STAGE_1") class TestStage1Transitions(unittest.TestCase): @@ -72,17 +72,17 @@ def setUp(self): def test_stage_1_to_stage_2_on_good_performance(self): metrics = make_metrics( - unignored_trials_per_session=[200], foraging_efficiency_per_session=[0.6], stage_name="stage_1" + unignored_trials_per_session=[200], foraging_efficiency_per_session=[0.6], stage_name="STAGE_1" ) updated = TRAINER.evaluate(self.trainer_state, metrics) - self.assertEqual(updated.stage.name, "stage_2") + self.assertEqual(updated.stage.name, "STAGE_2") def test_stage_1_no_transition_on_poor_performance(self): metrics = make_metrics( - unignored_trials_per_session=[100], foraging_efficiency_per_session=[0.4], stage_name="stage_1" + unignored_trials_per_session=[100], foraging_efficiency_per_session=[0.4], stage_name="STAGE_1" ) updated = TRAINER.evaluate(self.trainer_state, metrics) - self.assertEqual(updated.stage.name, "stage_1") + self.assertEqual(updated.stage.name, "STAGE_1") class TestStage2Transitions(unittest.TestCase): @@ -91,31 +91,34 @@ def setUp(self): def test_stage_2_to_stage_3_on_good_performance(self): metrics = make_metrics( - unignored_trials_per_session=[300], foraging_efficiency_per_session=[0.65], stage_name="stage_2" + unignored_trials_per_session=[300], + foraging_efficiency_per_session=[0.65], + consecutive_sessions_at_current_stage=3, + stage_name="STAGE_2", ) updated = TRAINER.evaluate(self.trainer_state, metrics) - self.assertEqual(updated.stage.name, "stage_3") + self.assertEqual(updated.stage.name, "STAGE_3") def test_stage_2_rollback_to_stage_1_on_poor_trials(self): metrics = make_metrics( - unignored_trials_per_session=[150], foraging_efficiency_per_session=[0.6], stage_name="stage_2" + unignored_trials_per_session=[150], foraging_efficiency_per_session=[0.6], stage_name="STAGE_2" ) updated = TRAINER.evaluate(self.trainer_state, metrics) - self.assertEqual(updated.stage.name, "stage_1") + self.assertEqual(updated.stage.name, "STAGE_1") def test_stage_2_rollback_to_stage_1_on_poor_efficiency(self): metrics = make_metrics( - unignored_trials_per_session=[199], foraging_efficiency_per_session=[0.5], stage_name="stage_2" + unignored_trials_per_session=[199], foraging_efficiency_per_session=[0.5], stage_name="STAGE_2" ) updated = TRAINER.evaluate(self.trainer_state, metrics) - self.assertEqual(updated.stage.name, "stage_1") + self.assertEqual(updated.stage.name, "STAGE_1") def test_stage_2_no_transition_on_middle_performance(self): metrics = make_metrics( - unignored_trials_per_session=[250], foraging_efficiency_per_session=[0.6], stage_name="stage_2" + unignored_trials_per_session=[250], foraging_efficiency_per_session=[0.6], stage_name="STAGE_2" ) updated = TRAINER.evaluate(self.trainer_state, metrics) - self.assertEqual(updated.stage.name, "stage_2") + self.assertEqual(updated.stage.name, "STAGE_2") class TestStage3Transitions(unittest.TestCase): @@ -124,31 +127,31 @@ def setUp(self): def test_stage_3_to_final_on_good_performance(self): metrics = make_metrics( - unignored_trials_per_session=[400], foraging_efficiency_per_session=[0.7], stage_name="stage_3" + unignored_trials_per_session=[400], foraging_efficiency_per_session=[0.7], stage_name="STAGE_3" ) updated = TRAINER.evaluate(self.trainer_state, metrics) - self.assertEqual(updated.stage.name, "final") + self.assertEqual(updated.stage.name, "STAGE_FINAL") def test_stage_3_rollback_to_stage_2_on_poor_trials(self): metrics = make_metrics( - unignored_trials_per_session=[250], foraging_efficiency_per_session=[0.7], stage_name="stage_3" + unignored_trials_per_session=[250], foraging_efficiency_per_session=[0.7], stage_name="STAGE_3" ) updated = TRAINER.evaluate(self.trainer_state, metrics) - self.assertEqual(updated.stage.name, "stage_2") + self.assertEqual(updated.stage.name, "STAGE_2") def test_stage_3_rollback_to_stage_2_on_poor_efficiency(self): metrics = make_metrics( - unignored_trials_per_session=[299], foraging_efficiency_per_session=[0.6], stage_name="stage_3" + unignored_trials_per_session=[299], foraging_efficiency_per_session=[0.6], stage_name="STAGE_3" ) updated = TRAINER.evaluate(self.trainer_state, metrics) - self.assertEqual(updated.stage.name, "stage_2") + self.assertEqual(updated.stage.name, "STAGE_2") def test_stage_3_no_transition_on_middle_performance(self): metrics = make_metrics( - unignored_trials_per_session=[350], foraging_efficiency_per_session=[0.67], stage_name="stage_3" + unignored_trials_per_session=[350], foraging_efficiency_per_session=[0.67], stage_name="STAGE_3" ) updated = TRAINER.evaluate(self.trainer_state, metrics) - self.assertEqual(updated.stage.name, "stage_3") + self.assertEqual(updated.stage.name, "STAGE_3") class TestFinalTransitions(unittest.TestCase): @@ -161,10 +164,10 @@ def test_final_to_graduated_on_excellent_performance(self): foraging_efficiency_per_session=[0.70] * 5, total_sessions=10, consecutive_sessions_at_current_stage=5, - stage_name="final", + stage_name="STAGE_FINAL", ) updated = TRAINER.evaluate(self.trainer_state, metrics) - self.assertEqual(updated.stage.name, "graduated") + self.assertEqual(updated.stage.name, "GRADUATED") def test_final_rollback_to_stage_3_on_poor_performance(self): metrics = make_metrics( @@ -172,10 +175,10 @@ def test_final_rollback_to_stage_3_on_poor_performance(self): foraging_efficiency_per_session=[0.55] * 5, total_sessions=10, consecutive_sessions_at_current_stage=5, - stage_name="final", + stage_name="STAGE_FINAL", ) updated = TRAINER.evaluate(self.trainer_state, metrics) - self.assertEqual(updated.stage.name, "stage_3") + self.assertEqual(updated.stage.name, "STAGE_3") def test_final_no_graduation_without_enough_sessions(self): metrics = make_metrics( @@ -183,10 +186,10 @@ def test_final_no_graduation_without_enough_sessions(self): foraging_efficiency_per_session=[0.70] * 5, total_sessions=5, consecutive_sessions_at_current_stage=3, - stage_name="final", + stage_name="STAGE_FINAL", ) updated = TRAINER.evaluate(self.trainer_state, metrics) - self.assertNotEqual(updated.stage.name, "graduated") + self.assertNotEqual(updated.stage.name, "GRADUATED") def test_graduated_is_absorbing(self): trainer_state = TRAINER.create_trainer_state(stage=make_s_stage_graduated()) @@ -195,10 +198,10 @@ def test_graduated_is_absorbing(self): foraging_efficiency_per_session=[0.9] * 5, total_sessions=20, consecutive_sessions_at_current_stage=10, - stage_name="final", + stage_name="GRADUATED", ) updated = TRAINER.evaluate(trainer_state, metrics) - self.assertEqual(updated.stage.name, "graduated") + self.assertEqual(updated.stage.name, "GRADUATED") if __name__ == "__main__": diff --git a/workspace/aind_behavior_dynamic_foraging_curricula/tests/test_metrics.py b/workspace/aind_behavior_dynamic_foraging_curricula/tests/test_metrics.py index c9861566..589bbb25 100644 --- a/workspace/aind_behavior_dynamic_foraging_curricula/tests/test_metrics.py +++ b/workspace/aind_behavior_dynamic_foraging_curricula/tests/test_metrics.py @@ -27,7 +27,7 @@ def _make_trial( def _patch_dataset( - trials: list[dict], is_baiting: bool = True, prev_metrics: Optional[dict] = None, stage_name: str = "stage_1" + trials: list[dict], is_baiting: bool = True, prev_metrics: Optional[dict] = None, stage_name: str = "STAGE_1" ): """Patch df_foraging_dataset with a mock matching the access pattern in metrics_from_dataset.""" @@ -100,7 +100,7 @@ def test_previous_metrics_accumulate(self): "unignored_trials_per_session": [10], "total_sessions": 1, "consecutive_sessions_at_current_stage": 1, - "stage_name": "stage_1_warmup", + "stage_name": "STAGE_1_WARMUP", } with _patch_dataset(trials, prev_metrics=metrics): result = metrics_from_dataset(self.tmp_path) diff --git a/workspace/aind_behavior_dynamic_foraging_curricula/tests/test_uncoupled.py b/workspace/aind_behavior_dynamic_foraging_curricula/tests/test_uncoupled.py index ea4024c3..4be63cc7 100644 --- a/workspace/aind_behavior_dynamic_foraging_curricula/tests/test_uncoupled.py +++ b/workspace/aind_behavior_dynamic_foraging_curricula/tests/test_uncoupled.py @@ -17,7 +17,7 @@ def make_metrics( unignored_trials_per_session: list[int] = None, total_sessions: int = 1, consecutive_sessions_at_current_stage: int = 1, - stage_name: str = "stage_1_warmup", + stage_name: str = "STAGE_1_WARMUP", ) -> DynamicForagingMetrics: return DynamicForagingMetrics( foraging_efficiency_per_session=foraging_efficiency_per_session or [0.0], @@ -32,16 +32,16 @@ class TestCurriculumStructure(unittest.TestCase): def test_all_stages_in_curriculum(self): stages = CURRICULUM.see_stages() stage_names = [s.name for s in stages] - self.assertIn("stage_1_warmup", stage_names) - self.assertIn("stage_1", stage_names) - self.assertIn("stage_2", stage_names) - self.assertIn("stage_3", stage_names) - self.assertIn("final", stage_names) - self.assertIn("graduated", stage_names) + self.assertIn("STAGE_1_WARMUP", stage_names) + self.assertIn("STAGE_1", stage_names) + self.assertIn("STAGE_2", stage_names) + self.assertIn("STAGE_3", stage_names) + self.assertIn("STAGE_FINAL", stage_names) + self.assertIn("GRADUATED", stage_names) def test_enrollment_starts_at_stage_1_warmup(self): trainer_state = TRAINER.create_enrollment() - self.assertEqual(trainer_state.stage.name, "stage_1_warmup") + self.assertEqual(trainer_state.stage.name, "STAGE_1_WARMUP") class TestWarmupTransitions(unittest.TestCase): @@ -50,20 +50,20 @@ def setUp(self): def test_warmup_to_stage_2_on_good_performance(self): metrics = make_metrics( - unignored_trials_per_session=[250], foraging_efficiency_per_session=[0.65], stage_name="stage_1_warmup" + unignored_trials_per_session=[250], foraging_efficiency_per_session=[0.65], stage_name="STAGE_1_WARMUP" ) updated = TRAINER.evaluate(self.trainer_state, metrics) - self.assertEqual(updated.stage.name, "stage_2") + self.assertEqual(updated.stage.name, "STAGE_2") def test_warmup_to_stage_1_after_first_session(self): metrics = make_metrics( unignored_trials_per_session=[100], foraging_efficiency_per_session=[0.4], consecutive_sessions_at_current_stage=1, - stage_name="stage_1_warmup", + stage_name="STAGE_1_WARMUP", ) updated = TRAINER.evaluate(self.trainer_state, metrics) - self.assertEqual(updated.stage.name, "stage_1") + self.assertEqual(updated.stage.name, "STAGE_1") class TestStage1Transitions(unittest.TestCase): @@ -72,17 +72,17 @@ def setUp(self): def test_stage_1_to_stage_2_on_good_performance(self): metrics = make_metrics( - unignored_trials_per_session=[200], foraging_efficiency_per_session=[0.6], stage_name="stage_1" + unignored_trials_per_session=[200], foraging_efficiency_per_session=[0.6], stage_name="STAGE_1" ) updated = TRAINER.evaluate(self.trainer_state, metrics) - self.assertEqual(updated.stage.name, "stage_2") + self.assertEqual(updated.stage.name, "STAGE_2") def test_stage_1_no_transition_on_poor_performance(self): metrics = make_metrics( - unignored_trials_per_session=[100], foraging_efficiency_per_session=[0.4], stage_name="stage_1" + unignored_trials_per_session=[100], foraging_efficiency_per_session=[0.4], stage_name="STAGE_1" ) updated = TRAINER.evaluate(self.trainer_state, metrics) - self.assertEqual(updated.stage.name, "stage_1") + self.assertEqual(updated.stage.name, "STAGE_1") class TestStage2Transitions(unittest.TestCase): @@ -91,31 +91,44 @@ def setUp(self): def test_stage_2_to_stage_3_on_good_performance(self): metrics = make_metrics( - unignored_trials_per_session=[300], foraging_efficiency_per_session=[0.65], stage_name="stage_2" + unignored_trials_per_session=[300], + foraging_efficiency_per_session=[0.65], + consecutive_sessions_at_current_stage=3, + stage_name="STAGE_2", + ) + updated = TRAINER.evaluate(self.trainer_state, metrics) + self.assertEqual(updated.stage.name, "STAGE_3") + + def test_stage_2_requires_two_sessions_before_stage_3(self): + metrics = make_metrics( + unignored_trials_per_session=[300], + foraging_efficiency_per_session=[0.65], + consecutive_sessions_at_current_stage=1, + stage_name="STAGE_2", ) updated = TRAINER.evaluate(self.trainer_state, metrics) - self.assertEqual(updated.stage.name, "stage_3") + self.assertEqual(updated.stage.name, "STAGE_2") def test_stage_2_rollback_to_stage_1_on_poor_trials(self): metrics = make_metrics( - unignored_trials_per_session=[150], foraging_efficiency_per_session=[0.6], stage_name="stage_2" + unignored_trials_per_session=[150], foraging_efficiency_per_session=[0.6], stage_name="STAGE_2" ) updated = TRAINER.evaluate(self.trainer_state, metrics) - self.assertEqual(updated.stage.name, "stage_1") + self.assertEqual(updated.stage.name, "STAGE_1") def test_stage_2_rollback_to_stage_1_on_poor_efficiency(self): metrics = make_metrics( - unignored_trials_per_session=[199], foraging_efficiency_per_session=[0.5], stage_name="stage_2" + unignored_trials_per_session=[199], foraging_efficiency_per_session=[0.5], stage_name="STAGE_2" ) updated = TRAINER.evaluate(self.trainer_state, metrics) - self.assertEqual(updated.stage.name, "stage_1") + self.assertEqual(updated.stage.name, "STAGE_1") def test_stage_2_no_transition_on_middle_performance(self): metrics = make_metrics( - unignored_trials_per_session=[250], foraging_efficiency_per_session=[0.6], stage_name="stage_2" + unignored_trials_per_session=[250], foraging_efficiency_per_session=[0.6], stage_name="STAGE_2" ) updated = TRAINER.evaluate(self.trainer_state, metrics) - self.assertEqual(updated.stage.name, "stage_2") + self.assertEqual(updated.stage.name, "STAGE_2") class TestStage3Transitions(unittest.TestCase): @@ -124,10 +137,10 @@ def setUp(self): def test_stage_3_to_final_one_trial_performance(self): metrics = make_metrics( - unignored_trials_per_session=[400], foraging_efficiency_per_session=[0.7], stage_name="stage_3" + unignored_trials_per_session=[400], foraging_efficiency_per_session=[0.7], stage_name="STAGE_3" ) updated = TRAINER.evaluate(self.trainer_state, metrics) - self.assertEqual(updated.stage.name, "final") + self.assertEqual(updated.stage.name, "STAGE_FINAL") class TestFinalTransitions(unittest.TestCase): @@ -140,10 +153,10 @@ def test_final_to_graduated_on_excellent_performance(self): foraging_efficiency_per_session=[0.70] * 5, total_sessions=10, consecutive_sessions_at_current_stage=5, - stage_name="final", + stage_name="STAGE_FINAL", ) updated = TRAINER.evaluate(self.trainer_state, metrics) - self.assertEqual(updated.stage.name, "graduated") + self.assertEqual(updated.stage.name, "GRADUATED") def test_final_rollback_to_stage_3_on_poor_performance(self): metrics = make_metrics( @@ -151,10 +164,10 @@ def test_final_rollback_to_stage_3_on_poor_performance(self): foraging_efficiency_per_session=[0.55] * 5, total_sessions=10, consecutive_sessions_at_current_stage=5, - stage_name="final", + stage_name="STAGE_FINAL", ) updated = TRAINER.evaluate(self.trainer_state, metrics) - self.assertEqual(updated.stage.name, "stage_3") + self.assertEqual(updated.stage.name, "STAGE_3") def test_final_no_graduation_without_enough_sessions(self): metrics = make_metrics( @@ -162,10 +175,10 @@ def test_final_no_graduation_without_enough_sessions(self): foraging_efficiency_per_session=[0.70] * 5, total_sessions=5, consecutive_sessions_at_current_stage=3, - stage_name="final", + stage_name="STAGE_FINAL", ) updated = TRAINER.evaluate(self.trainer_state, metrics) - self.assertNotEqual(updated.stage.name, "graduated") + self.assertNotEqual(updated.stage.name, "GRADUATED") def test_graduated_is_absorbing(self): trainer_state = TRAINER.create_trainer_state(stage=make_s_stage_graduated()) @@ -174,10 +187,10 @@ def test_graduated_is_absorbing(self): foraging_efficiency_per_session=[0.9] * 5, total_sessions=20, consecutive_sessions_at_current_stage=10, - stage_name="final", + stage_name="GRADUATED", ) updated = TRAINER.evaluate(trainer_state, metrics) - self.assertEqual(updated.stage.name, "graduated") + self.assertEqual(updated.stage.name, "GRADUATED") if __name__ == "__main__": diff --git a/workspace/aind_behavior_dynamic_foraging_curricula/tests/test_uncoupled_baiting.py b/workspace/aind_behavior_dynamic_foraging_curricula/tests/test_uncoupled_baiting.py index a84fcae5..aa51e03e 100644 --- a/workspace/aind_behavior_dynamic_foraging_curricula/tests/test_uncoupled_baiting.py +++ b/workspace/aind_behavior_dynamic_foraging_curricula/tests/test_uncoupled_baiting.py @@ -17,7 +17,7 @@ def make_metrics( unignored_trials_per_session: list[int] = None, total_sessions: int = 1, consecutive_sessions_at_current_stage: int = 1, - stage_name: str = "stage_1_warmup", + stage_name: str = "STAGE_1_WARMUP", ) -> DynamicForagingMetrics: return DynamicForagingMetrics( foraging_efficiency_per_session=foraging_efficiency_per_session or [0.0], @@ -32,16 +32,16 @@ class TestCurriculumStructure(unittest.TestCase): def test_all_stages_in_curriculum(self): stages = CURRICULUM.see_stages() stage_names = [s.name for s in stages] - self.assertIn("stage_1_warmup", stage_names) - self.assertIn("stage_1", stage_names) - self.assertIn("stage_2", stage_names) - self.assertIn("stage_3", stage_names) - self.assertIn("final", stage_names) - self.assertIn("graduated", stage_names) + self.assertIn("STAGE_1_WARMUP", stage_names) + self.assertIn("STAGE_1", stage_names) + self.assertIn("STAGE_2", stage_names) + self.assertIn("STAGE_3", stage_names) + self.assertIn("STAGE_FINAL", stage_names) + self.assertIn("GRADUATED", stage_names) def test_enrollment_starts_at_stage_1_warmup(self): trainer_state = TRAINER.create_enrollment() - self.assertEqual(trainer_state.stage.name, "stage_1_warmup") + self.assertEqual(trainer_state.stage.name, "STAGE_1_WARMUP") class TestWarmupTransitions(unittest.TestCase): @@ -50,20 +50,20 @@ def setUp(self): def test_warmup_to_stage_2_on_good_performance(self): metrics = make_metrics( - unignored_trials_per_session=[250], foraging_efficiency_per_session=[0.65], stage_name="stage_1_warmup" + unignored_trials_per_session=[250], foraging_efficiency_per_session=[0.65], stage_name="STAGE_1_WARMUP" ) updated = TRAINER.evaluate(self.trainer_state, metrics) - self.assertEqual(updated.stage.name, "stage_2") + self.assertEqual(updated.stage.name, "STAGE_2") def test_warmup_to_stage_1_after_first_session(self): metrics = make_metrics( unignored_trials_per_session=[100], foraging_efficiency_per_session=[0.4], consecutive_sessions_at_current_stage=1, - stage_name="stage_1_warmup", + stage_name="STAGE_1_WARMUP", ) updated = TRAINER.evaluate(self.trainer_state, metrics) - self.assertEqual(updated.stage.name, "stage_1") + self.assertEqual(updated.stage.name, "STAGE_1") class TestStage1Transitions(unittest.TestCase): @@ -72,17 +72,17 @@ def setUp(self): def test_stage_1_to_stage_2_on_good_performance(self): metrics = make_metrics( - unignored_trials_per_session=[200], foraging_efficiency_per_session=[0.6], stage_name="stage_1" + unignored_trials_per_session=[200], foraging_efficiency_per_session=[0.6], stage_name="STAGE_1" ) updated = TRAINER.evaluate(self.trainer_state, metrics) - self.assertEqual(updated.stage.name, "stage_2") + self.assertEqual(updated.stage.name, "STAGE_2") def test_stage_1_no_transition_on_poor_performance(self): metrics = make_metrics( - unignored_trials_per_session=[100], foraging_efficiency_per_session=[0.4], stage_name="stage_1" + unignored_trials_per_session=[100], foraging_efficiency_per_session=[0.4], stage_name="STAGE_1" ) updated = TRAINER.evaluate(self.trainer_state, metrics) - self.assertEqual(updated.stage.name, "stage_1") + self.assertEqual(updated.stage.name, "STAGE_1") class TestStage2Transitions(unittest.TestCase): @@ -91,31 +91,44 @@ def setUp(self): def test_stage_2_to_stage_3_on_good_performance(self): metrics = make_metrics( - unignored_trials_per_session=[300], foraging_efficiency_per_session=[0.65], stage_name="stage_2" + unignored_trials_per_session=[300], + foraging_efficiency_per_session=[0.65], + consecutive_sessions_at_current_stage=3, + stage_name="STAGE_2", + ) + updated = TRAINER.evaluate(self.trainer_state, metrics) + self.assertEqual(updated.stage.name, "STAGE_3") + + def test_stage_2_requires_two_sessions_before_stage_3(self): + metrics = make_metrics( + unignored_trials_per_session=[300], + foraging_efficiency_per_session=[0.65], + consecutive_sessions_at_current_stage=1, + stage_name="STAGE_2", ) updated = TRAINER.evaluate(self.trainer_state, metrics) - self.assertEqual(updated.stage.name, "stage_3") + self.assertEqual(updated.stage.name, "STAGE_2") def test_stage_2_rollback_to_stage_1_on_poor_trials(self): metrics = make_metrics( - unignored_trials_per_session=[150], foraging_efficiency_per_session=[0.6], stage_name="stage_2" + unignored_trials_per_session=[150], foraging_efficiency_per_session=[0.6], stage_name="STAGE_2" ) updated = TRAINER.evaluate(self.trainer_state, metrics) - self.assertEqual(updated.stage.name, "stage_1") + self.assertEqual(updated.stage.name, "STAGE_1") def test_stage_2_rollback_to_stage_1_on_poor_efficiency(self): metrics = make_metrics( - unignored_trials_per_session=[199], foraging_efficiency_per_session=[0.5], stage_name="stage_2" + unignored_trials_per_session=[199], foraging_efficiency_per_session=[0.5], stage_name="STAGE_2" ) updated = TRAINER.evaluate(self.trainer_state, metrics) - self.assertEqual(updated.stage.name, "stage_1") + self.assertEqual(updated.stage.name, "STAGE_1") def test_stage_2_no_transition_on_middle_performance(self): metrics = make_metrics( - unignored_trials_per_session=[250], foraging_efficiency_per_session=[0.6], stage_name="stage_2" + unignored_trials_per_session=[250], foraging_efficiency_per_session=[0.6], stage_name="STAGE_2" ) updated = TRAINER.evaluate(self.trainer_state, metrics) - self.assertEqual(updated.stage.name, "stage_2") + self.assertEqual(updated.stage.name, "STAGE_2") class TestStage3Transitions(unittest.TestCase): @@ -124,10 +137,10 @@ def setUp(self): def test_stage_3_to_final_one_trial_performance(self): metrics = make_metrics( - unignored_trials_per_session=[400], foraging_efficiency_per_session=[0.7], stage_name="stage_3" + unignored_trials_per_session=[400], foraging_efficiency_per_session=[0.7], stage_name="STAGE_3" ) updated = TRAINER.evaluate(self.trainer_state, metrics) - self.assertEqual(updated.stage.name, "final") + self.assertEqual(updated.stage.name, "STAGE_FINAL") class TestFinalTransitions(unittest.TestCase): @@ -140,10 +153,10 @@ def test_final_to_graduated_on_excellent_performance(self): foraging_efficiency_per_session=[0.70] * 5, total_sessions=10, consecutive_sessions_at_current_stage=5, - stage_name="final", + stage_name="STAGE_FINAL", ) updated = TRAINER.evaluate(self.trainer_state, metrics) - self.assertEqual(updated.stage.name, "graduated") + self.assertEqual(updated.stage.name, "GRADUATED") def test_final_rollback_to_stage_3_on_poor_performance(self): metrics = make_metrics( @@ -151,10 +164,10 @@ def test_final_rollback_to_stage_3_on_poor_performance(self): foraging_efficiency_per_session=[0.55] * 5, total_sessions=10, consecutive_sessions_at_current_stage=5, - stage_name="final", + stage_name="STAGE_FINAL", ) updated = TRAINER.evaluate(self.trainer_state, metrics) - self.assertEqual(updated.stage.name, "stage_3") + self.assertEqual(updated.stage.name, "STAGE_3") def test_final_no_graduation_without_enough_sessions(self): metrics = make_metrics( @@ -162,10 +175,10 @@ def test_final_no_graduation_without_enough_sessions(self): foraging_efficiency_per_session=[0.70] * 5, total_sessions=5, consecutive_sessions_at_current_stage=3, - stage_name="final", + stage_name="STAGE_FINAL", ) updated = TRAINER.evaluate(self.trainer_state, metrics) - self.assertNotEqual(updated.stage.name, "graduated") + self.assertNotEqual(updated.stage.name, "GRADUATED") def test_graduated_is_absorbing(self): trainer_state = TRAINER.create_trainer_state(stage=make_s_stage_graduated()) @@ -174,10 +187,10 @@ def test_graduated_is_absorbing(self): foraging_efficiency_per_session=[0.9] * 5, total_sessions=20, consecutive_sessions_at_current_stage=10, - stage_name="final", + stage_name="GRADUATED", ) updated = TRAINER.evaluate(trainer_state, metrics) - self.assertEqual(updated.stage.name, "graduated") + self.assertEqual(updated.stage.name, "GRADUATED") if __name__ == "__main__": From da3aad158401ecb4e94bcb1061551dd63476323e Mon Sep 17 00:00:00 2001 From: Micah Woodard Date: Wed, 9 Sep 2026 15:34:10 -0700 Subject: [PATCH 2/3] removes redundant step in cicd --- .github/workflows/dynamic-foraging-cicd.yml | 9 --------- 1 file changed, 9 deletions(-) diff --git a/.github/workflows/dynamic-foraging-cicd.yml b/.github/workflows/dynamic-foraging-cicd.yml index 270998c2..421fcc55 100644 --- a/.github/workflows/dynamic-foraging-cicd.yml +++ b/.github/workflows/dynamic-foraging-cicd.yml @@ -59,15 +59,6 @@ jobs: uv run python -m unittest discover tests uv run --directory .\workspace\aind_behavior_dynamic_foraging_curricula\ python -m unittest discover tests - - name: Run workspace python unit tests - shell: bash - run: | - for workspace in workspace/*; do - if [ -d "$workspace/tests" ]; then - uv run --directory "$workspace" python -m unittest discover tests - fi - done - - name: Regenerate all schemas run: | uv run dynamic-foraging regenerate From 6000210ffefc269028980e105823a6c37276122a Mon Sep 17 00:00:00 2001 From: Micah Woodard Date: Wed, 9 Sep 2026 15:38:37 -0700 Subject: [PATCH 3/3] tests boundries of transition --- .../tests/test_coupled_baiting.py | 2 +- .../tests/test_uncoupled.py | 2 +- .../tests/test_uncoupled_baiting.py | 2 +- 3 files changed, 3 insertions(+), 3 deletions(-) diff --git a/workspace/aind_behavior_dynamic_foraging_curricula/tests/test_coupled_baiting.py b/workspace/aind_behavior_dynamic_foraging_curricula/tests/test_coupled_baiting.py index 888a406a..94e853b7 100644 --- a/workspace/aind_behavior_dynamic_foraging_curricula/tests/test_coupled_baiting.py +++ b/workspace/aind_behavior_dynamic_foraging_curricula/tests/test_coupled_baiting.py @@ -93,7 +93,7 @@ def test_stage_2_to_stage_3_on_good_performance(self): metrics = make_metrics( unignored_trials_per_session=[300], foraging_efficiency_per_session=[0.65], - consecutive_sessions_at_current_stage=3, + consecutive_sessions_at_current_stage=1, stage_name="STAGE_2", ) updated = TRAINER.evaluate(self.trainer_state, metrics) diff --git a/workspace/aind_behavior_dynamic_foraging_curricula/tests/test_uncoupled.py b/workspace/aind_behavior_dynamic_foraging_curricula/tests/test_uncoupled.py index 4be63cc7..46e741c3 100644 --- a/workspace/aind_behavior_dynamic_foraging_curricula/tests/test_uncoupled.py +++ b/workspace/aind_behavior_dynamic_foraging_curricula/tests/test_uncoupled.py @@ -93,7 +93,7 @@ def test_stage_2_to_stage_3_on_good_performance(self): metrics = make_metrics( unignored_trials_per_session=[300], foraging_efficiency_per_session=[0.65], - consecutive_sessions_at_current_stage=3, + consecutive_sessions_at_current_stage=2, stage_name="STAGE_2", ) updated = TRAINER.evaluate(self.trainer_state, metrics) diff --git a/workspace/aind_behavior_dynamic_foraging_curricula/tests/test_uncoupled_baiting.py b/workspace/aind_behavior_dynamic_foraging_curricula/tests/test_uncoupled_baiting.py index aa51e03e..e7a1a648 100644 --- a/workspace/aind_behavior_dynamic_foraging_curricula/tests/test_uncoupled_baiting.py +++ b/workspace/aind_behavior_dynamic_foraging_curricula/tests/test_uncoupled_baiting.py @@ -93,7 +93,7 @@ def test_stage_2_to_stage_3_on_good_performance(self): metrics = make_metrics( unignored_trials_per_session=[300], foraging_efficiency_per_session=[0.65], - consecutive_sessions_at_current_stage=3, + consecutive_sessions_at_current_stage=2, stage_name="STAGE_2", ) updated = TRAINER.evaluate(self.trainer_state, metrics)