JAX 是一个 Python 数组计算与程序变换库,面向高性能数值计算和大规模机器学习场景。它围绕高效数组操作与程序变换展开,提供接近 NumPy 的使用方式,帮助研究人员和工程师降低采用成本。JAX 的同一份代码可以运行在 CPU、GPU 和 TPU 等多个后端上,并通过可组合的函数变换支持编译、批处理、自动微分和并行化。官网将 JAX 定位为一个范围较窄的核心工具,同时介绍了围绕它形成的神经网络、优化器、数据加载、概率编程、物理与仿真以及大语言模型等生态。
核心功能
NumPy 风格数组计算
JAX 提供熟悉的 NumPy-style API,核心文档从数组和 jax.numpy 开始,覆盖数组计算的表达方式。对于已经接触 NumPy 的研究人员和工程师,这种接口设计可以作为使用 JAX 的入门路径。官网将数组操作列为 JAX 本身的核心范围,而不是把它描述成一个包含完整机器学习应用层的产品。
程序编译与加速
JAX 提供可组合的函数变换,其中包括编译能力。文档在 JAX 201 层级介绍了使用 jax.jit() 进行编译,以及提前编译、控制流和编译时间诊断等内容。通过这些功能,用户可以进一步学习如何处理性能和扩展相关问题。
自动微分
自动微分是 JAX 文档体系中的重要能力,基础内容包括 jax.grad() 变换,进阶内容则覆盖 JVP、VJP、Jacobian、Hessian、自定义导数规则、带分片的自动微分和梯度检查点。文档还介绍了带可变状态的自动微分,适合需要对数值计算过程求导并继续扩展 JAX 的使用者。
批处理与向量化
JAX 支持批处理相关的函数变换,JAX 101 文档将 jax.vmap() 列为代表性内容。它与数组计算、追踪机制、pytrees、随机数和状态等基础主题一起构成计算表达的入门部分。相关内容用于说明如何在 JAX 中表达和变换计算,而不是提供某个特定模型的训练流程。
并行化、分片与多设备运行
JAX 包含并行化变换,并支持数据放置、分片、自动并行化以及使用 shard_map 进行按设备编程。官网还列出跨多个主机的多控制器 JAX、分布式数据加载、容错、导出与序列化、持久化编译缓存和传输保护等系统主题。相同代码可运行在 CPU、GPU 和 TPU 后端,但具体计算方式会涉及设备和分片相关设置。
自定义内核与外部代码调用
对于需要深入硬件或扩展运行能力的场景,JAX 401 文档介绍了使用 Pallas 编写自定义 GPU 和 TPU 内核,也介绍了通过外部函数接口调用外部代码。JAX 601 则进一步说明 jaxpr 语言、primitives 以及从头构建 JAX 核心的 Autodidax,面向需要理解内部机制的高级使用者。
使用方式
官网提供安装入口、JAX 101 入门教程和 API Reference。文档按 101、201、301、401、501、601 分层:初学者可以按顺序阅读,已经了解 JAX 的用户也可以直接查看与目标任务匹配的层级或 Newly documented 页面。正文没有说明需要注册、购买服务或自备模型 Key;同时没有介绍独立客户端、浏览器插件、命令行工具或单独的托管 API,主要使用方式是安装 Python 库并在代码中调用其 API。
适用人群与场景
研究人员和工程师可以使用 NumPy 风格接口表达数组计算,并在 CPU、GPU、TPU 上运行同一份代码。需要训练神经网络的用户可以将 JAX 作为底层计算与变换工具,并进一步查看官网提到的 JAX AI Stack。需要进行数值优化、求解器、概率编程、物理仿真或大语言模型开发的用户,可以从官网列出的相关 JAX 生态工具中选择配套项目。需要研究性能、自动微分、分片、设备编程或 JAX 内部机制的高级用户,则可以按 201 至 601 层级阅读对应文档。