traveler.core.array_tools.pytree_concat#

traveler.core.array_tools.pytree_concat(results, axis=0)[source]#

Concatenate matching leaves from a sequence of JAX pytrees.

Parameters:
  • results (sequence) – Pytrees that share the same structure and whose corresponding leaves can be concatenated.

  • axis (int or None, optional) – Axis passed to jax.numpy.concatenate().

Returns:

A pytree with the original structure and concatenated leaves.

Return type:

Any