JEP 9263:类型化密钥与可插拔 RNG#

Jake VanderPlas, Roy Frostig

2023 年 8 月

概述#

展望未来,JAX 中的 RNG 密钥将更具类型安全性和可定制性。单个 PRNG 密钥将不再由长度为 2 的 uint32 数组表示,而是表示为具有特殊 RNG 数据类型的标量数组,且满足 jnp.issubdtype(key.dtype, jax.dtypes.prng_key)

目前,旧式 RNG 密钥仍可通过 jax.random.PRNGKey() 创建。

>>> key = jax.random.PRNGKey(0)
>>> key
Array([0, 0], dtype=uint32)
>>> key.shape
(2,)
>>> key.dtype
dtype('uint32')

即日起,新型 RNG 密钥可通过 jax.random.key() 创建。

>>> key = jax.random.key(0)
>>> key
Array((), dtype=key<fry>) overlaying:
[0 0]
>>> key.shape
()
>>> key.dtype
key<fry>

这种(标量形状的)数组的行为与任何其他 JAX 数组相同,不同之处在于其元素类型是密钥(及相关元数据)。我们也可以创建非标量密钥数组,例如通过将 jax.vmap() 应用于 jax.random.key()

>>> key_arr = jax.vmap(jax.random.key)(jnp.arange(4))
>>> key_arr
Array((4,), dtype=key<fry>) overlaying:
[[0 0]
 [0 1]
 [0 2]
 [0 3]]
>>> key_arr.shape
(4,)

除了切换到新的构造函数外,大多数与 PRNG 相关的代码应继续按预期工作。您可以像以前一样在 jax.random API 中继续使用密钥;例如:

# split
new_key, subkey = jax.random.split(key)

# random number generation
data = jax.random.uniform(key, shape=(5,))

然而,并非所有的数值运算都适用于密钥数组。它们现在会刻意引发错误。

>>> key = key + 1  
Traceback (most recent call last):
TypeError: add does not accept dtypes key<fry>, int32.

如果由于某种原因您需要恢复底层缓冲区(旧式密钥),可以使用 jax.random.key_data() 来实现。

>>> jax.random.key_data(key)
Array([0, 0], dtype=uint32)

对于旧式密钥,key_data() 是恒等操作。

这对用户意味着什么?#

对于 JAX 用户,此更改目前不需要对代码进行任何修改,但我们希望您会发现升级很有价值并改用类型化密钥。要进行尝试,请将 jax.random.PRNGKey() 的使用替换为 jax.random.key()。这可能会导致您的代码在以下几种情况下出现中断:

  • 如果您的代码对密钥执行不安全/不支持的操作(例如索引、算术、转置等;请参阅下文“类型安全”部分),此更改将捕获这些操作。您可以更新代码以避免此类不支持的操作,或者使用 jax.random.key_data()jax.random.wrap_key_data() 以不安全的方式操作原始密钥缓冲区。

  • 如果您的代码包含关于 key.shape 的显式逻辑,您可能需要更新此逻辑,以说明末尾的密钥缓冲区维度不再是 shape 的显式组成部分这一事实。

  • 如果您的代码包含关于 key.dtype 的显式逻辑,则需要将其升级为使用新的公共 API 来判断 RNG 数据类型,例如 dtypes.issubdtype(dtype, dtypes.prng_key)

  • 如果您调用的基于 JAX 的库尚未处理类型化 PRNG 密钥,您可以暂时使用 raw_key = jax.random.key_data(key) 来恢复原始缓冲区,但请记录一个 TODO,以便在下游库支持类型化 RNG 密钥后将其移除。

未来某个时候,我们计划弃用 jax.random.PRNGKey() 并强制使用 jax.random.key()

检测新型类型化密钥#

要检查某个对象是否为新型类型化 PRNG 密钥,可以使用 jax.dtypes.issubdtypejax.numpy.issubdtype

>>> typed_key = jax.random.key(0)
>>> jax.dtypes.issubdtype(typed_key.dtype, jax.dtypes.prng_key)
True
>>> raw_key = jax.random.PRNGKey(0)
>>> jax.dtypes.issubdtype(raw_key.dtype, jax.dtypes.prng_key)
False

PRNG 密钥的类型注解#

推荐用于旧式和新型 PRNG 密钥的类型注解是 jax.Array。PRNG 密钥与其数组的区别在于其 dtype,目前无法在类型注解中指定 JAX 数组的数据类型。此前可以使用 jax.random.KeyArrayjax.random.PRNGKeyArray 作为类型注解,但它们在类型检查下始终被别名为 Any,因此 jax.Array 具有更高的特异性。

注意:jax.random.KeyArrayjax.random.PRNGKeyArray 在 JAX 0.4.16 版本中被弃用,并在 0.4.24 版本中被移除。.

