Macbook Pro M3 GPU上jaxopt运行报错:all placeholder ndarrays should have been allocated
问题背景
在MacBook Pro M3 GPU环境中,optax相关代码可正常运行,但jaxopt无法正常执行。测试以下简单代码:
from jax import numpy as jnp from jaxopt import LBFGS def loss(x): return x**2 opt = LBFGS(loss) print(opt.run(jnp.array(3.0)))
运行后出现如下错误:
/AppleInternal/Library/BuildRoots/0032d1ee-80fd-11ee-8227-6aecfccc70fe/Library/Caches/com.apple.xbs/Sources/MetalPerformanceShadersGraph/mpsgraph/MetalPerformanceShadersGraph/Runtimes/MPSRuntime/MPSRuntime.mm:1422:
failed assertion `MPSGraphKernelDAG: all placeholder ndarrays should have been allocated'
疑问:是操作有误,还是jax-metal的bug?毕竟jax-metal明确提示“JAX Apple GPU support is experimental and not all JAX functionality is correctly supported!”
解答
这大概率是jax-metal的实验性限制导致的bug,并非你的操作问题。
jax-metal对Apple GPU的支持仍处于实验阶段,官方已明确说明并非所有JAX生态的功能都能正常兼容。jaxopt的LBFGS实现可能用到了jax-metal当前未完全支持的底层操作,进而触发了MetalPerformanceShadersGraph的断言错误。
临时解决方案
- 强制使用CPU运行jaxopt代码,可通过配置JAX实现:
import jax jax.config.update('jax_platform_name', 'cpu') - 等待jax-metal后续版本更新,完善对JAX生态工具的支持。
内容的提问来源于stack exchange,提问作者mwlon

