JAX中高阶函数的自定义JVP与VJP实现方法问询
JAX高阶函数中自定义梯度的实现方法
要给你示例中高阶函数返回的child_func定义关于x和y的自定义梯度,需要结合JAX的custom_vjp(或custom_jvp)工具,同时注意闭包变量的梯度追踪问题。以下是两种可行的实现方式:
方式一:将x作为child_func的显式参数(推荐)
这种方式更清晰,避免闭包变量的梯度追踪限制,直接给包含x和y的child_func定义自定义梯度:
import jax import jax.numpy as jnp def parent_func(): @jax.custom_vjp def child_func(x, y): return x**2 * y # 前向传递:返回计算结果和反向需要的中间变量 def fwd(x, y): out = child_func(x, y) return out, (x, y) # 反向传递:自定义x和y的梯度 def bwd(res, cot): x_val, y_val = res # 自定义x的梯度:d(out)/dx = 2*x*y grad_x = cot * 2 * x_val * y_val # 自定义y的梯度:d(out)/dy = x² grad_y = cot * x_val ** 2 return (grad_x, grad_y) # 绑定VJP逻辑到child_func child_func.defvjp(fwd, bwd) return child_func # 使用示例 child = parent_func() x = jnp.array(2.0) y = jnp.array(3.0) # 计算x和y的自定义梯度 grad_x, grad_y = jax.grad(child, argnums=(0, 1))(x, y) print(f"自定义x梯度: {grad_x}, 自定义y梯度: {grad_y}")
方式二:保持child_func仅接收y(闭包捕获x)
如果要严格保留原高阶函数的结构(child_func只接收y),可以通过绑定二元自定义梯度函数的方式实现:
import jax import jax.numpy as jnp def parent_func(x): # 定义包含x和y的二元函数,并添加自定义VJP @jax.custom_vjp def combined_func(x_inner, y_inner): return x_inner**2 * y_inner def combined_fwd(x_inner, y_inner): return combined_func(x_inner, y_inner), (x_inner, y_inner) def combined_bwd(res, cot): x_in, y_in = res grad_x = cot * 2 * x_in * y_in grad_y = cot * x_in**2 return (grad_x, grad_y) combined_func.defvjp(combined_fwd, combined_bwd) # 使用jax.partial绑定x,返回仅接收y的child_func return jax.partial(combined_func, x) # 使用示例 child = parent_func(jnp.array(2.0)) y = jnp.array(3.0) # 计算y的自定义梯度 grad_y = jax.grad(child)(y) print(f"自定义y梯度: {grad_y}") # 计算x的自定义梯度 grad_x = jax.grad(lambda x_val: parent_func(x_val)(y))(jnp.array(2.0)) print(f"自定义x梯度: {grad_x}")
关键说明
jax.custom_vjp需要定义前向(fwd)和反向(bwd)函数:前向函数保存反向计算需要的中间变量,反向函数根据输出的余切值计算各输入参数的梯度。- JAX默认不会追踪闭包变量的梯度,因此方式二中通过将
x作为显式参数传入二元函数,再用jax.partial绑定的方式,既保留原函数结构,又能实现自定义梯度。
内容的提问来源于stack exchange,提问作者Jingyang Wang
相关产品推荐
相关产品推荐

