复数与微分#
JAX 对复数和微分提供了强大的支持。为了同时支持 全纯和非全纯微分,从 JVP(雅可比-向量积)和 VJP(向量-雅可比积)的角度来思考会很有帮助。
考虑一个复数到复数的函数 \(f: \mathbb{C} \to \mathbb{C}\),并将其等同于一个对应的函数 \(g: \mathbb{R}^2 \to \mathbb{R}^2\),
import jax.numpy as jnp
def f(z):
x, y = jnp.real(z), jnp.imag(z)
return u(x, y) + v(x, y) * 1j
def g(x, y):
return (u(x, y), v(x, y))
也就是说,我们将 \(f(z) = u(x, y) + v(x, y) i\) 分解,其中 \(z = x + y i\),并通过将 \(\mathbb{C}\) 等同于 \(\mathbb{R}^2\) 来得到 \(g\)。
由于 \(g\) 只涉及实数输入和输出,我们已经知道如何为它编写雅可比-向量积(JVP)。假设给定一个切向量 \((c, d) \in \mathbb{R}^2\),即:
\(\begin{bmatrix} \partial_0 u(x, y) & \partial_1 u(x, y) \\ \partial_0 v(x, y) & \partial_1 v(x, y) \end{bmatrix} \begin{bmatrix} c \\ d \end{bmatrix}\).
为了得到应用于切向量 \(c + di \in \mathbb{C}\) 的原始函数 \(f\) 的 JVP,我们只需使用相同的定义,并将结果识别为另一个复数:
\(\partial f(x + y i)(c + d i) = \begin{matrix} \begin{bmatrix} 1 & i \end{bmatrix} \\ ~ \end{matrix} \begin{bmatrix} \partial_0 u(x, y) & \partial_1 u(x, y) \\ \partial_0 v(x, y) & \partial_1 v(x, y) \end{bmatrix} \begin{bmatrix} c \\ d \end{bmatrix}\).
这就是我们对 \(\mathbb{C} \to \mathbb{C}\) 函数 JVP 的定义!请注意,\(f\) 是否全纯并不重要:JVP 的结果是明确无误的。
这里是一个验证
from jax import random, grad, jvp
def check(seed):
key = random.key(seed)
# random coeffs for u and v
key, subkey = random.split(key)
a, b, c, d = random.uniform(subkey, (4,))
def fun(z):
x, y = jnp.real(z), jnp.imag(z)
return u(x, y) + v(x, y) * 1j
def u(x, y):
return a * x + b * y
def v(x, y):
return c * x + d * y
# primal point
key, subkey = random.split(key)
x, y = random.uniform(subkey, (2,))
z = x + y * 1j
# tangent vector
key, subkey = random.split(key)
c, d = random.uniform(subkey, (2,))
z_dot = c + d * 1j
# check jvp
_, ans = jvp(fun, (z,), (z_dot,))
expected = (grad(u, 0)(x, y) * c +
grad(u, 1)(x, y) * d +
grad(v, 0)(x, y) * c * 1j+
grad(v, 1)(x, y) * d * 1j)
print(jnp.allclose(ans, expected))
check(0)
check(1)
check(2)
True
True
True
那么 VJP 呢?我们做类似的事情:对于余切向量 \(c + di \in \mathbb{C}\),我们将 \(f\) 的 VJP 定义为:
\((c + di)^* \; \partial f(x + y i) = \begin{matrix} \begin{bmatrix} c & -d \end{bmatrix} \\ ~ \end{matrix} \begin{bmatrix} \partial_0 u(x, y) & \partial_1 u(x, y) \\ \partial_0 v(x, y) & \partial_1 v(x, y) \end{bmatrix} \begin{bmatrix} 1 \\ -i \end{bmatrix}\).
为什么要用负号?它们只是为了处理复共轭,以及我们正在处理余向量这一事实。
这里是对 VJP 规则的验证
from jax import vjp
def check(seed):
key = random.key(seed)
# random coeffs for u and v
key, subkey = random.split(key)
a, b, c, d = random.uniform(subkey, (4,))
def fun(z):
x, y = jnp.real(z), jnp.imag(z)
return u(x, y) + v(x, y) * 1j
def u(x, y):
return a * x + b * y
def v(x, y):
return c * x + d * y
# primal point
key, subkey = random.split(key)
x, y = random.uniform(subkey, (2,))
z = x + y * 1j
# cotangent vector
key, subkey = random.split(key)
c, d = random.uniform(subkey, (2,))
z_bar = jnp.array(c + d * 1j) # for dtype control
# check vjp
_, fun_vjp = vjp(fun, z)
ans, = fun_vjp(z_bar)
expected = (grad(u, 0)(x, y) * c +
grad(v, 0)(x, y) * (-d) +
grad(u, 1)(x, y) * c * (-1j) +
grad(v, 1)(x, y) * (-d) * (-1j))
assert jnp.allclose(ans, expected, atol=1e-5, rtol=1e-5)
check(0)
check(1)
check(2)
那么像 jax.grad(), jax.jacfwd() 和 jax.jacrev() 这样的便捷封装呢?
对于 \(\mathbb{R} \to \mathbb{R}\) 函数,回顾一下,我们定义 grad(f)(x) 为 vjp(f, x)[1](1.0),这之所以有效,是因为将 VJP 应用于 1.0 值可以揭示梯度(即雅可比矩阵或导数)。我们对 \(\mathbb{C} \to \mathbb{R}\) 函数也可以这样做:我们仍然可以使用 1.0 作为余切向量,并得到一个总结了完整雅可比矩阵的复数结果。
def f(z):
x, y = jnp.real(z), jnp.imag(z)
return x**2 + y**2
z = 3. + 4j
grad(f)(z)
Array(6.-8.j, dtype=complex64)
对于一般的 \(\mathbb{C} \to \mathbb{C}\) 函数,雅可比矩阵具有 4 个实值自由度(如上方的 2x2 雅可比矩阵所示),因此我们无法指望用单个复数来表示所有这些信息。但对于全纯函数,我们可以做到!全纯函数正是那种具有特殊性质的 \(\mathbb{C} \to \mathbb{C}\) 函数,其导数可以表示为一个单一的复数。(柯西-黎曼方程确保了上述 2x2 雅可比矩阵具有复平面中缩放和旋转矩阵的特殊形式,即单个复数乘法的作用。)我们可以通过对 vjp 进行一次调用(余向量为 1.0)来揭示该复数。
由于这只适用于全纯函数,为了使用此技巧,我们需要向 JAX 保证我们的函数是全纯的;否则,当在具有复数输出的函数上使用 jax.grad() 时,JAX 将会报错。
def f(z):
return jnp.sin(z)
z = 3. + 4j
grad(f, holomorphic=True)(z)
Array(-27.034946-3.8511534j, dtype=complex64, weak_type=True)
holomorphic=True 的承诺仅仅是禁用输出为复数时的错误提示。即使函数不是全纯的,我们仍然可以写 holomorphic=True,但得到的结果将无法代表完整的雅可比矩阵。相反,它将是丢弃了输出虚部后的函数的雅可比矩阵。
def f(z):
return jnp.conjugate(z)
z = 3. + 4j
grad(f, holomorphic=True)(z) # f is not actually holomorphic!
Array(1.-0.j, dtype=complex64, weak_type=True)
这对于 jax.grad() 的工作方式有一些有用的结论。
我们可以在全纯的 \(\mathbb{C} \to \mathbb{C}\) 函数上使用
jax.grad()。我们可以使用
jax.grad()来优化 \(f : \mathbb{C} \to \mathbb{R}\) 函数,例如复数参数x的实值损失函数,方法是沿着grad(f)(x)的共轭方向迈步。如果我们有一个 \(\mathbb{R} \to \mathbb{R}\) 函数,它内部恰好使用了某些复数运算(其中一些必须是非全纯的,例如卷积中使用的 FFT),那么
jax.grad()仍然有效,并且我们得到的结果与仅使用实数实现的程序所给出的结果相同。
总之,JVP 和 VJP 总是明确无误的。如果我们想计算非全纯 \(\mathbb{C} \to \mathbb{C}\) 函数的完整雅可比矩阵,我们可以通过 JVP 或 VJP 来实现!
你应该预料到复数在 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)
Array([[-0.7534186 +0.j , -3.0509028 -10.940544j ,
5.9896846 +3.5423026j],
[-3.0509028 +10.940544j , -8.904491 +0.j ,
-5.1351523 -6.559373j ],
[ 5.9896846 -3.5423026j, -5.1351523 +6.559373j ,
0.01320427 +0.j ]], dtype=complex64)