JAX Enhancement Proposals(JEP)设计提案机制与核心 JEP 文档全解读

发布时间:2026/9/10 15:53:36

JAX Enhancement Proposals(JEP)设计提案机制与核心 JEP 文档全解读 JAX Enhancement ProposalsJEP设计提案机制与核心 JEP 文档全解读【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jaxJAX Enhancement ProposalsJEP是 JAX 社区沉淀重大设计决策的正式载体当一次改动涉及设计文档、需要长时间讨论时社区便以 JEP 的形式撰写长文、在 Pull Request 中迭代并接受评审。本文以 docs/jep/index.rst 为骨架完整梳理 JEP 的启动流程、文件命名规范与适用场景并对仓库中全部 19 个 JEP 的核心内容进行源码级解读——涵盖 PRNG 设计、shard_map、类型提升、版本管理等 JAX 关键技术决策帮助读者理解 JAX 关键 API 的来龙去脉与设计哲学。什么是 JEPJAX 的大多数改动通过普通的 issue、discussion 和 pull request 即可完成讨论。但部分改动范围更大或需要更深入的讨论这类改动应以 JEPJAX Enhancement Proposal的形式落地。JEP 允许撰写更长的设计文档并让文档本身在 Pull Request 中接受评审和迭代。从仓库结构看JEP 文档全部存放在 docs/jep 目录下以编号-短标题.md命名例如263-prng.md、14273-shard-map.md。这些文档由 docs/jep/index.rst 通过 Sphinxtoctree统一索引并最终呈现在 JAX 官方文档的 JEP 页面中。值得注意的是index.rst结尾还特别说明部分早期 JEP 是从其他文档、issue 和 pull request 事后转换而来的因此它们可能无法完全精确地反映上面描述的流程——也就是说JEP 体系的形成本身也是一个渐进演进的过程。何时应该使用 JEP根据 docs/jep/index.rst 的界定以下两类情况应当使用 JEP当你的改动需要一份设计文档design doc时。将设计文档收集为 JEP有利于后续的检索与引用——所有重大设计决策都能在同一处被找到。当你的改动需要大量讨论时。在 issue 或 pull request 上进行较短讨论是可以的但当讨论变长时后续消化这些讨论会变得不切实际。JEP 允许在主体文档中更新讨论摘要而这些更新本身又可以在添加 JEP 的 pull request 中被继续讨论。这份界定说明 JEP 的本质定位它是 JAX 重大设计的决策记录 活文档既用于决策当时的深度讨论也用于未来长期的可检索引用。如何启动一个 JEPindex.rst给出了完整的启动流程创建带JEP标签的 issue所有与该 JEP 相关的 pull request包括添加 JEP 本身的 PR以及后续实现该 JEP 的 PR都应链接到这个 issue 上。JEP 标签对应的 issue 列表可在 GitHub 上通过label:JEP搜索获得。创建 Pull Request 添加 JEP 文件文件名遵循%d-{short-title}.md格式其中%d为 issue 编号。例如 JEP 28845 对应 issue #28845其文件名为28845-stateful-rng.md。从仓库中实际的文件命名可以验证这一规范的一致性JEP 文件Issue 编号主题263-prng.md263JAX PRNG 设计2026-custom-derivatives.md2026自定义 JVP/VJP 规则4008-custom-vjp-update.md4008自定义 VJP 与nondiff_argnums更新4410-omnistaging.md4410Omnistaging9263-typed-keys.md9263Typed keys 与可插拔 RNG9407-type-promotion.md9407类型提升Type Promotion语义设计9419-jax-versioning.md9419jax 与 jaxlib 版本管理10657-sequencing-effects.md10657JAX 中的副作用排序11830-new-remat-checkpoint.md11830jax.remat/jax.checkpoint新实现12049-type-annotations.md12049JAX 类型标注路线图14273-shard-map.md14273shard_mapshmap15856-jex.md15856jax.extend扩展模块17111-shmap-transpose.md17111shard_map的高效转置18137-numpy-scipy-scope.md18137JAX NumPy/SciPy 包装器范围25516-effver.md25516基于努力的版本管理Effort-based versioning28661-jax-array-protocol.md28661__jax_array__协议28845-stateful-rng.md28845JAX 中的有状态随机数JEP 流程的设计特点以 issue 为锚点串联所有相关 PR使得一条设计决策的完整生命周期提案、评审、实现、回退都可以被追踪而活文档机制则允许讨论结论沉淀回文档本身。这正是 JAX 这种大型数值计算库管理核心 API 演进的核心机制。深入 JEP 的核心技术主题index.rst的toctree本身只是导航真正的技术内容分布在各个 JEP 文档中。以下按主题对仓库中 JEP 的核心技术内容进行解读。PRNG 设计三部曲263 → 9263 → 28845JAX 随机数体系的演进贯穿了三个 JEP这是理解jax.random的最佳路径。JEP 263JAX PRNG 设计263-prng.md该 JEP 奠定了 JAX 随机数的哲学基础。它提出一个理想的 PRNG 设计应当满足 7 条标准表达力强、可复现后端无关、语义不受jit编译边界与设备后端影响、支持 SIMD 向量化生成数组、可并行化不引入不必要的数据依赖顺序、可扩展到多副本/多核/分布式计算、契合 JAX 与 XLA 的功能性语义。基于这些标准该 JEP 逐一否定了两种传统模型有状态全局 PRNG如 NumPy 风格为保证可复现必须控制无关调用的求值顺序违背可并行性也与 XLA 功能性子表达式任意求值顺序的语义冲突显式线程化状态的功能性模型虽然把数据依赖显式化但每个随机函数都必须接受并返回状态表达力受限且无法避免顺序执行。最终结论JEP 的 TLDR 给出了精炼总结JAX PRNG Threefry 计数器 PRNG 功能性的可拆分splittable数组化模型计数器 PRNG 用哈希函数在整数区间[k1, …, ksample_size]上映射生成数组从而实现高效向量化split则从一个 key 派生出两个独立的新 key。当前仓库中jax/_src/random/core.py的splitL318、fold_inL288、keyL231等核心函数正是这一设计的实现。JAX 之所以坚持在软件中实现 PRNG也是出于当前硬件约束下的务实选择。JEP 9263Typed keys 与可插拔 RNG9263-typed-keys.md这是 JEP 263 设计的类型安全化升级RNG key 从长度为 2 的uint32数组变为一个带有特殊 RNG dtype 的标量数组满足jnp.issubdtype(key.dtype, jax.dtypes.prng_key)。该 JEP 对用户的影响包括新 API 用jax.random.key(0)创建 typed keydtypekeyfry底层以 Threefry 实现旧 APIjax.random.PRNGKey(0)仍可用但返回uint32数组对 key 进行算术如key 1、索引、转置等不安全操作会有意报错从而在类型层面拦截 PRNG 误用需要操作底层缓冲时可用jax.random.key_data(key)取出uint32数组旧式 key 下它是恒等操作以及jax.random.wrap_key_data包装回 typed key若代码中有显式的key.shape、key.dtype逻辑需要注意 key 缓冲区尾维不再属于 shape 的一部分。在仓库中jax/_src/random/core.py的key_dataL351与jax.random模块文档jax/random.py 顶部示例都体现了这一演进——如今的官方示例直接用jax.random.key(seed)创建 key。JEP 28845JAX 中的有状态随机数28845-stateful-rng.md这是最新的演进借助 JAX 新引入的可变引用mutable refs机制实现一个可选的有状态 PRNG用于与经典函数式 PRNG 互补。提案的 API 对齐 NumPy 最新的numpy.random.default_rng风格实现代码概要如下见 JEP 文档def stateful_rng(seed): Create a stateful PRNG Generator given an integer seed. return StatefulPRNG(jax.random.key(seed), jax.new_ref(0)) tree_util.register_dataclass dataclass(frozenTrue) class StatefulPRNG: base_key: jax.Array counter: jax.core.Ref def key(self): key jax.random.fold_in(self.base_key, self.counter[...]) jax.ref.addupdate(self.counter, ..., 1) # increment counter return key def random(self, size, dtypefloat): return random.uniform(self.key(), shapesize, dtypedtype)使用方式几乎与 NumPy 相同rng stateful_rng(1701)后反复调用rng.random((5,))每次都会因计数器自增而得到新的随机数。由于状态由 refs 追踪即使在jax.jit等通常要求纯函数语义的变换中使用随机状态也会正确更新。该提案从 jax.experimental.random 起步目标最终进入jax.random。仓库中 jax/_src/random/stateful_rng.py 已实现StatefulPRNG类L50与stateful_rng工厂函数L217并通过jax/experimental/random.py对外导出。同时该 JEP 也如实说明限制由于基于 refs在vmap/shard_map下会继承 refs 的限制例如无法在 vmapped 函数中直接使用未映射的rng。shard_map面向 per-device 代码的并行 APIJEP 14273shard_mapshmap14273-shard-map.md该 JEP 面向 JAX 多设备编程的第二条路线——让我写我想写的东西编写 per-device 代码并使用显式通信集合体collective。shard_map被定位为一个简单的多设备并行 API逻辑形状与每设备的物理 buffer 形状一致collective 精确对应跨设备通信xmap的裁剪特化scale-back 版本XLA SPMD Partitioner manual 模式的直接表面化。对pjit用户shmap是互补工具可在pjit计算内部临时切换到手动 collective模式作为编译器自动分区的逃生舱escape hatch对pmap用户则是严格升级——更具表达力、性能更好、与其他 JAX API 组合性更强。JEP 给出了经典的分块矩阵乘法示例from functools import partial import jax import jax.numpy as jnp from jax.sharding import Mesh, PartitionSpec as P from jax.experimental.shard_map import shard_map mesh jax.make_mesh((4, 2), (i, j)) a jnp.arange(8 * 16.).reshape(8, 16) b jnp.arange(16 * 32.).reshape(16, 32) partial(shard_map, meshmesh, in_specs(P(i, j), P(j, None)), out_specsP(i, None)) def matmul_basic(a_block, b_block): # a_block: f32[2, 8] # b_block: f32[8, 32] z_partialsum jnp.dot(a_block, b_block) z_block jax.lax.psum(z_partialsum, j) return z_block c matmul_basic(a, b) # c: f32[8, 32]该示例展示了shard_map相对pmap/xmap的诸多优势多轴并行无需嵌套或axis_index_groups调用方无需 reshape通过mesh精确控制设备放置逻辑轴名与物理轴名统一结果可直接高效传给pjit支持 eager 执行可中途pdb调试。JEP 还给出了全分片输出变体——用jax.lax.psum_scatter(c_partialsum, j, scatter_dimension1, tiledTrue)实现 reduce-scatter 风格的 matmul。配套的JEP 1711117111-shmap-transpose.md专门解决shard_map在反向模式自动微分中的高效转置问题说明shard_map设计从一开始就考虑了与自动微分的组合。其他代表性 JEP 速览JEP 9419jax 与 jaxlib 版本管理9419-jax-versioning.mdJAX 以两个独立 wheel 发布jax纯 Python与jaxlib主要为 C包含 XLA、LLVM 片段、MLIR/StableHLO 绑定以及快速 JIT 和 PyTree 的 C 库。分开发布让 Python 部分可以独立迭代无需每次重编 C。兼容性约束为jax版本x.y.z与jaxlib版本lx.ly.lz兼容当且仅当jaxlib≥ jax 声明的最低 jaxlib 版本且jax≥jaxlib因此jax可随时单独发布而每次发布新jaxlib必须同步发布一个jax版本约束在导入时由jax运行时检查而非 pip 包约束因为jaxlib面向 GPU、TPU 等多种软硬件组合分别提供 wheel不宜由 pip 自动安装平台相关的 extras如jax[cuda]会安装兼容的 jaxlib 版本。JEP 11830jax.checkpoint/jax.remat新实现11830-new-remat-checkpoint.md该 JEP 记录了在 JAX v0.3.17 定稿的jax.checkpoint新实现在 jax0.3.16 之前可用jax_new_checkpoint配置回退。其核心新特性是用户可自定义的重新物化策略policy参数精确控制前向传播中保存哪些中间值from functools import partial import jax def apply_layer(W, x): return jnp.sin(jnp.dot(W, x)) partial(jax.checkpoint, policyjax.checkpoint_policies.checkpoint_dots) def predict(params, x): for W in params[:-1]: x apply_layer(W, x) return jnp.dot(params[-1], x)例如checkpoint_dots策略只保存矩阵乘的结果sin/cos等逐元素运算的中间值在后向时重新计算——这类策略在 TPU 上尤其有效逐元素运算近乎免费而矩阵单元的结果值得保存。这在仓库 jax/_src/ad_checkpoint.py 中有完整实现checkpoint_dots dots_saveable DotsSaveable(False)L95、checkpoint_policies命名空间L203以及checkpoint(fun, *, prevent_cse..., ...)L224均可在源码中直接对应。JEP 9407类型提升Type Promotion语义设计9407-type-promotion.md该 JEP 论证了 JAX 为何不能照搬 NumPy 的类型提升规则NumPy 规则强烈偏向 64 位输出而 GPU/TPU 上 64 位浮点通常有明显性能惩罚部分加速器甚至不支持原生 64 位浮点。JEP 用格lattice表示来思考类型提升——任意两个节点间的上确界supremum就是它们提升到的类型并推导出面向加速器优化的提升语义。仓库中同时提供了 9407-type-promotion.ipynb 交互式笔记本版本便于读者动手验证。JEP 15856jax.extend扩展模块15856-jex.md针对许多项目把 JAX 内部当库用的现状该 JEP 提出引入jax.extend模块为 JAX 的部分内部组件提供库视图。它属于二级 API不承诺兼容性策略、没有弃用窗口每次发布都可能破坏现有调用方变更通过 changelog 公示。它与jax.experimental不同——后者是新功能的试验场最终会并入正式模块或被移除。仓库中 jax/extend 目录含random.py、core.py、mlir.py等子模块正是该提案落地的结果。JEP 4410Omnistaging4410-omnistaging.md这是一份升级指南性质的 JEP记录了 JAX 0.2.0 起默认开启的 tracing 基础设施变更改善内存性能与 trace 执行时间、简化内部实现但可能暴露既有代码的 bug。最常见的破坏点是用jax.numpy计算 shape 值或 trace-time 常量——正确做法是使用 Python 原生值或numpy。在 jax 0.2.00.2.11 中可临时通过环境变量JAX_OMNISTAGING、absl flagjax_omnistaging或jax.config.disable_omnistaging()禁用0.2.12 起不再支持禁用。其余 JEP 同样各有关键主题2026/4008定义自定义 JVP/VJP 规则与nondiff_argnums更新机制10657讨论副作用如打印、IO在 JAX 中的排序语义12049给出 JAX 类型标注路线图18137界定jax.numpy/jax.scipy包装器的支持范围与 NumPy/SciPy 的关系25516提出基于努力effort的版本管理28661设计__jax_array__协议以支持跨库数组互操作。感兴趣的读者可直接阅读对应文件并在仓库 docs/jep 中按需引用。JEP 的价值把设计决策变成可检索的资产纵观全部 JEP可以提炼出这套机制对 JAX 项目及读者的三层价值决策可追溯每个 JEP 以 issue 编号为锚从提案、评审到实现 PR 全部关联任何 API 的为什么都能找到原始设计文档文档即活文档JEP 允许在文档中沉淀讨论摘要且摘要更新本身可继续接受评审避免了长讨论失焦读者即受益者对使用 JAX 的开发者而言JEP 目录是一份设计决策索引——当某个 API 行为令人困惑时如 typed key 为何禁止算术、shard_map 与 pjit 如何分工、类型提升为何偏向 32 位在 docs/jep 中找到对应 JEP 即可获得第一手的设计动机与权衡分析其权威性远超二手解读。对于希望为 JAX 贡献重大设计的开发者遵循 docs/jep/index.rst 描述的流程带JEP标签建 issue → 以 issue 编号命名文件 → 提交 PR 迭代讨论即可让设计进入 JAX 的长期决策档案。【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
延伸阅读

