ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

JAX 状态化随机数生成器(jax.experimental.random):从函数式 PRNG 到隐式更新状态的有状态编程实践

JAX 状态化随机数生成器(jax.experimental.random):从函数式 PRNG 到隐式更新状态的有状态编程实践 JAX 状态化随机数生成器jax.experimental.random从函数式 PRNG 到隐式更新状态的有状态编程实践【免费下载链接】jaxComposable transformations of PythonNumPy 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.keyjax.random.split模式。jax.experimental.random模块提供了一套可选的、基于可变引用jax.Ref实现的状态化伪随机数生成器Stateful PRNGAPI 风格对齐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 经典函数式 PRNGjax.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_keyrandom.key(seed, implimpl), # 底层仍是标准的 typed PRNG key _counterref.new_ref(0) # 一个标量整数计数器包在 Ref 中 )即StatefulPRNG由一个固定的基础 key_base_key与一个可变的计数器引用_counter构成。2.2 最简示例from jax.experimental import random rng random.stateful_rng(42) rng # StatefulPRNG(_base_keyArray((), dtypekeyfry) overlaying: # [ 0 42], _counterRef(0, dtypeint32))重复采样会自动更新内部状态rng.uniform() # Array(0.5302608, dtypefloat32) rng.uniform() # Array(0.72766423, dtypefloat32)该行为在jax.jit变换下依然成立import jax jit_uniform jax.jit(rng.uniform) jit_uniform() # Array(0.6672406, dtypefloat32) jit_uniform() # Array(0.3890121, dtypefloat32)这正是该模块区别于普通 Python 类封装的关键点状态更新不是在 Python 层修改属性而是通过 JAX 的Ref机制见 jax/ref.py在编译期正确追踪的隐式更新。三、StatefulPRNG类的完整 APIStatefulPRNG是一个用dataclasses.dataclass(frozenTrue)修饰、并注册为 pytree 的冻结数据类见 jax/_src/random/stateful_rng.py包含两个字段字段类型说明_base_keyArraytyped PRNG key固定的基础 key构造后不再变化_countercore.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((), dtypekeyfry) overlaying: # [1797259609 2579123966] rng.key() # Array((), dtypekeyfry) overlaying: # [ 928981903 3453687069]shape可选形状用于一次返回多个 key。若基础 key 本身有形状即由split产生的已拆分生成器调用key()会抛出ValueError源码注释为 cannot operate on split stateful generator。底层实现非常简洁jax/_src/random/stateful_rng.pykey 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中对应的函数式采样器。这意味着每次采样都会消耗一个计数器步长。方法签名语义底层采样器randomrandom(sizeNone, dtypefloat)半开区间[0.0, 1.0)内的随机浮点数jax.random.uniformuniformuniform(low0, high1, sizeNone, *, dtypefloat)区间[low, high)的均匀分布jax.random.uniformnormalnormal(loc0, scale1, sizeNone, *, dtypefloat)均值为loc、标准差为scale的正态分布jax.random.normalintegersintegers(low, highNone, sizeNone, *, dtypeint)区间[low, high)的整数high省略时等价于[0, low)jax.random.randint其中size参数支持标量、形状元组如(5, 2)或NonesizeNone时自动根据其他参数如low/high/loc/scale的形状通过np.broadcast_shapes广播出输出形状见 jax/_src/random/stateful_rng.py 的_canonicalize_size辅助函数。示例rng random.stateful_rng(123) rng.uniform(low-1, high1, size(3, 2)) rng.normal(loc0.0, scale2.0, size5) 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 ], dtypefloat32)实现jax/_src/random/stateful_rng.pyreturn StatefulPRNG( _base_keyself.key(num), # 一次生成 num 个独立 key _counterref.new_ref(jnp.zeros(num, dtypeint)) )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], dtypeint32), Array([2, 1], dtypeint32)]每个子生成器拥有不同的_base_key且计数器从 0 开始彼此完全独立。四、底层原理基于Ref的隐式状态更新4.1 为什么需要RefJAX 变换要求函数纯净因此经典做法是把随机状态当作普通值显式传入传出。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)读出当前计数器值如 0random.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 函数中直接使用未拆分的rngrng 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.pytestVmapMapped验证 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, length10)[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(size10) # 顶层可以 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_inRef之上因此天然兼容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.mdRef机制详解 docs/array_refs.md完整测试用例可直接作为用法范本 tests/stateful_rng_test.py底层实现源码 jax/_src/random/stateful_rng.py【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表