Google TPU v3-8 VM中JAX 0.2.12/0.2.16无法识别TPU的求助
解决指定版本JAX在TPU v3-8上的识别问题
核心原因
旧版本JAX(0.2.12/0.2.16)对应的libtpu库版本与当前TPU虚拟机的驱动环境不兼容,导致无法自动识别TPU,需手动匹配对应版本的依赖并配置环境。
针对JAX 0.2.16的解决步骤
- 清理现有冲突依赖
pip uninstall -y jax jaxlib libtpu libtpu-nightly
- 安装指定版本JAX及对应兼容的libtpu
pip install "jax[tpu]==0.2.16" libtpu==0.1.0.dev20210806 -f https://storage.googleapis.com/jax-releases/libtpu_releases.html
- 设置必要的环境变量
export LIBTPU_INIT_ARGS="--xla_jf_max_parallelism=1"
- 验证TPU识别
import jax print(jax.device_count()) # 正常输出8(对应v3-8规格)
针对JAX 0.2.12的解决步骤
- 清理现有冲突依赖
pip uninstall -y jax jaxlib libtpu libtpu-nightly
- 安装匹配版本的JAX、jaxlib和libtpu
pip install jax==0.2.12 jaxlib==0.1.69+tpu libtpu==0.1.0.dev20210602 -f https://storage.googleapis.com/jax-releases/libtpu_releases.html
- 设置环境变量
export LIBTPU_INIT_ARGS="--xla_jf_max_parallelism=1"
- 验证TPU识别
import jax print(jax.device_count())
额外排查步骤
- 确认TPU已正确挂载:执行以下命令,若返回TPU节点IP则挂载正常
curl -H "Metadata-Flavor: Google" http://metadata.google.internal/computeMetadata/v1/instance/attributes/tpu-ip-address
- 若仍无法识别,可在代码中强制指定TPU后端:
import jax jax.config.update('jax_platform_name', 'tpu')
内容的提问来源于stack exchange,提问作者BlackHawk
相关产品推荐
相关产品推荐