更多相关文章

2026/9/10 15:53:36

Starlight文档平台接入Microsoft Clarity实践指南

1. 项目概述 去年接手公司文档平台重构时,我们选择了基于Astro构建的Starlight框架。这个轻量级的文档方案确实解决了多版本管理、搜索优化等痛点,但始终有个问题困扰着我们——无法直观了解用户的实际使用行为。直到引入了Microsoft Clarity这款免费的用…

2026/9/10 16:58:47

2026年家电企业客服软件怎么选?主流方案评测与3C家电出海案例解析

摘要:家电售后客服系统选型正从“可选项”变为“必选项”。本文围绕家电企业用什么在线客服系统、家电售后客服系统哪家好等问题,梳理2026年主流方案与选型要点,并介绍智齿科技“全球智能一体化客户联络中心”方案在3C/家电出海场景中的实践价…

2026/9/10 16:39:38

超人会飞不算本事:系统稳定依赖清晰规则与边界设计

开头先不绕弯子。“#斯坦李吐槽dc 所以超人是无缘无故会飞的嘛哈哈哈哈哈哈哈锤哥真是技术人才啊!#雷神 #复联”这类调侃式短标题,第一波冲击力在于它把两个宇宙的角色塞进同一个吐槽箱里,但细想一下就能发现,它真正碰到的根本不是…

