如何为多输出函数实现JAX custom_vjp自定义导数?
JAX多输出函数自定义特定输出的导数
在JAX中,针对多输出函数自定义某一输出(比如你的out_2)对输入的导数,可以通过jax.custom_vjp或jax.custom_jvp实现,核心是在导数规则中聚焦目标输出的梯度计算逻辑。以下是具体实现方案:
方法一:使用jax.custom_vjp(更适合反向传播梯度计算)
custom_vjp需要定义正向传播函数和反向传播函数,反向传播时我们可以仅基于目标输出的梯度权重计算输入的梯度。
示例代码
import jax import jax.numpy as jnp @jax.custom_vjp def test_func(*args, **kwargs): # 替换为你的实际业务逻辑 x, y = args out_1 = x + y out_2 = x * y return out_1, out_2 # 正向传播:返回函数输出+需保存的中间变量(用于反向梯度计算) def test_func_fwd(*args, **kwargs): x, y = args out_1 = x + y out_2 = x * y return (out_1, out_2), (x, y) # 反向传播:接收保存的中间变量、输出的梯度权重,返回输入的梯度 def test_func_bwd(res, cotangents): x, y = res _, cot_out2 = cotangents # 仅取out_2的梯度权重,忽略out_1的 # 自定义out_2对输入的梯度计算逻辑(此处为x*y的梯度) dx = y * cot_out2 dy = x * cot_out2 # 返回输入的梯度,kwargs的梯度返回None return (dx, dy), None # 注册VJP规则 test_func.defvjp(test_func_fwd, test_func_bwd) # 测试:计算out_2对x、y的梯度 x = jnp.array(2.0) y = jnp.array(3.0) grad_fn = jax.grad(lambda *args: test_func(*args)[1], argnums=(0, 1)) dx, dy = grad_fn(x, y) print(f"out_2对x的梯度: {dx}") # 输出3.0 print(f"out_2对y的梯度: {dy}") # 输出2.0
方法二:使用jax.custom_jvp(更适合正向模式自动微分)
custom_jvp定义正向传播的切线计算规则,我们可以针对目标输出单独定义切线逻辑。
示例代码
import jax import jax.numpy as jnp @jax.custom_jvp def test_func(*args, **kwargs): x, y = args out_1 = x + y out_2 = x * y return out_1, out_2 # 定义JVP规则:输入原始值+切线向量,返回输出值+输出切线向量 def test_func_jvp(primals, tangents): x, y = primals tx, ty = tangents # 计算原始输出 out_1 = x + y out_2 = x * y # 自定义out_2的切线(对应导数的正向传播) t_out2 = y * tx + x * ty # out_1的切线使用默认逻辑(也可自定义) t_out1 = tx + ty return (out_1, out_2), (t_out1, t_out2) # 注册JVP规则 test_func.defjvp(test_func_jvp) # 测试:计算out_2对x的导数(给x加切线1,y加切线0) x = jnp.array(2.0) y = jnp.array(3.0) outputs, tangents_out = jax.jvp(test_func, (x, y), (jnp.array(1.0), jnp.array(0.0))) print(f"out_2对x的梯度: {tangents_out[1]}") # 输出3.0
关键说明
- 若仅关心某一输出的导数,在反向传播(VJP)中只需提取对应输出的梯度权重(
cotangents中的对应元素),忽略其他输出即可。 - 使用
jax.grad时,通过lambda *args: test_func(*args)[1]指定对第二个输出求导,结合argnums参数指定要对哪些输入计算梯度。
内容的提问来源于stack exchange,提问作者Jingyang Wang
相关产品推荐
相关产品推荐

