JAX 中的前向和反向模式自动微分#

雅可比-向量积(JVPs,即前向模式自动微分)#

JAX 包含了高效且通用的前向和反向模式自动微分实现。我们熟悉的 jax.grad() 函数是基于反向模式构建的,但要解释这两种模式的区别以及各自的适用场景,需要一些数学背景知识。

数学中的 JVPs#

从数学上讲,给定函数 \(f : \mathbb{R}^n \to \mathbb{R}^m\)\(f\) 在输入点 \(x \in \mathbb{R}^n\) 处的雅可比矩阵(记为 \(\partial f(x)\))通常被视为 \(\mathbb{R}^m \times \mathbb{R}^n\) 中的一个矩阵

\(\qquad \partial f(x) \in \mathbb{R}^{m \times n}\).

但你也可以将 \(\partial f(x)\) 看作一个线性映射,它将 \(f\) 在点 \(x\) 处的定义域切空间(即另一个 \(\mathbb{R}^n\) 的副本)映射到 \(f\) 在点 \(f(x)\) 处的值域切空间(即 \(\mathbb{R}^m\) 的副本)

\(\qquad \partial f(x) : \mathbb{R}^n \to \mathbb{R}^m\).

这个映射被称为 \(f\)\(x\) 处的前推映射(pushforward map)。雅可比矩阵只是该线性映射在标准基下的矩阵表示。

如果不固定具体的输入点 \(x\),那么可以将函数 \(\partial f\) 视为先接收一个输入点,然后返回该输入点处对应的雅可比线性映射

\(\qquad \partial f : \mathbb{R}^n \to \mathbb{R}^n \to \mathbb{R}^m\).

特别地,你可以对函数进行去柯里化(uncurry),使得给定输入点 \(x \in \mathbb{R}^n\) 和切向量 \(v \in \mathbb{R}^n\),得到 \(\mathbb{R}^m\) 中的输出切向量。我们将这种从 \((x, v)\) 对到输出切向量的映射称为雅可比-向量积(Jacobian-vector product),记作

\(\qquad (x, v) \mapsto \partial f(x) v\)

JAX 代码中的 JVPs#

回到 Python 代码,JAX 的 jax.jvp() 函数模拟了这种变换。给定一个评估 \(f\) 的 Python 函数,JAX 的 jax.jvp() 提供了一种方法来获取一个用于评估 \((x, v) \mapsto (f(x), \partial f(x) v)\) 的 Python 函数。

import jax
import jax.numpy as jnp

key = jax.random.key(0)

# Initialize random model coefficients
key, W_key, b_key = jax.random.split(key, 3)
W = jax.random.normal(W_key, (3,))
b = jax.random.normal(b_key, ())

# Define a sigmoid function.
def sigmoid(x):
    return 0.5 * (jnp.tanh(x / 2) + 1)

# Outputs probability of a label being true.
def predict(W, b, inputs):
    return sigmoid(jnp.dot(inputs, W) + b)

# Build a toy dataset.
inputs = jnp.array([[0.52, 1.12,  0.77],
                   [0.88, -1.08, 0.15],
                   [0.52, 0.06, -1.30],
                   [0.74, -2.49, 1.39]])

# Isolate the function from the weight matrix to the predictions
f = lambda W: predict(W, b, inputs)

key, subkey = jax.random.split(key)
v = jax.random.normal(subkey, W.shape)

# Push forward the vector `v` along `f` evaluated at `W`
y, u = jax.jvp(f, (W,), (v,))

如果使用类似 Haskell 的类型签名,可以写成

jvp :: (a -> b) -> a -> T a -> (b, T b)

其中 T a 用于表示 a 的切空间类型。

换句话说,jvp 接收一个类型为 a -> b 的函数、一个类型为 a 的值,以及一个类型为 T a 的切向量作为参数。它返回一个包含类型为 b 的值和类型为 T b 的输出切向量的对(pair)。

