@jax.jit def matmul_jax_jit(x, y): return jnp.dot(x, y) # First call: JAX traces and compiles the function # This takes a moment, so we don't include it in the benchmark print("Compiling...") _ = matmul_jax_jit(x_jax, y_jax).block_until_ready() print("Done.") # Timed run (using the compiled version) start = time.perf_counter() result_jit = matmul_jax_jit(x_jax, y_jax).block_until_ready() jit_time = time.perf_counter() - start print(f"JAX time (with JIT): {jit_time:.4f} seconds") __ __