# matrix multiplication in jax import jax.numpy as jnp from jax import jit jxArrOne = jnp.random.rand(1000, 1000) jxArrTwo = jnp.random.rand(1000, 1000) @jit def matmul(jxArrOne, jxArrTwo): return jnp.dot(jxArrOne, jxArrTwo) start = time.time() jxArrThree = matmul(jxArrOne, jxArrOne) end = time.time() print("JAX time:", end - start) __ __