2026/9/10 11:16:38

超人VS蜘蛛侠:拆解超级IP的影响力与传播方法论

把“蜘蛛侠 vs 超人”放在 CSDN 上聊,可能很多人第一反应是走错片场了。但如果把这两个角色看成“两个持续运营了 80 多年的文化产品”,你会发现,这场比较本质上是两个不同 IP 策略的长期结果对比:超人赢在定义了整个超级英雄题材…

2026/9/9 16:31:09

基于CNN的调制信号识别:MATLAB实现时频图分类实战

简介:本资源是一套面向通信工程与信号处理方向学习者、研究者的深度学习实践方案,聚焦调制信号自动检测与识别这一典型无线通信任务,解决传统方法依赖人工特征、低信噪比下性能下降等痛点。压缩包共12个文件(10.73MB)&…

2026/9/10 0:00:55

目录对比去重实战:用哈希算法精准清理重复文件

我电脑里现在还有一块换了三次机的“数据墓地”硬盘,里面存着2016年以前所有旧笔记本的完整备份。平时不觉得有什么,直到前阵子想把它整理归档,发现同一个安装包、同一批照片、同一份论文草稿,在几个不同的备份目录里反复出现。更…

2026/9/10 0:00:55

Leaflet离线地图完整Demo合集:内网部署与坐标纠偏实战

