使用 jax.checkpoint (又名 jax.remat) 控制自动微分的保存值#
import jax
import jax.numpy as jnp
概述#
将 jax.checkpoint 装饰器(别名为 jax.remat)与 jax.grad 结合使用,可以控制在正向传播过程中保存哪些中间变量,以及在反向传播过程中重新计算哪些变量,从而权衡内存占用和计算量(FLOPs)。
请务必阅读实用注意事项,其中讨论了 jax.checkpoint 如何与 jax.jit 进行交互。
如果不使用 jax.checkpoint,jax.grad(f)(x) 的正向传播会保存雅可比系数和其他中间值以供反向传播使用。我们将这些保存的值称为残差 (residuals)。
def g(W, x):
y = jnp.dot(W, x)
return jnp.sin(y)
def f(W1, W2, W3, x):
x = g(W1, x)
x = g(W2, x)
x = g(W3, x)
return x
W1 = jnp.ones((5, 4))
W2 = jnp.ones((6, 5))
W3 = jnp.ones((7, 6))
x = jnp.ones(4)
# Inspect the 'residual' values to be saved on the forward pass
# if we were to evaluate `jax.grad(f)(W1, W2, W3, x)`
from jax.ad_checkpoint import print_saved_residuals
jax.ad_checkpoint.print_saved_residuals(f, W1, W2, W3, x)
f32[5,4] from the argument 'W1'
f32[6,5] from the argument 'W2'
f32[7,6] from the argument 'W3'
f32[4] from the argument 'x'
f32[5] output of sin from <ipython-input-4-f510dde58e22>:3 (g)
f32[5] output of cos from <ipython-input-4-f510dde58e22>:3 (g)
f32[6] output of sin from <ipython-input-4-f510dde58e22>:3 (g)
f32[6] output of cos from <ipython-input-4-f510dde58e22>:3 (g)
f32[7] output of cos from <ipython-input-4-f510dde58e22>:3 (g)
通过将 jax.checkpoint 应用于子函数(作为装饰器或在特定的应用位置),我们强制 JAX 不保存该子函数的任何残差。相反,只有 jax.checkpoint 装饰函数的输入才可能被保存,而反向传播过程中消耗的任何残差都会根据需要从这些输入中重新计算。
def f2(W1, W2, W3, x):
x = jax.checkpoint(g)(W1, x)
x = jax.checkpoint(g)(W2, x)
x = jax.checkpoint(g)(W3, x)
return x
jax.ad_checkpoint.print_saved_residuals(f2, W1, W2, W3, x)
f32[5,4] from the argument 'W1'
f32[6,5] from the argument 'W2'
f32[7,6] from the argument 'W3'
f32[4] from the argument 'x'
f32[5] output of sin from <ipython-input-4-f510dde58e22>:3 (g)
f32[6] output of sin from <ipython-input-4-f510dde58e22>:3 (g)
在这里,两个 sin 应用的结果被保存了下来,因为它们是后续 jax.checkpoint 装饰函数 g 的参数,且 jax.checkpoint 装饰函数的输入可能会被保存。但没有任何 cos 应用的值被保存。
为了在不编辑被微分函数定义的情况下控制哪些值可以保存,可以使用重计算(rematerialization)策略。以下是一个示例,它仅保存没有批次维度(batch dimensions)的 dot 操作结果(因为它们通常是计算密集型的,因此保存比重新计算更划算)。
f3 = jax.checkpoint(f, policy=jax.checkpoint_policies.dots_with_no_batch_dims_saveable)
jax.ad_checkpoint.print_saved_residuals(f3, W1, W2, W3, x)
f32[5,4] from the argument 'W1'
f32[6,5] from the argument 'W2'
f32[7,6] from the argument 'W3'
f32[4] from the argument 'x'
f32[5] output of dot_general from <ipython-input-4-f510dde58e22>:2 (g)
f32[6] output of dot_general from <ipython-input-4-f510dde58e22>:2 (g)
f32[7] output of dot_general from <ipython-input-4-f510dde58e22>:2 (g)
您还可以使用策略来引用使用 jax.ad_checkpoint.checkpoint_name 命名的中间值。
from jax.ad_checkpoint import checkpoint_name
def f4(W1, W2, W3, x):
x = checkpoint_name(g(W1, x), name='a')
x = checkpoint_name(g(W2, x), name='b')
x = checkpoint_name(g(W3, x), name='c')
return x
f4 = jax.checkpoint(f4, policy=jax.checkpoint_policies.save_only_these_names('a'))
jax.ad_checkpoint.print_saved_residuals(f4, W1, W2, W3, x)
f32[5,4] from the argument 'W1'
f32[6,5] from the argument 'W2'
f32[7,6] from the argument 'W3'
f32[4] from the argument 'x'
f32[5] named 'a' from <ipython-input-7-fc0ed1c14b8d>:4 (f4)
在尝试这些示例时,我们可以使用本 Notebook 中定义的 print_fwd_bwd 工具来更仔细地查看正在发生的事情。
from jax.tree_util import tree_flatten, tree_unflatten
from rich.console import Console
from rich.table import Table
import rich.text
def print_fwd_bwd(f, *args, **kwargs) -> None:
args, in_tree = tree_flatten((args, kwargs))
def f_(*args):
args, kwargs = tree_unflatten(in_tree, args)
return f(*args, **kwargs)
fwd = jax.make_jaxpr(lambda *args: jax.vjp(f_, *args))(*args).jaxpr
y, f_vjp = jax.vjp(f_, *args)
res, in_tree = tree_flatten(f_vjp)
def g_(*args):
*res, y = args
f_vjp = tree_unflatten(in_tree, res)
return f_vjp(y)
bwd = jax.make_jaxpr(g_)(*res, y).jaxpr
table = Table(show_header=False, show_lines=True, padding=(1, 2, 0, 2), box=None)
table.add_row("[bold green]forward computation:",
"[bold green]backward computation:")
table.add_row(rich.text.Text.from_ansi(str(fwd)),
rich.text.Text.from_ansi(str(bwd)))
console = Console(width=240, force_jupyter=True)
console.print(table)
def _renderable_repr(self):
return self.html
rich.jupyter.JupyterRenderable._repr_html_ = _renderable_repr
# no use of jax.checkpoint:
print_fwd_bwd(f, W1, W2, W3, x)
forward computation: backward computation: { lambda ; a:f32[5,4] b:f32[6,5] c:f32[7,6] d:f32[4]. let { lambda ; a:f32[7] b:f32[6] c:f32[7,6] d:f32[6] e:f32[5] f:f32[6,5] g:f32[5] h:f32[4] e:f32[5] = dot_general[dimension_numbers=(([1], [0]), ([], []))] a d i:f32[5,4] j:f32[7]. let f:f32[5] = sin e k:f32[7] = mul j a g:f32[5] = cos e l:f32[6] = dot_general[dimension_numbers=(([0], [0]), ([], []))] k c h:f32[6] = dot_general[dimension_numbers=(([1], [0]), ([], []))] b f m:f32[7,6] = dot_general[dimension_numbers=(([], []), ([], []))] k b i:f32[6] = sin h n:f32[6] = mul l d j:f32[6] = cos h o:f32[5] = dot_general[dimension_numbers=(([0], [0]), ([], []))] n f k:f32[7] = dot_general[dimension_numbers=(([1], [0]), ([], []))] c i p:f32[6,5] = dot_general[dimension_numbers=(([], []), ([], []))] n e l:f32[7] = sin k q:f32[5] = mul o g m:f32[7] = cos k r:f32[4] = dot_general[dimension_numbers=(([0], [0]), ([], []))] q i in (l, m, i, c, j, f, b, g, d, a) } s:f32[5,4] = dot_general[dimension_numbers=(([], []), ([], []))] q h in (s, p, m, r) }
# using jax.checkpoint with policy=jax.checkpoint_policies.dots_with_no_batch_dims_saveable:
print_fwd_bwd(f3, W1, W2, W3, x)
forward computation: backward computation: { lambda ; a:f32[5,4] b:f32[6,5] c:f32[7,6] d:f32[4]. let { lambda ; a:f32[5] b:f32[6] c:f32[7] d:f32[5,4] e:f32[6,5] f:f32[7,6] g:f32[4] h:f32[7]. let e:f32[5] = dot_general[dimension_numbers=(([1], [0]), ([], []))] a d i:f32[5,4] j:f32[6,5] k:f32[7,6] l:f32[4] = remat2[ f:f32[5] = sin e differentiated=True g:f32[6] = dot_general[dimension_numbers=(([1], [0]), ([], []))] b f jaxpr={ lambda ; m:f32[5] n:f32[6] o:f32[7] p:f32[5,4] q:f32[6,5] r:f32[7,6] h:f32[6] = sin g s:f32[4] t:f32[7]. let i:f32[7] = dot_general[dimension_numbers=(([1], [0]), ([], []))] c h u:f32[5] = sin m j:f32[7] = sin i v:f32[5] = cos m in (j, e, g, i, a, b, c, d) } w:f32[6] = sin n x:f32[6] = cos n y:f32[7] = cos o z:f32[7] = mul t y ba:f32[6] = dot_general[dimension_numbers=(([0], [0]), ([], []))] z r bb:f32[6] = mul ba x bc:f32[5] = dot_general[dimension_numbers=(([0], [0]), ([], []))] bb q bd:f32[5] = mul bc v be:f32[4] = dot_general[dimension_numbers=(([0], [0]), ([], []))] bd p bf:f32[5,4] = dot_general[dimension_numbers=(([], []), ([], []))] bd s bg:f32[6,5] = dot_general[dimension_numbers=(([], []), ([], []))] bb u bh:f32[7,6] = dot_general[dimension_numbers=(([], []), ([], []))] z w in (bf, bg, bh, be) } policy=<function dot_with_no_batch_dims at 0x7f5e469b1700> prevent_cse=True ] a b c d e f g h in (i, j, k, l) }
让我们一步步思考#
您可能需要先(重新)阅读自动微分实战指南(第一部分)。
jax.checkpoint 的基础#
在 jax.linearize 和 jax.vjp 中,计算某些值的方式和时间具有灵活性。不同的选择可以在内存使用和计算量(FLOPs)之间进行权衡。JAX 通过 jax.checkpoint 提供了对这些选择的控制。
其中一个选择是:是在正向传播时(输入可用时立即计算)进行雅可比系数计算,还是在反向传播时(紧接在需要系数之前)进行计算。考虑 sin_vjp 的示例。
def sin_vjp(x):
y = jnp.sin(x)
cos_x = jnp.cos(x)
return y, lambda y_bar: cos_x * y_bar
另一种有效的实现方式是在反向传播时计算 jnp.cos(x) 的值,而不是在正向传播时。
def sin_vjp2(x):
y = jnp.sin(x)
return y, lambda y_bar: jnp.cos(x) * y_bar
对于这个特定函数,两个版本使用的内存量相同,尽管我们减少了原始计算(即正向传播)的计算量,但增加了余切计算(即反向传播)的计算量。
当涉及函数组合时,还有另一个选择。回想一下我们组合两个函数的 VJP 规则。
def f(x):
y = g(x)
z = h(y)
return z
def f_vjp(x):
y, g_vjp = jax.vjp(g, x)
z, h_vjp = jax.vjp(h, y)
def f_bwd(z_bar):
y_bar, = h_vjp(z_bar)
x_bar, = g_vjp(y_bar)
return x_bar
return z, f_bwd
另一种选择是
def f_vjp_checkpoint(x):
y = g(x)
z, h_vjp = jax.vjp(h, y)
def f_bwd2(z_bar):
y_bar, = h_vjp(z_bar)
_, g_vjp = jax.vjp(g, x)
x_bar, = g_vjp(y_bar)
return x_bar
return z, f_bwd2
简单来说,这种替代实现不会在正向传播时计算 g_vjp 或其闭包中的残差值。相反,它只在反向传播 f_bwd2 中计算它们。这意味着 f_vjp_checkpoint 需要的内存更少:如果 g 和 h 的残差所需的内存量相似,且都远大于 x,那么 f_vjp_checkpoint(x) 产生的函数所需要的内存仅为 f_vjp(x) 的一半!
我们付出的代价是冗余的工作:在 f_bwd2 中,我们必须作为 jax.vjp(g, x) 的一部分重新评估 g(x),仅仅是为了丢弃它的值(在代码行 _, g_vjp = jax.vjp(g, x) 中的下划线变量处)。
我们可以在自动微分中实现这种 VJP 行为——而无需直接编写 VJP 函数——只需在原始函数 f 的替代定义中使用 jax.checkpoint。
def f_checkpoint(x):
y = jax.checkpoint(g)(x)
z = h(y)
return z
换句话说,我们将 jax.checkpoint 应用于 f 的第一阶段 g,而不是 f 本身。这样,当我们评估 jax.grad(f_checkpoint)(x) 时,我们将得到如下计算过程:
运行
g的正向传播,丢弃残差值;运行
h的正向传播,保存残差;运行
h的反向传播,消耗来自步骤 2 的残差;重新运行
g的正向传播,保存残差;运行
g的反向传播,消耗来自步骤 4 的残差。
也就是说,通过评估 jax.grad(f_checkpoint)(x),我们将得到与以下相同的计算结果:
def f_checkpoint_grad(x):
y = g(x) # step 1
_, h_vjp = jax.vjp(h)(y) # step 2
y_bar, = h_vjp(1.0) # step 3
_, g_vjp = jax.vjp(g, x) # step 4
x_bar, = g_vjp(y_bar) # step 5
return x_bar
通常情况下,jax.checkpoint(foo) 是一个新函数,它具有与 foo 相同的输入输出行为,但在自动微分下表现不同,特别是在 jax.linearize 和 jax.vjp(及其包装器,如 jax.grad)下,但在 jax.jvp 下则不然。当进行微分时,在正向传播过程中只存储 jax.checkpoint 微分函数的输入;在反向传播过程中,残差(即 foo 产生的中间值及其反向传播所需的雅可比系数)会被重新计算。
注意,如果 f = lambda x: h(g(x)) 是我们要微分的函数,即如果我们想应用 jax.grad(f),将 jax.checkpoint 应用于 f 本身并不会节省任何内存。这是因为评估 jax.grad(jax.checkpoint(f))(x) 会导致如下计算:
运行正向传播,丢弃所有残差;
立即重新运行正向传播,保存残差;
运行反向传播,消耗来自步骤 2 的残差。
也就是说,用代码表示的话大概是这样:
def f_grad_bad1(x):
_ = f(x) # step 1
_, f_vjp = jax.vjp(f, x) # step 2
x_bar, = f_vjp(1.0) # step 3
return x_bar
将 jax.checkpoint 应用于 f 的第二阶段 h 也不会节省内存。这是因为评估 jax.grad(lambda x: jax.checkpoint(h)(g(x))) 会导致如下计算:
运行
g的正向传播,保存残差;运行
h的正向传播,丢弃残差;立即重新运行
h的正向传播,保存残差;运行
h的反向传播,消耗来自步骤 3 的残差;运行
g的反向传播,消耗来自步骤 1 的残差。
也就是说,用代码表示的话大概是这样:
def f_grad_bad2(x):
y, g_vjp = jax.vjp(g, x) # step 1
_z = h(y) # step 2
_, h_vjp = jax.vjp(h, y) # step 3
y_bar, = h_vjp(1.0) # step 3
x_bar, = g_vjp(y_bar) # step 5
return x_bar
更概括地说,如果我们有一个函数链组合,例如 f = lambda x: f3(f2(f1(x))),并且我们有兴趣评估 jax.grad(f),我们可以说:
我们不应该将
jax.checkpoint应用于整个函数f,因为那不会节省任何内存(且会执行浪费的重新计算);我们不应该将
jax.checkpoint应用于最后一个子函数f3,因为那也不会节省任何内存(且会执行浪费的重新计算);我们可以将
jax.checkpoint应用于f1、f2或它们的组合lambda x: f2(f1(x)),因为其中任何一个都可能节省内存,并会体现不同的内存/重新计算权衡。
关于可保存内容的自定义策略#
如上所示,使用 jax.checkpoint 可以在两个极端之间进行切换:
没有
jax.checkpoint时,JAX 的自动微分倾向于在正向传播中计算一切可能的内容,并将其存储以备反向传播使用;使用
jax.checkpoint装饰器时,我们转而在正向传播中尽可能少地计算,并根据需要在反向传播中重新计算值。
为了在这两个极端之间运作——即保存一些内容而不保存其他内容——我们可以仔细地将 jax.checkpoint 装饰器放置在子函数上。但这需要编辑被微分的函数(例如模型代码),这可能不方便,而且很难尝试不同的组合。
因此,另一种方法是使用 jax.checkpoint 的 policy 参数。策略是一个可调用对象(即函数),它将一阶原语应用的类型级规范作为输入,并返回一个布尔值,指示对应的输出值是否允许被保存为残差(或者必须根据需要在(余)切计算中重新计算)。为了编写稳健的代码,策略应该从 jax.checkpoint_policies 的属性中选择,例如 jax.checkpoint_policies.dots_with_no_batch_dims_saveable,因为编写自定义策略可调用对象的 API 被认为是内部接口。
例如,考虑这个待微分函数:
def loss(params, x, y):
return jnp.sum((predict(params, x) - y)**2)
def predict(params, x):
*Ws, Wlast = params
for W in Ws:
x = layer(W, x)
x = jnp.dot(Wlast, x)
return x
def layer(W, x):
return jnp.sin(jnp.dot(W, x))
W1 = W2 = W3 = jnp.ones((4, 4))
params = [W1, W2, W3]
x = jnp.ones(4)
y = jnp.ones(4)
print_saved_residuals(loss, params, x, y)
f32[4,4] from the argument 'params'
f32[4,4] from the argument 'params'
f32[4,4] from the argument 'params'
f32[4] from the argument 'x'
f32[4] output of sin from <ipython-input-18-3808b5023c3d>:12 (layer)
f32[4] output of cos from <ipython-input-18-3808b5023c3d>:12 (layer)
f32[4] output of sin from <ipython-input-18-3808b5023c3d>:12 (layer)
f32[4] output of cos from <ipython-input-18-3808b5023c3d>:12 (layer)
f32[4] output of mul from <ipython-input-18-3808b5023c3d>:2 (loss)
与其在正向传播中保存那么多值,也许我们只想保存没有批次维度的矩阵乘法结果(因为它们可能是计算密集型而非内存密集型的)。我们可以使用策略 jax.checkpoint_policies.dots_with_no_batch_dims_saveable 来做到这一点。
loss_checkpoint = jax.checkpoint(loss, policy=jax.checkpoint_policies.dots_with_no_batch_dims_saveable)
print_saved_residuals(loss_checkpoint, params, x, y)
f32[4,4] from the argument 'params'
f32[4,4] from the argument 'params'
f32[4,4] from the argument 'params'
f32[4] from the argument 'x'
f32[4] from the argument 'y'
f32[4] output of dot_general from <ipython-input-18-3808b5023c3d>:12 (layer)
f32[4] output of dot_general from <ipython-input-18-3808b5023c3d>:12 (layer)
f32[4] output of dot_general from <ipython-input-18-3808b5023c3d>:8 (predict)
还要注意,通过提供策略,我们不需要编辑定义 loss、predict 或 layer 的代码。如果我们想在调用代码(例如训练脚本)中尝试策略而不更改库代码(例如神经网络库),这特别方便。
有些策略可以引用使用 jax.ad_checkpoint.checkpoint_name 命名的值。
def predict(params, x):
*Ws, Wlast = params
for i, W in enumerate(Ws):
x = layer(W, x)
x = checkpoint_name(x, name=f'layer{i}_output')
x = jnp.dot(Wlast, x)
return x
checkpoint_name 本身只是一个恒等函数。但由于某些策略函数知道如何寻找它们,我们可以利用这些名称来控制 checkpoint_name 输出的某些值是否被视为可保存的。
print_saved_residuals(loss, params, x, y)
f32[4,4] from the argument 'params'
f32[4,4] from the argument 'params'
f32[4,4] from the argument 'params'
f32[4] from the argument 'x'
f32[4] output of cos from <ipython-input-18-3808b5023c3d>:12 (layer)
f32[4] named 'layer0_output' from <ipython-input-22-e48aedf368ad>:7 (predict)
f32[4] output of cos from <ipython-input-18-3808b5023c3d>:12 (layer)
f32[4] named 'layer1_output' from <ipython-input-22-e48aedf368ad>:7 (predict)
f32[4] output of mul from <ipython-input-18-3808b5023c3d>:2 (loss)
loss_checkpoint2 = jax.checkpoint(loss, policy=jax.checkpoint_policies.save_any_names_but_these('layer1_output'))
print_saved_residuals(loss_checkpoint2, params, x, y)
f32[4,4] from the argument 'params'
f32[4,4] from the argument 'params'
f32[4,4] from the argument 'params'
f32[4] from the argument 'x'
f32[4] from the argument 'y'
另一个引用名称的策略是 jax.checkpoint_policies.save_only_these_names。
可以在此处找到策略列表。
策略仅指示什么可以保存;只有在反向传播确实需要时,值才会被实际保存。
进阶:递归 jax.checkpoint#
通过以正确的方式应用 jax.checkpoint,可以实现内存使用和(重新)计算之间的多种权衡。一个令人惊奇的例子是递归检查点,我们将 jax.checkpoint 应用于一个本身调用了 jax.checkpoint 装饰函数的函数,使得 \(D\) 个函数链组合的内存使用量按 \(\mathcal{O}(\log_2 D)\) 而非 \(\mathcal{O}(D)\) 的规模缩放。
作为一个简单的玩具示例,考虑多个 jnp.sin 函数的链式组合:
def chain_compose(funs):
def f(x):
for fun in funs:
x = fun(x)
return x
return f
f = chain_compose([jnp.sin] * 8)
print_saved_residuals(f, 3.)
f32[] output of cos from <ipython-input-25-46b5594773cb>:4 (f)
f32[] output of cos from <ipython-input-25-46b5594773cb>:4 (f)
f32[] output of cos from <ipython-input-25-46b5594773cb>:4 (f)
f32[] output of cos from <ipython-input-25-46b5594773cb>:4 (f)
f32[] output of cos from <ipython-input-25-46b5594773cb>:4 (f)
f32[] output of cos from <ipython-input-25-46b5594773cb>:4 (f)
f32[] output of cos from <ipython-input-25-46b5594773cb>:4 (f)
f32[] output of cos from <ipython-input-25-46b5594773cb>:4 (f)
通常情况下,存储的残差数量随链长线性增长。
f = chain_compose([jnp.sin] * 16)
print_saved_residuals(f, 3.)
f32[] output of cos from <ipython-input-25-46b5594773cb>:4 (f)
f32[] output of cos from <ipython-input-25-46b5594773cb>:4 (f)
f32[] output of cos from <ipython-input-25-46b5594773cb>:4 (f)
f32[] output of cos from <ipython-input-25-46b5594773cb>:4 (f)
f32[] output of cos from <ipython-input-25-46b5594773cb>:4 (f)
f32[] output of cos from <ipython-input-25-46b5594773cb>:4 (f)
f32[] output of cos from <ipython-input-25-46b5594773cb>:4 (f)
f32[] output of cos from <ipython-input-25-46b5594773cb>:4 (f)
f32[] output of cos from <ipython-input-25-46b5594773cb>:4 (f)
f32[] output of cos from <ipython-input-25-46b5594773cb>:4 (f)
f32[] output of cos from <ipython-input-25-46b5594773cb>:4 (f)
f32[] output of cos from <ipython-input-25-46b5594773cb>:4 (f)
f32[] output of cos from <ipython-input-25-46b5594773cb>:4 (f)
f32[] output of cos from <ipython-input-25-46b5594773cb>:4 (f)
f32[] output of cos from <ipython-input-25-46b5594773cb>:4 (f)
f32[] output of cos from <ipython-input-25-46b5594773cb>:4 (f)
但我们可以递归应用 jax.checkpoint 来改善这种缩放比例。
def recursive_checkpoint(funs):
if len(funs) == 1:
return funs[0]
elif len(funs) == 2:
f1, f2 = funs
return lambda x: f1(f2(x))
else:
f1 = recursive_checkpoint(funs[:len(funs)//2])
f2 = recursive_checkpoint(funs[len(funs)//2:])
return lambda x: f1(jax.checkpoint(f2)(x))
f = recursive_checkpoint([jnp.sin] * 8)
print_saved_residuals(f, 3.)
f32[] from the argument 'x'
f32[] output of sin from <ipython-input-27-86f83c871e81>:6 (<lambda>)
f32[] output of cos from <ipython-input-27-86f83c871e81>:6 (<lambda>)
f32[] output of cos from <ipython-input-27-86f83c871e81>:6 (<lambda>)
f = recursive_checkpoint([jnp.sin] * 16)
print_saved_residuals(f, 3.)
f32[] from the argument 'x'
f32[] output of sin from <ipython-input-27-86f83c871e81>:6 (<lambda>)
f32[] output of sin from <ipython-input-27-86f83c871e81>:6 (<lambda>)
f32[] output of cos from <ipython-input-27-86f83c871e81>:6 (<lambda>)
f32[] output of cos from <ipython-input-27-86f83c871e81>:6 (<lambda>)
通常,这里的代价是重新计算:特别是,我们最终执行的计算量(FLOPs)会增加 \(\mathcal{O}(\log_2 D)\) 倍。
f = chain_compose([jnp.sin] * 8)
print_fwd_bwd(f, 3.)
forward computation: backward computation: { lambda ; a:f32[]. let { lambda ; a:f32[] b:f32[] c:f32[] d:f32[] e:f32[] f:f32[] g:f32[] h:f32[] i:f32[]. let b:f32[] = sin a j:f32[] = mul i a c:f32[] = cos a k:f32[] = mul j b d:f32[] = sin b l:f32[] = mul k c e:f32[] = cos b m:f32[] = mul l d f:f32[] = sin d n:f32[] = mul m e g:f32[] = cos d o:f32[] = mul n f h:f32[] = sin f p:f32[] = mul o g i:f32[] = cos f q:f32[] = mul p h j:f32[] = sin h in (q,) } k:f32[] = cos h l:f32[] = sin j m:f32[] = cos j n:f32[] = sin l o:f32[] = cos l p:f32[] = sin n q:f32[] = cos n in (p, q, o, m, k, i, g, e, c) }
f = recursive_checkpoint([jnp.sin] * 8)
print_fwd_bwd(f, 3.)
forward computation: backward computation: { lambda ; a:f32[]. let { lambda ; a:f32[] b:f32[] c:f32[] d:f32[]. let b:f32[] = remat2[ e:f32[] = mul d a differentiated=False f:f32[] = mul e b jaxpr={ lambda ; c:f32[]. let d:f32[] = sin c; e:f32[] = sin d in (e,) } g:f32[] = remat2[ policy=None differentiated=True prevent_cse=True jaxpr={ lambda ; h:f32[] i:f32[]. let ] a j:f32[] = sin h f:f32[] = sin b k:f32[] = cos h g:f32[] = sin f l:f32[] = cos j h:f32[] = sin g m:f32[] = mul i l i:f32[] = sin h n:f32[] = mul m k j:f32[] = sin i in (n,) } k:f32[] = cos i policy=None l:f32[] = sin j prevent_cse=True m:f32[] = cos j ] c f in (l, m, k, g, a) } o:f32[] = remat2[ differentiated=True jaxpr={ lambda ; p:f32[] q:f32[]. let r:f32[] = sin p s:f32[] = sin r t:f32[] = sin s u:f32[] = cos s v:f32[] = cos t w:f32[] = mul q v x:f32[] = mul w u y:f32[] = remat2[ differentiated=True jaxpr={ lambda ; z:f32[] ba:f32[]. let bb:f32[] = sin z bc:f32[] = cos z bd:f32[] = cos bb be:f32[] = mul ba bd bf:f32[] = mul be bc in (bf,) } policy=None prevent_cse=True ] p x in (y,) } policy=None prevent_cse=True ] 3.0 g in (o,) }
实用注意事项#
当微分函数被分阶段(staged out)到 XLA 进行编译时(例如将 jax.jit 应用于包含 jax.grad 调用的函数),XLA 会自动优化计算,包括何时计算或重新物化值的决策。因此,对于 jax.jit 下的微分函数,通常不需要 jax.checkpoint。XLA 会为您完成优化。
一个例外是使用分阶段控制流时,例如 jax.lax.scan。跨多个控制流原语(例如跨正向传播的 scan 和相应的反向传播 scan)的自动编译器优化通常不够彻底。因此,在传递给 jax.lax.scan 的主体函数上使用 jax.checkpoint 通常是一个好主意。
例如,大型 Transformer 模型中的一个常见模式是将架构表示为各层上的 jax.lax.scan,以减少编译时间。也就是说,以简单的全连接网络为例,与其写成这样:
LayerParam = tuple[jnp.ndarray, jnp.ndarray] # weights, bias pair for a layer
ParamsList = list[LayerParam]
def net(params: ParamsList, x: jnp.ndarray):
for W, b in params:
x = jnp.maximum(jnp.dot(x, W) + b, 0.)
return x
我们转而使用 jax.lax.scan 来迭代层应用:
StackedWeights = jnp.ndarray # all weight matrices stacked together
StackedBiases = jnp.ndarray # all bias vectors stacked together
all_weights = jnp.stack([W for W, _ in params])
all_biases = jnp.stack([b for _, b in params])
def layer(x, W_b_pair):
W, b = W_b_pair
out = jnp.maximum(jnp.dot(x, W) + b, 0.)
return out, None
def net(all_weights, all_biases, x):
x, _ = jax.lax.scan(layer, x, (all_weights, all_biases))
return x
这种 scan-over-layers 版本减少了编译时间,但由于阻碍了某些编译器优化,它可能导致梯度的计算效率低下。为了缓解这个问题,我们会在被扫描的函数上使用 jax.checkpoint。
from functools import partial
@partial(jax.checkpoint,
policy=jax.checkpoint_policies.dots_with_no_batch_dims_saveable)
def layer(x, W_b_pair):
W, b = W_b_pair
out = jnp.maximum(jnp.dot(x, W) + b, 0.)
return out, None
通过以这种方式使用 jax.checkpoint,我们是在手动控制 JAX 的自动微分在正向和反向传播之间保存哪些值,因此不需要依赖 XLA 优化来替我们做决定。