def single_multiply(a, b): return a * b # Map over rows, then over columns double_batched = jax.vmap(jax.vmap(single_multiply)) matrix_a = jnp.array([[1, 2], [3, 4]]) matrix_b = jnp.array([[10, 20], [30, 40]]) result = double_batched(matrix_a, matrix_b) print(result) # [[10, 40], # [90, 160]] __ __