在 JAX 中编写自定义 Jaxpr 解释器#

Open in Colab Open in Kaggle

JAX 提供了多种可组合的函数变换(jitgradvmap 等),使编写简洁、加速的代码成为可能。

在这里,我们将展示如何通过编写自定义 Jaxpr 解释器,将你自己的函数变换添加到系统中。并且,我们将免费获得与其他所有变换的组合能力。

本示例使用了 JAX 内部 API,这些 API 可能会随时更改。任何不在 API 文档中的内容都应被视为内部 API。

import jax
import jax.numpy as jnp
from jax import jit, grad, vmap
from jax import random

JAX 在做什么?#

JAX 为数值计算提供了类似 NumPy 的 API,可以直接使用,但 JAX 的真正强大之处在于可组合的函数变换。以 jit 函数变换为例,它接收一个函数并返回一个语义相同但由 XLA 为加速器进行延迟编译的函数。

x = random.normal(random.key(0), (5000, 5000))
def f(w, b, x):
  return jnp.tanh(jnp.dot(x, w) + b)
fast_f = jit(f)

当我们调用 fast_f 时,会发生什么?JAX 会跟踪该函数并构建一个 XLA 计算图。随后,该图会被 JIT 编译并执行。其他变换的工作方式类似,即先跟踪函数,然后以某种方式处理输出轨迹。要了解更多关于 Jax 跟踪机制的信息,可以参考 README 中的 “How it works” 章节。

Jaxpr 跟踪器#

Jax 中一个特别重要的跟踪器是 Jaxpr 跟踪器,它将操作记录到 Jaxpr(Jax 表达式)中。Jaxpr 是一种可以像小型函数式编程语言一样求值的数据结构,因此 Jaxpr 是函数变换的有用中间表示。

要初步了解 Jaxpr,可以考虑 make_jaxpr 变换。make_jaxpr 本质上是一种“美化打印”变换:它将一个函数转换为在给定示例参数时,能够产生其计算的 Jaxpr 表示的函数。make_jaxpr 对于调试和内省很有用。让我们用它来看看一些示例 Jaxpr 是如何构建的。

def examine_jaxpr(closed_jaxpr):
  jaxpr = closed_jaxpr.jaxpr
  print("invars:", jaxpr.invars)
  print("outvars:", jaxpr.outvars)
  print("constvars:", jaxpr.constvars)
  for eqn in jaxpr.eqns:
    print("equation:", eqn.invars, eqn.primitive, eqn.outvars, eqn.params)
  print()
  print("jaxpr:", jaxpr)

def foo(x):
  return x + 1
print("foo")
print("=====")
examine_jaxpr(jax.make_jaxpr(foo)(5))

print()

def bar(w, b, x):
  return jnp.dot(w, x) + b + jnp.ones(5), x
print("bar")
print("=====")
examine_jaxpr(jax.make_jaxpr(bar)(jnp.ones((5, 10)), jnp.ones(5), jnp.ones(10)))
foo
=====
invars: [Var(id=128899284404800):int32[]]
outvars: [Var(id=128899030708416):int32[]]
constvars: []
equation: [Var(id=128899284404800):int32[], Literal(TypedInt(1, dtype=int32))] add [Var(id=128899030708416):int32[]] {}

jaxpr: { lambda ; a:i32[]. let b:i32[] = add a 1:i32[] in (b,) }

bar
=====
invars: [Var(id=128899030753280):float32[5,10], Var(id=128899030753152):float32[5], Var(id=128899030752832):float32[10]]
outvars: [Var(id=128899030844992):float32[5], Var(id=128899030752832):float32[10]]
constvars: []
equation: [Var(id=128899030753280):float32[5,10], Var(id=128899030752832):float32[10]] dot_general [Var(id=128899030759744):float32[5]] {'dimension_numbers': (((1,), (0,)), ((), ())), 'precision': None, 'preferred_element_type': dtype('float32'), 'out_sharding': None}
equation: [Var(id=128899030759744):float32[5], Var(id=128899030753152):float32[5]] add [Var(id=128899030844608):float32[5]] {}
equation: [Literal(1.0)] broadcast_in_dim [Var(id=128899030844672):float32[5]] {'shape': (5,), 'broadcast_dimensions': (), 'sharding': None}
equation: [Var(id=128899030844608):float32[5], Var(id=128899030844672):float32[5]] add [Var(id=128899030844992):float32[5]] {}

