在 JAX 中对副作用进行定序#

sharadmv@ 2022 年 5 月 9 日

概述#

当我们编写 JAX 代码时,通常可以假装自己是在编写单线程、立即执行的 Python 代码,尽管在底层,JAX 及其运行时可能会在后台异步执行这些代码。只要我们编写的是纯函数式(无副作用)的代码,这些性能优化通常对我们是不可见的,也不会干扰我们的单线程思维模型。异步执行非常棒——我们无需任何思考就能获得高性能的并行代码!

然而,在存在副作用的情况下,这种假象开始破裂,我们思维模型的缺陷也随之显现。具体来说,当我们思考副作用发生的顺序时,这些差异就会显现出来。

在本设计说明中,我们将探讨 JAX 的执行模型与副作用排序之间的相互作用。我们还将提供一种强制执行“单线程”副作用顺序的方法。

背景#

当我们编写以下 Python 代码时

def f():
  print("hello")
  return 2
def g():
  print("world")
  return 3
f()
g()

我们期望 "hello""world" 之前被打印。这看起来很显而易见,但请考虑以下 JAX 代码

@partial(jax.jit, device=<device 0>)
def f():
  return 2

@partial(jax.jit, device=<device 1>)
def g():
  return 3
f()
g()

在许多情况下,JAX 会并行执行 fg,将计算分发到不同的线程上——g 实际上可能比 f 先执行。并行执行是一种很好的性能优化,特别是如果与设备之间的数据拷贝非常昂贵时(详见异步分发说明)。然而在实践中,我们通常不需要考虑异步分发,因为我们编写的是纯函数,只关心函数的输入和输出——我们自然会在使用未来值(future values)时进行阻塞等待。

但是,现在想象我们有一个能在 JIT 编译的 JAX 函数内部运行的 jax.print 函数(host_callback.id_print 就是一个例子)。让我们回到上一个例子,但在其中加入打印操作。

@partial(jax.jit, device=<device 0>)
def f():
  jax.print("hello")
  return 2

@partial(jax.jit, device=<device 1>)
def g():
  jax.print("world")
  return 3
f()
g()

得益于异步分发,我们实际上可能会看到 "world""hello" 之前被打印出来。打印副作用的重排序打破了单线程执行模型的假象。

副作用可能“揭示”乱序执行的另一个例子是当我们编译 JAX 程序时。考虑以下 JAX 代码

@jax.jit
def f(x):
  jax.print("hello")
  jax.print("world")
  return x

即使在 Python 中,我们将 "hello" 的打印写在 "world" 的打印之前,像 XLA 这样的编译器也可以自由地对它们进行重排序,因为打印之间没有明确的数据依赖关系。

动机#

我们希望支持“有序”副作用。当我们说有序时,是指副作用发生的顺序与我们执行单线程 Python 程序时的顺序相同。这是我们的主要诉求。在存在诸如 pmap 或用户线程等显式并行的情况下,我们不需要维持这种行为,但至少在用户没有明确要求并行时,我们希望保留单线程的执行顺序。

在深入探讨之前,让我们先退一步,问问自己:为了性能而重排序副作用是否可以接受?反过来,我们是否确实需要强制执行副作用的顺序?在某些情况下,我们不需要排序。也许某些副作用不应该对 JAX 程序的性能产生负面影响。然而,对于其他副作用,我们可能希望强制执行单线程程序顺序,以免用户遇到违反直觉的行为。考虑一个日志记录(logging)副作用。

@jax.jit
def f(x, y):
  log_value(x)
  log_value(y)
f(1, 2)

如果 log 正在修改一个全局列表,我们可能期望先添加 x,再添加 y。对于更严格的副作用,我们可能希望有选项来对这些副作用进行排序。

强制执行有序副作用#

我们强制执行计算顺序的主要工具是数据依赖(data-dependence)。简单来说,如果函数 g 的输入是函数 f 的输出,那么 f 必须在 g 之前执行。

