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

Macbook Pro M3 GPU上jaxopt运行报错:all placeholder ndarrays should have been allocated

在MacBook Pro M3 GPU上运行jaxopt代码报错的问题

问题背景

在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.29 16:59:52