PositionSplitManager.axis_sizes

PositionSplitManager.axis_sizes#

property PositionSplitManager.axis_sizes: tuple[int, ...]#

Total axis sizes for each contained split.

Examples

>>> import jax.numpy as jnp
>>> from liesel.optim import PositionSplitManager, PositionSplit
>>> from liesel.optim.types import Position
>>> manager = PositionSplitManager(
...     [
...         PositionSplit(
...             Position({"x": jnp.arange(2)}),
...             Position({}),
...             Position({}),
...             2,
...             0,
...             0,
...         ),
...         PositionSplit(
...             Position({"y": jnp.arange(3)}),
...             Position({}),
...             Position({}),
...             3,
...             0,
...             0,
...         ),
...     ]
... )
>>> manager.axis_sizes
(2, 3)