jaxpr: { lambda ; a:f32[5,10] b:f32[5] c:f32[10]. let
    d:f32[5] = dot_general[
      dimension_numbers=(([1], [0]), ([], []))
      preferred_element_type=float32
    ] a c
    e:f32[5] = add d b
    f:f32[5] = broadcast_in_dim 1.0:f32[]
    g:f32[5] = add e f
  in (g, c) }
  • jaxpr.invars - Jaxpr 的 invars 是 Jaxpr 输入变量的列表,类似于 Python 函数中的参数。

  • jaxpr.outvars - Jaxpr 的 outvars 是由 Jaxpr 返回的变量。每个 Jaxpr 都有多个输出。

  • jaxpr.constvars - constvars 是一个变量列表,它们也是 Jaxpr 的输入,但对应于来自轨迹的常量(我们稍后会更详细地介绍这些内容)。

  • jaxpr.eqns - 等式列表,本质上是 let 绑定。每个等式都是输入变量列表、输出变量列表和一个用于对输入求值以产生输出的 原语 (primitive)。每个等式还有一个 params,即参数字典。

总而言之,Jaxpr 封装了一个简单的程序,可以用输入来求值以产生输出。我们稍后会介绍具体如何做到这一点。现在需要注意的是,Jaxpr 是一种可以按我们想要的任何方式进行操作和求值的数据结构。

为什么 Jaxpr 很有用?#

Jaxpr 是简单的程序表示,易于转换。由于 Jax 允许我们将 Python 函数暂存(stage out)为 Jaxpr,这为我们提供了一种转换用 Python 编写的数值程序的方法。

你的第一个解释器:invert#

让我们尝试实现一个简单的函数“求逆器 (inverter)”,它接收原始函数的输出并返回产生这些输出的输入。目前,我们专注于由其他可逆一元函数组成的简单一元函数。

目标

def f(x):
  return jnp.exp(jnp.tanh(x))
f_inv = inverse(f)
assert jnp.allclose(f_inv(f(1.0)), 1.0)

我们实现的方法是:(1)将 f 跟踪为 Jaxpr,然后(2)反向解释该 Jaxpr。在反向解释 Jaxpr 时,对于每个等式,我们将在表中查找原语的逆并应用它。

1. 跟踪函数#

让我们使用 make_jaxpr 将一个函数跟踪为 Jaxpr。

# Importing Jax functions useful for tracing/interpreting.
from functools import wraps

from jax import lax
from jax.extend import core
from jax._src.util import safe_map

jax.make_jaxpr 返回一个 封闭的 (closed) Jaxpr,即与来自轨迹的常量(literals)捆绑在一起的 Jaxpr。

def f(x):
  return jnp.exp(jnp.tanh(x))

closed_jaxpr = jax.make_jaxpr(f)(jnp.ones(5))
print(closed_jaxpr.jaxpr)
print(closed_jaxpr.literals)
{ lambda ; a:f32[5]. let b:f32[5] = tanh a; c:f32[5] = exp b in (c,) }
[]

2. 求值 Jaxpr#

在我们编写自定义 Jaxpr 解释器之前,先实现“默认”解释器 eval_jaxpr,它按原样求值 Jaxpr,计算出的值与原始、未变换的 Python 函数相同。

为此,我们首先创建一个环境来存储每个变量的值,并随着我们在 Jaxpr 中评估的每个等式更新环境。

def eval_jaxpr(jaxpr, consts, *args):
  # Mapping from variable -> value
  env = {}

  def read(var):
    # Literals are values baked into the Jaxpr
    if type(var) is core.Literal:
      return var.val
    return env[var]

  def write(var, val):
    env[var] = val

  # Bind args and consts to environment
  safe_map(write, jaxpr.invars, args)
  safe_map(write, jaxpr.constvars, consts)

  # Loop through equations and evaluate primitives using `bind`
  for eqn in jaxpr.eqns:
    # Read inputs to equation from environment
    invals = safe_map(read, eqn.invars)
    # `bind` is how a primitive is called
    outvals = eqn.primitive.bind(*invals, **eqn.params)
    # Primitives may return multiple outputs or not
    if not eqn.primitive.multiple_results:
      outvals = [outvals]
    # Write the results of the primitive into the environment
    safe_map(write, eqn.outvars, outvals)
  # Read the final result of the Jaxpr from the environment
  return safe_map(read, jaxpr.outvars)
