From 446a15024821358d4ab97a5531ddd14211fcf10c Mon Sep 17 00:00:00 2001 From: arjunsridhar12345 Date: Fri, 7 Aug 2026 12:13:26 -0700 Subject: [PATCH 1/2] fix: add check for uniform distribution --- .../processing/_trial_table.py | 13 ++++++++++--- 1 file changed, 10 insertions(+), 3 deletions(-) diff --git a/src/dynamic_foraging_processing/processing/_trial_table.py b/src/dynamic_foraging_processing/processing/_trial_table.py index 9daaab9..8158162 100644 --- a/src/dynamic_foraging_processing/processing/_trial_table.py +++ b/src/dynamic_foraging_processing/processing/_trial_table.py @@ -233,7 +233,9 @@ def _distribution_stats( ``beta`` is the scale of an exponential distribution (``1 / rate``); it is ``None`` for non-exponential families (e.g. the scalar quiescent - duration). ``min``/``max`` come from the truncation parameters when set. + duration). ``min``/``max`` come from the truncation parameters when set, + except for a uniform distribution, whose bounds are its own ``min`` and + ``max`` distribution parameters. Parameters ---------- @@ -247,11 +249,16 @@ def _distribution_stats( """ beta: t.Optional[float] = None params = distribution.distribution_parameters - if params.family == DistributionFamily.EXPONENTIAL and params.rate: - beta = 1.0 / params.rate truncation = distribution.truncation_parameters minimum = truncation.min if truncation is not None else None maximum = truncation.max if truncation is not None else None + if params.family == DistributionFamily.EXPONENTIAL and params.rate: + beta = 1.0 / params.rate + elif params.family == DistributionFamily.UNIFORM: + # A uniform distribution carries its bounds in the distribution + # parameters rather than the truncation parameters. + minimum = params.min + maximum = params.max return beta, minimum, maximum @staticmethod From 30a0bbd78c9549f3e06f082cdebb9645cf6ac098 Mon Sep 17 00:00:00 2001 From: arjunsridhar12345 Date: Fri, 7 Aug 2026 12:13:34 -0700 Subject: [PATCH 2/2] test: update tests --- tests/test_processing/test_trial_table.py | 51 +++++++++++++++++++---- 1 file changed, 42 insertions(+), 9 deletions(-) diff --git a/tests/test_processing/test_trial_table.py b/tests/test_processing/test_trial_table.py index e5ebe3f..ebadf13 100644 --- a/tests/test_processing/test_trial_table.py +++ b/tests/test_processing/test_trial_table.py @@ -23,6 +23,8 @@ Scalar, ScalarDistributionParameter, TruncationParameters, + UniformDistribution, + UniformDistributionParameters, ) from dynamic_foraging_processing.processing import TrialConfig, TrialTableBuilder @@ -139,16 +141,28 @@ def _outcome( } -def _task_logic(quiescent_scalar=True): - """Build a coupled-generator task logic with known distribution parameters.""" - quiescent = ( - Scalar(distribution_parameters=ScalarDistributionParameter(value=0.0)) - if quiescent_scalar - else ExponentialDistribution( - distribution_parameters=ExponentialDistributionParameters(rate=1.0), - truncation_parameters=TruncationParameters(min=0.0, max=1.0), +def _quiescent_distribution(kind): + """Build the quiescent-duration distribution for the requested family. + + ``"uniform"`` deliberately also carries truncation parameters that differ + from its own bounds, so tests can pin down which pair is reported. + """ + if kind == "scalar": + return Scalar(distribution_parameters=ScalarDistributionParameter(value=0.0)) + if kind == "uniform": + return UniformDistribution( + distribution_parameters=UniformDistributionParameters(min=0.25, max=0.75), + truncation_parameters=TruncationParameters(min=9.0, max=99.0), ) + return ExponentialDistribution( + distribution_parameters=ExponentialDistributionParameters(rate=1.0), + truncation_parameters=TruncationParameters(min=0.0, max=1.0), ) + + +def _task_logic(quiescent="scalar"): + """Build a coupled-generator task logic with known distribution parameters.""" + quiescent = _quiescent_distribution(quiescent) spec = CoupledTrialGeneratorSpec( quiescent_duration=quiescent, inter_trial_interval_duration=ExponentialDistribution( @@ -348,6 +362,9 @@ def test_build_full_dataset(): assert first["ITI_min"] == 1.0 and first["ITI_max"] == 10.0 assert first["block_beta"] == pytest.approx(20.0) assert pd.isna(first["delay_beta"]) # scalar quiescent distribution + # Scalar has neither a scale nor truncation parameters -> null bounds. + assert pd.isna(first["delay_min"]) + assert pd.isna(first["delay_max"]) assert pd.isna(first["min_reward_each_block"]) # removed from generator schema assert first["base_reward_probability_sum"] == pytest.approx(0.8) @@ -418,7 +435,7 @@ def test_build_exponential_quiescent_sets_delay_beta(): """A non-scalar quiescent distribution populates delay beta/min/max.""" behavior = _full_dataset().children["Behavior"] behavior.children["InputSchemas"].children["TaskLogic"] = _Stream( - _task_logic(quiescent_scalar=False) + _task_logic(quiescent="exponential") ) table = TrialTableBuilder(_Node({"Behavior": behavior})).build() assert table.iloc[0]["delay_beta"] == pytest.approx(1.0) @@ -426,6 +443,22 @@ def test_build_exponential_quiescent_sets_delay_beta(): assert table.iloc[0]["delay_max"] == 1.0 +def test_build_uniform_quiescent_sets_delay_bounds_from_parameters(): + """A uniform quiescent distribution reports its own bounds and no beta.""" + behavior = _full_dataset().children["Behavior"] + behavior.children["InputSchemas"].children["TaskLogic"] = _Stream( + _task_logic(quiescent="uniform") + ) + table = TrialTableBuilder(_Node({"Behavior": behavior})).build() + first = table.iloc[0] + # Uniform has no scale parameter, so beta stays null. + assert pd.isna(first["delay_beta"]) + # The bounds come from the distribution parameters (0.25/0.75), not from the + # truncation parameters the fixture also sets (9.0/99.0). + assert first["delay_min"] == pytest.approx(0.25) + assert first["delay_max"] == pytest.approx(0.75) + + def test_build_warns_on_misaligned_streams(caplog): """A per-trial stream shorter than TrialOutcome warns but still builds.""" table = TrialTableBuilder(_misaligned_dataset()).build()