def matmul_jax(x, y): return jnp.dot(x, y) # Warmup _ = matmul_jax(x_jax, y_jax).block_until_ready() # Timed run start = time.perf_counter() result_jax = matmul_jax(x_jax, y_jax).block_until_ready() jax_time = time.perf_counter() - start print(f"JAX time (no JIT): {jax_time:.4f} seconds") __ __