致 JAX 库作者的说明#

如果您维护一个基于 JAX 的库,您的用户也是 JAX 用户。请知悉,JAX 目前将在 jax.random 中继续支持“原始”旧式密钥,因此调用者可能期望它们在任何地方都被接受。如果您希望在库中强制要求使用新型类型化密钥,则可能需要按照以下思路进行检查:

from jax import dtypes

def ensure_typed_key_array(key: Array) -> Array:
  if dtypes.issubdtype(key.dtype, dtypes.prng_key):
    return key
  else:
    raise TypeError("New-style typed JAX PRNG keys required")

动机#

推动这一变化的两个主要因素是可定制性和安全性。

自定义 PRNG 实现#

JAX 目前使用单个全局配置的 PRNG 算法运行。PRNG 密钥是一个无符号 32 位整数向量,jax.random API 使用它来生成伪随机流。任何更高秩的 uint32 数组都被解释为此类密钥缓冲区的数组,其中末尾维度表示密钥。

随着我们引入必须通过设置全局或局部配置标志来选择的替代 PRNG 实现,这种设计的弊端变得更加明显。不同的 PRNG 实现具有不同大小的密钥缓冲区,以及用于生成随机位的不同算法。使用全局标志确定此行为容易出错,特别是在整个进程中同时使用多种密钥实现时。

我们的新方法是将实现作为 PRNG 密钥类型的一部分,即作为密钥数组的元素类型携带。使用新的密钥 API,下面是一个在默认 threefry2x32 实现(由纯 Python 实现并由 JAX 编译)和非默认 rbg 实现(对应于单个 XLA 随机位生成操作)下生成伪随机值的示例。

>>> key = jax.random.key(0, impl='threefry2x32')  # this is the default impl
>>> key
Array((), dtype=key<fry>) overlaying:
[0 0]
>>> jax.random.uniform(key, shape=(3,))
Array([0.947667  , 0.9785799 , 0.33229148], dtype=float32)

>>> key = jax.random.key(0, impl='rbg')
>>> key
Array((), dtype=key<rbg>) overlaying:
[0 0 0 0]
>>> jax.random.uniform(key, shape=(3,))
Array([0.39904642, 0.8805201 , 0.73571277], dtype=float32)

安全地使用 PRNG 密钥#

从原则上讲,PRNG 密钥实际上仅旨在支持少数几种操作,即密钥派生(例如拆分)和随机数生成。PRNG 的设计初衷是生成独立的伪随机数,前提是正确地拆分密钥且每个密钥仅被消耗一次。

以其他方式操作或消耗密钥数据的代码通常表示出现了意外的错误,而将密钥数组表示为原始 uint32 缓冲区导致了此类滥用的易发性。以下是我们实际遇到的一些误用示例:

密钥缓冲区索引#

对底层整数缓冲区的访问使得用户很容易尝试以非标准方式派生密钥,有时会产生意外的严重后果。

# Incorrect
key = random.PRNGKey(999)
new_key = random.PRNGKey(key[1])  # identical to the original key!
# Correct
key = random.PRNGKey(999)
key, new_key = random.split(key)

如果此密钥是使用 random.key(999) 创建的新型类型化密钥,则对密钥缓冲区进行索引会报错。

密钥算术运算#

密钥算术运算是另一种从现有密钥中派生密钥的危险方式。通过直接操作密钥数据来绕过 jax.random.split()jax.random.fold_in() 派生密钥,会产生一批密钥,这些密钥根据 PRNG 实现的不同,随后在该批次内生成相关的随机数。

# Incorrect
key = random.PRNGKey(0)
batched_keys = key + jnp.arange(10, dtype=key.dtype)[:, None]
# Correct
key = random.PRNGKey(0)
batched_keys = random.split(key, 10)

使用 random.key(0) 创建的新型类型化密钥通过禁止对密钥进行算术运算解决了这个问题。

无意中对密钥缓冲区进行转置#

使用“原始”旧式密钥数组,很容易无意中交换批处理(前导)维度和密钥缓冲区(末尾)维度。同样,这可能会产生导致生成相关伪随机数的密钥。我们长期以来观察到的一个模式可以归结为:

# Incorrect
keys = random.split(random.PRNGKey(0))
data = jax.vmap(random.uniform, in_axes=1)(keys)
# Correct
keys = random.split(random.PRNGKey(0))
data = jax.vmap(random.uniform, in_axes=0)(keys)

这里的错误很微妙。通过对 in_axes=1 进行映射,此代码通过组合每个批处理密钥缓冲区中的单个元素来制作新密钥。生成的密钥彼此不同,但实际上是以非标准方式“派生”的。同样,PRNG 并未设计或测试为从这种密钥批次中产生独立的随机流。

