@jax.jit def broken_random(): return np.random.randn(5) # Using NumPy's random result1 = broken_random() result2 = broken_random() print(f"First call: {result1}") print(f"Second call: {result2}") __ __