然而,我们可能有一些像打印这样完全没有输入的副作用,所以简单地看,我们无法对它们进行排序。因此,我们使用令牌(tokens)作为向计算中注入人工数据依赖的手段。

什么是令牌?令牌只是一个可以传入和传出计算的虚拟值。通过将同一个令牌传入和传出多个计算,我们强制要求它们必须按特定顺序发生。让我们以之前的打印示例为例,看看在加入令牌后它会是什么样子

@jax.jit
def f(token, x):
  token = jax.print(token, "hello")
  token = jax.print(token, "world")
  return token, x

如果我们重写 jax.print 以接收并返回一个令牌,我们就对两次打印进行了定序,因为第二次打印的输入依赖于第一次打印的输出。令牌的实际值其实可以是任何东西,但在实践中我们会看到令牌对用户是不可见的。

运行时令牌与编译器令牌#

在这里,我们将开始讨论实现细节。在实践中,我们需要两种不同类型的令牌来对副作用进行定序:分别针对上述提到的两种重排序来源。我们需要运行时令牌来对异步分发的副作用计算进行定序,还需要编译器令牌来对计算内部的副作用进行定序。

在实践中,我们的计算将被重写为如下形式

@jax.jit
def f(runtime_token, x):
  compiler_token = new_compiler_token()
  compiler_token = jax.print(compiler_token, "hello")
  compiler_token = jax.print(compiler_token, "world")
  return runtime_token, x

请注意,运行时令牌仅在 JIT 边界处使用,而编译器令牌仅在编译后的代码内使用。编译器令牌是在“降低”(lowering,即将 Python 代码转换为 HLO 或 StableHLO 等底层表示)过程中创建的,但运行时令牌需要在 Python 中进行管理,因为它们需要在 JIT 编译后的函数之间传入传出。

此外,请注意运行时令牌与编译器令牌是“断开连接”的,这意味着它们之间没有数据依赖。这可能存在潜在危险,因为如果丢失了两个已分发函数调用主体之间的数据依赖关系,就会出问题。但是,如果我们假设“严格执行”——即分发后的函数仅在所有输入准备就绪时才开始执行,并且其所有输出将同时准备就绪——那么创建新的编译器令牌并返回与输出无关的运行时令牌是安全的。

管理运行时令牌#

为了代表用户管理运行时令牌,我们需要挂钩到 JAX 的分发机制中。每当我们调用一个 JIT 编译后的函数时,最终都会落到底层的一个类似如下的函数中

def _execute(compiled_computation, *args):
  outputs = compiled_computation.execute(*args)
  return outputs

在这一点上,我们需要将运行时令牌“注入”到计算中,并从计算的输出中将其“提取”出来

def _execute(compiled_computation, *args):
  runtime_token = get_runtime_token() # Grab global token
  runtime_token, *outputs = compiled_computation.execute(runtime_token, *args)
  update_runtime_token(runtime_token) # Update global token
  return outputs

runtime_token 到底是什么?我们需要能够将其传递到 compiled_computation 中,这意味着它需要是某种类型的数组(目前如此,因为在编译后的 JAX 代码内部和外部没有共享的令牌表示)。在实践中,我们可以使用一个 (0,) 形状的数组来最小化开销。

我们还需要考虑多设备使用场景,例如第一个例子:我们先在设备 0 上调用一个 JIT 编译函数,然后在设备 1 上调用另一个。在这种情况下,我们还需要将从第一个计算返回的运行时令牌(存在于设备 0 上)拷贝到设备 1,以便将其传入第二个计算。如果两个连续的计算共享同一个设备,则无需进行此拷贝。

添加编译器令牌#

当我们把 Python 代码降低到 HLO 或 StableHLO 时,我们需要在计算开始时创建一个令牌,并确保在进行需要定序的副作用计算时,令牌是可用的。这些副作用计算将把令牌作为输入并作为输出返回。

这种令牌传递的实现涉及升级 JAX 的降低机制,以自动完成此账本工作。主要挑战在于处理高阶原语(primitives),如调用原语和控制流原语。我们不会在本设计说明中详细讨论如何处理这些情况。

