import jax import jax.numpy as jnp from jax import random # Basic array operations (same as NumPy) x = jnp.array([1, 2, 3]) y = jnp.zeros((3, 3)) z = jnp.dot(a, b) # Updating arrays (immutable style) new_x = x.at[0].set(99) new_x = x.at[1:].add(10) # JIT compilation @jax.jit def fast_function(x): return jnp.dot(x, x.T) # Explicit randomness key = random.PRNGKey(0) key, subkey = random.split(key) samples = random.normal(subkey, shape=(100,)) # Block for accurate timing result = fast_function(x).block_until_ready() __ __