import jax import jax.numpy as jnp # A function that works on ONE number def square(x): return x ** 2 # Transform it to work on MANY numbers batched_square = jax.vmap(square) # Now use it numbers = jnp.array([1, 2, 3, 4, 5]) result = batched_square(numbers) print(result) # [1, 4, 9, 16, 25] __ __