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

在Metal上运行JAX报错:PJRT API版本不匹配求助

问题:Metal版JAX安装后验证失败,出现PJRT版本不匹配错误

问题背景

按照官方步骤在虚拟环境安装适用于Metal的JAX,执行验证命令时触发版本不匹配错误。

安装步骤

python3 -m venv ~/jax-metal
source ~/jax-metal/bin/activate
python -m pip install -U pip
python -m pip install numpy wheel
python -m pip install jax-metal

验证命令

python -c 'import jax; print(jax.numpy.arange(10))'

报错信息

Platform 'METAL' is experimental and not all JAX functionality may be correctly supported!
Traceback (most recent call last):
  File "/Users/ridhibandaru/jax-metal/lib/python3.11/site-packages/jax/_src/xla_bridge.py", line 887, in backends
    backend = _init_backend(platform)
              ^^^^^^^^^^^^^^^^^^^^^^^
  File "/Users/ridhibandaru/jax-metal/lib/python3.11/site-packages/jax/_src/xla_bridge.py", line 973, in _init_backend
    backend = registration.factory()
              ^^^^^^^^^^^^^^^^^^^^^^
  File "/Users/ridhibandaru/jax-metal/lib/python3.11/site-packages/jax/_src/xla_bridge.py", line 667, in factory
    xla_client.initialize_pjrt_plugin(plugin_name)
  File "/Users/ridhibandaru/jax-metal/lib/python3.11/site-packages/jaxlib/xla_client.py", line 177, in initialize_pjrt_plugin
    _xla.initialize_pjrt_plugin(plugin_name)
jaxlib.xla_extension.XlaRuntimeError: INVALID_ARGUMENT: Mismatched PJRT plugin PJRT API version (0.47) and framework PJRT API version 0.54).

During handling of the above exception, another exception occurred:

Traceback (most recent call last):
  File "<string>", line 1, in <module>
  File "/Users/ridhibandaru/jax-metal/lib/python3.11/site-packages/jax/_src/numpy/lax_numpy.py", line 4264, in arange
    output = _arange(start, stop=stop, step=step, dtype=dtype)
             ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
  File "/Users/ridhibandaru/jax-metal/lib/python3.11/site-packages/jax/_src/numpy/lax_numpy.py", line 4304, in _arange
    return lax.iota(dtype, start)  # type: ignore[arg-type]
           ^^^^^^^^^^^^^^^^^^^^^^
  File "/Users/ridhibandaru/jax-metal/lib/python3.11/site-packages/jax/_src/lax/lax.py", line 1338, in iota
    return broadcasted_iota(dtype, (size,), 0)
           ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
  File "/Users/ridhibandaru/jax-metal/lib/python3.11/site-packages/jax/_src/lax/lax.py", line 1348, in broadcasted_iota
    return iota_p.bind(*dynamic_shape, dtype=dtype, shape=tuple(static_shape),
           ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
  File "/Users/ridhibandaru/jax-metal/lib/python3.11/site-packages/jax/_src/core.py", line 429, in bind
    return self.bind_with_trace(find_top_trace(args), args, params)
           ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
  File "/Users/ridhibandaru/jax-metal/lib/python3.11/site-packages/jax/_src/core.py", line 433, in bind_with_trace
    out = trace.process_primitive(self, map(trace.full_raise, args), params)
          ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
  File "/Users/ridhibandaru/jax-metal/lib/python3.11/site-packages/jax/_src/core.py", line 939, in process_primitive
    return primitive.impl(*tracers, **params)
           ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
  File "/Users/ridhibandaru/jax-metal/lib/python3.11/site-packages/jax/_src/dispatch.py", line 87, in apply_primitive
    outs = fun(*args)
           ^^^^^^^^^^
RuntimeError: Unable to initialize backend 'METAL': INVALID_ARGUMENT: Mismatched PJRT plugin PJRT API version (0.47) and framework PJRT API version 0.54). (you may need to uninstall the failing plugin package, or set JAX_PLATFORMS=cpu to skip this backend.)
--------------------
For simplicity, JAX has removed its internal frames from the traceback of the following exception. Set JAX_TRACEBACK_FILTERING=off to include these.

错误原因

核心问题是PJRT API版本不兼容:jax-metal插件依赖的PJRT API版本(0.47)与JAX框架所依赖的jaxlib的PJRT API版本(0.54)不一致,导致Metal后端初始化失败。直接执行pip install jax-metal会自动拉取最新版jax/jaxlib,但jax-metal的更新速度跟不上jax主版本,出现版本错位。

解决方法

方法一:安装版本匹配的jax、jaxlib和jax-metal

  1. 清理旧环境并重新创建:

    deactivate
    rm -rf ~/jax-metal
    python3 -m venv ~/jax-metal
    source ~/jax-metal/bin/activate
    
  2. 安装指定兼容版本(示例组合:jax-metal 0.0.7 对应 jax/jaxlib 0.4.16,可根据官方最新说明调整版本号):

    python -m pip install -U pip
    python -m pip install numpy wheel
    python -m pip install jax==0.4.16 jaxlib==0.4.16
    python -m pip install jax-metal==0.0.7
    
  3. 验证安装:

    python -c 'import jax; print(jax.numpy.arange(10)); print("Metal后端已初始化:", jax.default_backend() == "metal")'
    

方法二:临时使用CPU后端(无需修改版本)

如果只是想临时运行代码,可通过环境变量强制JAX使用CPU后端:

export JAX_PLATFORMS=cpu

之后再执行验证命令即可正常运行。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.18 23:00:02