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

使用Jax优化变量求函数最大值遇类型转换错误求助

解决Jax结合Scipy L-BFGS-B优化时的数组类型错误

错误根源

Scipy的L-BFGS-B优化器要求梯度返回普通NumPy数组,但Jax的jax.grad返回的是Jax DeviceArray,类型不兼容导致转换失败。

修复步骤

  • 包装梯度函数转换数组类型:将梯度输出从Jax数组转为NumPy数组,确保Scipy能识别。
  • 简化目标函数冗余操作:目标函数中jnp.array(x)这类转换是多余的,xy切片后本身就是Jax数组,无需重复转换。
  • 统一初始变量类型:初始变量直接用NumPy数组即可,Jax会自动兼容处理,避免类型混淆。

修复后完整代码

import jax.numpy as jnp
import jax 
import scipy
import numpy as np

def temp_func(x,y,z):
    tmp = x + jnp.dot(jnp.power(y, 3), jnp.tanh(z))
    return -tmp

def obj_func(xy, z):
    x, y = xy[:2], xy[2:].reshape(2,2)
    return jnp.sum(temp_func(x, y, z))

# 包装梯度函数,将输出转为NumPy数组
def grad_func(xy, z):
    return jax.grad(obj_func, argnums=0)(xy, z).numpy()

# 初始变量用NumPy数组,Jax自动兼容
xy = np.concatenate([np.random.rand(2), np.random.rand(2*2)])
z = np.random.rand(2,2)

print(obj_func(jnp.array(xy), z))

result = scipy.optimize.minimize(
    obj_func,
    xy,
    args=(z,),
    method='L-BFGS-B',
    jac=grad_func
)

补充说明:你通过返回-tmp将最大化问题转为最小化问题的逻辑是正确的;若后续使用GPU上的Jax数组,转换为NumPy数组时会自动同步到CPU,Scipy仅支持处理CPU数组。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.01 23:47:20