Jax 和 Jaxlib 版本控制#

为什么 jaxjaxlib 是独立的软件包?#

我们将 JAX 发布为两个独立的 Python wheel:一个是纯 Python 的 jax,另一个主要是 C++ 的 jaxlib,其中包含如下库:

  • XLA,

  • XLA 使用的 LLVM 组件,

  • MLIR 基础设施,例如 StableHLO Python 绑定。

  • 用于快速 JIT 和 PyTree 操作的 JAX 专用 C++ 库。

我们将 jaxjaxlib 分开分发,因为这样可以轻松地对 JAX 的 Python 部分进行开发,而无需构建 C++ 代码,甚至无需安装 C++ 工具链。jaxlib 是一个大型库,许多用户难以构建,但 JAX 的大部分更改仅涉及 Python 代码。通过允许独立更新 Python 部分和 C++ 部分,我们提高了 Python 代码的开发速度。

此外,构建 jaxlib 的成本很高,而我们希望能够在没有大量 CPU 资源的各种环境中(例如在 Github Actions 中或在笔记本电脑上)迭代并运行 JAX 测试。我们的许多 CI 构建直接使用预构建的 jaxlib,而无需在每个 PR 上重新构建 JAX 的 C++ 组件。

正如我们将看到的,将 jaxjaxlib 分开分发是有代价的,因为它要求对 jaxlib 的更改必须保持 API 向后兼容。然而,我们认为从总体上讲,让 Python 的更改变得简单是更可取的,即使这会使 C++ 的更改稍微困难一些。

jaxjaxlib 是如何进行版本控制的?#

摘要:jaxjaxlib 在 JAX 源代码树中共享相同的版本号,但作为独立的 Python 软件包发布。安装时,jax 软件包的版本必须大于或等于 jaxlib 的版本,且 jaxlib 的版本必须大于或等于 jax 指定的最低 jaxlib 版本。

jaxjaxlib 的发布版本号均为 x.y.z,其中 x 是主版本号,y 是次版本号,z 是可选的补丁版本号。版本号必须遵循 PEP 440。版本号比较是基于整数元组的字典序比较。

每个 jax 版本都有一个关联的最低 jaxlib 版本 mx.my.mzjax 版本 x.y.z 的最低 jaxlib 版本不得超过 x.y.z

若要使 jax 版本 x.y.zjaxlib 版本 lx.ly.lz 兼容,必须满足以下条件:

  • jaxlib 版本(lx.ly.lz)必须大于或等于最低 jaxlib 版本(mx.my.mz)。

  • jax 版本(x.y.z)必须大于或等于 jaxlib 版本(lx.ly.lz)。

这些限制条件对发布版本隐含了以下规则:

  • jax 可以随时独立发布,而无需更新 jaxlib

  • 如果发布了新的 jaxlib,则必须同时发布相应的 jax 版本。

这些 版本限制 目前由 jax 在导入时检查,而不是表示为 Python 软件包的版本限制。jax 在运行时检查 jaxlib 版本,而不是使用 pip 软件包版本限制,因为我们 为各种硬件和软件版本(例如 GPU、TPU 等)提供了独立的 jaxlib wheel。由于我们不知道对于任何给定的用户哪一个是正确的选择,因此我们不希望 pip 自动为我们安装 jaxlib 软件包。

将来,我们希望将 jaxlib 中特定于硬件的部分拆分为独立的插件,届时最低版本可以表示为 Python 软件包依赖项。目前,我们提供了特定于平台的额外需求,用于安装兼容的 jaxlib 版本,例如 jax[cuda]

