ARTICLE DETAIL

资讯详情

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

JAX jax.random 模块完全指南:PRNG Key 体系、采样器 API 与底层实现解析

JAX jax.random 模块完全指南:PRNG Key 体系、采样器 API 与底层实现解析 JAX jax.random 模块完全指南PRNG Key 体系、采样器 API 与底层实现解析【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax导读jax.random是 JAX 中确定性伪随机数生成PRNG的核心模块它以显式传递的key密钥取代传统库的全局隐式随机状态提供从均匀分布、正态分布到 categorical、permutation 等四十余种采样函数并内置 Six 套可切换的 PRNG 实现Threefry / Philox / XLA RBG。读完本文你将掌握 key 的创建、分裂与格式转换 API理解类型化 keykeyfry与传统uint32key 的区别并能依据官方对比表为 TPU、多分片pjit等场景正确选择 PRNG 实现。1. 模块定位纯函数式随机数生成jax.random的模块说明见 jax/random.py明确指出该包提供一系列确定性生成伪随机数序列的例程。与 NumPy / SciPy 用户习惯的有状态PRNG 不同JAX 的所有随机函数都要求将一个显式的 PRNG 状态作为第一个参数传入这个状态是一种特殊的数组元素类型称为key通常由jax.random.key生成from jax import random key random.key(0) key # Array((), dtypekeyfry) overlaying: # [0 0]该 key 可以直接用于 JAX 的任何随机数生成例程random.uniform(key) # Array(0.947667, dtypefloat32)一个关键设计是使用 key 不会修改它。同一个 key 重复使用会得到完全相同的结果random.uniform(key) # Array(0.947667, dtypefloat32) # 与上次相同如果需要新的随机数必须用jax.random.split派生新的子 keykey, subkey random.split(key) random.uniform(subkey) # Array(0.00729382, dtypefloat32)这种key 即值的纯函数设计从源码角度同样可验证jax.random只是一个薄薄的导出层所有函数均从 jax/_src/random/core.py 以import ... as ...形式重导出见 jax/random.py保证采样函数不读取任何全局可变状态。1.1 为什么不用全局生成器在 JAX 的 101 教程 docs/101/random.md 中对此有系统阐述numpy.random的全局生成器意味着隐藏状态被推进你的结果取决于程序各处采样调用的确切顺序和次数。而 JAX 的编译器需要自由地重排、跨设备并行化计算因此随机数生成必须满足三点可复现reproducible结果只由 key 值决定可并行parallelizable多副本、多核计算中随机函数调用之间不存在顺序约束可向量化vectorizable生成数组值时能自由使用vmap等变换。这正是jax.random文档 Design and background 一节jax/random.py所强调的设计目标也是JAX PRNG Threefry 计数器 PRNG 函数式数组化分裂模型这一 TLDR 的由来。2. API ReferenceKey 创建与操作jax.random的 API 文档将密钥相关函数归为 Key Creation Manipulation 一组共 8 个成员见 docs/jax.random.rst函数用途key由整数种子创建新式类型化PRNG keykey_dtype获取某个 PRNG 实现对应的 key dtypekey_data取回 PRNG key 底层的数据位uint32数组wrap_key_data将 key 数据位数组包装回 PRNG key 数组fold_in将数据折叠进 key 以派生新 keysplit将 key 分裂为num个新 key增加前导轴clone克隆 keyPRNGKey由整数种子创建旧式uint32key2.1key创建类型化 keyjax.random.key(seed, *, implNone, dtypeNone)接受一个 64 位或 32 位整数种子返回一个标量 PRNG key 数组其 dtype 指示所使用的 PRNG 实现源码定义见 jax/_src/random/core.py。impl与dtype二者只能指定其一且在新版本中impl已被dtype取代dtype可直接传实现名字符串如threefry2x32。若都不传则采用全局配置jax_default_prng_impl决定的默认实现。2.2 新旧两种 key 格式模块文档特别提示jax/random.py类型化 key 数组元素类型如keyfry于 JAX v0.4.16 引入。在此之前key 习惯上用uint32数组表示其最后一维代表 key 的比特级表示。新式jax.random.key创建dtype 是jax.dtypes.prng_key的子类型如keyfry推荐用于新代码旧式jax.random.PRNGKey创建是uint32数组。旧式 key 有三个缺陷模块文档原文多一个尾部维度具有数值 dtypeuint32允许对 key 做整数运算等本不该做的操作不携带 RNG 实现信息传给jax.random函数时由全局配置决定实现。两者可用key_data/wrap_key_data互转。旧格式在需要与 JAX 外部系统对接时如导出为可序列化格式仍有必要。源码层面PRNGKey的实现jax/_src/random/core.py会先创建类型化 key再通过_return_prng_keys解包成uint32数组。key_data/wrap_key_data的可逆性有官方 docstring 示例佐证jax/_src/random/core.pykey jax.random.key(42) data jax.random.key_data(key) new_key jax.random.wrap_key_data(data, dtypekey.dtype) key new_key # Array(True, dtypebool)2.3split与fold_in派生新 key 的两种途径split(key, num2)把 key 分裂成num个新 key通过新增前导轴实现源码 jax/_src/random/core.py。num可以是正整数或整数元组指定分裂结果的形状。fold_in(key, data)把 32 位整数data折叠进 key 形成新 key源码 jax/_src/random/core.py。它要求data为标量非常适合在不贯穿 key 线程的前提下为每步训练或每个样本派生 key。教程 docs/101/random.md 给出了典型用法并强调一个重要实践准则保持 key 树宽而不是深。key random.key(42) subkeys random.split(key, num4) # 一次分裂出 4 个 for step in range(3): step_key random.fold_in(key, step) # 直接从父 key 派生深链式分裂每步 key 由上一步 key 分裂而来有两个问题其一串行化——百万步意味着百万次顺序哈希而宽派生是一次可向量化、可并行化的批量操作其二易碰撞——对固定 key底层哈希是其输入上的伪随机置换故单次split产生的 key 保证互异但将哈希视作key 的函数时它并非置换而是随机函数每多一跳就多一次碰撞机会在默认 64 位 key 空间中长链会累积到生日界。少量链式分裂无害但与训练步数或数据集规模成正比时应从公共父 key 宽派生。3. API Reference随机采样器Random Samplersjax.random通过autosummary自动生成采样函数清单见 docs/jax.random.rst完整列表如下ball、bernoulli、beta、binomial、bits、categorical、cauchy、chisquare、choice、dirichlet、double_sided_maxwell、exponential、f、gamma、generalized_normal、geometric、gumbel、laplace、loggamma、logistic、lognormal、maxwell、multinomial、multivariate_normal、normal、orthogonal、pareto、permutation、poisson、rademacher、randint、rayleigh、t、triangular、truncated_normal、uniform、wald、weibull_min所有这些采样器都遵循同一约定key 作为第一个参数。它们的全部实现都位于 jax/_src/random/core.py如uniform在 L469、randint在 L592、normal在 L911、categorical在 L2339、binomial在 L3632并在 jax/random.py 统一导出。3.1 采样器通用参数约定以uniform为例源码 jax/_src/random/core.py其签名与默认值展示了采样器的通用结构uniform(key, shape(), dtypeNone, minval0., maxval1., *, out_shardingNone)shape结果形状默认()dtype浮点 dtype默认在jax_enable_x64开启时为float64否则为float32minval/maxval区间[minval, maxval)需与shape广播兼容默认[0, 1)out_sharding多设备计算中输出数组的分片规范可为NamedSharding或PartitionSpecP默认None主要在显式分片模式下使用。uniform内部先做 dtype 规范化与core.canonicalize_shape形状规整再经maybe_auto_axes包装后调用jit(static_argnums(3, 4))修饰的内部实现_uniformjax/_src/random/core.py其中把minval/maxval转换为目标 dtype 并广播到目标 rank。也就是说每个采样器最终都编译为 XLA 计算天然融入 JIT 流程。3.2 低层工具bitsbits(key, shape(), dtypeNone, *, out_shardingNone)jax/_src/random/core.py是唯一直接暴露原始随机比特的采样器返回无符号整数形式的均匀随机位dtype必须是无符号整型默认jax_enable_x64开启时为uint64否则uint32bit_width dtype.itemsize * 8最终送入底层prng.random_bits。它相当于其它分布采样器的地基——所有连续/离散分布最终都从这类基础随机位经变换与拒绝采样算法产生。4. 高级 RNG 配置六套 PRNG 实现JAX 提供多套 PRNG 实现模块文档 Advanced RNG configuration 一节jax/random.py。选择方式有二在jax.random.key创建时传impl/dtype关键字或在创建时省略该参数由全局配置jax_default_prng_impl决定。可用实现的名字如下实现名说明threefry2x32默认基于 Threefry 哈希变体的计数器 PRNG64 位 key 空间 64 位计数器空间threefry4x32同上128 位 key 空间 128 位计数器空间philox2x32基于 Philox 哈希变体的计数器 PRNG32 位 key 空间 64 位计数器空间philox4x32同上64 位 key 空间 128 位计数器空间rbg实验性基于 XLA Random Bit Generator (RBG) 算法生成用 RBGkey 派生复用threefry2x32的方法unsafe_rbg实验性生成与 key 派生均用 XLA RBG三种基础算法Threefry / Philox 的 2x32、4x32 变体与 RBG均有独立实现文件位于 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 的PRNGImpl机制register_prng注册实现resolve_prng_impl负责按名字解析见 jax/_src/random/core.py。4.1 实验性 RBG 实现的注意点模块文档明确警示jax/random.pyrbg与unsafe_rbg生成的随机数未经BigCrush 等经验随机性测试unsafe_rbg的 key 派生质量也未经验证名字中的 unsafe 即强调其 key 派生与生成质量尚不明确两者在jax.vmap下行为异常对一批 key 做vmap时输出值可能与对同一批 key 做真实映射的结果不同——整批输出随机数只由输入 key 批的第一个 key生成。例如对 8 个 key 向量jax.vmap(jax.random.normal)(keys)等价于jax.random.normal(keys[0], shape(8,))。这是对 XLA RBG 有限批处理支持的绕行方案。4.2 为何需要替代默认 RNG更换默认实现的主要原因文档原文默认的 Threefry 在TPU 上编译慢、执行相对较慢。因此 TPU 上的大规模训练通常需要切换到rbg/unsafe_rbg或 Philox。4.3 自动分区Automatic partitioning相关标志为了让jax.jit高效地对生成分片随机数数组或 key 数组的函数做自动分区各实现依赖额外标志jax/random.pythreefry2x32及rbg的 key 派生需jax_threefry_partitionableTrueJAX v0.5.0 起为默认unsafe_rbg及rbg的随机生成需设置 XLA 标志--xla_tpu_spmd_rng_bit_generator_unsafe1通过环境变量生效XLA_FLAGS--xla_tpu_spmd_rng_bit_generator_unsafe14.4 官方属性对比表模块文档末尾给出完整对比表jax/random.py是选择实现的权威依据属性ThreefryThreefry*Philoxrbgunsafe_rbgrbg**unsafe_rbg**TPU 上最快✅✅✅✅可用 pjit 高效分片✅✅✅✅不同分片下结果一致✅✅✅✅✅CPU/GPU/TPU 结果一致✅✅✅对 key 的jax.vmap结果精确✅✅✅*设置jax_threefry_partitionable1JAX v0.5.0 起默认 **设置XLA_FLAGS--xla_tpu_spmd_rng_bit_generator_unsafe1。默认实现是绝大多数场景的正确选择除非 PRNG 生成出现在性能剖析的热点中。5. 最佳实践速查结合 docs/101/random.md 教程与模块文档总结实践要点新代码一律用jax.random.key创建类型化 key避免uint32旧式 key 的误用仅在对接外部系统时用key_data/wrap_key_data转换。不要复用 key——除非你刻意想要相同输出。相同 key 喂给不同采样器会产生相关结果。split一次多分fold_in按步派生保持 key 树宽而非深兼顾并行度与碰撞安全。不要期待顺序等价JAX 不承诺逐个采样 N 个数与一次性采样 N 个数得到相同序列教程 docs/101/random.md 有专门演示。放弃顺序等价正是生成可自由向量化与分片的前提用默认实现时jax.vmap(random.normal)(subkeys)与逐 key 调用严格等价。TPU 场景评估rbg/unsafe_rbg与 Philox并按第 4.3 节设置对应分区标志实验性实现的正确性边界vmap 语义、未经 BigCrush 测试需自行权衡。如需深入了解 PRNG 设计背景可继续阅读仓库中的设计 JEP 文档 docs/jep/263-prng.md 与类型化 key 的 JEP docs/jep/9263-typed-keys.md以及jax.random的入门教程 docs/101/random.md。【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表