JAX 状态化随机数生成器(jax.experimental.random):从函数式 PRNG 到隐式更新状态的有状态编程实践
【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax
导读
JAX 以其纯函数式编程范式著称:jax.jit、jax.vmap、jax.grad等变换都要求被包装的函数是"纯函数",随机数状态因此必须由使用者显式创建、显式更新、显式传递(即经典的jax.random.key+jax.random.split模式)。jax.experimental.random模块提供了一套可选的、基于可变引用(jax.Ref)实现的状态化伪随机数生成器(Stateful PRNG),API 风格对齐numpy.random.default_rng:重复调用rng.uniform()会自动推进内部状态,无需手动管理 key。本文以 docs/jax.experimental.random.rst 所指代的jax.experimental.random模块为核心,结合其底层实现 jax/_src/random/stateful_rng.py、JEP 设计文档 docs/jep/28845-stateful-rng.md 与测试套件 tests/stateful_rng_test.py,系统讲解该模块的 API、底层原理、与jax.jit/jax.vmap等变换的交互方式以及使用边界,帮助你快速上手这一 JAX 原生的有状态随机数编程风格。
一、模块总览:jax.experimental.random
在 JAX 的 API 文档体系中,docs/jax.experimental.random.rst 通过 Sphinx 的automodule指令挂载了jax.experimental.random模块的完整 docstring,并通过autosummary声明了该模块对外暴露的两个核心对象:
stateful_rng:工厂函数,用于创建状态化随机数生成器;StatefulPRNG:状态化随机数生成器类。
从源码看,jax/experimental/random.py 是一个极薄的公共接口层,实际实现位于jax/_src/random/stateful_rng.py:
# jax/experimental/random.py from jax._src.random.stateful_rng import ( stateful_rng as stateful_rng, StatefulPRNG as StatefulPRNG, )也就是说,jax.experimental.random.stateful_rng与jax.experimental.random.StatefulPRNG的真正实现在 jax/_src/random/stateful_rng.py 中,模块 docstring 将其定位为:"Stateful, implicitly-updated PRNG implementation based on mutable refs."(基于可变引用的、隐式更新状态的状态化 PRNG 实现)。
该 API 是可选的:它是对 JAX 经典函数式 PRNG(jax.random)的便捷封装,供"有状态更顺手"的场景使用;对于性能敏感的生产级应用,官方仍推荐使用显式管理 key 的函数式方案(详见下文"适用边界"一节)。
二、快速上手:stateful_rng()工厂函数
2.1 函数签名
stateful_rng(seed: ArrayLike | None = None, *, impl: PRNGSpecDesc | None = None) -> StatefulPRNG参数说明:
seed:可选,一个 64 位或 32 位整数,作为 key 的种子。如果生成器是在被 JAX 变换(如jax.jit)包装的代码内部实例化的,则必须显式指定;在程序顶层使用时可以省略,此时 RNG 会使用 NumPy 默认的种子生成方式(基于np.random.SeedSequence()的熵)自动播种。impl:可选,指定 PRNG 实现的字符串,例如'threefry2x32'(即默认的 Threefry2x32 算法)。
其实现(见 jax/_src/random/stateful_rng.py)本质上是两件事的组合:
return StatefulPRNG( _base_key=random.key(seed, impl=impl), # 底层仍是标准的 typed PRNG key _counter=ref.new_ref(0) # 一个标量整数计数器,包在 Ref 中 )即:StatefulPRNG由一个固定的基础 key(_base_key)与一个可变的计数器引用(_counter)构成。
2.2 最简示例
from jax.experimental import random rng = random.stateful_rng(42) rng # StatefulPRNG(_base_key=Array((), dtype=key<fry>) overlaying: # [ 0 42], _counter=Ref(0, dtype=int32))重复采样会自动更新内部状态:
rng.uniform() # Array(0.5302608, dtype=float32) rng.uniform() # Array(0.72766423, dtype=float32)该行为在jax.jit变换下依然成立:
import jax jit_uniform = jax.jit(rng.uniform) jit_uniform() # Array(0.6672406, dtype=float32) jit_uniform() # Array(0.3890121, dtype=float32)这正是该模块区别于普通 Python 类封装的关键点:状态更新不是"在 Python 层修改属性",而是通过 JAX 的Ref机制(见 jax/ref.py)在编译期正确追踪的隐式更新。
三、StatefulPRNG类的完整 API
StatefulPRNG是一个用dataclasses.dataclass(frozen=True)修饰、并注册为 pytree 的冻结数据类(见 jax/_src/random/stateful_rng.py),包含两个字段:
| 字段 | 类型 | 说明 |
|---|---|---|
_base_key | Array(typed PRNG key) | 固定的基础 key,构造后不再变化 |
_counter | core.Ref(标量整数) | 每次生成 key 时自动加 1 的计数器 |
注意它是**冻结(frozen)**数据类——字段本身不可变,状态更新完全发生在_counter这个Ref内部,这也从数据结构层面保证了它与 JAX 变换兼容的纯函数式语义(从变换的视角看,每次调用只是"读取并写入了一个被追踪的 Ref")。
3.1key(shape=()):生成新的 JAX PRNG key
生成一个新的、与_base_key同实现同 dtype 的独立 PRNG key,同时隐式推进内部状态:
rng = random.stateful_rng(0) rng.key() # Array((), dtype=key<fry>) overlaying: # [1797259609 2579123966] rng.key() # Array((), dtype=key<fry>) overlaying: # [ 928981903 3453687069]shape:可选形状,用于一次返回多个 key。- 若基础 key 本身有形状(即由
split产生的"已拆分生成器"),调用key()会抛出ValueError(源码注释为 "cannot operate on split stateful generator")。
底层实现非常简洁(jax/_src/random/stateful_rng.py):
key = random.fold_in(self._base_key, ref_primitives.ref_get(self._counter)) ref_primitives.ref_addupdate(self._counter, ..., 1) shape_tuple = _canonicalize_size(shape) return random.split(key, shape_tuple) if shape_tuple else key即:用fold_in将固定的基础 key与当前计数器值结合派生出一个新 key,随后通过ref_addupdate把计数器加 1。这样做的好处(详见 JEP 文档 docs/jep/28845-stateful-rng.md 的"Statistical Considerations"一节)是:生成器会完整遍历 32 位或 64 位 key 空间后才循环回初始状态,避免了"反复 split 基础 key"带来的统计相关性隐患。
3.2 随机数采样方法
所有采样方法都遵循同一模式:内部调用self.key()取一个新 key,再交给jax.random中对应的函数式采样器。这意味着每次采样都会消耗一个计数器步长。
| 方法 | 签名 | 语义 | 底层采样器 |
|---|---|---|---|
random | random(size=None, dtype=float) | 半开区间[0.0, 1.0)内的随机浮点数 | jax.random.uniform |
uniform | uniform(low=0, high=1, size=None, *, dtype=float) | 区间[low, high)的均匀分布 | jax.random.uniform |
normal | normal(loc=0, scale=1, size=None, *, dtype=float) | 均值为loc、标准差为scale的正态分布 | jax.random.normal |
integers | integers(low, high=None, size=None, *, dtype=int) | 区间[low, high)的整数(high省略时等价于[0, low)) | jax.random.randint |
其中size参数支持标量、形状元组(如(5, 2))或None;size=None时自动根据其他参数(如low/high/loc/scale的形状)通过np.broadcast_shapes广播出输出形状(见 jax/_src/random/stateful_rng.py 的_canonicalize_size辅助函数)。
示例:
rng = random.stateful_rng(123) rng.uniform(low=-1, high=1, size=(3, 2)) rng.normal(loc=0.0, scale=2.0, size=5) rng.integers(0, 10, 4) # [0, 10) 内的 4 个整数 rng.integers(10, 4) # high 省略:等价于 [0, 10) 内的 4 个整数3.3split(num):拆分出"可映射"的生成器
split生成一个批量化的StatefulPRNG(基础 key 形状为num,计数器为同形状的 Ref),专门用于配合jax.vmap等逐元素映射变换:
import jax import jax.numpy as jnp rng = random.stateful_rng(123) x = jnp.zeros(3) def f(rng, x): return x + rng.uniform() jax.vmap(f)(rng.split(3), x) # Array([0.35525954, 0.21937883, 0.5336956 ], dtype=float32)实现(jax/_src/random/stateful_rng.py):
return StatefulPRNG( _base_key=self.key(num), # 一次生成 num 个独立 key _counter=ref.new_ref(jnp.zeros(num, dtype=int)) )split与spawn的区别:
split(num)返回一个批量化StatefulPRNG对象,适合作为vmap的映射参数(in_axes自动识别);spawn(n_children)返回一个长度为n_children的 Python 列表,每个元素是独立的标量StatefulPRNG,适合在普通 Python 循环或列表推导中使用。
3.4spawn(n_children):生成一组独立子生成器
rng = random.stateful_rng(123) child_rngs = rng.spawn(2) [r.integers(0, 10, 2) for r in child_rngs] # [Array([4, 5], dtype=int32), Array([2, 1], dtype=int32)]每个子生成器拥有不同的_base_key且计数器从 0 开始,彼此完全独立。
四、底层原理:基于Ref的隐式状态更新
4.1 为什么需要Ref
JAX 变换要求函数纯净,因此经典做法是把随机状态当作普通值显式传入传出。StatefulPRNG之所以能"看起来有状态",是因为它把计数器放进了jax.Ref(见 jax/ref.py)——这是 JAX 引入的一种受限可变引用机制,允许在变换内部以受控方式就地更新(详见 docs/array_refs.md 与 docs/stateful-computations.md)。Ref通过 JAX 的**效果系统(effect system)**被编译器正确追踪:读取用ref_get,就地累加用ref_addupdate,变换会把这些操作编译为副作用,而不是在 Python 层"偷偷改属性"。
4.2 一次key()调用发生了什么
ref_get(self._counter)读出当前计数器值(如 0);random.fold_in(self._base_key, counter)派生新 key —— 基础 key 不变,只做数学上的折叠(fold-in),统计质量有保证;ref_addupdate(self._counter, ..., 1)将计数器原地加 1;- 返回新 key(若指定
shape则先split)。
因此,连续调用产生的是fold_in(base_key, 0)、fold_in(base_key, 1)、fold_in(base_key, 2)…… 序列,天然互不相关,且不会被"用过的 key 再次出现"所困扰。
4.3 与函数式 PRNG 的关系
StatefulPRNG的所有采样最终都委托给jax.random的函数式采样器(uniform/normal/randint等),模块 docstring 也明确指出:它是经典无状态 PRNG 的便捷封装。这种设计意味着:
- 通过
rng.key()可以直接拿到标准 typed key,随时"切换回"纯函数式模式; - 状态推进与采样解耦,
fold_in+ 计数器的方案避免了迭代 split 的统计陷阱。
五、与 JAX 变换的交互
5.1jax.jit:支持
状态更新是"隐式"的,在 JIT 下仍能正确推进:
rng = random.stateful_rng(42) jit_uniform = jax.jit(rng.uniform) jit_uniform() # 第一次 jit_uniform() # 第二次,结果不同测试 tests/stateful_rng_test.py 中的testRepeatedDrawsJIT验证了这一点。
5.2jax.vmap:必须先split
由于Ref的限制,不能在 vmapped 函数中直接使用未拆分的rng:
rng = random.stateful_rng(0) def f(x): return x + rng.uniform() jax.vmap(f)(jnp.arange(10)) # Exception: performing an addupdate operation with vmapped value on an # unbatched array reference of type Ref{int32[]}. Move the array # reference to be an argument to the vmapped function?正确用法是把split后的生成器作为参数传入:
def f(x, rng): return x + rng.uniform() jax.vmap(f)(jnp.arange(5), rng.split(5))对应测试 tests/stateful_rng_test.py:testVmapMapped验证 split 用法与 spawn + 列表推导的结果逐元素一致;testVmapUnmapped验证未 split 直接使用会抛出 addupdate 错误。这一限制与shard_map等映射类变换同理(JEP 文档 docs/jep/28845-stateful-rng.md 的 "Interaction with vmap and shard_map" 一节对此有专述)。
5.3jax.lax.scan等控制流:通过闭包捕获使用
StatefulPRNG对象不能作为carry值传入scan/while_loop,但可以在 scan 的函数体内通过闭包捕获:
def f1(seed): rng = random.stateful_rng(seed) def scan_f(_, __): return None, rng.uniform() return jax.lax.scan(scan_f, None, length=10)[1]测试 tests/stateful_rng_test.py 的testScanClosure验证了该用法与 Python 列表推导逐次采样结果一致。
六、适用边界与注意事项
6.1 明确的限制(模块 docstring 原文)
- 不能作为变换函数的返回值:
StatefulPRNG对象不能出现在 JIT 或其他 JAX 变换包装函数的返回值中;尤其意味着不能作为jax.lax.scan、jax.lax.while_loop等控制流原语的carry值。 - 不能与
jax.checkpoint/jax.remat共用:因为Ref依赖 JAX 的效果系统,而remat当前不支持效果;此类场景应改用rng.key()生成标准 key 走函数式路径。
6.2 变换内实例化必须显式给 seed
在变换代码内部调用stateful_rng()且不传seed会直接报错(源码中通过core.trace_ctx.is_top_level()判断是否处于变换追踪上下文):
def f(): return random.stateful_rng().uniform(size=10) # 顶层可以 jax.jit(f)() # TypeError: When used within transformed code, ...测试 tests/stateful_rng_test.py 的testDefaultSeedErrorUnderJIT/Grad/Vmap覆盖了这三种变换下的报错路径。
6.3 工程上的取舍(来自 JEP 的讨论)
- 顺序依赖:有状态采样在程序中引入了天然的串行依赖,编译器无法重排依赖随机数的操作,使用者也不容易在不改变后续采样序列的前提下重构代码(如更换神经网络某一层会"消耗"一个 key,从而改变后续所有层的随机数)。
- 性能:对于性能关键的路径,官方推荐回归
jax.random.key显式管理状态,以获得更充分的编译优化空间与批量生成能力。 - 在
jax.vmap/shard_map等多设备场景下,split未来可能需要增加sharding参数以支持分片语义(JEP 文档中已预告该方向)。
七、结论与延伸阅读
jax.experimental.random为 JAX 提供了一个"低门槛"的状态化随机数编程入口:API 形似numpy.random.default_rng,但底层完全构建在 JAX 原生机制(typed PRNG key +fold_in+Ref)之上,因此天然兼容jax.jit,并可通过split/spawn优雅地适配jax.vmap。对于初入 JAX 的开发者,它显著降低了理解"纯函数式随机状态管理"的陡峭学习曲线,同时保留了通过rng.key()随时切换回函数式范式的通道。
该 API 目前位于experimental命名空间,按 JEP 28845 的规划(见 docs/jep/28845-stateful-rng.md),未来可能正式进入jax.random模块,并可能以default_rng别名出现在jax.numpy.random中。建议继续阅读:
- 函数式 PRNG 的完整 API: docs/jax.random.rst 与 docs/random-numbers.md
- 有状态计算的一般讨论: docs/stateful-computations.md
Ref机制详解: docs/array_refs.md- 完整测试用例(可直接作为用法范本): tests/stateful_rng_test.py
- 底层实现源码: jax/_src/random/stateful_rng.py
【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考