ITADN
jax-ml/jax
README.md
以下内容由 AI 翻译,如有问题请点此提交 issue 反馈
logo

大规模可转换数值计算

Continuous integration PyPI version

变换 | 缩放 | 安装指南 | 变更日志 | 参考文档

什么是 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.gradjax.jitjax.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.vmapjax.gradjax.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_64Linux aarch64Mac aarch64Windows x86_64Windows WSL2 x86_64
CPU
NVIDIA GPU不适用实验性支持
Google TPU不适用不适用不适用不适用
AMD GPU不适用实验性支持
Apple GPU不适用实验性支持不适用不适用
Intel GPU实验性支持不适用不适用

说明

平台说明
CPUpip install -U jax
NVIDIA GPUpip install -U "jax[cuda13]"
Google TPUpip 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 开发者入门的信息,请参阅 开发者文档