如何使用Autograd对嵌套定义的函数求解各变量的偏导数?
实现方法
你定义的是柯里化结构的嵌套函数,用autograd计算指定点位偏导可以按以下步骤操作:
1. 依赖准备
首先确保安装了autograd,未安装可以执行命令:pip install autograd,之后导入梯度计算工具:
from autograd import grad
2. 保留原有函数定义
不需要修改你写的嵌套函数逻辑:
def f(x, y): def g(w): def h(z): return z * (w ** 2) * (x + y) return h return g
3. 封装平层函数(推荐)
为了方便对不同变量求偏导,我们可以把嵌套的柯里化函数,封装为同时接收4个参数的普通函数:
def f_flat(x, y, w, z): return f(x, y)(w)(z)
4. 定义各变量偏导函数
通过grad的argnum参数指定要对第几个参数求导(参数索引从0开始计数):
df_dx = grad(f_flat, argnum=0) # 对x求偏导 df_dy = grad(f_flat, argnum=1) # 对y求偏导 df_dw = grad(f_flat, argnum=2) # 对w求偏导 df_dz = grad(f_flat, argnum=3) # 对z求偏导
5. 代入点位计算
注意autograd要求输入为浮点型,不要传入整数类型:
# 替换为你需要的x0、y0、w0、z0即可 x0, y0, w0, z0 = 1.0, 2.0, 3.0, 4.0 print("对x的偏导:", df_dx(x0, y0, w0, z0)) print("对y的偏导:", df_dy(x0, y0, w0, z0)) print("对w的偏导:", df_dw(x0, y0, w0, z0)) print("对z的偏导:", df_dz(x0, y0, w0, z0))
结果验证
你可以用理论推导的偏导公式核对结果:
- ∂f/∂x = z * w²
- ∂f/∂y = z * w²
- ∂f/∂w = 2 * z * w * (x + y)
- ∂f/∂z = w² * (x + y)
上面的示例点位计算输出依次为36.0、36.0、72.0、27.0,和理论值完全一致。
替代方案(无需封装平层函数)
如果你不想修改原有函数结构,也可以针对单个变量单独构造单参数函数求导,以x为例:
x0, y0, w0, z0 = 1.0, 2.0, 3.0, 4.0 # 固定y、w、z,仅保留x为变量 def fx(x): return f(x, y0)(w0)(z0) print("对x的偏导:", grad(fx)(x0))
其余变量的偏导按相同逻辑调整固定参数即可。
内容的提问来源于stack exchange,提问作者Daniel B.
相关产品推荐
相关产品推荐

