如何通过requirements.txt安装支持CUDA的JAX版本?
解决JAX GPU版本在requirements.txt中无法正确安装的问题
方法1:使用--find-links规范格式
requirements.txt中不能直接混用命令行参数和包名,需将源地址单独成行声明:
--find-links https://storage.googleapis.com/jax-releases/jax_cuda_releases.html jax[cuda]
执行安装时带上升级参数即可:
pip install -r requirements.txt --upgrade
这种写法会让pip优先从指定源拉取CUDA兼容的JAX包。
方法2:直接锁定GPU版本号
先通过命令查看当前已安装的GPU版JAX和jaxlib版本:
pip freeze | grep -E "jax|jaxlib"
比如输出可能是:
jax==0.4.23 jaxlib==0.4.23+cuda12.cudnn89
将这两行直接写入requirements.txt:
jax==0.4.23 jaxlib==0.4.23+cuda12.cudnn89
这样pip会精确安装指定的GPU兼容版本,完全避免拉取CPU版。
方法3:用环境变量强制绑定CUDA后端
安装前设置环境变量,让JAX优先启用CUDA支持:
export JAX_PLATFORM_NAME=cuda pip install -r requirements.txt
配合requirements.txt的正确源配置,可确保安装GPU版本。
注意:不要在requirements.txt中加入--pre参数,除非你明确需要预发布版本,否则可能导致拉取未适配CUDA的包。
内容的提问来源于stack exchange,提问作者MathiesW
相关产品推荐
相关产品推荐

