# update an element in numpy npArr = np.array([1, 2, 3]) npArr[0] = 4 print(npArr) # output: [4, 2, 3] # update an element in jax jxArr = jnp.array([1, 2, 3]) jxArr = jnp.index_update(b, 0, 4) print(jxArr) # output: [4, 2, 3] __ __