Skip to content

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 integer N, the array will be divided into N equal 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).