SplitManager

Contents

SplitManager#

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

Bases: object

Wraps multiple Split objects for multi-branch splitting.

Split stays scalar: each instance assumes one axis size. SplitManager coordinates several such scalar splitters and returns a PositionSplitManager with merged train/validation/test positions.

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

  • passthrough_position_keys (Sequence[str], default: <factory>) – Keyword-only names included unchanged in train, validate, and test. They are not assigned to a child Split or batched automatically. Use this for shared lookup tables or constants, not per-observation data.

Examples

>>> import jax.numpy as jnp
>>> from liesel.optim import SplitManager, Split
>>> manager = SplitManager(
...     [
...         Split(["x"], axis_size=10, validate_axis_size=2, shuffle=False),
...         Split(["y"], axis_size=6, validate_axis_size=1, shuffle=False),
...     ]
... )
>>> split = manager.split_position({"x": jnp.arange(10), "y": jnp.arange(6)})
>>> split.train_axis_sizes
(8, 5)
>>> split.validate["x"].tolist(), split.validate["y"].tolist()
([8, 9], [5])

Automatically group model observations by axis size:

>>> import liesel.model as lsl
>>> x = lsl.Var.new_obs(jnp.arange(8.0), name="x")
>>> y = lsl.Var.new_obs(jnp.arange(5.0), name="y")
>>> model = lsl.Model([x, y])
>>> manager = SplitManager.from_model(
...     model,
...     position_keys=[["x"], ["y"]],
...     validate_axis_share=0.2,
...     shuffle=True,
...     seed=42,
... )
>>> manager.axis_sizes
(8, 5)

Methods

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

Builds a SplitManager from inferred or explicit groups.

split_position(position)

Splits a position with every child and merges the result.

Attributes

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 splits.

test_axis_sizes

Test axis sizes for each contained split.

train_axis_sizes

Training axis sizes for each contained split.

validate_axis_sizes

Validation axis sizes for each contained split.

splits

Non-empty sequence of Split objects.

passthrough_position_keys

Keyword-only names included unchanged in train, validate, and test.