You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

安装JAX、jaxlib与Optax后GPU支持失效问题求助

JAX GPU支持因Optax安装失效的解决办法

问题根源

Optax部分版本的依赖配置会自动安装CPU版的jax/jaxlib,直接覆盖你原本安装的CUDA兼容版本,这就是GPU支持消失的核心原因。

具体解决步骤

  • 安装兼容的Optax版本:针对JAX 0.4.23,Optax 0.1.7是完全匹配的版本,不会触发依赖冲突。执行命令:

    pip install optax==0.1.7
    
  • 修复被覆盖的jaxlib(若已安装过不兼容Optax):先卸载所有相关包,再重新安装CUDA版JAX和兼容的Optax:

    pip uninstall -y jax jaxlib optax
    # 根据你的CUDA版本选择对应jaxlib版本,例如CUDA 11.8用jaxlib==0.4.23+cuda11.cudnn86,CUDA 12.x用jaxlib==0.4.23+cuda12.cudnn89
    pip install jax==0.4.23 jaxlib==0.4.23+cuda11.cudnn86
    pip install optax==0.1.7
    
  • 验证GPU恢复情况:运行以下代码确认:

    from jax.lib import xla_bridge
    print(xla_bridge.get_backend().platform)
    

    输出gpu即表示GPU支持正常。

版本匹配说明

Optax和JAX的版本必须对应,核心是Optax的依赖声明要包含你使用的JAX版本(0.4.23)。旧版本如Optax 0.1.5的依赖范围不包含JAX 0.4.23,会强制安装低版本JAX,从而替换掉CUDA版jaxlib。

内容的提问来源于stack exchange,提问作者Alessandro Castelli

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.16 00:03:13