MacBook M3-Pro运行jnp.linalg.qr(A)触发XlaRuntimeError错误
解决JAX在M3-Pro上执行QR分解时的XlaRuntimeError错误
- 核心问题:版本不兼容
你安装的JAX 0.0.5和jaxlib 4.20版本完全不匹配,JAX当前稳定版本为0.4.x系列,jaxlib版本必须与jax主版本严格对应(如jax 0.4.20对应jaxlib 0.4.20),版本不一致会导致底层算子调用失败,这就是你遇到mhlo.custom_call无法合法化的原因。
修复步骤
- 卸载现有不兼容版本
pip uninstall -y jax jaxlib
- 重新安装适配Apple Silicon的匹配版本
直接安装对应版本的JAX和jaxlib:
pip install jax==0.4.20 jaxlib==0.4.20 -f https://storage.googleapis.com/jax-releases/jax_macOS_releases.html
或者安装自动适配的CPU版本(M系列芯片默认支持):
pip install "jax[cpu]" -f https://storage.googleapis.com/jax-releases/jax_macOS_releases.html
- 验证版本一致性
运行以下代码确认版本匹配:
import jax import jax.numpy as jnp print(f"JAX version: {jax.__version__}") print(f"jaxlib version: {jax.lib.xla_extension.__version__}")
输出的两个版本号必须完全相同。
- 重新测试QR分解
执行你的测试代码:
import jax.numpy as jnp A = jnp.array([[1, 2], [3, 4]]) Q, R = jnp.linalg.qr(A) print("Q matrix:\n", Q) print("R matrix:\n", R)
临时 workaround(若版本修复后仍有问题)
强制切换到CPU后端运行,避免M系列芯片后端的算子支持问题:
import jax jax.config.update('jax_platform_name', 'cpu') # 再执行QR分解代码
内容的提问来源于stack exchange,提问作者Benjamin Evans
相关产品推荐
相关产品推荐

