大规模可转换数值计算
什么是 JAX?
JAX 是一个用于加速器导向的数组计算和程序变换的 Python 库, 专为高性能数值计算和大规模机器学习而设计。
JAX 可以自动对原生
Python 和 NumPy 函数进行微分。它可以对循环、分支、
递归和闭包进行微分,并且可以计算导数的导数的
导数。它支持通过 jax.grad 进行反向模式微分(也称为反向传播),
以及前向模式微分,
并且两者可以任意组合到任意阶数。
JAX 使用 XLA
在 TPU、GPU 和其他硬件加速器上编译和扩展你的 NumPy 程序。
你可以使用 jax.jit 编译你自己的纯函数。
编译和自动微分可以任意组合。
深入挖掘,你会发现 JAX 实际上是一个可扩展的系统,用于 可组合的函数变换 在 规模 上。
这是一个研究项目,不是官方的 Google 产品。请做好 遇到棘手问题 的准备。 请通过试用、报告 bug 来帮助我们, 并告诉我们你的想法!
import jax
import jax.numpy as jnp
def predict(params, inputs):
for W, b in params:
outputs = jnp.dot(inputs, W) + b
inputs = jnp.tanh(outputs) # inputs to the next layer
return outputs # no activation on last layer
def loss(params, inputs, targets):
preds = predict(params, inputs)
return jnp.sum((preds - targets)**2)
grad_loss = jax.jit(jax.grad(loss)) # compiled gradient evaluation function
perex_grads = jax.jit(jax.vmap(grad_loss, in_axes=(None, 0, 0))) # fast per-example grads
目录
变换
在其核心,JAX 是一个用于变换数值函数的可扩展系统。
以下是三个:jax.grad、jax.jit 和 jax.vmap。
使用 grad 进行自动微分
使用 jax.grad
高效计算反向模式梯度:
import jax
import jax.numpy as jnp
def tanh(x):
y = jnp.exp(-2.0 * x)
return (1.0 - y) / (1.0 + y)
grad_tanh = jax.grad(tanh)
print(grad_tanh(1.0))
# prints 0.4199743
您可以使用 grad 进行任意阶的求导:
print(jax.grad(jax.grad(jax.grad(tanh)))(1.0))
# prints 0.62162673
您可以自由地在 Python 控制流中使用微分:
def abs_val(x):
if x > 0:
return x
else:
return -x
abs_val_grad = jax.grad(abs_val)
print(abs_val_grad(1.0)) # prints 1.0
print(abs_val_grad(-1.0)) # prints -1.0 (abs_val is re-evaluated)
请参阅 JAX 自动微分 手册 以及关于自动 微分的参考文档 以获取更多信息。
使用 jit 进行编译
使用 XLA 端到端编译你的函数,
通过 jit,
将其用作 @jit 装饰器或高阶函数。
import jax
import jax.numpy as jnp
def slow_f(x):
# Element-wise ops see a large benefit from fusion
return x * x + x * 2.0
x = jnp.ones((5000, 5000))
fast_f = jax.jit(slow_f)
%timeit -n10 -r3 fast_f(x)
%timeit -n10 -r3 slow_f(x)
使用 jax.jit 会限制函数可以使用的 Python 控制流类型;
请参阅 使用 JIT 的控制流和逻辑运算符教程
以获取更多信息。
使用 vmap 进行自动向量化
vmap 沿数组轴映射
一个函数。
但它不仅仅是循环应用函数,而是将循环下推
到函数的原始操作上,例如将矩阵-向量乘法转换为
矩阵-矩阵乘法以获得更好的性能。
使用 vmap 可以省去在代码中
携带批次维度的麻烦:
import jax
import jax.numpy as jnp
def l1_distance(x, y):
assert x.ndim == y.ndim == 1 # only works on 1D inputs
return jnp.sum(jnp.abs(x - y))
def pairwise_distances(dist1D, xs):
return jax.vmap(jax.vmap(dist1D, (0, None)), (None, 0))(xs, xs)
xs = jax.random.normal(jax.random.key(0), (100, 3))
dists = pairwise_distances(l1_distance, xs)
dists.shape # (100, 100)
通过将 jax.vmap 与 jax.grad 和 jax.jit 组合,我们可以得到高效的
Jacobian 矩阵,或逐样本梯度:
per_example_grads = jax.jit(jax.vmap(jax.grad(loss), in_axes=(None, 0, 0)))
扩展
要在数千个设备上扩展你的计算,你可以使用以下任意组合:
- 基于编译器的自动并行化 你像使用单台全局机器一样进行编程,由编译器选择 如何分片数据并划分计算(带有一些用户提供的约束);
- 显式分片与自动划分
你仍然拥有全局视图,但数据分片
在 JAX 类型中是显式的,可以使用
jax.typeof进行检查; - 手动逐设备编程 你拥有数据和计算的逐设备视图, 并可以通过显式集合通信进行通信。
| 模式 | 视图? | 显式分片? | 显式集合通信? |
|---|---|---|---|
| 自动 | 全局 | ❌ | ❌ |
| 显式 | 全局 | ✅ | ❌ |
| 手动 | 逐设备 | ✅ | ✅ |
from jax.sharding import set_mesh, AxisType, PartitionSpec as P
mesh = jax.make_mesh((8,), ('data',), axis_types=(AxisType.Explicit,))
set_mesh(mesh)
# parameters are sharded for FSDP:
for W, b in params:
print(f'{jax.typeof(W)}') # f32[512@data,512]
print(f'{jax.typeof(b)}') # f32[512]
# shard data for batch parallelism:
inputs, targets = jax.device_put((inputs, targets), P('data'))
# evaluate gradients, automatically parallelized!
gradfun = jax.jit(jax.grad(loss))
param_grads = gradfun(params, (inputs, targets))
注意事项与易错点
请参阅注意事项 笔记本。
安装
支持的平台
| Linux x86_64 | Linux aarch64 | Mac aarch64 | Windows x86_64 | Windows WSL2 x86_64 | |
|---|---|---|---|---|---|
| CPU | 是 | 是 | 是 | 是 | 是 |
| NVIDIA GPU | 是 | 是 | 不适用 | 否 | 实验性支持 |
| Google TPU | 是 | 不适用 | 不适用 | 不适用 | 不适用 |
| AMD GPU | 是 | 否 | 不适用 | 否 | 实验性支持 |
| Apple GPU | 不适用 | 否 | 实验性支持 | 不适用 | 不适用 |
| Intel GPU | 实验性支持 | 不适用 | 不适用 | 否 | 否 |
说明
| 平台 | 说明 |
|---|---|
| CPU | pip install -U jax |
| NVIDIA GPU | pip install -U "jax[cuda13]" |
| Google TPU | pip install -U "jax[tpu]" |
| AMD GPU (Linux) | pip install -U "jax[rocm7-local]" |
| Intel GPU | 遵循 Intel 的说明. |
参见 文档 以获取有关替代安装策略的信息。这些包括从源代码编译、 使用 Docker 安装、使用其他版本的 CUDA、 社区支持的 conda 构建,以及一些常见问题的解答。
引用 JAX
要引用此仓库:
@software{jax2018github,
author = {James Bradbury and Roy Frostig and Peter Hawkins and Matthew James Johnson and Yash Katariya and Chris Leary and Dougal Maclaurin and George Necula and Adam Paszke and Jake Vander{P}las and Skye Wanderman-{M}ilne and Qiao Zhang},
title = {{JAX}: composable transformations of {P}ython+{N}um{P}y programs},
url = {http://github.com/jax-ml/jax},
version = {0.3.13},
year = {2018},
}
在上述 bibtex 条目中,姓名按字母顺序排列,版本号 旨在取自 jax/version.py,并且 年份对应于项目的开源发布。
JAX 的一个早期版本,仅支持自动微分和 编译到 XLA,在 一篇发表于 SysML 2018 的论文 中进行了描述。我们目前正在 撰写一篇更全面且最新的论文,以涵盖 JAX 的理念和能力。
参考文档
有关 JAX API 的详细信息,请参阅 参考文档。
有关作为 JAX 开发者入门的信息,请参阅 开发者文档。