# vmap in Jax def add(x, y): return x + y jxArrOne = jnp.array([[1, 2], [3, 4]]) jxArrTwo = jnp.array([[10, 20], [30, 40]]) # apply add to multiple inputs using broadcasting jxArrThree = jnp.vmap(add)(jxArrOne, jxArrTwo) print(jxArrThree) # output: [[11, 22], [33, 44]] __ __