jvp 变换后的函数评估过程与原函数非常相似,但它会伴随每个类型为 a 的原始值(primal value),同时推进类型为 T a 的切值。对于原函数中应用的每一个基本数值操作,jvp 变换后的函数都会执行该原语的“JVP 规则”,既评估原语的原始值,又在这些原始值上应用原语的 JVP。

这种评估策略对计算复杂度有一些直接影响。由于我们是在进行过程中评估 JVP,因此不需要存储任何内容以备后用,因此内存成本与计算深度无关。此外,jvp 变换后的函数的浮点运算(FLOP)成本大约是仅评估原函数成本的 3 倍(例如,评估原函数如 sin(x) 需要一个单位的工作量;线性化如 cos(x) 需要一个单位;将线性化函数应用于向量如 cos_x * v 又需要一个单位)。换句话说,对于固定的原始点 \(x\),评估 \(v \mapsto \partial f(x) \cdot v\) 的边际成本与评估 \(f\) 的成本大致相同。

这种内存复杂度听起来非常吸引人!那么为什么我们在机器学习中不常看到前向模式呢?

要回答这个问题,首先想一下如何使用 JVP 构建完整的雅可比矩阵。如果我们对一个独热(one-hot)切向量应用 JVP,它会揭示雅可比矩阵的一列,对应于我们输入的非零条目。因此,我们可以一次一列地构建完整的雅可比矩阵,而获取每一列的成本与一次函数评估大致相同。这对具有“高瘦型”雅可比矩阵的函数来说是有效的,但对“宽短型”雅可比矩阵则效率低下。

如果你在机器学习中进行基于梯度的优化,通常希望最小化一个从 \(\mathbb{R}^n\) 中的参数映射到 \(\mathbb{R}\) 中标量损失值的损失函数。这意味着该函数的雅可比矩阵是一个非常宽的矩阵:\(\partial f(x) \in \mathbb{R}^{1 \times n}\),我们通常将其等同于梯度向量 \(\nabla f(x) \in \mathbb{R}^n\)。一次一列地构建该矩阵,且每次调用所需的浮点运算次数与评估原函数相似,看起来确实效率低下!特别是对于训练神经网络(其中 \(f\) 是训练损失函数,\(n\) 可达数百万甚至数十亿),这种方法显然无法扩展。

为了在此类函数上取得更好的效果,只需要使用反向模式。

向量-雅可比积(VJPs,即反向模式自动微分)#

前向模式返回一个用于评估雅可比-向量积的函数,可用于逐列构建雅可比矩阵;而反向模式提供了一种评估向量-雅可比积(等同于雅可比转置-向量积)的函数,可用于逐行构建雅可比矩阵。

数学中的 VJPs#

再次考虑函数 \(f : \mathbb{R}^n \to \mathbb{R}^m\)。沿用 JVPs 的符号,VJPs 的符号非常简单

\(\qquad (x, v) \mapsto v \partial f(x)\),

其中 \(v\)\(f\)\(x\) 处的余切空间元素(同构于 \(\mathbb{R}^m\) 的另一个副本)。严谨地说,我们应将 \(v\) 视为线性映射 \(v : \mathbb{R}^m \to \mathbb{R}\),当我们写 \(v \partial f(x)\) 时,指的是函数复合 \(v \circ \partial f(x)\),由于 \(\partial f(x) : \mathbb{R}^n \to \mathbb{R}^m\),类型是匹配的。但在常见情况下,我们可以将 \(v\) 等同于 \(\mathbb{R}^m\) 中的一个向量,并几乎可以互换使用,就像我们有时会在“列向量”和“行向量”之间切换而不做过多说明一样。

基于这种等同,我们也可以将 VJP 的线性部分视为 JVP 线性部分的转置(或伴随共轭)

\(\qquad (x, v) \mapsto \partial f(x)^\mathsf{T} v\).

对于给定的点 \(x\),我们可以写出如下签名

\(\qquad \partial f(x)^\mathsf{T} : \mathbb{R}^m \to \mathbb{R}^n\).

