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),按以下步骤修复:
确认系统实际加载的CuDNN版本
执行命令查看当前优先加载的CuDNN库:find /usr/local/cuda -name "libcudnn*.so" | head -n 1 | xargs ls -l检查环境变量是否包含旧版本路径:
echo $LD_LIBRARY_PATH清理旧版本并设置正确路径优先级
- 删除系统中残留的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生效。
- 删除系统中残留的CuDNN 8.5.0文件(通常在
检查conda环境依赖
查看环境中是否存在旧版本cudnn包:conda list | grep cudnn卸载旧版本并安装匹配版本:
conda remove cudnn conda install cudnn=8.6.0=cuda11.7_0验证修复结果
运行以下代码确认版本匹配: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
相关产品推荐
相关产品推荐

