Split

Contents

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: object

Defines how observed position entries are split into train, validation, and test.

Split stores a vector of observation indices. The first train_axis_size indices become the training split, the next validate_axis_size indices become the validation split, and the final test_axis_size indices become the test split. If shuffle=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 to None in split_axes are included unchanged in train, validate, and test and 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 at None, it is computed as axis_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 to None makes it passthrough data: it is included unchanged in train, validate, and test, 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 use default_split_axis.

  • default_split_axis (int, default: 0) – Split axis for all position keys not listed in split_axes.

  • shuffle (bool, default: True) – Whether to shuffle observations during initialization; defaults to True. 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 to 0. Explicit None uses 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 resulting PositionSplit.

  • keep_in_train (Sequence[int] | None, default: None) – Positional row indices that must belong to the training partition. Reserved rows count toward train_axis_size; this controls membership, not order. This is useful when a random split must retain rare categories in training.

Raises:

ValueError – If axis_size is not positive, if any split size is negative, or if train_axis_size + validate_axis_size + test_axis_size is not exactly equal to axis_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 Split from 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

axis_size

Number of observations along each split axis.

default_split_axis

Split axis for all position keys not listed in split_axes.

has_test

Whether this splitter assigns observations to testing.

has_validation

Whether this splitter assigns observations to validation.

indices_test

Observation indices for the test split.

indices_train

Observation indices for the training split.

indices_validate

Observation indices for the validation split.

keep_in_train

Positional row indices that must belong to the training partition.

passthrough_position_keys

Position keys copied unchanged into every split part.

position_keys

Names of position entries that should be included.

sample_sizes

Optional effective sample sizes passed to the resulting PositionSplit.

seed

Seed or JAX pseudo-random key used for shuffled holdouts.

shuffle

Whether to shuffle observations during initialization; defaults to True.

split_position_keys

Position keys that are partitioned by this splitter.

test_axis_share

Share of observations assigned to testing.

test_axis_size

Number of test observations.

train_axis_size

Number of training observations.

validate_axis_share

Share of observations assigned to validation.

validate_axis_size

Number of validation observations.

split_axes

Optional mapping from position key to split axis.

indices

Observation order, optionally shuffled once during initialization.