在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
清理旧环境并重新创建:
deactivate rm -rf ~/jax-metal python3 -m venv ~/jax-metal source ~/jax-metal/bin/activate安装指定兼容版本(示例组合: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验证安装:
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
相关产品推荐
相关产品推荐

