PositionSplitManager

PositionSplitManager#

class liesel.optim.PositionSplitManager(splits, *, passthrough=<factory>)[source]#

Bases: object

Coordinates multiple PositionSplit objects as one split interface.

PositionSplitManager is the split-side counterpart to BatchManager. It is useful when a model has observed branches with different axis sizes. Each child PositionSplit stores the split data for one branch, while the manager exposes merged train, validate, and test positions.

Parameters:
  • splits (Sequence[PositionSplit]) – Non-empty sequence of PositionSplit objects. Their PositionSplit.position_keys must not overlap. Either all children must contain validation data or none may contain validation data; the same rule applies to test data.

  • passthrough (Position (dict[str, Any]), default: <factory>) – Keyword-only shared position entries included unchanged in the manager’s merged train, validate, and test positions. They are not split or batched automatically. Use this for shared lookup tables or constants, not per-observation data.

Raises:

ValueError – If splits is empty, if position keys overlap, or if validation/test availability differs across children.

Notes

Branch-specific sizes are available as axis_sizes, train_axis_sizes, validate_axis_sizes, and test_axis_sizes. Branch-specific likelihood sizes are available as sample_sizes, train_sample_sizes, validate_sample_sizes, and test_sample_sizes. Scalar aliases such as train_axis_size, validate_axis_share, and validate_sample_scale are available only when all children have the same value. Use sample_size() for the total likelihood sample size of a split part across all branches.

Examples

Merge two branches with different axis sizes:

>>> import jax.numpy as jnp
>>> from liesel.optim import PositionSplitManager, Split
>>> position = {"x": jnp.arange(10), "y": jnp.arange(6)}
>>> split_x = Split(
...     ["x"], axis_size=10, validate_axis_size=2, shuffle=False
... ).split_position(position)
>>> split_y = Split(
...     ["y"], axis_size=6, validate_axis_size=1, shuffle=False
... ).split_position(position)
>>> manager = PositionSplitManager([split_x, split_y])
>>> manager.position_keys
['x', 'y']
>>> manager.train_axis_sizes
(8, 5)
>>> manager.train["x"].shape, manager.train["y"].shape
((8,), (5,))

Unequal scalar aliases raise and direct users to the plural property:

>>> try:
...     manager.train_axis_size
... except ValueError as error:
...     print("train_axis_sizes" in str(error))
True

Build a manager directly from a model with two observation sizes:

>>> import liesel.model as lsl
>>> y1 = lsl.Var.new_obs(jnp.arange(10.0), name="y1")
>>> y2 = lsl.Var.new_obs(jnp.arange(6.0), name="y2")
>>> model = lsl.Model([y1, y2])
>>> managed = PositionSplitManager.from_model(
...     model,
...     position_keys=[["y1"], ["y2"]],
...     validate_axis_share=0.2,
...     shuffle=True,
...     seed=42,
... )
>>> managed.validate_axis_sizes
(2, 1)

Methods

from_model(model[, position_keys, ...])

Builds grouped position splits from a model.

sample_scale(part)

Common sample scale, available only when all branches agree.

sample_scales(part)

Sample scaling factors for each contained split.

sample_size(part)

Total effective likelihood sample size for one split part.

scaled_log_lik(model, model_state[, part])

Returns the log likelihood with branch-specific split scaling.

Attributes

axis_size

Common total axis size, available only when all branches agree.

axis_sizes

Total axis sizes for each contained split.

has_test

Whether all child splits contain test data.

has_validation

Whether all child splits contain validation data.

position_keys

Position keys claimed by all contained split objects.

sample_sizes

Raw effective sample-size mappings for each contained split.

split_position_keys

Position keys that are partitioned rather than passed through.

test

Merged test position.

test_axis_share

Common test share, available only when all branches agree.

test_axis_shares

Test shares for each contained split.

test_axis_size

Common test axis size, available only when all branches agree.

test_axis_sizes

Test axis sizes for each contained split.

test_sample_size

Common test sample size, available only when all branches agree.

test_sample_sizes

Test sample sizes for each contained split.

train

Merged training position.

train_axis_size

Common training axis size, available only when all branches agree.

train_axis_sizes

Training axis sizes for each contained split.

train_sample_size

Common training sample size, available only when all branches agree.

train_sample_sizes

Training sample sizes for each contained split.

validate

Merged validation position.

validate_axis_share

Common validation share, available only when all branches agree.

validate_axis_shares

Validation shares for each contained split.

validate_axis_size

Common validation axis size, available only when all branches agree.

validate_axis_sizes

Validation axis sizes for each contained split.

validate_sample_scale

Common validation likelihood scale.

validate_sample_scales

Validation likelihood scales for each contained split.

validate_sample_size

Common validation sample size, available only when all branches agree.

validate_sample_sizes

Validation sample sizes for each contained split.

splits

Non-empty sequence of PositionSplit objects.

passthrough

Keyword-only shared position entries included unchanged in the manager's merged train, validate, and test positions.