# Matrix size size = 3000 # Create random matrices with NumPy x_np = np.random.normal(size=(size, size)).astype(np.float32) y_np = np.random.normal(size=(size, size)).astype(np.float32) # Convert to JAX arrays x_jax = jnp.array(x_np) y_jax = jnp.array(y_np) print(f"Matrix shape: {x_np.shape}") print(f"Total elements per matrix: {size * size:,}") print(f"Operations for multiplication: {size ** 3:,}") __ __