JAX 对复数函数求导怎么做:解析函数与非解析函数的 JVP 和 VJP
【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax
如果你需要用 JAX 对复数输入的函数求导,会遇到一个具体分叉:函数输出是实数还是复数?函数是解析的(holomorphic)还是非解析的(non-holomorphic)?JAX 的jax.grad只对实值输出函数直接可用;对复值输出,要么函数是解析的并显式传入holomorphic=True,要么改用jax.jvp/jax.vjp直接处理实线性导数。这篇文章覆盖三条可执行路径:对非解析复函数验证 JVP/VJP 并求出完整 Jacobian;对解析函数用grad(f, holomorphic=True)取复导数;对ℂ → ℝ损失函数用grad的共轭做梯度下降。
本文内容来自 JAX 文档 复数与求导 和 自动微分手册。准备环境只需要 CPU 版 JAX:
pip install --upgrade pip pip install --upgrade jax先明确:JVP 和 VJP 对任何可微复函数都是良定义的
JAX 中复值函数的微分是定义在底层实导数上的。把f: ℂ → ℂ按f(z) = u(x, y) + v(x, y)·1j分解后,它对应实函数F: ℝ² → ℝ²,其导数是实 2×2 Jacobian 矩阵:
J = [[∂₀u, ∂₁u], [∂₀v, ∂₁v]]JVP 就是把这个实线性映射作用到切向量上,复数只是实数对的表示;这个定义不要求解析性。所以非解析函数没有歧义——jvp和vjp始终可用,也是文档给出的兜底建议:"When in doubt about what a complex derivative means, usejvpandvjpdirectly: they are always well-defined, for any function."
下面用文档中一个非解析函数验证 JVP。u、v的选取使fun不满足 Cauchy–Riemann 方程,因此不是解析函数:
import jax import jax.numpy as jnp from jax import grad, jvp, vjp def u(x, y): return x**2 + jnp.sin(y) def v(x, y): return x * y def fun(z): # not holomorphic! x, y = jnp.real(z), jnp.imag(z) return u(x, y) + v(x, y) * 1j z = 1.5 + 0.5j x, y = jnp.real(z), jnp.imag(z) J = jnp.array([[grad(u, 0)(x, y), grad(u, 1)(x, y)], [grad(v, 0)(x, y), grad(v, 1)(x, y)]]) t = 0.7 - 0.3j _, t_out = jvp(fun, (z,), (t,)) t_pair = J @ jnp.array([jnp.real(t), jnp.imag(t)]) print(jnp.allclose(t_out, t_pair[0] + t_pair[1] * 1j))预期输出(文档示例)为True:jvp的结果等于实 Jacobian 作用在切向量(t₁, t₂)上再拼回复数。
VJP 的配对约定:JAX 用双线性配对
vjp是导数的对偶映射。由于f一般只对ℝ可微,余切是ℝ线性泛函,但vjp返回的余切与原始值同类型,所以泛函必须用复数来表示,这就涉及一个配对约定的选择。在ℂ ≅ ℝ²上有两种标准实值配对:
- 双线性(bilinear)配对:
⟨w, t⟩ = Re(wt) = w₁t₁ - w₂t₂; - 半双线性(sesquilinear)配对:
⟨w, t⟩ = Re(w̄t) = w₁t₁ + w₂t₂,即ℝ²上标准欧氏内积。
两者相差第一参数上的共轭,因此它们诱导的转置相差一次逐元素共轭。JAX 采用双线性配对。在这个约定下,vjp由下面的恒等式刻画,注意全程是普通复数乘积、不显式取共轭:
Re(w · jvp(t)) = Re(vjp(w) · t) 对任意 t, w 成立沿用上文的fun、z和切向量t,可以核对这一点(代码沿用上一节的变量):
w = -0.2 + 1.1j _, fun_vjp = vjp(fun, z) w_out, = fun_vjp(w) print(jnp.allclose(jnp.real(w * t_out), jnp.real(w_out * t))) # True print(jnp.allclose(jnp.real(jnp.conj(w) * t_out), jnp.real(jnp.conj(w_out) * t))) # False!第一行输出True、第二行输出False!是文档给出的示例结果:如果误用了带共轭的 sesquilinear 形式,核对会失败。
求非解析函数的完整 Jacobian:两次 JVP 或 VJP
一般ℝ可微的ℂ → ℂ映射的导数有 4 个实数自由度,单次 JVP 或 VJP 只是它的二维投影。对两个线性无关的切向量(例如1和1j)各求一次,就能恢复完整 Jacobian。对实域或实值码域则一次就够:一次jvp确定ℝ → ℂ函数的导数,一次vjp(或grad)确定ℂ → ℝ函数的导数。
上一节例子里的J就是这样用 4 次grad按分量拼出来的;在正式代码中按同样思路对1和1j两个方向求jvp即可得到完整 Jacobian 的两个列。
解析函数:grad(f, holomorphic=True) 直接给出 f'(z)
grad默认要求输出为实数:复值输出会直接报错。错误信息本身就在源码里给出了两条出路(见 jax/_src/api.py):
grad requires real-valued outputs (output dtype that is a sub-dtype of np.floating), but got <complex dtype>. For holomorphic differentiation, pass holomorphic=True. For differentiation of non-holomorphic functions involving complex outputs, use jax.vjp directly.函数是解析的,意味着 Cauchy–Riemann 方程把 2×2 实 Jacobian 限制成复平面上的一个缩放旋转,导数完全由单个复数f'(z)刻画。此时jvp和vjp都退化为普通复数乘法:
jvp(t) = f'(z)·t, vjp(w) = f'(z)·wgrad(f, holomorphic=True)做的就是用共轭向量1.0调一次 VJP,返回f'(z):
print(grad(jnp.sin, holomorphic=True)(3. + 4j)) print(jnp.cos(3. + 4j))两行输出相同,即文档示例中jnp.cos(3. + 4j)的值。
holomorphic=True只做一件事:关掉对复值输出的报错检查,它不验证函数是否真的解析。文档同时提醒:对非解析函数传入该参数仍然可以运行,但返回值不是完整 Jacobian,而是"丢弃输出虚部之后"的函数(实部)的 Jacobian:
def f(z): return jnp.conjugate(z) # not holomorphic! grad(f, holomorphic=True)(3. + 4j)另外,holomorphic=True要求输入和输出都必须是复数 dtype,否则分别抛TypeError(检查逻辑同样在 jax/_src/api.py 与 jax/_src/api.py)。
复数在 JAX 的变换和线性代数中是普遍支持的,文档给出的一个例子是对复矩阵 Cholesky 分解求导:
A = jnp.array([[5., 2.+3j, 5j], [2.-3j, 7., 1.+7j], [-5j, 1.-7j, 12.]]) def f(X): L = jnp.linalg.cholesky(X) return jnp.sum((L - jnp.sin(L))**2) grad(f, holomorphic=True)(A)优化 ℂ → ℝ 损失:必须沿 grad 的共轭方向走
这是最容易出错的一处。对f: ℂ → ℝ,JAX 定义grad(f)(x)为vjp(f, x)1;代入双线性转置公式得:
grad(f)(z) = ∂₀u(x, y) - ∂₁u(x, y)·i它是实梯度向量(∂₀u, ∂₁u)的复共轭,而不是梯度向量本身。因此方向导数由普通乘积的实部给出:
lim_{ε→0} (f(z + εt) - f(z))/ε = Re(grad(f)(z) · t)复平面上的最速上升方向是conj(grad(f)(z)),梯度下降更新必须写成z ← z - η·conj(grad(f)(z))。文档用一个最小化点在原点的|z|²演示了两种写法的差别:
def f(z): x, y = jnp.real(z), jnp.imag(z) return x**2 + y**2 # |z|^2, minimized at z = 0 print(grad(f)(3. + 4j)) # 6 - 8j: conjugate of the steepest-ascent 6 + 8jz = 3. + 4j for _ in range(100): z = z - 0.05 * jnp.conj(grad(f)(z)) # with the conjugate: descends print(f(z)) z = 3. + 4j for _ in range(100): z = z - 0.05 * grad(f)(z) # without: the imaginary part grows! print(f(z))取共轭的循环会下降,不取的循环虚部分量会变大(文档示例注释)。
由此得到两条直接可用的结论:
- 用复参数优化实值损失时,沿
conj(grad(f)(z))迈步。为实参数编写的优化器库不会替你加这次共轭,用在复参数上时必须自行核对。 - 一阶 Taylor 近似和方向导数用不取共轭的乘积:
f(z + t) ≈ f(z) + Re(grad(f)(z) · t)。
使用建议与文档边界
按函数类型选择求导接口:
ℂ → ℝ实值损失:直接用grad,但更新方向取conj(grad(...));ℂ → ℂ且确实是解析函数:用grad(f, holomorphic=True)取f'(z),并自行保证解析性——JAX 不检查;ℂ → ℂ非解析函数:不要用grad,直接用jvp/vjp,两次独立方向的求值恢复完整 Jacobian。
与 Wirtinger 记号的关系文档也有对照:JVP 可写成jvp(t) = (∂f/∂z)·t + (∂f/∂z̄)·t̄,函数解析当且仅当∂f/∂z̄ = 0。对实值函数,JAX 的grad计算的是2·∂f/∂z,而最速上升向量是2·∂f/∂z̄ = conj(grad(f)(z));PyTorch 和 TensorFlow 采用的是把共轭吸收进返回导数的 sesquilinear 约定(∂L/∂z* = 2·∂L/∂z̄)。两套约定表示的是同一个底层实导数,只是共轭出现的位置不同——这也是从其他框架迁移过来时最容易踩的坑。
相关文档可继续深入:docs/complex-differentiation.md、docs/301/cookbook.md(jax-301-complex一节),以及 docs/301/custom-jvp-vjp.md 中自定义 JVP/VJP 规则的接口。
【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考