SplitManager#
- class liesel.optim.SplitManager(splits, *, passthrough_position_keys=<factory>)[source]#
Bases:
objectWraps multiple
Splitobjects for multi-branch splitting.Splitstays scalar: each instance assumes one axis size.SplitManagercoordinates several such scalar splitters and returns aPositionSplitManagerwith merged train/validation/test positions.- Parameters:
splits (
Sequence[Split]) – Non-empty sequence ofSplitobjects. Theirposition_keysmust 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 intrain,validate, andtest. They are not assigned to a childSplitor 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
SplitManagerfrom inferred or explicit groups.split_position(position)Splits a position with every child and merges the result.
Attributes
Total axis sizes for each contained split.
Whether all child splits contain test data.
Whether all child splits contain validation data.
Position keys claimed by all contained splits.
Test axis sizes for each contained split.
Training axis sizes for each contained split.
Validation axis sizes for each contained split.
Non-empty sequence of
Splitobjects.Keyword-only names included unchanged in
train,validate, andtest.