如何安全地对 jaxlib 的 API 进行更改?#

  • jax 可以随时放弃与旧版 jaxlib 的兼容性,前提是将最低 jaxlib 版本提高到兼容版本。但请注意,即使是未发布的 jax 版本,最低 jaxlib 版本也必须是一个已发布的版本!这使我们能够在 CI 构建中使用已发布的 jaxlib wheel,并允许 Python 开发人员在 HEAD 上进行 JAX 开发,而无需构建 jaxlib

    例如,要删除 jax Python 代码中的旧向后兼容路径,只需提高最低 jaxlib 版本,然后删除该兼容路径即可。

  • jaxlib 可以放弃对低于其自身发布版本号的旧版 jax 的兼容性。jax 执行的版本限制将禁止使用不兼容的 jaxlib

    例如,若 jaxlib 要删除旧版 jax 使用的 Python 绑定 API,则必须增加 jaxlib 的次版本号或主版本号。

  • 如果可能,对 jaxlib 的更改应以向后兼容的方式进行。

    通常情况下,只要遵循关于 jax 必须与所有至少等于最低版本的 jaxlib 兼容的规则,jaxlib 就可以自由地更改其 API。这意味着 jax 必须始终与至少两个版本的 jaxlib 兼容,即上一个发布版本和 HEAD 处(实际上是下一个发布版本)的版本。如果保持兼容性,这更容易做到,尽管可以使用 jax 中的版本测试来进行不兼容的更改;请见下文。

    例如,向 jaxlib 添加新函数通常是安全的,但如果当前的 jax 仍在使用该函数,则删除现有函数或更改其签名是不安全的。对 jax 的更改必须在所有从最低版本到 HEAD 的 jaxlib 发布版本上正常工作或优雅降级。

请注意,此处的兼容性规则仅适用于已发布的 jaxjaxlib 版本。它们不适用于未发布的版本;也就是说,如果 API 从未发布过,或者没有任何已发布的 jax 版本使用该 API,则引入并随后删除 jaxlib 中的 API 是可以的。

jaxlib 的源代码是如何布局的?#

jaxlib 分布在两个主要仓库中,即 JAX 主仓库中的 jaxlib/ 子目录 以及 位于 XLA 仓库内部的 XLA 源代码树。XLA 中 JAX 特定的部分主要位于 xla/python 子目录

JAX 的 C++ 组件(例如 Python 绑定和运行时组件)位于 XLA 树中的原因既有历史原因,也有技术原因。

历史原因是,最初设想 xla/python 绑定是可能与其他框架共享的通用 Python 绑定。在实践中,情况已大不相同,xla/python 合并了许多 JAX 特定的部分,并且很可能还会合并更多。因此,简单地将 xla/python 视为 JAX 的一部分可能是最好的。

技术原因是 XLA C++ API 不稳定。通过将 XLA:Python 绑定保留在 XLA 树中,其 C++ 实现可以与 XLA 的 C++ API 原子地更新。维护 Python API 的向后和向前兼容性比 C++ API 更容易,因此 xla/python 公开 Python API,并负责在 Python 层面上保持向后兼容性。

jaxlib 是使用 Bazel 从 jax 仓库构建的。XLA 仓库中 jaxlib 的组件作为 Bazel 子模块 合并到构建中。要更新构建期间使用的 XLA 版本,必须更新 Bazel WORKSPACE 中的固定版本。这是根据需要手动完成的,但可以在每次构建时覆盖。

我们如何在版本发布之间处理 jaxjaxlib 边界上的更改?#

jaxlib 版本是一个粗糙的工具:它只允许我们考虑*发布版本*。

然而,由于 jaxjaxlib 代码分布在无法在一次更改中原子更新的仓库中,我们需要以比发布周期更细的粒度管理兼容性。为了管理细粒度的兼容性,我们有独立于 jaxlib 发布版本号的额外版本控制。

我们在 XLA 仓库的 xla_client.py 维护了一个额外的版本号(_version)。该版本号与 JAX 的 C++ 部分一起在 xla/python 中定义,同时作为 jax._src.lib.jaxlib_extension_version 对 JAX Python 可见。每当对 XLA/Python 代码进行对 jax 有向后兼容性影响的更改时,都必须增加此版本号。JAX Python 代码随后可以使用此版本号来维护向后兼容性,例如:

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.

请注意,此版本号是*在*已发布版本号的限制条件*之外*的,即此版本号的存在是为了帮助管理未发布代码在开发期间的兼容性。发布版本也必须遵循上述给出的兼容性规则。