复数与微分

复数与微分#

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() 的工作方式有一些有用的结论。

  1. 我们可以在全纯的 \(\mathbb{C} \to \mathbb{C}\) 函数上使用 jax.grad()

  2. 我们可以使用 jax.grad() 来优化 \(f : \mathbb{C} \to \mathbb{R}\) 函数,例如复数参数 x 的实值损失函数,方法是沿着 grad(f)(x) 的共轭方向迈步。

  3. 如果我们有一个 \(\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)