拓冰建站拓冰建站
首页 / 资讯中心 / 正文

JAX 对复数函数求导怎么做:解析函数与非解析函数的 JVP 和 VJP

JAX 对复数函数求导怎么做解析函数与非解析函数的 JVP 和 VJP【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax如果你需要用 JAX 对复数输入的函数求导会遇到一个具体分叉函数输出是实数还是复数函数是解析的holomorphic还是非解析的non-holomorphicJAX 的jax.grad只对实值输出函数直接可用对复值输出要么函数是解析的并显式传入holomorphicTrue要么改用jax.jvp/jax.vjp直接处理实线性导数。这篇文章覆盖三条可执行路径对非解析复函数验证 JVP/VJP 并求出完整 Jacobian对解析函数用grad(f, holomorphicTrue)取复导数对ℂ → ℝ损失函数用grad的共轭做梯度下降。本文内容来自 JAX 文档 复数与求导 和 自动微分手册。准备环境只需要 CPU 版 JAXpip 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))预期输出文档示例为Truejvp的结果等于实 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, holomorphicTrue) 直接给出 f(z)grad默认要求输出为实数复值输出会直接报错。错误信息本身就在源码里给出了两条出路见 jax/_src/api.pygrad requires real-valued outputs (output dtype that is a sub-dtype of np.floating), but got complex dtype. For holomorphic differentiation, pass holomorphicTrue. 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, holomorphicTrue)做的就是用共轭向量1.0调一次 VJP返回f(z)print(grad(jnp.sin, holomorphicTrue)(3. 4j)) print(jnp.cos(3. 4j))两行输出相同即文档示例中jnp.cos(3. 4j)的值。holomorphicTrue只做一件事关掉对复值输出的报错检查它不验证函数是否真的解析。文档同时提醒对非解析函数传入该参数仍然可以运行但返回值不是完整 Jacobian而是丢弃输出虚部之后的函数实部的 Jacobiandef f(z): return jnp.conjugate(z) # not holomorphic! grad(f, holomorphicTrue)(3. 4j)另外holomorphicTrue要求输入和输出都必须是复数 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, holomorphicTrue)(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, holomorphicTrue)取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.mdjax-301-complex一节以及 docs/301/custom-jvp-vjp.md 中自定义 JVP/VJP 规则的接口。【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门