阻塞输出令牌#

为副作用计算添加对运行时和编译器令牌的支持对于定序很重要,但令牌还有另一个细微的用例,即在副作用计算上进行阻塞。即使我们不要求副作用计算必须有序,我们也可能希望等待它完成。目前我们有 jax.block_until_ready,它会等待未来值的结果准备好。然而,对于副作用计算,我们可能会有函数没有返回值但在执行副作用。看看这个简单的例子

@jax.jit
def f():
  jax.print("hello world")
  return
f() # Executed asynchronously

这个编译后的计算没有显式输入,也没有显式输出。如果这是一个有序的打印副作用,我们可以阻塞在返回的运行时令牌上。然而,当这是一个无序计算时,我们不会进行任何令牌传递。当我们没有可供调用 block_until_ready 的输出值时,该如何等待 f() 执行结束呢?好吧,我们可以应用同样的令牌策略,只是我们只返回运行时令牌,而不把它们作为输入。这将给我们一个可以阻塞等待的值,它只有在 f() 执行完毕后才会准备就绪。我们将这些令牌称为输出令牌。我们最终得到的函数看起来像这样

@jax.jit
def f():
  jax.print("hello world")
  return new_runtime_token()
f() # Executed asynchronously

在底层,我们将以管理运行时令牌相同的方式管理输出令牌,但为用户提供一种方法来阻塞当前的输出令牌集合。与运行时令牌不同,输出令牌需要是设备特定(device-specific)的。考虑单设备使用场景

@jax.jit
def f():
  jax.print("hello")

@jax.jit
def g():
  jax.print("world")

f()
g()

由于 f()g() 在同一个设备上执行,阻塞在 g() 的输出令牌上实际上也就阻塞了 f(),因为(目前为止!)JAX 运行时不会交错执行在同一设备上运行的计算。当然,如果这一点发生变化,我们将不得不修改整个设计。

然而,考虑双设备使用场景

@partial(jax.jit, device=<device 0>)
def f():
  jax.print("hello")

@partial(jax.jit, device=<device 1>)
def g():
  jax.print("world")

f()
g()

这里我们不想显式地对 f()g() 进行定序,但希望等待它们两者都执行完毕。我们需要一个针对 f() 的输出令牌和一个针对 g() 的输出令牌,并阻塞在这两个令牌上

@partial(jax.jit, device=<device 0>)
def f():
  jax.print("hello")
  return new_runtime_token()

@partial(jax.jit, device=<device 1>)
def g():
  jax.print("world")
  return new_runtime_token()

t0 = f()
t1 = g()
block_until_ready((t0, t1))

因此,我们需要一个按设备划分的输出令牌,这样我们既可以避免对不同设备上的计算进行定序,又能提供阻塞副作用计算的能力。我们最终对 JAX 分发机制做了以下(近似的)修改

def _execute(compiled_computation, *args):
  output_token, *outputs = compiled_computation.execute(runtime_token, *args)
  update_output_token(output_token, compiled_computation.device)
  return outputs

我们还需要公开一个可以阻塞在输出令牌上的函数

def effects_barrier():
  output_token.block_until_ready()

请注意,阻塞输出令牌可能并不常见,因为大多数 JAX 计算会返回一个可以用于阻塞的值。然而,输出令牌对于测试和性能分析很有帮助,支持它们有助于构建一个一致且内聚的副作用系统。

更多细节#

  • 上述所有令牌管理基础设施都将是线程局部(thread-local)的。这意味着每个用户线程都将拥有自己独立的运行时令牌流。定序仅在用户线程级别保证。

  • 在实践中,每个副作用对应一个运行时令牌。该副作用的不同实例将被定序。这是为了避免对彼此之间可能没有关系的副作用计算进行定序。从技术上讲,这违背了我们最初强制执行单线程 Python 程序顺序的目标,但这是一种可以通过同时拥有“特定副作用”令牌和“全局”令牌来权衡的方案。