JAX 伪随机数生成:基于 Key 的纯函数 PRNG 设计、源码实现与实战指南
【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax
JAX 的随机数系统没有任何隐藏的全局生成器:每一个jax.random采样函数都是其 key 参数的纯函数,同样的 key 永远给出同样的值。本文基于 JAX 官方教程文档 docs/101/random.md 展开,结合 jax/random.py、jax/_src/random/core.py 等仓库源码,完整讲解 key 的创建与派生、"宽树不深链"的 key 管理规则、typed key 与 legacy key 的边界转换,以及底层 Threefry 计数式 PRNG 的设计原理与多实现对比。读完后你将掌握在 JIT、vmap 和多设备分片场景下写出可复现、可并行随机数代码的全部关键实践。
为什么不用全局生成器
NumPy 的numpy.random是典型的有状态设计:一个全局生成器,seed 一次,之后每次采样都悄悄推进内部状态:
import numpy as np np.random.seed(0) print(np.random.uniform()) print(np.random.uniform()) # 不同:隐藏状态已经前进问题在于"隐藏状态前进了"这句话——它让你的结果依赖于程序中每次采样调用的精确顺序和数量:
np.random.seed(0) def bar(): return np.random.uniform() def baz(): return np.random.uniform() def foo(): return bar() + 2 * baz() print(foo())这个值之所以可复现,仅因为 NumPy 承诺bar()一定在baz()之前执行。而这类顺序承诺恰恰是 JAX 必须避免的:优化编译器应当可以自由地重排工作、并在多台设备间并行化。JAX 需要的是可复现(reproducible)、可并行(parallelizable)、可向量化(vectorizable)的随机数生成机制,这直接排除了任何秘密读写共享状态的采样函数。
解决方案并不是把生成器状态变成一个显式参数在函数间来回传递。JAX 的做法是:采样本身就是一个 key 值的纯函数,不产生任何需要穿线带出的更新状态。这一点在 PRNG 设计 JEP docs/jep/263-prng.md 中有完整论证:有状态模型同时违反可复现性、可并行性、编译边界不变性等多条设计要求,因此设计必须走向函数式。
Key 就是值:创建与使用
用整数种子创建 key,调用 {jax.random.key}(实现在 jax/_src/random/core.py 的key函数中):
import jax from jax import random key = random.key(42) key一个 key 是秩为 0 的数组(shape 为()),带有特殊的元素类型:它的 JAX 类型是key<fry>[],其中key<fry>指明默认的 Threefry PRNG 实现:
print(jax.typeof(key)) # 输出 key<fry>[]把 key 传给采样函数不会修改它、也不会以任何物理意义上的方式"消耗"它——样本就是 key 的确定性函数:
print(random.normal(key)) print(random.normal(key)) # 同一个 key,同一个值——必然如此这意味着可复现性是自动获得的:结果只依赖于你的程序构造出的 key 值,永远不依赖执行顺序、调用次数或哪个设备跑了什么。
反过来,不同的随机数需要不同的key。这引出了 JAX 随机数的唯一规则:
绝不复用 key(除非你就是要相同的输出)。把同一个 key 喂给两个不同的采样器会产生相关(correlated)的结果,剥夺你程序的"救命混沌"。
从源码看,这个"纯函数"承诺是有强制校验的:jax/_src/random/core.py 中的_check_prng_key会检查传入的 key 是否为prng_key子类型;对于 legacyuint32裸数组,会根据jax_legacy_prng_key配置发出警告甚至抛出ValueError,防止误用。
源码细节:64 位种子如何变成 key
默认的threefry2x32实现中,key 的内容就是两个uint32(即 64 位)。jax/_src/random/threefry2x32.py 中的threefry_seed/_threefry_seed展示了构造过程:一个 64 位(或 32 位)整数种子被按位拆成高 32 位和低 32 位,分别转成uint32后拼接成 shape(2,)的 key 内容;32 位种子会先补零。这个函数还带有参数校验——seed 必须是标量整数,否则会抛出TypeError(提示批处理请用jax.vmap)。
派生新 key:split 与 fold_in
要获得全新的 key,就从已有的 key 派生。jax.random.split确定性地产生任意数量的新 key,每一个都可以用来生成统计独立的样本(实现见 jax/_src/random/core.py):
key = random.key(42) key, subkey = random.split(key) print(random.normal(subkey))split的num参数默认是 2,也可以一次要多个:
key = random.key(42) subkeys = random.split(key, num=4) [float(random.normal(k)) for k in subkeys]jax.random.fold_in则从一个 key 加一个整数派生新 key,特别适合生成"按步"或"按样本"的 key,而无需在循环里穿线传递 key:
key = random.key(42) for step in range(3): step_key = random.fold_in(key, step) print(f"step {step}: {random.normal(step_key)}")注意这个模式的形状:每个step_key都直接从同一个父 key 派生,而不是从前一个 key 派生。这是刻意设计的。
让 key 树宽而扁,不要又长又深
你程序中的 key 构成一棵树:根是种子,通过split或fold_in生长。糟糕的模式是把树长成长链——每一步的 key 都从上一步的 key split 出来:
for step in range(num_steps): key, subkey = random.split(key) # 每个 key 都派生自前一个:避免! ...应优先使用宽树:所有步 key 都挂在同一个父 key 下,通过一次split(key, num_steps)或者如上文的fold_in(key, step)完成。链式版本有两个问题,一个计算层面、一个统计层面:
- 它强制串行化。每个 key 都依赖于前一个 key,一百万步的链就意味着一百万次顺序执行的哈希运算。而宽式 key 派生是单批操作,可以自由向量化和并行化。
- 它在自找碰撞。对于固定的key,PRNG 底层的哈希是其输入的伪随机置换(permutation),所以一次
split(或对不同整数做fold_in)产生的 key 保证互不相同。但作为key 的函数,该哈希并不是置换,而表现像随机函数。因此树上每一跳派生都是两个 key 可能重合的一次独立机会,在默认的 64 位 key 空间中,长链会把碰撞概率累积到生日界(birthday bound)附近。一旦发生碰撞,意味着从碰撞点开始两条随机流完全相同。
少量链式 split 是无害的:碰撞数学只在大规模时才咬人,大量正确代码都会连续 split 几次。但任何与训练长度或数据集规模成正比的东西,都应从一个公共父 key 宽式派生 key。
无顺序等价性(No sequential equivalence)
NumPy 保证"一次取一个地采 N 个数"与"一次性采 N 个"得到相同序列。JAX 刻意不做这种承诺:
key = random.key(42) subkeys = random.split(key, 3) print("individually:", np.stack([random.normal(k) for k in subkeys])) key = random.key(42) print("all at once: ", random.normal(key, shape=(3,)))顺序等价性会施加的正是 JAX 设计要规避的那种顺序约束。放弃它之后,从独立 key 采出的样本之间不存在任何顺序依赖,生成过程可以自由向量化和分片(sharding)。
由于 key 就是普通数组,它们与 JAX 中的一切都组合。你可以用vmap对一批 key 做向量化的采样:
import jax jax.vmap(random.normal)(subkeys)在默认 PRNG 实现下,这与对每个 key 分别调用random.normal完全等价——对 key 做向量化不会改变数值(对比之下,rbg/unsafe_rbg实验性实现在 vmap 下有特殊行为,详见下文实现对比)。
Typed key 与 legacy key 的边界转换
教程文档使用由jax.random.key创建的 typed key。你可能还会遇到用jax.random.PRNGKey的旧代码,它产生的是裸的uint32数组——仍然可用,但容易被误用(任何东西都拦不住你拿它做算术),而且不记录 key 属于哪个 PRNG 实现。新代码应优先使用jax.random.key,在与需要裸数组的系统交互时用jax.random.key_data和jax.random.wrap_key_data在边界处转换。完整背景见仓库内的 typed PRNG keys JEP docs/jep/9263-typed-keys.md。
两种 key 的具体差异在 jax/random.py 的模块文档中列出:
- typed key 数组(如
key<fry>)自 JAX v0.4.16 引入;此前 key 惯例上表示为uint32数组,其最后一维表示 key 的比特级表示; - 两种形式至今都能创建和使用,typed 用
jax.random.key创建,legacy 用jax.random.PRNGKey创建; - legacy key 的坑:多一个尾随维度;dtype 是数值型(
uint32),允许对 key 做本不该做的整数运算;不携带 RNG 实现信息——把 legacy key 传给jax.random函数时,由全局配置决定使用哪个 RNG 实现。
对应源码接口(jax/_src/random/core.py):
key_data(keys):调用prng.random_unwrap,取回 key 数组底层的比特数据;wrap_key_data(key_bits_array, impl=.../dtype=...):把比特数组重新包装为 typed key,impl与dtype二选一(dtype是推荐写法,impl已被标记为弃用方向),同时指定两者会抛出ValueError;PRNGKey(seed, impl=...):创建 legacy key,同样受jax_default_prng_impl全局配置约束。
wrap_key_data的 docstring 中还给出了往返示例:data = key_data(key),new_key = wrap_key_data(data, dtype=key.dtype),则key == new_key为True。
底层设计:Threefry 计数式 PRNG + 函数式切分模型
JAX 的 PRNG 是基于计数器的 Threefry 哈希与函数式切分模型的组合,选择它的目的就是让生成过程完全没有任何顺序约束。设计动机全文见 docs/jep/263-prng.md,其中列出了七条设计要求:可复现、可并行(采样调用之间无顺序约束)、jit编译边界与后端不变、SIMD 向量化、可扩展到多副本多核分布式计算、与 JAX/XLA 语义契合等。JEP 的 TLDR 概括为:
JAX PRNG = Threefry counter PRNG + a functional array-oriented splitting model
实现层面,jax.random提供 uniform、normal、categorical、permutation 等覆盖很广的分布采样器,全部以 key 作为第一个参数,导出清单见 jax/random.py(uniform、normal、bits、gamma、categorical、permutation、split、fold_in等 40 余个函数均从jax._src.random.core导出)。
可用的 PRNG 实现与选型对比
除默认的 Threefry 外,JAX 还提供多套实现,可通过jax.random.key的impl/dtype参数按 key 选择,或通过全局jax_default_prng_impl配置选择(配置定义见 jax/_src/config.py)。各实现的字符串名称与语义(摘自 jax/random.py 模块文档):
"threefry2x32"(默认)与"threefry4x32":基于 Threefry 哈希变体的计数式 PRNG。threefry2x32有 64 位 key 空间和 64 位计数器空间;threefry4x32有 128 位 key 空间和 128 位计数器空间。"philox2x32"与"philox4x32":基于 Philox 哈希变体的计数式 PRNG。philox2x32有 32 位 key 空间、64 位计数器空间;philox4x32有 64 位 key 空间、128 位计数器空间。"rbg"与"unsafe_rbg"(实验性):构建在 XLA 的 Random Bit Generator 算法之上。"rbg"用 XLA RBG 做采样、用与"threefry2x32"相同的方法做 key 派生;"unsafe_rbg"两者都用 XLA RBG。这类实验方案未经经验随机性测试(如 BigCrush),且两者在vmap下行为特殊:对一批 key 做jax.vmap(jax.random.normal)(keys)时,整批输出只从输入批次中的第一个 key 生成,即等于jax.random.normal(keys[0], shape=(8,))——这是对 XLA RBG 批处理支持有限的一种 workaround,与默认实现"vmap 后逐 key 精确等价"的语义形成鲜明对比。
仓库中的对应实现文件分别位于 jax/_src/random/threefry2x32.py、jax/_src/random/threefry4x32.py、jax/_src/random/philox2x32.py、jax/_src/random/philox4x32.py 和 jax/_src/random/rbg.py,公共框架在 jax/_src/random/prng.py。
选择非默认实现的主要理由是:在 TPU 上编译和执行相对较慢的默认实现可能被优化。模块文档给出的属性对比表如下:
| 属性 | threefry2x32 | threefry2x32* | threefry4x32 | philox | rbg | unsafe_rbg | unsafe_rbg** |
|---|---|---|---|---|---|---|---|
| TPU 上最快 | 支持 | 支持 | 支持 | 支持 | |||
| 可高效分片(配合 pjit) | 支持 | 支持 | 支持 | ||||
| 跨分片方式结果一致 | 支持 | 支持 | 支持 | 支持 | 支持 | ||
| 跨 CPU/GPU/TPU 结果一致 | 支持 | 支持 | 支持 | ||||
对 key 的vmap精确等价 | 支持 | 支持 | 支持 |
*:需要设置jax_threefry_partitionable=1(自 JAX v0.5.0 起默认开启); **:需要设置XLA_FLAGS=--xla_tpu_spmd_rng_bit_generator_unsafe=1。
关于自动分片(automatic partitioning):为了让jax.jit高效地对生成分片随机数组(或 key 数组)的函数做自动分区,"threefry2x32"及"rbg"的 key 派生依赖jax_threefry_partitionable=True(自 v0.5.0 起默认开启);"unsafe_rbg"及"rbg"的采样则需通过XLA_FLAGS环境变量设置上述 XLA 标志。
默认实现是正确的选择,除非你在 profile 中看到了 PRNG 生成的瓶颈。
从源码看 Threefry 哈希与采样路径
默认的 Threefry 哈希在 jax/_src/random/threefry2x32.py 中以apply_round实现经典轮函数:v[0] += v[1]、v[1] = rotate_left(v[1], rot)、v[1] ^= v[0],配合 32 位左旋(_make_rotate_left,用左移与逻辑右移的按位或实现)。抽象求值_threefry2x32_abstract_eval要求全部入参为uint32,输出两个uint32值——这与"key 内容 = 2×uint32"的模型一致。
采样侧的路径同样值得了解:以uniform为例(jax/_src/random/core.py),它先经_check_prng_key校验并规范化 key,再走bits(同文件 L420-L457)按目标 dtype 的位宽生成无符号整数比特流,最后做线性映射到[minval, maxval)。bits还接受可选的out_sharding参数,用于多设备显式分片场景下指定输出分片方式。这些行为在 tests/random_test.py 中有系统性的回归测试覆盖。
小结与下一步
Key 机制让每一个函数保持纯函数:随机性成为显式的、可推导的值,而非隐藏在调用顺序里的状态。核心实践可以浓缩为三条:
- 用
jax.random.key(seed)创建 typed key,新代码不要再写PRNGKey; - 用
split(一次性批量)或fold_in(按步/按样本)从单一父 key 宽式派生,避免长链式 split; - 牢记无顺序等价性:一次性
normal(key, shape=(n,))与逐 key 采样结果不同,这是特性不是缺陷;与裸数组系统交互时用key_data/wrap_key_data在边界转换。
下一篇教程主题是状态(state)——随程序运行而演化、真正原地修改的值,见 docs/101/state.md。
【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考