You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何为多输出函数实现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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.30 19:50:23