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)