余切空间上的相应映射通常被称为 \(f\)\(x\) 处的拉回(pullback)。对我们而言,关键在于它从看起来像 \(f\) 输出的东西转变为看起来像 \(f\) 输入的东西,正如我们对转置线性函数所期望的那样。

JAX 代码中的 VJPs#

回到 Python,JAX 函数 vjp 可以接收一个用于评估 \(f\) 的 Python 函数,并返回一个用于评估 VJP \((x, v) \mapsto (f(x), v^\mathsf{T} \partial f(x))\) 的 Python 函数。

from jax import vjp

# Isolate the function from the weight matrix to the predictions
f = lambda W: predict(W, b, inputs)

y, vjp_fun = vjp(f, W)

key, subkey = jax.random.split(key)
u = jax.random.normal(subkey, y.shape)

# Pull back the covector `u` along `f` evaluated at `W`
v = vjp_fun(u)

如果使用类似 Haskell 的类型签名,可以写成

vjp :: (a -> b) -> a -> (b, CT b -> CT a)

其中 CT a 用于表示 a 的余切空间类型。简而言之,vjp 接收一个类型为 a -> b 的函数和一个类型为 a 的点,返回一个包含类型为 b 的值和类型为 CT b -> CT a 的线性映射的对。

这非常棒,因为它允许我们逐行构建雅可比矩阵,且评估 \((x, v) \mapsto (f(x), v^\mathsf{T} \partial f(x))\) 的浮点运算成本仅为评估 \(f\) 成本的三倍左右。特别地,如果我们想要函数 \(f : \mathbb{R}^n \to \mathbb{R}\) 的梯度,只需一次调用即可完成。这就是 jax.grad() 能够高效进行基于梯度优化的原因,即使对于拥有数百万或数十亿参数的神经网络训练损失函数也是如此。

不过这里有一个代价:尽管浮点运算成本很友好,但内存开销会随着计算深度线性增加。此外,其实现逻辑比前向模式更为复杂,尽管 JAX 有一些巧妙的处理方式(这将在未来的笔记本中介绍!)。

有关反向模式如何工作的更多信息,请参阅 2017 年深度学习暑期学校的这篇教程视频

使用 VJPs 计算向量值梯度#

如果你对计算向量值梯度(如 tf.gradients)感兴趣

def vgrad(f, x):
  y, vjp_fn = jax.vjp(f, x)
  return vjp_fn(jnp.ones(y.shape))[0]

print(vgrad(lambda x: 3*x**2, jnp.ones((2, 2))))
[[6. 6.]
 [6. 6.]]

结合使用前向和反向模式的黑塞矩阵-向量积#

在前面的章节中,你实现了仅使用反向模式的黑塞矩阵-向量积函数(假设存在连续二阶导数)

def hvp(f, x, v):
    return jax.grad(lambda x: jnp.vdot(jax.grad(f)(x), v))(x)

这是有效的,但通过结合使用前向模式和反向模式,你可以做得更好,并节省内存。

从数学上讲,给定待微分函数 \(f : \mathbb{R}^n \to \mathbb{R}\)、进行函数线性化的点 \(x \in \mathbb{R}^n\) 以及向量 \(v \in \mathbb{R}^n\),我们要得到的黑塞矩阵-向量积函数为

\((x, v) \mapsto \partial^2 f(x) v\)

考虑辅助函数 \(g : \mathbb{R}^n \to \mathbb{R}^n\),定义为 \(f\) 的导数(或梯度),即 \(g(x) = \partial f(x)\)。你只需要它的 JVP,因为这会给出

\((x, v) \mapsto \partial g(x) v = \partial^2 f(x) v\).

我们可以将其几乎直接转换为代码

# forward-over-reverse
def hvp(f, primals, tangents):
  return jax.jvp(jax.grad(f), primals, tangents)[1]

更好的是,由于不必直接调用 jnp.dot(),这个 hvp 函数适用于任何形状的数组和任意容器类型(例如存储为嵌套列表/字典/元组的向量),甚至不依赖于 jax.numpy

