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.
Noneflattens the inputs first.dtype – The desired data type of the result.
- Returns:
The concatenated array.