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

CuDNN版本不匹配引发XlaRuntimeError错误求助

问题描述

2023-07-31 01:53:45.016563: E external/xla/xla/stream_executor/cuda/cuda_dnn.cc:427] 运行时加载的CuDNN库版本为8.5.0,但源码编译时使用的版本为8.6.0。CuDNN库需要主版本匹配且次版本大于等于编译版本。如果是二进制安装,请升级CuDNN库;如果是从源码构建,请确保运行时加载的库与编译配置指定的版本兼容。

错误回溯

XlaRuntimeError 回溯(最近的调用最后)
Cell In[4], 第29行

26 model = trainer.make_model(nmask)
28 lr_fn, opt = trainer.make_optimizer(steps_per_epoch=len(train_dl))
---> 29 state = trainer.create_train_state(jax.random.PRNGKey(0), model, opt)
30 state = checkpoints.restore_checkpoint(ckpt.parent, state)

文件 /mnt/data/miniconda/envs/energy_transformer_117/lib/python3.11/site-packages/jax/_src/random.py:137, 在 PRNGKey(seed)中

134 if np.ndim(seed):
135 raise TypeError("PRNGKey accepts a scalar seed, but was given an array of"
136 f"shape {np.shape(seed)} != (). Use jax.vmap for batching")
---> 137 key = prng.seed_with_impl(impl, seed)
138 return _return_prng_keys(True, key)

文件 /mnt/data/miniconda/envs/energy_transformer_117/lib/python3.11/site-packages/jax/_src/prng.py:320, 在 seed_with_impl(impl, seed)中

319 def seed_with_impl(impl: PRNGImpl, seed: Union[int, Array]) -> PRNGKeyArrayImpl:
---> 320 return random_seed(seed, impl=impl)

文件 /mnt/data/miniconda/envs/energy_transformer_117/lib/python3.11/site-packages/jax/_src/prng.py:734, 在 random_seed(seeds, impl)中

732 else:
733 seeds_arr = jnp.asarray(seeds)
---> 734 return random_seed_p.bind(seeds_arr, impl=impl)

文件 /mnt/data/miniconda/envs/energy_transformer_117/lib/python3.11/site-packages/jax/_src/core.py:380, 在 Primitive.bind(self, *args, **params)中

377 def bind(self, *args, **params):
378 assert (not config.jax_enable_checks or
379 all(isinstance(arg, Tracer) or valid_jaxtype(arg) for arg in args)), args
---> 380 return self.bind_with_trace(find_top_trace(args), args, params)

文件 /mnt/data/miniconda/envs/energy_transformer_117/lib/python3.11/site-packages/jax/_src/core.py:383, 在 Primitive.bind_with_trace(self, trace, args, params)中

382 def bind_with_trace(self, trace, args, params):
---> 383 out = trace.process_primitive(self, map(trace.full_raise, args), params)
384 return map(full_lower, out) if self.multiple_results else full_lower(out)

文件 /mnt/data/miniconda/envs/energy_transformer_117/lib/python3.11/site-packages/jax/_src/core.py:790, 在 EvalTrace.process_primitive(self, primitive, tracers, params)中

789 def process_primitive(self, primitive, tracers, params):
---> 790 return primitive.impl(*tracers, **params)

文件 /mnt/data/miniconda/envs/energy_transformer_117/lib/python3.11/site-packages/jax/_src/prng.py:746, 在 random_seed_impl(seeds, impl)中

744 @random_seed_p.def_impl
745 def random_seed_impl(seeds, *, impl):
---> 746 base_arr = random_seed_impl_base(seeds, impl=impl)
747 return PRNGKeyArrayImpl(impl, base_arr)

文件 /mnt/data/miniconda/envs/energy_transformer_117/lib/python3.11/site-packages/jax/_src/prng.py:751, 在 random_seed_impl_base(seeds, impl)中

749 def random_seed_impl_base(seeds, *, impl):
750 seed = iterated_vmap_unary(seeds.ndim, impl.seed)
---> 751 return seed(seeds)

文件 /mnt/data/miniconda/envs/energy_transformer_117/lib/python3.11/site-packages/jax/_src/prng.py:980, 在 threefry_seed(seed)中

968 def threefry_seed(seed: typing.Array) -> typing.Array:
969 """Create a single raw threefry PRNG key from an integer seed.
970
971 Args:
(...)
978 first padding out with zeros).
979 """
---> 980 return _threefry_seed(seed)

[... 跳过隐藏的12帧]
文件 /mnt/data/miniconda/envs/energy_transformer_117/lib/python3.11/site-packages/jax/_src/dispatch.py:463, 在 backend_compile(backend, module, options, host_callbacks)中

458 return backend.compile(built_c, compile_options=options,
459 host_callbacks=host_callbacks)
460 # Some backends don't have host_callbacks option yet
461 # TODO(sharadmv): remove this fallback when all backends allow compile
462 # to take in host_callbacks
---> 463 return backend.compile(built_c, compile_options=options)

XlaRuntimeError: FAILED_PRECONDITION: DNN库初始化失败。查看上方错误获取更多详情。

环境信息

  • jax版本:0.4.10
  • jaxlib版本:0.4.10+cuda11.cudnn86
  • 加速器:GPU
  • 系统信息:Python 3.11.4、Ubuntu 22.04、CUDA 11.7、CuDNN 8.6(用户声明安装版本)

修复方案

核心问题是运行时加载的CuDNN版本(8.5.0)低于jaxlib编译依赖的版本(8.6.0),按以下步骤修复:

  1. 确认系统实际加载的CuDNN版本
    执行命令查看当前优先加载的CuDNN库:

    find /usr/local/cuda -name "libcudnn*.so" | head -n 1 | xargs ls -l
    

    检查环境变量是否包含旧版本路径:

    echo $LD_LIBRARY_PATH
    
  2. 清理旧版本并设置正确路径优先级

    • 删除系统中残留的CuDNN 8.5.0文件(通常在/usr/local/cuda/include和/usr/local/cuda/lib64目录下)。
    • 将CuDNN 8.6的路径添加到LD_LIBRARY_PATH最前端:
      export LD_LIBRARY_PATH=/usr/local/cuda-11.7/lib64:$LD_LIBRARY_PATH
      
      若要永久生效,将上述命令写入~/.bashrc或~/.zshrc,再执行source ~/.bashrc生效。
  3. 检查conda环境依赖
    查看环境中是否存在旧版本cudnn包:

    conda list | grep cudnn
    

    卸载旧版本并安装匹配版本:

    conda remove cudnn
    conda install cudnn=8.6.0=cuda11.7_0
    
  4. 验证修复结果
    运行以下代码确认版本匹配:

    import jax
    print(jax.lib.xla_bridge.get_backend().platform)
    print(jax.lib.xla_bridge.get_backend().cuda_version)
    print(jax.lib.xla_bridge.get_backend().cudnn_version)
    

    确保输出的cudnn_version为8.6.0及以上。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.14 10:34:58