import jax if __name__ == "__main__": print(jax.devices())