closed_jaxpr = jax.make_jaxpr(f)(jnp.ones(5))
eval_jaxpr(closed_jaxpr.jaxpr, closed_jaxpr.literals, jnp.ones(5))
[Array([2.1416876, 2.1416876, 2.1416876, 2.1416876, 2.1416876], dtype=float32)]

请注意,eval_jaxpr 总是返回一个扁平列表,即使原始函数不是这样。

此外,该解释器不处理高阶原语(如 jitpmap),我们将在本指南中跳过它们。你可以参考 core.eval_jaxpr (链接) 来查看该解释器未涵盖的边界情况。

自定义 inverse Jaxpr 解释器#

inverse 解释器看起来与 eval_jaxpr 没有太大区别。我们将首先建立将原语映射到其逆的注册表。然后,我们将编写一个在注册表中查找原语的自定义解释器。

事实证明,这个解释器看起来也类似于反向模式自动微分中使用的“转置”解释器,可以在此处找到

inverse_registry = {}

现在,我们将为一些原语注册逆。按照惯例,Jax 中的原语以 _p 结尾,许多流行的原语都位于 lax 中。

inverse_registry[lax.exp_p] = jnp.log
inverse_registry[lax.tanh_p] = jnp.arctanh

inverse 将首先跟踪函数,然后自定义解释 Jaxpr。让我们建立一个简单的骨架。

def inverse(fun):
  @wraps(fun)
  def wrapped(*args, **kwargs):
    # Since we assume unary functions, we won't worry about flattening and
    # unflattening arguments.
    closed_jaxpr = jax.make_jaxpr(fun)(*args, **kwargs)
    out = inverse_jaxpr(closed_jaxpr.jaxpr, closed_jaxpr.literals, *args)
    return out[0]
  return wrapped

现在我们只需要定义 inverse_jaxpr,它将反向遍历 Jaxpr,并在可能的情况下反转原语。

def inverse_jaxpr(jaxpr, consts, *args):
  env = {}

  def read(var):
    if type(var) is core.Literal:
      return var.val
    return env[var]

  def write(var, val):
    env[var] = val
  # Args now correspond to Jaxpr outvars
  safe_map(write, jaxpr.outvars, args)
  safe_map(write, jaxpr.constvars, consts)

  # Looping backward
  for eqn in jaxpr.eqns[::-1]:
    #  outvars are now invars
    invals = safe_map(read, eqn.outvars)
    if eqn.primitive not in inverse_registry:
      raise NotImplementedError(
          f"{eqn.primitive} does not have registered inverse.")
    # Assuming a unary function
    outval = inverse_registry[eqn.primitive](*invals)
    safe_map(write, eqn.invars, [outval])
  return safe_map(read, jaxpr.invars)

就是这样!

def f(x):
  return jnp.exp(jnp.tanh(x))

f_inv = inverse(f)
assert jnp.allclose(f_inv(f(1.0)), 1.0)

重要的是,你可以通过 Jaxpr 解释器进行跟踪。

jax.make_jaxpr(inverse(f))(f(1.))
{ lambda ; a:f32[]. let b:f32[] = log a; c:f32[] = atanh b in (c,) }

这就是将新变换添加到系统的全部内容,并且你可以免费获得与其他所有变换的组合能力!例如,我们可以将 jitvmapgradinverse 一起使用!

jit(vmap(grad(inverse(f))))((jnp.arange(5) + 1.) / 5.)
Array([-3.1440797, 15.584931 ,  2.2551253,  1.3155028,  1.       ],      dtype=float32, weak_type=True)

读者练习#

  • 处理具有多个参数且输入部分已知的情况,例如 lax.add_plax.mul_p

  • 处理 xla_callxla_pmap 原语,这些原语在目前的 eval_jaxprinverse_jaxpr 实现中将无法工作。