# indexing in numpy npArr = np.array([1, 2, 3]) print(npArr[0]) # output: 1 # indexing in jax jxArr = jnp.array([1, 2, 3]) print(jxArr[0]) # output: 1 __ __