如何在WSL中禁用JAX的TPU自动调用,改用CPU运行?
解决JAX自动调用TPU、改用CPU运行的问题
以下是几种可行的解决方案,按优先级排序:
1. 用环境变量强制指定CPU后端
这是最快的临时解决方法,不需要重新安装包:
- 终端运行脚本前设置:在启动Python脚本前执行这条命令,会让当前会话的JAX只使用CPU:
export JAX_PLATFORM_NAME=cpu - 代码内硬设置:如果不想每次终端都输命令,直接在Python代码最开头添加:
这个变量会跳过TPU后端的初始化,强制JAX只加载CPU设备。import os os.environ['JAX_PLATFORM_NAME'] = 'cpu' import jax
2. 彻底替换为纯CPU版JAX
如果卸载重装没生效,大概率是之前装的TPU版jaxlib没清干净,或者安装时默认拉了TPU源的包,按以下步骤操作:
- 完全卸载所有JAX相关包:
pip uninstall -y jax jaxlib cloud-tpu-client - 安装官方纯CPU版本:
默认PyPI源的jaxlib就是CPU版本,不会带TPU相关依赖。pip install jax jaxlib
3. 清理WSL中的TPU残留配置
检查WSL的shell配置文件(比如~/.bashrc、~/.zshrc),如果里面有设置过JAX_PLATFORMS、TPU_NAME这类和TPU相关的环境变量,直接删除对应的行,然后重启终端生效。
验证是否成功
运行以下代码,输出如果是[CpuDevice(id=0)]就说明已经切换到CPU运行:
import jax print(jax.devices())
内容的提问来源于stack exchange,提问作者Shereo
相关产品推荐
相关产品推荐

