import jax import jax.numpy as jnp import numpy as np import time print(f"JAX version: {jax.__version__}") print(f"Available devices: {jax.devices()}") __ __