CUDA与JAX版本不兼容引发PTX版本错误求助
问题:CUDA与JAX/PTX版本不兼容导致运行报错
环境信息
- CUDA版本:11.6
- JAX版本:0.4.16
- jaxlib版本:0.4.16+cuda11.cudnn86
核心报错信息
第一个错误:
W external/xla/xla/service/gpu/buffer_comparator.cc:1054] INTERNAL: ptxas exited with non-zero error code 65280, output: ptxas /tmp/tempfile-meiji-993e158e-113566-608cf99230264, line 10; fatal : Unsupported .version 7.8; current version is '7.6'
第二个错误:
jaxlib.xla_extension.XlaRuntimeError: INTERNAL: Failed to load PTX text as a module: CUDA_ERROR_UNSUPPORTED_PTX_VERSION: the provided PTX was compiled with an unsupported toolchain.
详细错误日志
WARNING: All log messages before absl::InitializeLog() is called are written to STDERR I0000 00:00:1698537556.277868 113566 tfrt_cpu_pjrt_client.cc:349] TfrtCpuClient created. random key: [0 0] 2023-10-28 19:59:41.023080: W external/xla/xla/service/gpu/buffer_comparator.cc:1054] INTERNAL: ptxas exited with non-zero error code 65280, output: ptxas /tmp/tempfile-meiji-993e158e-113566-608cf99230264, line 10; fatal : Unsupported .version 7.8; current version is '7.6' ptxas fatal : Ptx assembly aborted due to errors Relying on driver to perform ptx compilation. Setting XLA_FLAGS=--xla_gpu_cuda_data_dir=/path/to/cuda or modifying $PATH can be used to set the location of ptxas This message will only be logged once. 2023-10-28 19:59:41.077542: E external/xla/xla/stream_executor/cuda/cuda_driver.cc:857] failed to load PTX text as a module: CUDA_ERROR_UNSUPPORTED_PTX_VERSION: the provided PTX was compiled with an unsupported toolchain. 2023-10-28 19:59:41.077572: E external/xla/xla/stream_executor/cuda/cuda_driver.cc:862] error log buffer (98 bytes): ptxas application ptx input, line 10; fatal : Unsupported .version 7.8; current version is '7.6 jax.errors.SimplifiedTraceback: For simplicity, JAX has removed its internal frames from the traceback of the following exception. Set JAX_TRACEBACK_FILTERING=off to include these. The above exception was the direct cause of the following exception: Traceback (most recent call last): File "/home/shuting/nucleotide_transformer/nucleotide_transformer_test.py", line 76, in <module> outs = forward_fn.apply(parameters, random_key, tokens) File "/home/shuting/anaconda3/envs/transformer5/lib/python3.9/site-packages/haiku/_src/transform.py", line 183, in apply_fn out, state = f.apply(params, None, *args, **kwargs) File "/home/shuting/anaconda3/envs/transformer5/lib/python3.9/site-packages/haiku/_src/transform.py", line 456, in apply_fn out = f(*args, **kwargs) File "/home/shuting/anaconda3/envs/transformer5/lib/python3.9/site-packages/nucleotide_transformer/model.py", line 363, in nucleotide_transformer_fn outs = encoder( File "/home/shuting/anaconda3/envs/transformer5/lib/python3.9/site-packages/haiku/_src/module.py", line 458, in wrapped out = f(*args, **kwargs) File "/home/shuting/anaconda3/envs/transformer5/lib/python3.9/contextlib.py", line 79, in inner return func(*args, **kwds) File "/home/shuting/anaconda3/envs/transformer5/lib/python3.9/site-packages/haiku/_src/module.py", line 299, in run_interceptors return bound_method(*args, **kwargs) File "/home/shuting/anaconda3/envs/transformer5/lib/python3.9/site-packages/nucleotide_transformer/model.py", line 310, in __call__ x, outs = self.apply_attention_blocks( File "/home/shuting/anaconda3/envs/transformer5/lib/python3.9/site-packages/nucleotide_transformer/model.py", line 226, in apply_attention_blocks output = layer( File "/home/shuting/anaconda3/envs/transformer5/lib/python3.9/site-packages/haiku/_src/module.py", line 458, in wrapped out = f(*args, **kwargs) File "/home/shuting/anaconda3/envs/transformer5/lib/python3.9/contextlib.py", line 79, in inner return func(*args, **kwds) File "/home/shuting/anaconda3/envs/transformer5/lib/python3.9/site-packages/haiku/_src/module.py", line 299, in run_interceptors return bound_method(*args, **kwargs) File "/home/shuting/anaconda3/envs/transformer5/lib/python3.9/site-packages/nucleotide_transformer/layers.py", line 281, in __call__ output = self.self_attention( File "/home/shuting/anaconda3/envs/transformer5/lib/python3.9/site-packages/nucleotide_transformer/layers.py", line 240, in self_attention return self.sa_layer(x, x, x, attention_mask=attention_mask) File "/home/shuting/anaconda3/envs/transformer5/lib/python3.9/site-packages/haiku/_src/module.py", line 458, in wrapped out = f(*args, **kwargs) File "/home/shuting/anaconda3/envs/transformer5/lib/python3.9/contextlib.py", line 79, in inner return func(*args, **kwds) File "/home/shuting/anaconda3/envs/transformer5/lib/python3.9/site-packages/haiku/_src/module.py", line 299, in run_interceptors return bound_method(*args, **kwargs) File "/home/shuting/anaconda3/envs/transformer5/lib/python3.9/site-packages/nucleotide_transformer/layers.py", line 149, in __call__ attention_weights = self.attention_weights(query, key, attention_mask) File "/home/shuting/anaconda3/envs/transformer5/lib/python3.9/site-packages/nucleotide_transformer/layers.py", line 83, in attention_weights query_heads = self._linear_projection_he_init(query, self.key_size, "query") File "/home/shuting/anaconda3/envs/transformer5/lib/python3.9/site-packages/nucleotide_transformer/layers.py", line 173, in _linear_projection_he_init y = hk.Linear( File "/home/shuting/anaconda3/envs/transformer5/lib/python3.9/site-packages/haiku/_src/module.py", line 458, in wrapped out = f(*args, **kwargs) File "/home/shuting/anaconda3/envs/transformer5/lib/python3.9/contextlib.py", line 79, in inner return func(*args, **kwds) File "/home/shuting/anaconda3/envs/transformer5/lib/python3.9/site-packages/haiku/_src/module.py", line 299, in run_interceptors return bound_method(*args, **kwargs) File "/home/shuting/anaconda3/envs/transformer5/lib/python3.9/site-packages/haiku/_src/basic.py", line 181, in __call__ out = jnp.dot(inputs, w, precision=precision) jaxlib.xla_extension.XlaRuntimeError: INTERNAL: Failed to load PTX text as a module: CUDA_ERROR_UNSUPPORTED_PTX_VERSION: the provided PTX was compiled with an unsupported toolchain. I0000 00:00:1698537581.636191 113566 tfrt_cpu_pjrt_client.cc:352] TfrtCpuClient destroyed.
解决方案
原因分析
报错提示PTX版本7.8不被当前环境支持,当前仅支持7.6。根源在于:
- 安装的jaxlib版本编译时使用的PTX版本,高于系统中
ptxas工具(CUDA 11.6自带)支持的版本 - CUDA 11.6对应的PTX版本为7.6,而jaxlib 0.4.16默认可能采用了更高版本的PTX(对应CUDA 11.7+的7.8)
修复步骤
指定XLA使用正确的CUDA路径
运行代码前设置环境变量,让XLA定位到CUDA 11.6的ptxas工具:export XLA_FLAGS=--xla_gpu_cuda_data_dir=/usr/local/cuda-11.6(替换为你实际的CUDA 11.6安装路径)
降级jaxlib到适配CUDA 11.6的版本
若上述方法无效,卸载当前jaxlib并安装完全兼容CUDA 11.6的版本:pip uninstall jaxlib -y pip install jaxlib==0.4.16+cuda11.cudnn86若仍有问题,可尝试jaxlib 0.4.15及以下版本,这类版本通常对PTX 7.6兼容性更好。
升级CUDA环境
若允许调整环境,可将CUDA升级到11.7或更高版本(对应PTX 7.8),同时确保显卡驱动版本同步兼容。
内容的提问来源于stack exchange,提问作者Rowan
相关产品推荐
相关产品推荐

