Split Bijector¤
distreqx.bijectors.Split(distreqx.bijectors.AbstractForwardInverseBijector, distreqx.bijectors.AbstractInvLogDetJacBijector, distreqx.bijectors.AbstractFwdLogDetJacBijector)
¤
A bijector that splits a single array into a tuple of arrays along an axis.
This operates as a wrapper around jax.numpy.split.
__init__(indices_or_sections: int | tuple | list, axis: int = -1)
¤
Initializes a Split bijector.
Arguments:
indices_or_sections: If an integerN, the array will be divided intoNequal arrays along axis. If a tuple/list of sorted integers, the entries indicate where along axis the array is split.axis: The axis along which to split. Defaults to -1 (last axis).