JAX中实现任意变量数特定多变量导数求解的递归错误解决方法
解决JAX中循环生成高阶偏导数的递归错误
问题背景
手动编写代码可针对数组x的三个输入求函数的高阶偏导数:
import jax.numpy as jnp import jax def f(x): return jnp.prod(x) x=jnp.array([1.0,2.0,3.0]) funs = [f] funs.append(lambda x: jax.grad(funs[0])(x)[0]) funs.append(lambda x: jax.grad(funs[1])(x)[1]) funs.append(lambda x: jax.grad(funs[2])(x)[2]) print(funs[3](x))
但尝试用循环推广至任意变量数量时,出现递归错误:
--------------------------------------------------------------------------- RecursionError Traceback (most recent call last) /tmp/ipykernel_331066/1225421420.py in <module> 19 funs.append(z) 20 ---> 21 funs[3](x) /tmp/ipykernel_331066/1225421420.py in <lambda>(x) 16 funs = [f] 17 for i in range(n): ---> 18 z = lambda x: jax.grad(funs[i])(x)[i] 19 funs.append(z) 20 [... skipping hidden 10 frame] ... last 11 frames repeated, from the frame below ... /tmp/ipykernel_331066/1225421420.py in <lambda>(x) 16 funs = [f] 17 for i in range(n): ---> 18 z = lambda x: jax.grad(funs[i])(x)[i] 19 funs.append(z) 20 RecursionError: maximum recursion depth exceeded
错误原因
循环中定义的lambda捕获的是变量i的引用,而非循环迭代时的i值。当循环结束,i的最终值为n-1,所有lambda都会引用这个最终值。调用时,funs[i]指向最后一个lambda,导致无限递归调用,触发RecursionError。
解决方案
方法1:用默认参数绑定当前i值
给lambda添加默认参数i=i,将循环迭代时的i值直接绑定到lambda内部,避免引用循环变量:
import jax.numpy as jnp import jax def f(x): return jnp.prod(x) x = jnp.array([1.0, 2.0, 3.0]) n = 3 # 可替换为任意变量数量 funs = [f] for i in range(n): # 通过默认参数绑定当前循环的i值 z = lambda x, i=i: jax.grad(funs[i])(x)[i] funs.append(z) print(funs[3](x)) # 输出1.0,对应三阶偏导数
方法2:用闭包封装当前i值
编写辅助函数,通过闭包固定每次循环的i值和对应的funs元素:
import jax.numpy as jnp import jax def f(x): return jnp.prod(x) x = jnp.array([1.0, 2.0, 3.0]) n = 3 def make_grad_func(i, funs): # 闭包固定i和funs[i]的引用 return lambda x: jax.grad(funs[i])(x)[i] funs = [f] for i in range(n): z = make_grad_func(i, funs) funs.append(z) print(funs[3](x)) # 输出1.0
内容的提问来源于stack exchange,提问作者Jose Arcadio Buendia
相关产品推荐
相关产品推荐

