import jax import jax.numpy as jnp # Check available devices print(jax.devices()) # Use GPU for computation device = jax.devices('gpu')[0] jax.jit(function, device=device) # Use TPU for computation device = jax.devices('tpu')[0] jax.jit(function, device=device) __ __