jax_concatenate

probly.representation.jax_functions.jax_concatenate(arrays: Sequence[JaxArrayLike], /, axis: int | None = 0, dtype: DTypeLike | None = None) jax.Array[source]

Join a sequence of arrays along an existing axis, mirroring jax.numpy.concatenate.

Parameters:
  • arrays – The arrays to concatenate.

  • axis – The axis to concatenate along. None flattens the inputs first.

  • dtype – The desired data type of the result.

Returns:

The concatenated array.