Ubuntu 18.04+CUDA12.4环境JAX GPU版本安装失败求助
求适配以下服务器规格的GPU版JAX可靠安装命令
系统规格
- 操作系统:Ubuntu 18.04 LTS
- GPU:8×NVIDIA Quadro RTX 8000
- NVIDIA驱动:550.144.03
- 驱动支持的CUDA版本:12.4
- Python:3.10(由Conda管理)
已尝试的方法
每次尝试均创建全新conda环境,确保无依赖冲突。
尝试1:官方推荐标准方法
使用官方提供的命令:
conda create -n jax_test python=3.10 -y conda activate jax_test pip install "jax[cuda12_pip]" -f https://storage.googleapis.com/jax-releases/jax_cuda_releases.html --no-cache-dir
- 预期结果:安装包含CUDA库的大体积
jaxlibwheel包(数GB) - 实际结果:
pip忽略[cuda12_pip]指令,下载CPU版本jaxlib(89.9MB)并告警,验证失败:
WARNING: jax 0.6.2 does not provide the extra 'cuda12-pip' Downloading jaxlib-0.6.2-cp310-cp310-manylinux2014_x86_64.whl (89.9 MB) ... $ python -c "import jax; print(jax.devices())" WARNING: An NVIDIA GPU may be present on this machine, but a CUDA-enabled jaxlib is not installed. Falling back to cpu. [CpuDevice(id=0)]
尝试2:直接URL安装法
强制安装特定GPU wheel包:
# 在干净环境中执行 pip install "https://storage.googleapis.com/jax-releases/cuda12/jaxlib-0.4.23+cuda12.cudnn88-cp310-cp310-manylinux2014_x86_64.whl" pip install jax==0.4.23 "numpy<2.0"
- 预期结果:安装指定GPU版本
jaxlib - 实际结果:URL失效,谷歌已移除旧文件,安装失败:
ERROR: HTTP error 404 while getting https://.../jaxlib-0.4.23...whl
依赖旧URL并非稳定方案。
尝试3:分阶段插件安装法
先安装CUDA插件再安装JAX:
# 在干净环境中执行 pip install --upgrade "jax-cuda12-plugin" pip install jax
- 预期结果:JAX使用插件提供的
jaxlib - 实际结果:安装
jax时会覆盖插件库,重新安装CPU版本jaxlib,最终回退到CPU模式
问题求助
目前陷入困境:标准安装无法选择GPU包、URL安装文件丢失、插件安装引发依赖冲突。求适配上述服务器规格的明确可行的JAX GPU版本安装命令?
内容的提问来源于stack exchange,提问作者PowerPoint Trenton
相关产品推荐
相关产品推荐