简介:这是一份面向Web GIS开发者的LeafLet离线地图示例合集,帮助开发者快速掌握离线地图从搭建到交互的完整流程。压缩包共723个文件,大小14.06MB,以319个js脚本、175个html页面和29个css样式文件为主体,配合png/svg图…

2026/9/10 0:00:55

MATLAB读取Rinex 3.02观测文件:多系统GNSS数据解析实战

简介:基于MATLAB开发的Rinex3.02版观测文件(o文件)读取代码包,面向卫星定位导航方向的学习者与研究人员,用于解决新版观测文件的数据解析、历元提取与时间转换问题。压缩包共4个文件,包含两个m脚本、一个19…

2026/9/10 12:32:02

USB Type-C PCB布局分区设计:电源、高速信号与PD协议全攻略

做硬件这行,Type-C接口算是典型的“看着简单,做起来全坑”的东西。光引脚就24个,高低速信号、电源、控制线全部塞在一个小小的连接器里,如果PCB布局不做规划,打样回来基本就是“插上没反应”、“高速掉线”、“静电一打…

2026/9/10 15:19:50

系统编程学习原型如何补齐稳定性边界

系统编程学习原型如何补齐稳定性边界预算有限时&#xff0c;我先优化明显多余的复制&#xff0c;而不是猜测性地换容器。用借用传递只读数据通常就能减少分配&#xff1a; fn parse(line: &str) -> Result<Item, Error> { /* ... */ }用基准确认热点确实在分配&am…

2026/9/10 15:49:53

雨花区哪家财务公司代理记账比较好?

在雨花区&#xff0c;企业处理财税事务常常面临诸多挑战&#xff0c;选择一家靠谱的财务公司至关重要。湖南巨勤财务管理咨询有限公司就是本地正规实体财税服务机构&#xff0c;深耕本地工商财税行业多年&#xff0c;熟悉当地工商局、税务局最新政策与申报流程。主营公司注册、…

还想了解更多?直接咨询顾问

免费诊断 + 免费方案 + 透明报价。

全国咨询热线400-8866-253
免费获取方案
咨询二维码