JAX 自定义求导规则实战:用 `jax.custom_jvp` 与 `jax.custom_vjp` 掌控 JVP/VJP
导读
在 JAX 中,jax.grad、jax.jvp、jax.vjp 等变换会沿着 jax.numpy 与 jax.lax 原语的默认微分规则自动求导,但默认规则并不总能满足数值稳定性、特定定义域约定或工程需求(如梯度裁剪)。本文以仓库文档 Custom_derivative_rules_for_Python_code.md 为主体,系统讲解 jax.custom_jvp 与 jax.custom_vjp 两套 API:先给出可复制的 TL;DR 示例,再深入 5 个真实问题场景(数值稳定性、求导约定、梯度裁剪、反向传播调试、迭代算法的隐函数微分),随后逐项解析 API 签名、nondiff_argnums、pytree 支持等细节,最后结合仓库源码剖析其底层 core.Primitive 实现。读完本文,你将能够为任意 JAX 可变换函数定制求导行为,同时保留其 jit、vmap 等其余变换能力。
一、JAX 中定义求导规则的两种方式
JAX 提供了两条定义微分规则的路径:
- 使用
jax.custom_jvp与jax.custom_vjp,为本身已经可被 JAX 变换的 Python 函数定制求导规则——这是本文的主题; - 定义全新的
core.Primitive实例并为其编写全部变换规则,用于对接求解器、模拟器等外部系统(可参考仓库中的 jax-primitives 文档)。
此外,仓库还提供了统一两种思路的实验性 API——hijax 原语:一个 Python 实现携带自定义微分(及其他变换)规则,相关内容见 hijax_custom_derivatives.md,其内容与本 notebook 互为镜像。
阅读本文前,建议先了解 jax.jvp、jax.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)
考虑函数 。用 jax.numpy 直接实现:
def log1pexp(x):
return jnp.log(1. + jnp.exp(x))
log1pexp(3.)
因为完全由 jax.numpy 构成,它天然可被 jit、grad、vmap 变换:
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
数学上,其导数为 ,对大 应趋近 1。用 make_jaxpr 检查梯度计算的 jaxpr,可以看到求导过程实际上在计算 lambda x: (1 / (1 + jnp.exp(x))) * jnp.exp(x):当 很大时,浮点会把两个因子分别舍入为 0 和 inf,等效于计算 0. * jnp.inf,结果必然出错。
我们真正想要的是直接求值数学上等价的表达式 ——全程不出现需要"抵消"的大数。这正是 jax.custom_jvp 的用武之地:把 log1pexp 作为一个整体指定求导规则,同时保留原始 Python 定义供 jit、vmap 等其他变换使用:
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 问题二:强制求导约定(定义域边界)
考虑定义在 上的函数 。作为整个实数轴上的函数, 在 0 处不可导(左极限不存在),因此默认自动微分给出 nan:
def f(x):
return x / (1 + jnp.sqrt(x))
print(grad(f)(0.)) # nan
但若将 视为 上的函数,它在 0 处可导(可理解为从右侧取方向导数),合理值应为 1.0。JAX 默认假设函数定义在整个 上,所以不会自动给出该值。用自定义 JVP 规则把导数函数 直接写进去即可:
@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 的导数在 外会超出 的裁剪窗口):
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]
注意:lo、hi 是普通数组参数(整数 dtype 也算),不需要用 nondiff_argnums 声明;直接用 None 返回零余切即可。
3.4 问题四:Python 调试(在反向传播中下断点)
当需要定位 nan 来源、或仔细检查反向传播中传播的余切(梯度)值时,可以在与主计算中某个特定点对应的反向步骤处插入 pdb 调试器——jax.custom_vjp 让这变得很自然(完整示例见第四节)。
3.5 问题五:迭代实现的隐函数微分(fixed_point)
这是数学上最深入的一个例子。jax.custom_vjp 的另一个应用场景是:函数可以被 jit、vmap 变换,却无法被高效求导——典型如包含 lax.while_loop 的迭代算法(XLA HLO 无法表达一个需要无界内存的反向 While 程序,因此无法高效计算 XLA HLO While 循环的反向导数)。
考虑用 while_loop 迭代求解方程 的 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)的数学处理。推导如下:
设 在 邻域成立,两边对 求导:
令 、,整理得:
于是向量-雅可比积可以写成:
其中 ,即 是映射 的不动点——这提示我们:fixed_point 的 VJP 可以递归地用 fixed_point 本身来写!展开 、 后可见只需在 处计算 的 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_jvp 与 jax.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 与未加装饰器时行为完全一致(defjvps 中 primal_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:若只求值、或做 jit、vmap 等非微分变换,只有 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。
反向传播调试实战:把 pdb 与 custom_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 支持
list、tuple、namedtuple、dict 等标准 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 的正反例:那里 lo、hi 是整数 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_jvp、custom_vjp 提升为顶层 API,所以可以直接 from jax import custom_jvp, custom_vjp。
从实现结构看,两个装饰器最终各自把规则打包成带 subfuns(fun, jvp 或 fun, fwd, bwd)参数的 core.Primitive 调用:
CustomJVPCallPrimitive(jax/_src/custom_derivatives.py)声明了multiple_results = True与skip_canonicalization = True,并通过bind_with_trace把fun与jvp分发给当前 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 选项(配合 CustomVJPPrimal 与 custom_vjp_primal_tree_values 使用,用于把带扰动信息的前向参数树还原为原始形式),并提供了带日志的 defvjp_with_logs 变体。
仓库对这两个 API 的专门测试位于 tests/custom_api_test.py,此外 tests/api_test.py 等测试也大量覆盖了与 jit、vmap、grad 组合的行为,可作为自定义规则正确性的验证模板。
七、进一步阅读
- 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.odeint(ode.py):需要闭包变量参与微分的场景。
atomcodeClaude Code 的开源替代方案。连接任意大模型,编辑代码,运行命令,自动验证 — 全自动执行。用 Rust 构建,极致性能。 | An open-source alternative to Claude Code. Connect any LLM, edit code, run commands, and verify changes — autonomously. Built in Rust for speed. Get StartedRust4.21 K635- DDeepSeek-V4.1-FlashDeepSeek-V4.1-Flash 是一个多模态混合专家(MoE)模型,拥有 5520 亿骨干参数,并支持最多一百万 token 的上下文长度。该模型原生支持图像和文本输入,并以自回归方式生成文本Python70
jforgamejforgame是一个一站式游戏服务器开发框架。包含游戏服务器开发所需要的各种组件,比如网关,socket服务端与客户端,自定义高效消息编解码,游戏热更新,游戏通用工具等等。包含游戏服,跨服,匹配服,后台管理系统等实现,同时提供大量业务案例以供学习。亦可用于其他socket应用,例如及时聊天等。Java161
fizz-gateway-nodeAn Aggregation API Gateway in Java . FizzGate 是一个基于 Java开发的微服务聚合网关,是拥有自主知识产权的应用网关国产化替代方案,能够实现热服务编排聚合、自动授权选择、线上服务脚本编码、在线测试、高性能路由、API审核管理、回调管理等目的,拥有强大的自定义插件系统可以自行扩展,并且提供友好的图形化配置界面,能够快速帮助企业进行API服务治理、减少中间层胶水代码以及降低编码投入、提高 API 服务的稳定性和安全性。Java90
certd开源SSL证书管理工具;全自动证书申请、更新、续期;通配符证书,泛域名证书申请;证书自动化部署到阿里云、腾讯云、主机、群晖、宝塔;https证书,pfx证书,der证书,TLS证书,nginx证书自动续签自动部署JavaScript120
Spark-X2.5-1.7BSpark-X2.5-1.7B 旨在让强大的 AI 更加实用、高效且易于获取。这些模型在广泛的日常任务中表现出色,涵盖对话、写作、翻译、推理、编程、工具调用和智能体工作流,并在同等规模的开源模型中取得领先结果。Spark-X2.5 将面向效率的架构与最高 1M tokens 的原生上下文窗口相结合,并支持 200 多种语言。Python300