以下是使用方法示例

def f(X):
  return jnp.sum(jnp.tanh(X)**2)

key, subkey1, subkey2 = jax.random.split(key, 3)
X = jax.random.normal(subkey1, (30, 40))
V = jax.random.normal(subkey2, (30, 40))

def hessian(f):
    return jax.jacfwd(jax.jacrev(f))

ans1 = hvp(f, (X,), (V,))
ans2 = jnp.tensordot(hessian(f)(X), V, 2)

print(jnp.allclose(ans1, ans2, 1e-4, 1e-4))
True

你可能考虑的另一种写法是使用反向模式套用前向模式

# Reverse-over-forward
def hvp_revfwd(f, primals, tangents):
  g = lambda primals: jax.jvp(f, primals, tangents)[1]
  return jax.grad(g)(primals)

但这并不是最好的,因为前向模式的开销比反向模式小,而且由于这里的外部微分算子必须对比内部计算更复杂的计算过程进行微分,因此将前向模式保留在外部效果最好。

# Reverse-over-reverse, only works for single arguments
def hvp_revrev(f, primals, tangents):
  x, = primals
  v, = tangents
  return jax.grad(lambda x: jnp.vdot(jax.grad(f)(x), v))(x)


print("Forward over reverse")
%timeit -n10 -r3 hvp(f, (X,), (V,))
print("Reverse over forward")
%timeit -n10 -r3 hvp_revfwd(f, (X,), (V,))
print("Reverse over reverse")
%timeit -n10 -r3 hvp_revrev(f, (X,), (V,))

print("Naive full Hessian materialization")
%timeit -n10 -r3 jnp.tensordot(jax.hessian(f)(X), V, 2)
Forward over reverse
2.93 ms ± 64.1 μs per loop (mean ± std. dev. of 3 runs, 10 loops each)
Reverse over forward
The slowest run took 4.22 times longer than the fastest. This could mean that an intermediate result is being cached.
9.82 ms ± 7.19 ms per loop (mean ± std. dev. of 3 runs, 10 loops each)
Reverse over reverse
11.4 ms ± 5.79 ms per loop (mean ± std. dev. of 3 runs, 10 loops each)
Naive full Hessian materialization
42.2 ms ± 437 μs per loop (mean ± std. dev. of 3 runs, 10 loops each)

组合使用 VJPs、JVPs 和 jax.vmap#

雅可比-矩阵积和矩阵-雅可比积#

现在你已经拥有了 jax.jvp()jax.vjp() 变换,它们可以分别实现单个向量的前推或拉回,你可以使用 JAX 的 jax.vmap() 变换来一次性前推或拉回整个基。特别地,你可以利用这一点来编写快速的矩阵-雅可比积和雅可比-矩阵积。

# Isolate the function from the weight matrix to the predictions
f = lambda W: predict(W, b, inputs)

# Pull back the covectors `m_i` along `f`, evaluated at `W`, for all `i`.
# First, use a list comprehension to loop over rows in the matrix M.
def loop_mjp(f, x, M):
    y, vjp_fun = jax.vjp(f, x)
    return jnp.vstack([jnp.asarray(vjp_fun(mi)) for mi in M])

# Now, use vmap to build a computation that does a single fast matrix-matrix
# multiply, rather than an outer loop over vector-matrix multiplies.
def vmap_mjp(f, x, M):
    y, vjp_fun = jax.vjp(f, x)
    outs, = jax.vmap(vjp_fun)(M)
    return outs

key = jax.random.key(0)
num_covecs = 128
U = jax.random.normal(key, (num_covecs,) + y.shape)

loop_vs = loop_mjp(f, W, M=U)
print('Non-vmapped Matrix-Jacobian product')
%timeit -n10 -r3 loop_mjp(f, W, M=U)

print('\nVmapped Matrix-Jacobian product')
vmap_vs = vmap_mjp(f, W, M=U)
%timeit -n10 -r3 vmap_mjp(f, W, M=U)

