news 2026/9/10 6:57:59

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

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
JAX 状态化随机数生成器(jax.experimental.random):从函数式 PRNG 到隐式更新状态的有状态编程实践

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.jitjax.vmapjax.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_rngjax.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_keyArray(typed 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((), 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中对应的函数式采样器。这意味着每次采样都会消耗一个计数器步长

方法签名语义底层采样器
randomrandom(size=None, dtype=float)半开区间[0.0, 1.0)内的随机浮点数jax.random.uniform
uniformuniform(low=0, high=1, size=None, *, dtype=float)区间[low, high)的均匀分布jax.random.uniform
normalnormal(loc=0, scale=1, size=None, *, dtype=float)均值为loc、标准差为scale的正态分布jax.random.normal
integersintegers(low, high=None, size=None, *, dtype=int)区间[low, high)的整数(high省略时等价于[0, low)jax.random.randint

其中size参数支持标量、形状元组(如(5, 2))或Nonesize=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)) )

splitspawn的区别:

  • 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()调用发生了什么

  1. ref_get(self._counter)读出当前计数器值(如 0);
  2. random.fold_in(self._base_key, counter)派生新 key —— 基础 key 不变,只做数学上的折叠(fold-in),统计质量有保证;
  3. ref_addupdate(self._counter, ..., 1)将计数器原地加 1;
  4. 返回新 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.scanjax.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),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/10 6:57:51

医院温湿度监控系统全流程解析:从需求到运维

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/10 6:57:26

TVBoxOSC 完全配置指南:从安装电视盒子播放器到日常使用

TVBoxOSC 完全配置指南&#xff1a;从安装电视盒子播放器到日常使用 【免费下载链接】TVBoxOSC TVBoxOSC - 一个基于第三方项目的代码库&#xff0c;用于电视盒子的控制和管理。 项目地址: https://gitcode.com/GitHub_Trending/tv/TVBoxOSC TVBoxOSC 是一款开源免费的电…

作者头像 李华
网站建设 2026/9/10 6:57:26

车载Android串口开发实战:从UART/RS485到Modbus协议解析

做车载 Android 开发这几年&#xff0c;串口这块踩过的坑比写的代码还多。从最初在调试板上拿 USB 转串口线测 UART&#xff0c;到后来在量产车机上调 RS485 多设备组网&#xff0c;中间经历过电平不匹配烧板子、SELinux 权限搞不定一直打不开设备、Modbus 帧解析各种乱码丢包&…

作者头像 李华
网站建设 2026/9/10 6:56:17

前端版本信息Tags实现:静态注入与动态拉取方案详解

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/10 6:56:01

微信好友申请也能交给程序处理?个人微信二次开发功能介绍

好友申请处理不是一个接口的事&#xff0c;而是一条完整的自动化管线。从收到申请到完成处理&#xff0c;分三个环节。一、感知环节——程序怎么知道有人申请靠好友事件回调。用户发起好友申请时&#xff0c;Eyun 通过 Webhook 推送事件通知&#xff0c;回调数据里带申请人标识…

作者头像 李华
网站建设 2026/9/10 6:49:43

基于STM32与RC522的智能门禁卡系统设计与实现

简介&#xff1a;一套面向电子/嵌入式方向学生的智能家居门禁卡综合管理系统毕业设计与课程设计资料包&#xff0c;覆盖原理图、源码、部署到演示的完整流程。系统基于STM32单片机与RC522射频模块&#xff0c;实现密码开锁、指纹开锁、刷卡开锁&#xff1b;管理员可通过密码进入…

作者头像 李华