from jax import random key = random.PRNGKey(42) print(random.normal(key, shape=(3,))) # Always the same print(random.normal(key, shape=(3,))) # Still the same! __ __