jax.remat / jax.checkpoint 变更:你需要了解的内容#

本文档讨论了在 2022 年 8 月发布的 JAX v0.3.17 中最终确定的 jax.checkpoint(又名 jax.remat)的变更。

目录#

发生了什么?#

#11830 开始,我们启用了 jax.checkpoint()(又名 jax.remat(),两者互为别名)的新实现。对于大多数代码,没有任何变化。但在极端情况下可能会有一些可观察到的差异;请参阅 升级后可能会出现什么问题?

我该如何禁用此变更,暂时回到旧的行为?#

如果您在使用此变更时遇到问题,jax==0.3.16 版本及之前,可以通过将 jax_new_checkpoint 配置选项设置为 False 来关闭新实现,方法如下:

  1. 设置 shell 环境变量 JAX_NEW_CHECKPOINT=0

  2. 执行 jax.config.update('jax_new_checkpoint', False)

  3. 如果您使用 absl 解析标志,请传递 --jax_new_checkpoint=False 选项。

如果您需要恢复到旧实现,请在 GitHub 上提交 issue 联系我们,以便我们能让新实现适配您的需求。

jax==0.3.17 开始,不再提供 jax_new_checkpoint 配置选项。如果您遇到问题,请在 issue 跟踪器上联系我们,以便我们为您解决!

为什么要这样做?#

在撰写本文时,JAX 有两套并行的 jax.checkpoint 实现。新实现已在可选的基础上使用了数月(例如在 Pax 和 Flaxformer/T5X 中),但之前并未默认启用。

我们希望将新实现作为默认设置,然后删除旧实现。使用新实现并删除旧实现可为用户带来多项好处。

用户可自定义的重新物化策略#

新实现的主要优点是对应 policy 参数的新功能。其理念是让用户能够精确控制在自动微分的前向传递过程中保存(或重新物化)哪些中间结果。通过对内存使用与重新计算之间的权衡进行精细控制,用户可以获得显著的性能提升,特别是在大型模型和我们的 MLPerf LLM 提交中!

该功能的完整文档即将推出,但这里有一个简单的示例:

from functools import partial
import jax

def apply_layer(W, x):
  return jnp.sin(jnp.dot(W, x))

@partial(jax.checkpoint, policy=jax.checkpoint_policies.checkpoint_dots)
def predict(params, x):
  for W in params[:-1]:
    x = apply_layer(W, x)
  return jnp.dot(params[-1], x)

通过在此处应用 jax.checkpoint 并设置 policy=jax.checkpoint_policies.checkpoint_dots,我们确保在前向传递期间仅保存矩阵乘法的结果。计算 cos 应用产生的雅可比系数,以及计算它们所需的 sin 应用值,不会在前向传递中保存,而是在反向传递过程中重新计算。(此类策略在 TPU 上非常有效,因为 TPU 上的逐元素计算几乎没有成本,但矩阵运算单元的结果值得保存。)

能够重新物化常量,而不仅仅是那些对参数有数据依赖的操作#

旧的 jax.checkpoint 实现实际上无法在不依赖被装饰函数参数的情况下重新物化计算。考虑这个小例子:

@jax.checkpoint
def f(x):
  a = some_function(jnp.arange(10_000_000))  # `a` does not depend on `x`
  return a * x

旧的 jax.checkpoint 实现被迫保存 a 的值,这可能需要大量内存。新的 jax.checkpoint 实现可以重新物化而不是保存 a 的值。

在某些情况下显著降低 Python 开销#

在某些情况下,新的 jax.checkpoint 产生的 Python 开销显著降低。简单的开销基准测试显示速度提升了 10 倍。这些开销仅出现在即时逐操作执行中,因此在 jax.jit 或类似环境下使用 jax.checkpoint 的常见情况中,这些加速并不相关。但无论如何,这很棒!

通过简化内部结构启用新的 JAX 功能#

此变更还为未来的用户功能带来了重大收益,例如自定义批处理规则(custom_vjpvmap 类比)以及对 custom_vjp 的前向微分升级。它还显著降低了 JAX 代码库部分内容的复杂性,这将有利于整体的可维护性和错误修复。

升级后可能会出现什么问题?#

细微的数值变化#

由于新实现可以重新物化更多的计算,包括可能很大的常量,某些代码可能会出现微小的数值变化。任何数值变化的幅度都应在我们预期编译器优化更改(例如浮点运算重排序)的范围内。但某些过于严格的测试容差可能需要适当放宽。

移除了 concrete=True 选项。#

旧的 jax.checkpoint 实现有一个布尔值 concrete 选项,它允许在具体的 Python 值上进行跟踪(而不是延迟所有计算并仅在抽象值上进行跟踪)。该选项很少使用,且在某些使用场景下有更简单的替代方案。因此,我们在新的 jax.checkpoint 中移除了该选项。

例如,Google 代码中绝大多数使用 concrete=True 的场景是为了支持传递类似 is_training 这样的参数。

@partial(jax.checkpoint, concrete=True)  # OLD jax.checkpoint API
def foo(x, is_training):
  if is_training:
    return g(x)
  else:
    return h(x)

使用新的 jax.checkpoint 实现,我们可以通过 static_argnums 选项达到同样的效果。

@partial(jax.checkpoint, static_argnums=(1,))  # NEW jax.checkpoint API
def foo(x, is_training):
  if is_training:
    ...

如果需要在静态参数上执行 jax.numpy 操作,并在 Python 跟踪过程中计算其数值结果而不是延迟计算,我们可以结合使用 static_argnumsjax.ensure_compile_time_eval()。但似乎不太可能需要这样做!