使用 random.key(0) 创建的新型类型化密钥通过隐藏单个密钥的缓冲区表示,转而将密钥视为密钥数组的不透明元素,解决了此问题。密钥数组没有可供索引、转置或映射的末尾“缓冲区”维度。

密钥重用#

与基于状态的 PRNG API(如 numpy.random)不同,JAX 的函数式 PRNG 在密钥被使用后不会隐式更新它。

# Incorrect
key = random.PRNGKey(0)
x = random.uniform(key, (100,))
y = random.uniform(key, (100,))  # Identical values!
# Correct
key = random.PRNGKey(0)
key1, key2 = random.split(random.key(0))
x = random.uniform(key1, (100,))
y = random.uniform(key2, (100,))

我们正在积极开发工具来检测和防止意外的密钥重用。这仍在开发中,但它依赖于类型化的密钥数组。现在升级到类型化密钥,为我们后续在构建这些安全功能时奠定了基础。

类型化 PRNG 密钥的设计#

类型化 PRNG 密钥在 JAX 中作为扩展数据类型 (extended dtypes) 的实例实现,其中新的 PRNG 数据类型是其子数据类型。

扩展数据类型#

从用户角度来看,扩展数据类型 dt 具有以下用户可见属性:

  • jax.dtypes.issubdtype(dt, jax.dtypes.extended) 返回 True:这是应该用于检测数据类型是否为扩展数据类型的公共 API。

  • 它具有类级属性 dt.type,返回 numpy.generic 层次结构中的类型类。这类似于 np.dtype('int32').type 返回 numpy.int32 的方式,后者不是数据类型,而是标量类型,也是 numpy.generic 的子类。

  • 与 numpy 标量类型不同,我们不允许实例化 dt.type 标量对象:这符合 JAX 将标量值表示为零维数组的决定。

从非公开的实现角度来看,扩展数据类型具有以下属性:

  • 其类型是私有基类 jax._src.dtypes.ExtendedDtype 的子类,这是用于扩展数据类型的非公开基类。 ExtendedDtype 的一个实例类似于 np.dtype 的实例,如 np.dtype('int32')

  • 它具有私有 _rules 属性,允许数据类型定义其在特定操作下的行为。例如,当 dtype 是扩展数据类型时,jax.lax.full(shape, fill_value, dtype) 将委托给 dtype._rules.full(shape, fill_value, dtype)

为什么要广泛引入扩展数据类型,而不仅仅限于 PRNG?我们在内部其他地方也复用了这种相同的扩展数据类型机制。例如,jax._src.core.bint 对象(一种用于动态形状实验工作的有界整数类型)是另一种扩展数据类型。在最近的 JAX 版本中,它满足上述属性(参见 jax/_src/core.py#L1789-L1802)。

PRNG 数据类型#

PRNG 数据类型被定义为扩展数据类型的一种特殊情况。具体而言,此更改引入了一个新的公共标量类型类 jax.dtypes.prng_key,它具有以下属性:

>>> jax.dtypes.issubdtype(jax.dtypes.prng_key, jax.dtypes.extended)
True

PRNG 密钥数组因此具有具有以下属性的数据类型:

>>> key = jax.random.key(0)
>>> jax.dtypes.issubdtype(key.dtype, jax.dtypes.extended)
True
>>> jax.dtypes.issubdtype(key.dtype, jax.dtypes.prng_key)
True

除了如上所述针对扩展数据类型一般的 key.dtype._rules 外,PRNG 数据类型还定义了 key.dtype._impl,其中包含定义 PRNG 实现的元数据。PRNG 实现目前由非公开的 jax._src.prng.PRNGImpl 类定义。目前,PRNGImpl 并非公共 API,但我们可能会很快重新审视这一点,以允许实现完全自定义的 PRNG。

进度#

以下是实现上述设计的关键拉取请求 (Pull Request) 的非完整列表。主要追踪 issue 是 #9263

  • 通过 PRNGImpl 实现可插拔 PRNG:#6899

  • 实现不带数据类型的 PRNGKeyArray#11952

  • 为带有 _rules 属性的 PRNGKeyArray 添加“自定义元素”数据类型属性:#12167

  • 将“自定义元素类型”重命名为“不透明数据类型 (opaque dtype)”:#12170

  • 重构 bint 以使用不透明数据类型基础架构:#12707

  • 添加 jax.random.key 以直接创建类型化密钥:#16086

  • keyPRNGKey 添加 impl 参数:#16589

  • 将“不透明数据类型”重命名为“扩展数据类型”并定义 jax.dtypes.extended#16824

  • 引入 jax.dtypes.prng_key 并将 PRNG 数据类型与扩展数据类型统一:#16781

  • 添加 jax_legacy_prng_key 标志,以在处理遗留(原始)PRNG 密钥时支持发出警告或报错:#17225