from jax import random @jax.jit def correct_random(key): return random.normal(key, shape=(5,)) # Create a PRNG key key = random.PRNGKey(42) # Split the key for each use key, subkey1 = random.split(key) result1 = correct_random(subkey1) key, subkey2 = random.split(key) result2 = correct_random(subkey2) print(f"First call: {result1}") print(f"Second call: {result2}") __ __