assert jnp.allclose(loop_vs, vmap_vs), 'Vmap and non-vmapped Matrix-Jacobian Products should be identical'
Non-vmapped Matrix-Jacobian product
74.5 ms ± 93.7 μs per loop (mean ± std. dev. of 3 runs, 10 loops each)

Vmapped Matrix-Jacobian product
3.41 ms ± 48.6 μs per loop (mean ± std. dev. of 3 runs, 10 loops each)
def loop_jmp(f, W, M):
    # jvp immediately returns the primal and tangent values as a tuple,
    # so we'll compute and select the tangents in a list comprehension
    return jnp.vstack([jax.jvp(f, (W,), (mi,))[1] for mi in M])

def vmap_jmp(f, W, M):
    _jvp = lambda s: jax.jvp(f, (W,), (s,))[1]
    return jax.vmap(_jvp)(M)
num_vecs = 128
S = jax.random.normal(key, (num_vecs,) + W.shape)

loop_vs = loop_jmp(f, W, M=S)
print('Non-vmapped Jacobian-Matrix product')
%timeit -n10 -r3 loop_jmp(f, W, M=S)
vmap_vs = vmap_jmp(f, W, M=S)
print('\nVmapped Jacobian-Matrix product')
%timeit -n10 -r3 vmap_jmp(f, W, M=S)

assert jnp.allclose(loop_vs, vmap_vs), 'Vmap and non-vmapped Jacobian-Matrix products should be identical'
Non-vmapped Jacobian-Matrix product
86.9 ms ± 180 μs per loop (mean ± std. dev. of 3 runs, 10 loops each)

Vmapped Jacobian-Matrix product
1.26 ms ± 29.7 μs per loop (mean ± std. dev. of 3 runs, 10 loops each)

jax.jacfwdjax.jacrev 的实现#

既然我们已经了解了快速的雅可比-矩阵积和矩阵-雅可比积,就不难猜出如何编写 jax.jacfwd()jax.jacrev() 了。我们只需使用同样的技术,一次性前推或拉回整个标准基(同构于单位矩阵)。

from jax import jacrev as builtin_jacrev

def our_jacrev(f):
    def jacfun(x):
        y, vjp_fun = jax.vjp(f, x)
        # Use vmap to do a matrix-Jacobian product.
        # Here, the matrix is the Euclidean basis, so we get all
        # entries in the Jacobian at once.
        J, = jax.vmap(vjp_fun, in_axes=0)(jnp.eye(len(y)))
        return J
    return jacfun

assert jnp.allclose(builtin_jacrev(f)(W), our_jacrev(f)(W)), 'Incorrect reverse-mode Jacobian results!'
from jax import jacfwd as builtin_jacfwd

def our_jacfwd(f):
    def jacfun(x):
        _jvp = lambda s: jax.jvp(f, (x,), (s,))[1]
        Jt = jax.vmap(_jvp, in_axes=1)(jnp.eye(len(x)))
        return jnp.transpose(Jt)
    return jacfun

assert jnp.allclose(builtin_jacfwd(f)(W), our_jacfwd(f)(W)), 'Incorrect forward-mode Jacobian results!'

有趣的是,Autograd 库无法做到这一点。Autograd 中反向模式 jacobian实现必须通过外部循环 map 逐个拉回向量。在计算过程中逐个前推向量的效率远低于使用 jax.vmap() 将它们全部批处理起来。

Autograd 无法做到的另一件事是 jax.jit()。有趣的是,无论你在待微分函数中使用多少 Python 动态特性,我们始终可以在计算的线性部分上使用 jax.jit()。例如

def f(x):
    try:
        if x < 3:
            return 2 * x ** 3
        else:
            raise ValueError
    except ValueError:
        return jnp.pi * x

y, f_vjp = jax.vjp(f, 4.)
print(jax.jit(f_vjp)(1.))
(Array(3.1415927, dtype=float32, weak_type=True),)