使用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
相关产品推荐
相关产品推荐

