import jax.numpy as jnp # create a 1D array jxArrOne = jnp.array([1, 2, 3]) print(jxArrOne) # create a 2D array jxArrTwo = jnp.array([[1, 2], [3, 4]]) print(jxArrTwo) __ __