PositionSplit.from_model()

PositionSplit.from_model()#

static PositionSplit.from_model(model, position_keys=None, axis_size=None, validate_axis_share=0.0, test_axis_share=0.0, split_axes=None, default_split_axis=0, shuffle=True, seed=0, multi_size='error', sample_sizes=None, infer_sample_sizes=True)[source]#

Builds a PositionSplit from the observed variables in a model.

Parameters:
  • model (Model) – Model containing the observed variables to split.

  • position_keys (Sequence[str] | Sequence[Sequence[str]] | None, default: None) – Names of observed position entries to include. If None, strong observed variables in model are used. Weak observations are recomputed from their strong inputs and cannot be selected directly. Flat keys are grouped by axis length; nested keys specify exact groups, including equal-sized groups. Each group must have matching lengths along its configured axes. Use split_axes={key: None} to include a selected entry unchanged in train, validate, and test without splitting or batching it automatically.

  • axis_size (int | None, default: None) – Number of observations along the split axis. If None, the number is guessed from model along default_split_axis.

  • validate_axis_share (float, default: 0.0) – Share of observations assigned to the validation split.

  • test_axis_share (float, default: 0.0) – Share of observations assigned to the test split.

  • split_axes (dict[str, int | None] | None, default: None) – 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 observations are shuffled before splitting; 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.

  • multi_size (Literal['error', 'manager'], default: 'error') – How to handle multiple inferred or explicit observation groups. The default "error" keeps PositionSplit scalar and raises a helpful error. Use "manager" to return a PositionSplitManager when multiple groups are detected, even if they have equal axis sizes. One group still returns a scalar split.

  • sample_sizes (Mapping[Literal['train', 'validate', 'test'], int | float] | None, default: None) – Optional effective sample sizes for train, validation, and test scaling. If supplied, these values are used instead of automatic inference.

  • infer_sample_sizes (bool, default: True) – Whether to infer effective sample sizes from pointwise observed log-probability arrays. Inference counts log-probability scalars, not observed value elements; for multivariate observation distributions, one observed event may have several value dimensions but one pointwise log-probability scalar.

Return type:

PositionSplit | PositionSplitManager

Returns:

PositionSplit or PositionSplitManager – Split observed position entries extracted from model.

Examples

>>> import jax.numpy as jnp
>>> import liesel.model as lsl
>>> from liesel.optim import PositionSplit
>>> y = lsl.Var.new_obs(jnp.arange(10.0), name="y")
>>> model = lsl.Model([y])
>>> split = PositionSplit.from_model(
...     model,
...     position_keys=[["y"]],
...     validate_axis_share=0.2,
...     test_axis_share=0.1,
...     shuffle=False,
... )
>>> split.train_axis_size, split.validate_axis_size, split.test_axis_size
(7, 2, 1)
>>> split.train["y"].tolist()
[0.0, 1.0, 2.0, 3.0, 4.0, 5.0, 6.0]

Multi-size observed data must opt into the manager API:

>>> y1_multi = lsl.Var.new_obs(jnp.arange(10.0), name="y1_multi")
>>> y2_multi = lsl.Var.new_obs(jnp.arange(6.0), name="y2_multi")
>>> model = lsl.Model([y1_multi, y2_multi])
>>> managed = PositionSplit.from_model(
...     model,
...     position_keys=[["y1_multi"], ["y2_multi"]],
...     validate_axis_share=0.2,
...     multi_size="manager",
...     shuffle=True,
...     seed=42,
... )
>>> type(managed).__name__
'PositionSplitManager'