在Jax中计算极坐标转笛卡尔坐标雅可比矩阵得到全零数组,求原因
问题原因及解决方法
你的代码得到全零雅可比矩阵的核心问题是:x和y的计算定义在函数f外部,并未依赖函数的输入参数var。
当你定义x = var[0]*jnp.cos(var[1])和y = var[0]*jnp.sin(var[1])时,这些值是基于初始的var数组计算的固定常量。后续调用f(var)时,函数只是返回这两个预先计算好的常量,而非根据传入的var重新计算。JAX计算雅可比矩阵时,是对常量求导,结果自然全为0。
修正后的代码
将x和y的计算逻辑移到函数f内部,让它们依赖于函数的输入参数:
import jax.numpy as jnp import numpy as np theta = np.pi/4 r = 4.0 var = np.array([r, theta]) def f(var): x = var[0] * jnp.cos(var[1]) y = var[0] * jnp.sin(var[1]) return jnp.array([x, y]) jac = jax.jacobian(f)(var) print(jac)
预期输出
运行修正后的代码,你会得到正确的雅可比矩阵:
DeviceArray([[ 0.70710677, -2.8284271 ], [ 0.70710677, 2.8284271 ]], dtype=float32)
内容的提问来源于stack exchange,提问作者David
相关产品推荐
相关产品推荐

