JAX 与 jaxlib 版本管理机制详解:双包架构、版本约束与跨仓库兼容策略
【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax
导读
JAX 以两个独立 Python 包(jax与jaxlib)的形式发布,二者共享同一版本号却在源码树与发布节奏上完全解耦,这一设计直接影响了安装方式、CI 流程与 API 演进策略。本文以官方 JEP 文档(docs/jep/9419-jax-versioning.md)为核心,结合本仓库源码,系统讲解jax/jaxlib为何分离、版本约束如何生效、跨仓库修改如何保持兼容,以及开发者应如何安全地演进jaxlib的 API。
一、为什么jax和jaxlib是两个独立包?
1.1 双包架构的定位
JAX 以两个 wheel 形式发布:
jax:纯 Python wheel,包含绝大多数 Python 层代码(jax/目录),改动只涉及 Python 时无需重新编译任何 C++ 代码。jaxlib:以 C++ 为主的 wheel,包含:- XLA 编译器;
- XLA 依赖的部分 LLVM 组件;
- MLIR 基础设施(如 StableHLO 的 Python 绑定);
- JAX 专属的 C++ 库,用于快速 JIT 与 PyTree 操作(本仓库中对应 jaxlib/ 目录下的
jax_jit.cc、pytree.cc等实现)。
1.2 分离发布的动机
文档给出的核心理由有三点:
- 降低开发门槛:绝大部分 JAX 改动只触及 Python 代码,分离后开发者可以在没有 C++ 工具链的环境下直接工作,无需每次构建
jaxlib。 - 加快 CI 迭代:
jaxlib构建昂贵,而 CI 构建可以直接复用预构建的jaxlibwheel。本仓库的 ci/run_pytest_cpu.sh 注释即明确说明"Runs Pyest CPU tests. Requires a jaxlib wheel to be present",并在测试前安装预构建的 jaxlib wheel,而不是在每个 PR 上重新编译 C++ 部分。 - 代价可控:分离带来的成本是
jaxlib必须维护向后兼容的 API,但权衡之下,让 Python 改动更轻量仍然更划算。
二、jax与jaxlib的版本规则
2.1 版本号格式
两个包均使用x.y.z三段式版本号:x为主版本、y为次版本、z为可选补丁版本。版本号遵循 PEP 440,比较方式是整数元组的字典序比较。
在源码树中,二者共享同一个版本定义文件 jax/version.py,当前仓库中_version = "0.11.2"。该文件被jax和jaxlib两个包共同包含,且会被setup.py以eval()方式读取,因此它不能有任何外部依赖。
2.2 兼容性约束
每个jax发布版都关联一个最低 jaxlib 版本mx.my.mz,且该最低版本不得高于jax自身的版本号。对于jax版本x.y.z与jaxlib版本lx.ly.lz,二者兼容需同时满足:
| 约束 | 含义 |
|---|---|
lx.ly.lz >= mx.my.mz | jaxlib 版本不低于 jax 声明的最低 jaxlib 版本 |
x.y.z >= lx.ly.lz | jax 版本不低于 jaxlib 版本(jaxlib 不得比 jax 更新) |
由此推出的发布规则:
jax可以随时单独发布,无需同步更新jaxlib;- 一旦发布新版
jaxlib,必须同时发布一个对应的jax版本。
本仓库 jax/version.py 中即可看到实例:_minimum_jaxlib_version = '0.11.1',即当前jax(0.11.2)要求 jaxlib 至少为 0.11.1。
2.3 为什么用运行时检查而不是 pip 约束?
这些版本约束由jax在import 时检查,而不是写成 Python 包的依赖声明。原因在于:JAX 为不同硬件/软件组合(GPU、TPU 等)提供了多种jaxlibwheel,pip无法替用户判断该安装哪一个,因此不自动安装jaxlib。
具体实现位于 jax/_src/lib/init.py 的check_jaxlib_version()函数。它用正则提取 PEP 440 版本号的前缀数字部分并转成整数元组,然后执行两条核心判断:
if _jaxlib_version < _minimum_jaxlib_version: raise RuntimeError( f'jaxlib is version {jaxlib_version}, but this version ' f'of jax requires version >= {minimum_jaxlib_version}.') if _jaxlib_version > _jax_version: raise RuntimeError( f'jaxlib version {jaxlib_version} is newer than and ' f'incompatible with jax version {jax_version}. Please ' 'update your jax and/or jaxlib packages.')对应的报错信息源文件在 jax/_src/lib/init.py。值得注意的是,setup.py 的install_requires中实际上也写明了jaxlib >= 最低版本, <= jax 版本的区间,作为包管理层面的兜底;而运行时检查则是双保险。
2.4 平台特定 extras
文档提到当前通过平台特定的 extra 安装兼容的 jaxlib,例如jax[cuda]。这在 setup.py 中有完整体现:
jax[cpu]:空 extra,纯 CPU 安装无需额外依赖(仅为了兼容保留);jax[cuda]/jax[cuda12]:安装jaxlib与jax-cuda12-plugin[with-cuda];jax[cuda13]:安装jax-cuda13-plugin[with-cuda];jax[tpu]:安装jaxlib、libtpu与requests(requests为jax.distributed.initialize所需);jax[minimum-jaxlib]:固定安装最低 jaxlib 版本,专用于测试;jax[ci]:固定安装 PyPI 上的最新 jaxlib,用于从 GitHub HEAD 构建的 CI 场景。
这些 extras 中的版本区间均写为jaxlib>=当前jaxlib版本,<=jax版本,与运行时检查的约束保持一致。文档也展望了未来:一旦将硬件相关部分拆分为独立插件,最低版本约束就可以表达为常规的 Python 包依赖。
三、如何安全地修改jaxlib的 API?
JEP 文档给出了三条明确的演进规则,全部围绕"jax必须兼容从最低版本到 HEAD 的所有 jaxlib"这一核心展开:
3.1 规则一:jax可以随时放弃对旧 jaxlib 的支持
只要把最低 jaxlib 版本提升到兼容版本即可。例如要删除jaxPython 代码中的旧向后兼容路径,只需提升最低 jaxlib 版本号,然后删除该兼容路径。
关键限制:即使对于未发布的jax版本,其最低 jaxlib 版本也必须是已发布版本。这保证了 CI 构建可以直接使用已发布的 jaxlib wheel,也允许 Python 开发者在 HEAD 上工作而无需构建 jaxlib。
3.2 规则二:jaxlib只能放弃对"低于自身版本号"的旧 jax 的支持
由于jax强制执行的版本约束(jax 版本 >= jaxlib 版本)会直接禁止不兼容组合,jaxlib若要删除某个旧jax使用的 Python 绑定 API,必须递增 jaxlib 的次版本或主版本号。
3.3 规则三:优先向后兼容,必要时用版本检测
jaxlib可以自由修改 API,但必须遵守"jax兼容所有不低于最低版本的 jaxlib"这一总规则。这意味着:
jax必须始终兼容至少两个版本的 jaxlib:最近一个发布版与 tip-of-tree(即下一个发布版);- 新增函数通常是安全的,删除现有函数或修改当前
jax仍在使用的函数签名则不安全; - 对
jax的改动必须在"最低版本到 HEAD"的所有 jaxlib 上正常工作或优雅降级。
注意:兼容规则只适用于已发布版本。未发布版本之间可以随意增删 API——只要该 API 从未发布过,或没有已发布的jax版本使用它。
四、jaxlib源码的布局与构建
4.1 跨两个仓库的源码分布
jaxlib的源码横跨两个主仓库:
- 主 JAX 仓库的 jaxlib/ 子目录:包含 JAX 专属的 C++ 代码、Python 绑定(如
jax_jit.cc、pytree.cc、py_array.cc)以及各硬件后端(cuda/、gpu/、tpu相关、mosaic/、triton/等); - XLA 仓库:XLA 本体及其中的
xla/python子目录,承载大量 Python 绑定与运行时组件。
4.2 为什么 JAX 的 C++ 代码在 XLA 树里?
文档给出了历史与技术两方面的原因:
- 历史原因:
xla/python最初被设想为可与其他框架共享的通用 Python 绑定,但实践中它已包含越来越多 JAX 专属内容,本质上可以视为 JAX 的一部分; - 技术原因:XLA 的 C++ API 不稳定。把 XLA:Python 绑定留在 XLA 树内,其 C++ 实现就能与 XLA 的 C++ API原子地同步更新。相比 C++ API,Python API 更容易维护向后/向前兼容,因此
xla/python以 Python API 形式对外暴露,并负责在 Python 层面维持兼容性。
4.3 Bazel 构建与 XLA 版本锁定
jaxlib使用 Bazel 从 JAX 仓库构建,XLA 部分以Bazel 模块依赖的方式引入。本仓库的 MODULE.bazel 中可以看到 XLA 的锁定方式:
# TODO: use a released version when available bazel_dep(name = "xla") archive_override( module_name = "xla", integrity = "sha256-...", strip_prefix = "xla-9ca0e4af7fb10440783117311fb400e63ab8a3b0", urls = ["https://github.com/openxla/xla/archive/9ca0e4af7fb10440783117311fb400e63ab8a3b0.tar.gz"], )更新构建所用的 XLA 版本,就是手动修改MODULE.bazel中锁定的 commit(按需进行),也可以在单次构建时用 override 覆盖。XLA 带来的 LLVM、StableHLO、Triton 等依赖同样通过MODULE.bazel的 module extensions 引入(见 MODULE.bazel)。
五、细粒度跨仓库兼容:jaxlib_extension_version
5.1 问题的来源
jaxlib的发布版本号是"粗粒度"工具,只能描述发布版之间的兼容关系。但jax与jaxlib的代码分布在两个无法原子更新的仓库中,开发期间需要在比发布周期更细的粒度上管理兼容性——例如某个 XLA/Python 改动尚未发布,却已经要求jax侧同步调整。
5.2 机制:一个独立于发布版本号的扩展版本号
解决方案是在xla/python/xla_client.py中维护一个额外的版本号_version,满足以下特征:
- 它定义在
xla/python(与 JAX 的 C++ 部分同处一地); - 每次对 XLA/Python 代码做出影响
jax向后兼容性的改动,都必须递增它; jax侧可通过jax._src.lib.jaxlib_extension_version读取。
本仓库 jax/_src/lib/init.py 正是这样实现的:
# Jaxlib code is split between the Jax and the XLA repositories. # Only for the internal usage of the JAX developers, we expose a version # number that can be used to perform changes without breaking the main # branch on the Jax github. jaxlib_extension_version: int = getattr(xla_client, '_version', 0) ifrt_version: int = getattr(xla_client, '_ifrt_version', 0)5.3 实际使用模式
JAX Python 代码通过该版本号在运行时做分支,实现新旧代码路径的共存。JEP 文档给出的模板如下:
from jax._src.lib import jaxlib_extension_version # 123 is the new version number for _version in xla_client.py if jaxlib_extension_version >= 123: # Use new code path ... else: # Use old code path.本仓库中有多个真实用例,例如:
- jax/_src/array.py:根据
jaxlib_extension_version >= 489决定调用xc.array_result_handler时是否传入_skip_checks=True; - jax/_src/distributed.py:根据
jaxlib_extension_version < 483决定分布式服务初始化时是否支持 mTLS 参数(旧版 jaxlib 下直接弹出 mTLS 相关 kwargs 并报错提示需要 jaxlib 0.11.2+)。
这种模式让jax主分支可以同时兼容"已发布的旧 jaxlib"与"尚未发布的 HEAD jaxlib",从而支撑"jax 至少兼容两个 jaxlib 版本"的总体策略。需要强调的是,该扩展版本号是在发布版本约束之外的额外机制,专门服务于未发布代码的开发期兼容;正式发布仍需遵守上文第三节的版本兼容规则。
六、实践要点小结
| 场景 | 正确做法 |
|---|---|
| 安装 JAX | pip install jax(纯 CPU,自动带 jaxlib);GPU/TPU 使用pip install "jax[cuda12]"、"jax[cuda13]"、"jax[tpu]"等平台 extras(见 setup.py) |
| 手工指定 jaxlib | 保持最低 jaxlib 版本 <= jaxlib <= jax 版本,否则 import 时报 RuntimeError(jax/_src/lib/init.py) |
单独发布jax | 允许,无需同步发布jaxlib |
发布新版jaxlib | 必须同时发布对应jax版本 |
删除jax中旧兼容路径 | 提升_minimum_jaxlib_version(jax/version.py)后删除对应代码 |
删除jaxlib中旧绑定 API | 递增 jaxlib 次/主版本号,由jax的版本检查禁止不兼容组合 |
| 开发期跨仓库改动 | 递增xla_client.py的_version,并在jax侧用jaxlib_extension_version分支(jax/_src/lib/init.py) |
| 更新构建所用 XLA | 修改 MODULE.bazel 中锁定的 XLA commit,可按构建覆盖 |
理解这套版本体系,对于排查"jaxlib 版本不匹配"类报错、参与 JAX 或 jaxlib 的 API 演进,以及在企业内部镜像/固定版本依赖时都至关重要:它保证了 Python 层的高迭代速度与 C++ 层的稳定兼容二者可以兼得。
【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考