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
13 changes: 10 additions & 3 deletions src/dynamic_foraging_processing/processing/_trial_table.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
----------
Expand All @@ -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
Expand Down
51 changes: 42 additions & 9 deletions tests/test_processing/test_trial_table.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,8 @@
Scalar,
ScalarDistributionParameter,
TruncationParameters,
UniformDistribution,
UniformDistributionParameters,
)

from dynamic_foraging_processing.processing import TrialConfig, TrialTableBuilder
Expand Down Expand Up @@ -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):

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

maybe type kind arg

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

oh wait this is tests nevermind

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

Expand Down Expand Up @@ -418,14 +435,30 @@ 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)
assert table.iloc[0]["delay_min"] == 0.0
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()
Expand Down