# JAX: The correct way arr = jnp.zeros(5) new_arr = arr.at[0].set(42) print(arr) # [0. 0. 0. 0. 0.] — Original unchanged print(new_arr) # [42. 0. 0. 0. 0.] — New array with the update __ __