traveler.core.multithread.pytree_concat#
- traveler.core.multithread.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: