首页
/ JAX 自定义求导规则实战:用 `jax.custom_jvp` 与 `jax.custom_vjp` 掌控 JVP/VJP

JAX 自定义求导规则实战:用 `jax.custom_jvp` 与 `jax.custom_vjp` 掌控 JVP/VJP

2026-09-09 15:58:35作者:苗圣禹Peter

导读

在 JAX 中,jax.gradjax.jvpjax.vjp 等变换会沿着 jax.numpyjax.lax 原语的默认微分规则自动求导,但默认规则并不总能满足数值稳定性、特定定义域约定或工程需求(如梯度裁剪)。本文以仓库文档 Custom_derivative_rules_for_Python_code.md 为主体,系统讲解 jax.custom_jvpjax.custom_vjp 两套 API:先给出可复制的 TL;DR 示例,再深入 5 个真实问题场景(数值稳定性、求导约定、梯度裁剪、反向传播调试、迭代算法的隐函数微分),随后逐项解析 API 签名、nondiff_argnums、pytree 支持等细节,最后结合仓库源码剖析其底层 core.Primitive 实现。读完本文,你将能够为任意 JAX 可变换函数定制求导行为,同时保留其 jitvmap 等其余变换能力。

一、JAX 中定义求导规则的两种方式

JAX 提供了两条定义微分规则的路径:

  1. 使用 jax.custom_jvpjax.custom_vjp,为本身已经可被 JAX 变换的 Python 函数定制求导规则——这是本文的主题;
  2. 定义全新的 core.Primitive 实例并为其编写全部变换规则,用于对接求解器、模拟器等外部系统(可参考仓库中的 jax-primitives 文档)。

此外,仓库还提供了统一两种思路的实验性 API——hijax 原语:一个 Python 实现携带自定义微分(及其他变换)规则,相关内容见 hijax_custom_derivatives.md,其内容与本 notebook 互为镜像。

阅读本文前,建议先了解 jax.jvpjax.grad 及 JVP/VJP 的数学含义(入门材料见 autodiff_cookbook.md)。理解 pytree 概念有助于掌握后文容器示例(见 pytrees.md)。

二、TL;DR:30 秒上手两个 API

2.1 用 jax.custom_jvp 定义前向规则

import jax.numpy as jnp
from jax import custom_jvp

@custom_jvp
def f(x, y):
  return jnp.sin(x) * y

@f.defjvp
def f_jvp(primals, tangents):
  x, y = primals
  x_dot, y_dot = tangents
  primal_out = f(x, y)
  tangent_out = jnp.cos(x) * x_dot * y + jnp.sin(x) * y_dot
  return primal_out, tangent_out

验证(前向模式 jvp 与反向模式 grad 均可用):

from jax import jvp, grad

print(f(2., 3.))
y, y_dot = jvp(f, (2., 3.), (1., 0.))
print(y)
print(y_dot)
print(grad(f)(2., 3.))

等价地,可以使用 defjvps 便捷包装器,为每个参数分别定义 JVP,结果会自动相加:

@custom_jvp
def f(x, y):
  return jnp.sin(x) * y

f.defjvps(lambda x_dot, primal_out, x, y: jnp.cos(x) * x_dot * y,
          lambda y_dot, primal_out, x, y: jnp.sin(x) * y_dot)

2.2 用 jax.custom_vjp 定义纯反向规则

from jax import custom_vjp

@custom_vjp
def f(x, y):
  return jnp.sin(x) * y

def f_fwd(x, y):
  # 返回前向输出,以及供反向传播使用的残差(residuals)
  return f(x, y), (jnp.cos(x), jnp.sin(x), y)

def f_bwd(res, g):
  cos_x, sin_x, y = res  # 取出 f_fwd 中保存的残差
  return (cos_x * g * y, sin_x * g)

f.defvjp(f_fwd, f_bwd)

print(grad(f)(2., 3.))

三、为什么需要自定义求导规则:五个典型问题

3.1 问题一:数值稳定性(log1pexp

考虑函数 xlog(1+ex)x \mapsto \log(1 + e^x)。用 jax.numpy 直接实现:

def log1pexp(x):
  return jnp.log(1. + jnp.exp(x))

log1pexp(3.)

因为完全由 jax.numpy 构成,它天然可被 jitgradvmap 变换:

from jax import jit, grad, vmap

print(jit(log1pexp)(3.))
print(jit(grad(log1pexp))(3.))
print(vmap(jit(grad(log1pexp)))(jnp.arange(3.)))

但数值稳定性隐患潜伏其中——对大输入求导结果错误:

print(grad(log1pexp)(100.))  # 期望约 1.0

数学上,其导数为 ex1+ex\frac{e^x}{1 + e^x},对大 xx 应趋近 1。用 make_jaxpr 检查梯度计算的 jaxpr,可以看到求导过程实际上在计算 lambda x: (1 / (1 + jnp.exp(x))) * jnp.exp(x):当 xx 很大时,浮点会把两个因子分别舍入为 0inf,等效于计算 0. * jnp.inf,结果必然出错。

我们真正想要的是直接求值数学上等价的表达式 111+ex1 - \frac{1}{1 + e^x}——全程不出现需要"抵消"的大数。这正是 jax.custom_jvp 的用武之地:把 log1pexp 作为一个整体指定求导规则,同时保留原始 Python 定义供 jitvmap 等其他变换使用:

from jax import custom_jvp

@custom_jvp
def log1pexp(x):
  return jnp.log(1. + jnp.exp(x))

@log1pexp.defjvp
def log1pexp_jvp(primals, tangents):
  x, = primals
  x_dot, = tangents
  ans = log1pexp(x)
  ans_dot = (1 - 1/(1 + jnp.exp(x))) * x_dot
  return ans, ans_dot

print(grad(log1pexp)(100.))  # 正确得到接近 1 的值
print(jit(log1pexp)(3.))
print(jit(grad(log1pexp))(3.))
print(vmap(jit(grad(log1pexp)))(jnp.arange(3.)))

defjvps 便捷写法:

@custom_jvp
def log1pexp(x):
  return jnp.log(1. + jnp.exp(x))

log1pexp.defjvps(lambda t, ans, x: (1 - 1/(1 + jnp.exp(x))) * t)

3.2 问题二:强制求导约定(定义域边界)

考虑定义在 R+=[0,)\mathbb{R}_+ = [0, \infty) 上的函数 f(x)=x1+xf(x) = \frac{x}{1 + \sqrt{x}}。作为整个实数轴上的函数,ff 在 0 处不可导(左极限不存在),因此默认自动微分给出 nan

def f(x):
  return x / (1 + jnp.sqrt(x))

print(grad(f)(0.))  # nan

但若将 ff 视为 R+\mathbb{R}_+ 上的函数,它在 0 处可导(可理解为从右侧取方向导数),合理值应为 1.0。JAX 默认假设函数定义在整个 R\mathbb{R} 上,所以不会自动给出该值。用自定义 JVP 规则把导数函数 xx+22(x+1)2x \mapsto \frac{\sqrt{x} + 2}{2(\sqrt{x} + 1)^2} 直接写进去即可:

@custom_jvp
def f(x):
  return x / (1 + jnp.sqrt(x))

@f.defjvp
def f_jvp(primals, tangents):
  x, = primals
  x_dot, = tangents
  ans = f(x)
  ans_dot = ((jnp.sqrt(x) + 2) / (2 * (jnp.sqrt(x) + 1)**2)) * x_dot
  return ans, ans_dot

print(grad(f)(0.))  # 1.0

便捷写法:

@custom_jvp
def f(x):
  return x / (1 + jnp.sqrt(x))

f.defjvps(lambda t, ans, x: ((jnp.sqrt(x) + 2) / (2 * (jnp.sqrt(x) + 1)**2)) * t)

3.3 问题三:梯度裁剪(反向模式)

有时我们需要刻意偏离数学,调整自动微分实际执行的计算——反向模式梯度裁剪是经典例子。利用 jnp.clip 配合 jax.custom_vjp 实现一个"恒等但会裁剪梯度"的函数:

from functools import partial
from jax import custom_vjp

@custom_vjp
def clip_gradient(lo, hi, x):
  return x  # 前向是恒等函数

def clip_gradient_fwd(lo, hi, x):
  return x, (lo, hi)  # 把边界保存为残差

def clip_gradient_bwd(res, g):
  lo, hi = res
  return (None, None, jnp.clip(g, lo, hi))  # 用 None 表示 lo、hi 的零余切

clip_gradient.defvjp(clip_gradient_fwd, clip_gradient_bwd)

对比效果(jnp.sin 的导数在 [π/2,π/2][-\pi/2, \pi/2] 外会超出 [0.75,0.75][-0.75, 0.75] 的裁剪窗口):

import matplotlib.pyplot as plt
from jax import vmap

t = jnp.linspace(0, 10, 1000)

plt.plot(jnp.sin(t))
plt.plot(vmap(grad(jnp.sin))(t))          # 未裁剪

def clip_sin(x):
  x = clip_gradient(-0.75, 0.75, x)
  return jnp.sin(x)

plt.plot(clip_sin(t))
plt.plot(vmap(grad(clip_sin))(t))         # 梯度被裁剪到 [-0.75, 0.75]

注意:lohi 是普通数组参数(整数 dtype 也算),不需要nondiff_argnums 声明;直接用 None 返回零余切即可。

3.4 问题四:Python 调试(在反向传播中下断点)

当需要定位 nan 来源、或仔细检查反向传播中传播的余切(梯度)值时,可以在与主计算中某个特定点对应的反向步骤处插入 pdb 调试器——jax.custom_vjp 让这变得很自然(完整示例见第四节)。

3.5 问题五:迭代实现的隐函数微分(fixed_point

这是数学上最深入的一个例子。jax.custom_vjp 的另一个应用场景是:函数可以被 jitvmap 变换,却无法被高效求导——典型如包含 lax.while_loop 的迭代算法(XLA HLO 无法表达一个需要无界内存的反向 While 程序,因此无法高效计算 XLA HLO While 循环的反向导数)。

考虑用 while_loop 迭代求解方程 x=f(a,x)x = f(a, x)fixed_point

from jax.lax import while_loop

def fixed_point(f, a, x_guess):
  def cond_fun(carry):
    x_prev, x = carry
    return jnp.abs(x_prev - x) > 1e-6

  def body_fun(carry):
    _, x = carry
    return x, f(a, x)

  _, x_star = while_loop(cond_fun, body_fun, (x_guess, f(a, x_guess)))
  return x_star

用它实现"只靠加、乘、除"的牛顿法开平方:

def newton_sqrt(a):
  update = lambda a, x: 0.5 * (x + a / x)
  return fixed_point(update, a, a)

print(newton_sqrt(2.))
print(jit(vmap(newton_sqrt))(jnp.array([1., 2., 3., 4.])))

while_loop 无法做反向自动微分;即便能做也不划算——与其穿透 fixed_point 的所有迭代求导,不如利用隐函数定理做更省内存(本例也更省 FLOP)的数学处理。推导如下:

x(a)=f(a,x(a))x^*(a) = f(a, x^*(a))a0a_0 邻域成立,两边对 aa 求导:

x(a)=0f(a,x(a))+1f(a,x(a))x(a)\partial x^*(a) = \partial_0 f(a, x^*(a)) + \partial_1 f(a, x^*(a))\, \partial x^*(a)

A=1f(a0,x(a0))A = \partial_1 f(a_0, x^*(a_0))B=0f(a0,x(a0))B = \partial_0 f(a_0, x^*(a_0)),整理得:

x(a0)=(IA)1B\partial x^*(a_0) = (I - A)^{-1} B

于是向量-雅可比积可以写成:

vx(a0)=v(IA)1B=wBv^\top \partial x^*(a_0) = v^\top (I - A)^{-1} B = w^\top B

其中 w=v+wAw^\top = v^\top + w^\top A,即 ww^\top 是映射 uv+uAu^\top \mapsto v^\top + u^\top A不动点——这提示我们:fixed_point 的 VJP 可以递归地用 fixed_point 本身来写!展开 AABB 后可见只需在 (a0,x(a0))(a_0, x^*(a_0)) 处计算 ff 的 VJP。实现如下:

from jax import vjp

@partial(custom_vjp, nondiff_argnums=(0,))
def fixed_point(f, a, x_guess):
  def cond_fun(carry):
    x_prev, x = carry
    return jnp.abs(x_prev - x) > 1e-6

  def body_fun(carry):
    _, x = carry
    return x, f(a, x)

  _, x_star = while_loop(cond_fun, body_fun, (x_guess, f(a, x_guess)))
  return x_star

def fixed_point_fwd(f, a, x_init):
  x_star = fixed_point(f, a, x_init)
  return x_star, (a, x_star)

def fixed_point_rev(f, res, x_star_bar):
  a, x_star = res
  _, vjp_a = vjp(lambda a: f(a, x_star), a)
  a_bar, = vjp_a(fixed_point(partial(rev_iter, f),
                             (a, x_star, x_star_bar),
                             x_star_bar))
  return a_bar, jnp.zeros_like(x_star)

def rev_iter(f, packed, u):
  a, x_star, x_star_bar = packed
  _, vjp_x = vjp(lambda x: f(a, x), x_star)
  return x_star_bar + vjp_x(u)[0]

fixed_point.defvjp(fixed_point_fwd, fixed_point_rev)

print(newton_sqrt(2.))
print(grad(newton_sqrt)(2.))
print(grad(grad(newton_sqrt))(2.))  # 二阶导也能算

用完全不同的实现 jnp.sqrt 交叉验证:

print(grad(jnp.sqrt)(2.))
print(grad(grad(jnp.sqrt))(2.))

该方法的局限:参数 f 不能闭包捕获任何参与微分的量(这正是我们把 a 显式留在 fixed_point 参数列表里的原因)。若需要在闭包变量上求导,应使用底层原语 lax.custom_root(支持自定义求根函数的同时允许闭包变量参与求导)。

四、jax.custom_jvpjax.custom_vjp API 详解

4.1 jax.custom_jvp:定义前向(并间接反向)规则

规范的最小示例(注释采用 Haskell 风格类型签名):

from jax import custom_jvp
import jax.numpy as jnp

# f :: a -> b
@custom_jvp
def f(x):
  return jnp.sin(x)

# f_jvp :: (a, T a) -> (b, T b)
def f_jvp(primals, tangents):
  x, = primals
  t, = tangents
  return f(x), jnp.cos(x) * t

f.defjvp(f_jvp)

签名约定f_jvp 接收一对输入——类型 a 的 primal 输入与类型 T a 的对应 tangent 输入,返回一对输出——类型 b 的 primal 输出与类型 T b 的 tangent 输出。tangent 输出必须是 tangent 输入的线性函数,否则自动转置会报错。f.defjvp 也可用作装饰器:

@custom_jvp
def f(x):
  ...

@f.defjvp
def f_jvp(primals, tangents):
  ...

只写 JVP 也能用 grad:尽管只定义了 JVP 规则,JAX 会自动转置 tangent 值上的线性计算,其 VJP 效率与手写规则相当:

from jax import grad

print(grad(f)(3.))
print(grad(grad(f))(3.))

多参数版本

@custom_jvp
def f(x, y):
  return x ** 2 * y

@f.defjvp
def f_jvp(primals, tangents):
  x, y = primals
  x_dot, y_dot = tangents
  primal_out = f(x, y)
  tangent_out = 2 * x * y * x_dot + x ** 2 * y_dot
  return primal_out, tangent_out

print(grad(f)(2., 3.))

defjvps 逐参数定义(各自结果分别计算后求和):

@custom_jvp
def f(x, y):
  return x ** 2 * y

f.defjvps(lambda x_dot, primal_out, x, y: 2 * x * y * x_dot,
          lambda y_dot, primal_out, x, y: x ** 2 * y_dot)

print(grad(f)(2., 3.))
print(grad(f, 0)(2., 3.))  # 同上
print(grad(f, 1)(2., 3.))

defjvps 还支持用 None 表示某参数的 JVP 恒为零:

@custom_jvp
def f(x, y):
  return x ** 2 * y

f.defjvps(lambda x_dot, primal_out, x, y: 2 * x * y * x_dot,
          None)

print(grad(f)(2., 3.))
print(grad(f, 0)(2., 3.))
print(grad(f, 1)(2., 3.))  # 关于 y 的梯度为 0

关键字/默认参数:以关键字调用 jax.custom_jvp 函数,或在定义中使用默认参数,都是允许的——只要能被标准库 inspect.signature 无歧义地映射到位置参数。

非微分调用不受影响:不求导时,f 与未加装饰器时行为完全一致(defjvpsprimal_out 由框架提供,无需自己计算):

@custom_jvp
def f(x):
  print('called f!')  # 无害的副作用
  return jnp.sin(x)

@f.defjvp
def f_jvp(primals, tangents):
  print('called f_jvp!')  # 无害的副作用
  x, = primals
  t, = tangents
  return f(x), jnp.cos(x) * t

from jax import vmap, jit

print(f(3.))                    # 只调用 f
print(vmap(f)(jnp.arange(3.)))  # 只调用 f
print(jit(f)(3.))               # 只调用 f

求导时才触发规则(前向与反向都会):

y, y_dot = jvp(f, (3.,), (1.,))
print(y_dot)

print(grad(f)(3.))

高阶微分的要点:注意 f_jvp 是通过调用 f 来计算 primal 输出的。在高阶微分中,每一次微分变换当且仅当规则内部调用原始 f 来计算 primal 时才会使用自定义 JVP 规则——这是一个根本性权衡:既想在规则中复用 f 求值过程的中间量,又想规则在所有阶微分中都生效,二者不可兼得。

grad(grad(f))(3.)

支持 Python 控制流

@custom_jvp
def f(x):
  if x > 0:
    return jnp.sin(x)
  else:
    return jnp.cos(x)

@f.defjvp
def f_jvp(primals, tangents):
  x, = primals
  x_dot, = tangents
  ans = f(x)
  if x > 0:
    return ans, 2 * x_dot
  else:
    return ans, 3 * x_dot

print(grad(f)(1.))   # 2.0
print(grad(f)(-1.))  # 3.0

4.2 jax.custom_vjp:定义纯反向规则

当需要直接控制 VJP 规则(如 3.3、3.5 两例)时,使用 jax.custom_vjp

from jax import custom_vjp
import jax.numpy as jnp

# f :: a -> b
@custom_vjp
def f(x):
  return jnp.sin(x)

# f_fwd :: a -> (b, c)
def f_fwd(x):
  return f(x), jnp.cos(x)

# f_bwd :: (c, CT b) -> CT a
def f_bwd(cos_x, y_bar):
  return (cos_x * y_bar,)

f.defvjp(f_fwd, f_bwd)

from jax import grad

print(f(3.))
print(grad(f)(3.))

签名约定

  • f_fwd 描述前向过程,既做主计算,也决定保存哪些值供反向使用。其输入签名与原始 f 一致;输出为一对:第一个元素是 primal 输出 b,第二个元素是任意"残差"数据 c(类似 PyTorch 的 save_for_backward 机制),反向时传给 f_bwd
  • f_bwd 描述反向过程:接收两个输入——f_fwd 产生的残差 c,以及对应 primal 输出的输出余切 CT b;输出 CT a 是对应 primal 输入的余切。输出必须是长度等于 primal 函数参数个数的序列(如元组)

多参数版本

from jax import custom_vjp

@custom_vjp
def f(x, y):
  return jnp.sin(x) * y

def f_fwd(x, y):
  return f(x, y), (jnp.cos(x), jnp.sin(x), y)

def f_bwd(res, g):
  cos_x, sin_x, y = res
  return (cos_x * g * y, sin_x * g)

f.defvjp(f_fwd, f_bwd)

print(grad(f)(2., 3.))

custom_jvp 相同,custom_vjp 函数也支持关键字参数与默认参数(通过 inspect.signature 无歧义映射)。

非微分调用只走 f:若只求值、或做 jitvmap 等非微分变换,只有 f 被调用,f_fwd/f_bwd 不会被触发。若使用 vjp 显式构造反向传播,则 f_fwd(构造阶段)与 f_bwd(应用阶段)都会执行:

@custom_vjp
def f(x):
  print("called f!")
  return jnp.sin(x)

def f_fwd(x):
  print("called f_fwd!")
  return f(x), jnp.cos(x)

def f_bwd(cos_x, y_bar):
  print("called f_bwd!")
  return (cos_x * y_bar,)

f.defvjp(f_fwd, f_bwd)

print(f(3.))          # 只打印 "called f!"

y, f_vjp = vjp(f, 3.)  # 打印 "called f!" 与 "called f_fwd!"
print(y)
print(f_vjp(1.))       # 打印 "called f_bwd!"

限制:前向模式不可用。对 jax.custom_vjp 函数使用 jax.jvp 会直接报错:

from jax import jvp

try:
  jvp(f, (3.,), (1.,))
except TypeError as e:
  print('ERROR! {}'.format(e))

若需要同时支持前向与反向模式,请改用 jax.custom_jvp

反向传播调试实战:把 pdbcustom_vjp 结合,在反向步骤中下断点:

import pdb

@custom_vjp
def debug(x):
  return x  # 相当于恒等函数

def debug_fwd(x):
  return x, x

def debug_bwd(x, g):
  pdb.set_trace()
  return g

debug.defvjp(debug_fwd, debug_bwd)

def foo(x):
  y = x ** 2
  y = debug(y)  # 在对应反向步骤插入 pdb
  return jnp.sin(y)

运行 jax.grad(foo)(3.) 即会在反向传播经过该点时暂停,可检查中间值:

> <ipython-input-113-b19a2dc1abf7>(12)debug_bwd()
-> return g
(Pdb) p x
Array(9., dtype=float32)
(Pdb) p g
Array(-0.91113025, dtype=float32)
(Pdb) q

五、更多特性与细节

5.1 容器与 pytree 支持

listtuplenamedtupledict 等标准 Python 容器及其嵌套结构均开箱即用——广义上,任何结构一致的 pytree 都允许。custom_jvp 示例(参数与输出均为 pytree):

from collections import namedtuple
Point = namedtuple("Point", ["x", "y"])

@custom_jvp
def f(pt):
  x, y = pt.x, pt.y
  return {'a': x ** 2,
          'b': (jnp.sin(x), jnp.cos(y))}

@f.defjvp
def f_jvp(primals, tangents):
  pt, = primals
  pt_dot, =  tangents
  ans = f(pt)
  ans_dot = {'a': 2 * pt.x * pt_dot.x,
             'b': (jnp.cos(pt.x) * pt_dot.x, -jnp.sin(pt.y) * pt_dot.y)}
  return ans, ans_dot

def fun(pt):
  dct = f(pt)
  return dct['a'] + dct['b'][0]

pt = Point(1., 2.)
print(f(pt))
print(grad(fun)(pt))

custom_vjp 版本(注意 f_bwd 需要按 pytree 结构拆解余切,并还原输出结构):

@custom_vjp
def f(pt):
  x, y = pt.x, pt.y
  return {'a': x ** 2,
          'b': (jnp.sin(x), jnp.cos(y))}

def f_fwd(pt):
  return f(pt), pt

def f_bwd(pt, g):
  a_bar, (b0_bar, b1_bar) = g['a'], g['b']
  x_bar = 2 * pt.x * a_bar + jnp.cos(pt.x) * b0_bar
  y_bar = -jnp.sin(pt.y) * b1_bar
  return (Point(x_bar, y_bar),)

f.defvjp(f_fwd, f_bwd)

def fun(pt):
  dct = f(pt)
  return dct['a'] + dct['b'][0]

pt = Point(1., 2.)
print(f(pt))
print(grad(fun)(pt))

5.2 非可微参数:nondiff_argnums

某些场景(如 3.5 的 fixed_point)需要把函数值参数等不可微参数传给自定义规则。类似的场景还有 jax.experimental.odeint(实现见 ode.py)。

jax.custom_jvp 中的 nondiff_argnums

from functools import partial

@partial(custom_jvp, nondiff_argnums=(0,))
def app(f, x):
  return f(x)

@app.defjvp
def app_jvp(f, primals, tangents):
  x, = primals
  x_dot, = tangents
  return f(x), 2. * x_dot

print(app(lambda x: x ** 3, 3.))
print(grad(app, 1)(lambda x: x ** 3, 3.))

注意事项(gotcha):无论不可微参数在原始签名中的位置如何,在 JVP 规则中它们总是被置于签名最前面。多不可微参数示例:

@partial(custom_jvp, nondiff_argnums=(0, 2))
def app2(f, x, g):
  return f(g((x)))

@app2.defjvp
def app2_jvp(f, g, primals, tangents):
  x, = primals
  x_dot, = tangents
  return f(g(x)), 3. * x_dot

print(app2(lambda x: x ** 3, 3., lambda y: 5 * y))
print(grad(app2, 1)(lambda x: x ** 3, 3., lambda y: 5 * y))

jax.custom_vjp 中的 nondiff_argnums:约定类似——不可微参数一律作为 _bwd 规则的前几个参数传入,无论其在原签名中的位置;_fwd 签名则与原始函数保持一致:

@partial(custom_vjp, nondiff_argnums=(0,))
def app(f, x):
  return f(x)

def app_fwd(f, x):
  return f(x), x

def app_bwd(f, x, g):
  return (5 * g,)

app.defvjp(app_fwd, app_bwd)

print(app(lambda x: x ** 2, 4.))
print(grad(app, 1)(lambda x: x ** 2, 4.))

边界约束nondiff_argnums 不应用于数组值参数(包括整数 dtype 的数组)。它只应标记不对应 JAX 类型(即非数组类型)的值,如 Python 可调用对象或字符串。若 JAX 检测到 nondiff_argnums 标记的位置出现了 JAX Tracer,会直接报错(见 3.3 节 clip_gradient 的正反例:那里 lohi 是整数 dtype 数组,但正确做法是返回 None 零余切而非标记为不可微)。

六、源码视角:两个 API 的底层实现

custom_jvp / custom_vjp 的公开入口定义于 jax/custom_derivatives.py,它从 jax._src.custom_derivatives 导入并再导出(# noqa 注释表明这些名称是刻意保留的公开符号):

from jax._src.custom_derivatives import (
  _sum_tangents,
  _zeros_like_pytree,
  closure_convert as closure_convert,
  custom_gradient as custom_gradient,
  custom_jvp as custom_jvp,
  custom_jvp_call_p as custom_jvp_call_p,
  custom_vjp as custom_vjp,
  custom_vjp_call_p as custom_vjp_call_p,
  ...
)

jax/__init__.py 再将 custom_jvpcustom_vjp 提升为顶层 API,所以可以直接 from jax import custom_jvp, custom_vjp

从实现结构看,两个装饰器最终各自把规则打包成带 subfunsfun, jvpfun, fwd, bwd)参数的 core.Primitive 调用:

  • CustomJVPCallPrimitivejax/_src/custom_derivatives.py)声明了 multiple_results = Trueskip_canonicalization = True,并通过 bind_with_tracefunjvp 分发给当前 trace 的 process_custom_jvp_call
  • CustomVJPCallPrimitive 同样声明多结果,把 fun, fwd, bwd 分发给 process_custom_vjp_call,反向规则会经 _handle_consts_in_bwd 处理常量闭包。

自动转置jax.custom_jvp 之所以"只写 JVP 就能用 grad",是因为反向微分 trace 会对你定义的 JVP 规则中的线性 tangent 计算做转置(transposition),等价于自动生成一个 VJP;前提正是文档强调的"输出 tangents 必须是输入 tangents 的线性函数",否则触发转置错误。

Tracer 防护_check_for_tracers(同文件)会递归遍历 nondiff_argnums 标记位置的所有叶子,一旦发现 core.Tracer 即抛出 UnexpectedTracerError,其错误消息明确说明:nondiff_argnums 只应标记函数值等不可能包含 Tracer 的参数,数组值参数通常不应标记——这与 5.2 节的边界约束完全对应。

便捷包装器defjvps 在源码中明确标注"不能与 nondiff_argnums 同时使用";defvjp 还支持 symbolic_zeros 选项(配合 CustomVJPPrimalcustom_vjp_primal_tree_values 使用,用于把带扰动信息的前向参数树还原为原始形式),并提供了带日志的 defvjp_with_logs 变体。

仓库对这两个 API 的专门测试位于 tests/custom_api_test.py,此外 tests/api_test.py 等测试也大量覆盖了与 jitvmapgrad 组合的行为,可作为自定义规则正确性的验证模板。

七、进一步阅读

  • autodiff_cookbook.md:JAX 自动微分 API 入门(JVP/VJP 的数学含义)。
  • hijax_custom_derivatives.md:hijax 原语——统一 custom rules 与 Primitive 的实验性 API,内容与本 notebook 镜像。
  • jax-primitives 文档:定义新 core.Primitive 及全部变换规则的第二种方式。
  • pytrees.md:pytree 概念与容器约定。
  • 底层求根原语 lax.custom_root(见 jax.lax.rst)与 jax.experimental.odeintode.py):需要闭包变量参与微分的场景。
热门项目推荐
相关项目推荐

项目优选

收起
ops-transformerops-transformer
本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。
C++
1.15 K
2.77 K
kernelkernel
deepin linux kernel
C
34
18
docsdocs
暂无描述
Markdown
900
5.83 K
ops-nnops-nn
本项目是CANN提供的神经网络类计算算子库,实现网络在NPU上加速计算。
C++
929
1.85 K
pytorchpytorch
作为 Ascend for PyTorch 社区的核心组件,TorchNPU 是昇腾专为 PyTorch 打造的深度学习适配插件,使 PyTorch 框架能够直接调用昇腾 NPU,为开发者提供昇腾 AI 处理器的超强算力。
Python
860
1.36 K
jiuwenswarmjiuwenswarm
JiuwenSwarm 是一款基于openJiuwen开发的智能AI Agent,它能够将大语言模型的强大能力,通过你日常使用的各类通讯应用,直接延伸至你的指尖。
Python
3.94 K
1.03 K
ops-mathops-math
本项目是CANN提供的数学类基础计算算子库,实现网络在NPU上加速计算。
C++
1.37 K
1.47 K
kernelkernel
openEuler内核是openEuler操作系统的核心,既是系统性能与稳定性的基石,也是连接处理器、设备与服务的桥梁。
C
534
603
AscendNPU-IRAscendNPU-IR
AscendNPU-IR是基于MLIR(Multi-Level Intermediate Representation)构建的,面向昇腾亲和算子编译时使用的中间表示,提供昇腾完备表达能力,通过编译优化提升昇腾AI处理器计算效率,支持通过生态框架使能昇腾AI处理器与深度调优
C++
548
398
leetcodeleetcode
🔥LeetCode solutions in any programming language | 多种编程语言实现 LeetCode、《剑指 Offer(第 2 版)》、《程序员面试金典(第 6 版)》题解
Markdown
77
23