Split#
- class liesel.optim.Split(position_keys=None, axis_size=0, validate_axis_size=0, test_axis_size=0, train_axis_size=None, split_axes=<factory>, default_split_axis=0, shuffle=True, seed=0, sample_sizes=None, keep_in_train=None)[source]#
Bases:
objectDefines how observed position entries are split into train, validation, and test.
Splitstores a vector of observation indices. The firsttrain_axis_sizeindices become the training split, the nextvalidate_axis_sizeindices become the validation split, and the finaltest_axis_sizeindices become the test split. Ifshuffle=True, the index vector is permuted once during initialization when a validation or test part is nonempty. Full-data splits preserve row order and do not use a seed.- Parameters:
position_keys (
Sequence[str] |None, default:None) – Names of position entries that should be included. If omitted,split_position()uses all keys in the supplied position. Entries mapped toNoneinsplit_axesare included unchanged intrain,validate, andtestand are not batched automatically.axis_size (
int, default:0) – Number of observations along each split axis. Must be positive.validate_axis_size (
int, default:0) – Number of validation observations.test_axis_size (
int, default:0) – Number of test observations.train_axis_size (
int|None, default:None) – Number of training observations. If left atNone, it is computed asaxis_size - validate_axis_size - test_axis_size.split_axes (
dict[str,int|None] |None, default:<factory>) – Optional mapping from position key to split axis. Mapping a key toNonemakes it passthrough data: it is included unchanged intrain,validate, andtest, is not split, and is excluded from automatically derived batches. Use this for shared lookup tables or constants, not per-observation data. Keys missing from this mapping usedefault_split_axis.default_split_axis (
int, default:0) – Split axis for all position keys not listed insplit_axes.shuffle (
bool, default:True) – Whether to shuffle observations during initialization; defaults toTrue. Full-data splits preserve order regardless of this setting.seed (
Array|int|None, default:0) – Seed or JAX pseudo-random key used for shuffled holdouts. Defaults to0. ExplicitNoneuses Unix time in whole seconds. Ignored for full-data splits.sample_sizes (
Mapping[Literal['train','validate','test'],int|float] |None, default:None) – Optional effective sample sizes passed to the resultingPositionSplit.keep_in_train (
Sequence[int] |None, default:None) – Positional row indices that must belong to the training partition. Reserved rows count towardtrain_axis_size; this controls membership, not order. This is useful when a random split must retain rare categories in training.
- Raises:
ValueError – If
axis_sizeis not positive, if any split size is negative, or iftrain_axis_size + validate_axis_size + test_axis_sizeis not exactly equal toaxis_size.
Examples
Split one vector without shuffling:
>>> import jax.numpy as jnp >>> from liesel.optim import Split >>> splitter = Split( ... ["x"], axis_size=10, validate_axis_size=2, test_axis_size=1, shuffle=False ... ) >>> splitter Split(train=7, validate=2, test=1) >>> splitter.indices_train.tolist() [0, 1, 2, 3, 4, 5, 6] >>> split = splitter.split_position({"x": jnp.arange(10)}) >>> ( ... split.train["x"].tolist(), ... split.validate["x"].tolist(), ... split.test["x"].tolist(), ... ) ([0, 1, 2, 3, 4, 5, 6], [7, 8], [9])
Split different entries along different split_axes:
>>> splitter = Split( ... ["x", "y"], ... axis_size=4, ... validate_axis_size=1, ... test_axis_size=1, ... split_axes={"x": 1}, ... shuffle=False, ... ) >>> position = { ... "x": jnp.arange(8).reshape(2, 4), ... "y": jnp.arange(12).reshape(4, 3), ... } >>> split = splitter.split_position(position) >>> split.train["x"].tolist() [[0, 1], [4, 5]] >>> split.validate["y"].tolist() [[6, 7, 8]]
Methods
from_axis_shares(position_keys, axis_size[, ...])Builds a
Splitfrom validation and test proportions.from_model(model[, position_keys, ...])permute_indices(key)Returns a random permutation of the current index vector.
split_position(position)Splits position entries into train, validation, and test positions.
Attributes
Number of observations along each split axis.
Split axis for all position keys not listed in
split_axes.Whether this splitter assigns observations to testing.
Whether this splitter assigns observations to validation.
Observation indices for the test split.
Observation indices for the training split.
Observation indices for the validation split.
Positional row indices that must belong to the training partition.
Position keys copied unchanged into every split part.
Names of position entries that should be included.
Optional effective sample sizes passed to the resulting
PositionSplit.Seed or JAX pseudo-random key used for shuffled holdouts.
Whether to shuffle observations during initialization; defaults to
True.Position keys that are partitioned by this splitter.
Share of observations assigned to testing.
Number of test observations.
Number of training observations.
Share of observations assigned to validation.
Number of validation observations.
Optional mapping from position key to split axis.
Observation order, optionally shuffled once during initialization.