与安装实战)
JAX 完全指南PythonNumPy 程序的可组合变换grad / jit / vmap / pmap与安装实战【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/gh_mirrors/jax/jax导读JAX 是一个面向加速器GPU/TPU的高性能数值计算与大规模机器学习 Python 库其核心能力是可组合的函数变换grad自动微分、jitXLA 即时编译、vmap自动向量化与pmap多设备 SPMD 并行四大变换可以任意嵌套组合让你在纯 Python 中写出兼具 NumPy 可读性与编译内核性能的代码。本文以仓库根目录 README.md 为骨架结合 jax/_src/api.py 等源码实现系统讲解 JAX 的定位、四大核心变换的用法与底层原理、常见坑位gotchas以及多平台安装方案读完后你可以直接上手编写可微分、可编译、可并行的高性能数值程序。一、JAX 是什么Autograd 的更新版 XLA 编译器JAX 定位为面向加速器的数组计算与程序变换库由两部分关键能力构成自动微分Autograd 的更新版可以对原生 Python 与 NumPy 函数自动求导且能穿过循环、分支、递归与闭包支持反向模式grad与前向模式并且可以对导数再求导、任意阶求导derivatives of derivatives of derivatives。XLA 编译与加速底层使用 XLA 将 NumPy 程序编译后在 GPU/TPU 上运行。库调用默认在幕后被 JIT 编译执行同时用户也可以通过jit这一个函数 API 把自己的 Python 函数即时编译为 XLA 优化内核。更本质地说JAX 是一个可扩展的复合函数变换系统grad、jit只是其中两个变换实例vmap自动向量化与pmap多加速器 SPMD 并行是另外两个四者可以任意组合。需要说明的是这是一个研究项目而非 Google 官方产品存在已知的sharp edges详见下文常见坑位一节。在仓库结构中这些变换的公开入口集中在 jax/_src/api.pyjit在 L142、grad在 L566、vmap在 L1035、pmap在 L1299其底层分别由 jax/_src/interpreters/ad.py反向模式vjp、linearize在 L120/L141与 jax/_src/interpreters/batching.pybatch在 L609等解释器支撑。开箱即用的核心示例README 给出的一段典型代码同时用到了三种变换——先grad求梯度再jit编译最后vmap批量计算逐样本梯度import jax.numpy as jnp from jax import grad, jit, vmap def predict(params, inputs): for W, b in params: outputs jnp.dot(inputs, W) b inputs jnp.tanh(outputs) # 作为下一层的输入 return outputs # 最后一层不加激活 def loss(params, inputs, targets): preds predict(params, inputs) return jnp.sum((preds - targets)**2) grad_loss jit(grad(loss)) # 编译后的梯度计算函数 perex_grads jit(vmap(grad_loss, in_axes(None, 0, 0))) # 快速逐样本梯度这段代码演示了 JAX 的核心心智模型变换是惰性包裹的jit(grad(...))把求导与编译叠成一层vmap再在外层批量展开全部在 Python 内完成。二、自动微分grad与高阶求导基本用法grad的 API 与 Autograd 大致相同用于反向模式梯度from jax import grad import jax.numpy as jnp def tanh(x): # 定义一个函数 y jnp.exp(-2.0 * x) return (1.0 - y) / (1.0 y) grad_tanh grad(tanh) # 获得它的梯度函数 print(grad_tanh(1.0)) # 在 x 1.0 处求值 # prints 0.4199743任意阶求导grad可以嵌套任意次print(grad(grad(grad(tanh)))(1.0)) # prints 0.62162673对 Python 控制流的支持与 Autograd 一致grad可以自由穿过 Python 控制结构且会按当前分支重新求值def abs_val(x): if x 0: return x else: return -x abs_val_grad grad(abs_val) print(abs_val_grad(1.0)) # prints 1.0 print(abs_val_grad(-1.0)) # prints -1.0abs_val 被重新求值源码视角grad的实现要点从 jax/_src/api.py 看grad实际委托给value_and_grad通过_vjp计算反向模式向量-雅可比积VJP再对结果乘上初始伴随值lax_internal._one(ans)得到梯度。其关键参数包括argnums对第几个位置参数求导默认 0可传整数或整数序列has_aux若fun返回(输出, 辅助数据)二元组则置 True此时返回(梯度, 辅助数据)holomorphic承诺函数是全纯函数时置 True要求输入输出为复数allow_int是否允许对整数输入求导梯度为 float0 平凡向量空间 dtype。注意grad要求被求导函数输出标量shape 为()的数组否则会抛出Gradient only defined for scalar-output functions错误相关检查见 jax/_src/api.py 的_check_scalar。更高级的原语vjp/jvp/ 雅可比 / 海森对于更高级的自动微分可以使用jax.vjp反向模式向量-雅可比积返回(输出, 反向函数)jax.jvp前向模式雅可比-向量积两者可互相组合也可与其它变换组合。例如下面的代码用jacfwd与jacrev组合出高效计算完整海森矩阵的函数from jax import jit, jacfwd, jacrev def hessian(fun): return jit(jacfwd(jacrev(fun)))对应实现分别为 jax/_src/api.pyjacfwd前向模式逐列求雅可比、jax/_src/api.pyjacrev反向模式逐行求雅可比与 jax/_src/api.pyhessian稠密海森矩阵。反向模式的底层由 jax/_src/interpreters/ad.py 的linearize/vjp提供 jaxpr 级线性化与转置支撑。更完整的原理讲解可继续阅读仓库内 docs/notebooks/autodiff_cookbook.md 与 docs/_tutorials/advanced-autodiff.md。三、编译jit与 XLA 融合内核jit可以用作装饰器也可以作为高阶函数使用将整个函数端到端编译到 XLAimport jax.numpy as jnp from jax import jit def slow_f(x): # 逐元素运算能从算子融合中获得巨大收益 return x * x x * 2.0 x jnp.ones((5000, 5000)) fast_f jit(slow_f) %timeit -n10 -r3 fast_f(x) # ~ 4.5 ms / loop on Titan X原文档示例数据 %timeit -n10 -r3 slow_f(x) # ~ 14.5 ms / loop同样经由 JAX 在 GPU 上运行jit可与grad以及任何其它变换任意混用。需要留意的是jit会对函数内可用的 Python 控制流形式施加约束详见下文的 Gotchas 一节。jit的关键参数源码级说明从 jax/_src/api.py 看jit的完整签名包括in_shardings、out_shardings、static_argnums、static_argnames、donate_argnums、donate_argnames、keep_unused、device、backend、inline、abstracted_axes等核心语义static_argnums/static_argnames把指定位置/名字参数视为编译期常量static它们必须可哈希且不可变换值会触发重新编译。非数组类参数如字符串、配置对象必须标记为 static。donate_argnums/donate_argnames声明哪些输入缓冲区允许被计算覆盖buffer donationXLA 可以借此复用输入缓冲区存储结果以减少显存占用捐赠后不得再复用这些缓冲区否则会报错。keep_unusedFalse默认未被函数使用的参数会被裁剪不再传输到设备端。in_shardings/out_shardings指定输入/输出的分片策略与jax.sharding.Sharding体系打通详见 jax/_src/sharding.py。device/backend指定执行设备如cpu、gpu、tpu或具体Device实例。在编译缓存方面JAX 对函数对象持有弱引用并以之作为编译缓存键因此传入jit的函数必须是可弱引用的。关于编译缓存与 AOT 的进一步讨论可参考 docs/persistent_compilation_cache.md 与 docs/aot.md。四、自动向量化vmap把循环压进原语vmap是向量化映射语义上等价于沿数组轴映射函数但不是把循环放在外层而是把循环下推到函数内部的原语运算中从而获得更好的性能省去在代码中手动搬运 batch 维度的麻烦。考虑一个仅适用于单输入向量的简单预测函数def predict(params, input_vec): assert input_vec.ndim 1 activations input_vec for W, b in params: outputs jnp.dot(W, activations) b # activations 在右侧 activations jnp.tanh(outputs) # 作为下一层输入 return outputs # 最后一层不加激活通常我们会写成jnp.dot(activations, W)以便在activations左侧加 batch 维但上面的函数只接受单个向量。要对一批输入同时计算语义上可以from functools import partial predictions jnp.stack(list(map(partial(predict, params), input_batch)))但逐个样本过网络太慢——更好的做法是把计算向量化让每一层都做矩阵-矩阵乘法而非矩阵-向量乘法。vmap自动完成这个变换from jax import vmap predictions vmap(partial(predict, params))(input_batch) # 或者等价地 predictions vmap(predict, in_axes(None, 0))(params, input_batch)执行后机器最终执行的正是矩阵-矩阵乘法效果与手工 batch 完全一致。in_axes(None, 0)表示params不映射None、input_batch沿第 0 轴映射。典型场景逐样本梯度per-example gradients手动向量化简单神经网络并不难但很多场景下手动向量化不现实甚至不可能例如高效计算逐样本梯度固定参数、对 batch 中每个样本分别求损失梯度。用vmap一行搞定per_example_gradients vmap(partial(grad(loss), params))(inputs, targets)vmap可以任意地与jit、grad及其它变换组合jax.jacfwd、jax.jacrev、jax.hessian中正是用它结合前向/反向自动微分完成快速雅可比与海森矩阵计算。源码视角vmap的参数语义从 jax/_src/api.py 看vmap的签名包含in_axes默认 0整数 /None/ 序列指定对哪些输入轴做映射轴序号必须在[-ndim, ndim)范围内若参数是 pytree 容器则in_axes需要是对应树前缀结构关键字参数总是映射其首轴轴 0。out_axes默认 0指定映射轴出现在输出的哪个位置。axis_name给映射轴命名以便在函数内部使用并行集合通信与pmap的axis_name语义一致。axis_size显式指定映射轴大小不提供时从参数推断。其底层实现位于 jax/_src/interpreters/batching.py 的batch函数L609它把被映射轴以批处理维度形式下推进每个原语而非逐元素循环。向量化的更深层原理可参考 docs/notebooks/How_JAX_primitives_work.md。五、多设备 SPMD 并行pmappmap用于对多个加速器如多张 GPU做并行编程用它写单程序多数据SPMD程序支持快速的并行集合通信。应用pmap后函数会像jit一样被 XLA 编译然后复制到各设备上并行执行。pmap与vmap语义上相似都是沿数组轴映射函数但区别在于vmap把映射轴压进原语做向量化而pmap是复制函数、让每个副本在各 XLA 设备上并行执行见 jax/_src/api.py 的 docstring。在 8 卡 GPU 机器上的例子from jax import random, pmap import jax.numpy as jnp # 创建 8 个随机 5000 x 6000 矩阵每张 GPU 一个 keys random.split(random.PRNGKey(0), 8) mats pmap(lambda key: random.normal(key, (5000, 6000)))(keys) # 在每个设备上并行执行本地矩阵乘无数据搬移 result pmap(lambda x: jnp.dot(x, x.T))(mats) # result.shape 为 (8, 5000, 5000) # 在每个设备上并行求均值并打印 print(pmap(jnp.mean)(result)) # prints [1.1566595 1.1805978 ... 1.2321935 1.2015157]集合通信lax.psum等除了纯映射pmap内部还可以使用设备间快速集合通信算子对应jax.lax中的并行算子详见 jax/_src/lax/init.pyfrom functools import partial from jax import lax partial(pmap, axis_namei) def normalize(x): return x / lax.psum(x, i) print(normalize(jnp.arange(4.))) # prints [0. 0.16666667 0.33333334 0.5 ]这里的axis_namei正是前文vmap中axis_name参数的同一机制——集合算子按轴名定位跨设备的归约目标。嵌套pmap与对并行计算求导pmap可以嵌套以获得更复杂的通信模式更关键的是它可以与自动微分任意组合from jax import grad pmap def f(x): y jnp.sin(x) pmap def g(z): return jnp.cos(z) * jnp.tan(y.sum()) * jnp.tanh(x).sum() return grad(lambda w: jnp.sum(g(w)))(x) print(f(x)) # [[ 0. , -0.7170853 ], # [-3.1085174 , -0.4824318 ], # [10.366636 , 13.135289 ], # [ 0.22163185, -0.52112055]] print(grad(lambda x: jnp.sum(f(x)))(x)) # [[ -3.2369726, -1.6356447], # [ 4.7572474, 11.606951 ], # [-98.524414 , 42.76499 ], # [ -1.6007166, -1.2568436]]对pmap函数做反向模式求导时如外层套grad反向传播会被前向过程一样并行化。从源码看pmap还支持devices指定参与设备、static_broadcasted_argnums、donate_argnums、global_arg_shapes等参数见 jax/_src/api.py其映射轴大小必须不超过jax.local_device_count()返回的本地设备数嵌套时乘积不超过设备总数。更深入的内容可继续阅读 cloud_tpu_colabs/Pmap_Cookbook.ipynb 与完整的 SPMD MNIST 示例 examples/spmd_mnist_classifier_fromscratch.py。若你的需求是按分片规范自动并行而非手工pmap可以进一步了解jax.experimental.pjit入口 jax/experimental/pjit.py文档见 docs/jax.experimental.pjit.rst。六、常见坑位Gotchas以下是从 README 提炼的几类高频坑更完整的示例与解释见 docs/notebooks/Common_Gotchas_in_JAX.md变换只作用于纯函数JAX 变换要求函数无副作用且满足引用透明性对象同一性判断is不保留。对不纯函数做变换可能报Exception: Cant lift Traced...或Exception: Different traces at same level。不支持就地修改x[i] y这类就地可变更新不受支持但有函数式替代方案jax.numpy的at系列如x.at[i].add(y)对应 jax/_src/numpy/init.py。在jit下这些函数式替代会自动复用缓冲区buffer reuse。随机数体系不同JAX 的随机数基于显式 PRNG key 的可分叉设计random.PRNGKey/random.key/random.split/random.fold_in见 jax/_src/random.py背后有充分理由设计文档见 docs/jep/263-prng.md。卷积算子在jax.lax包中需要卷积算子时从jax.lax导入相关教程见 docs/notebooks/convolutions.md。默认单精度32-bitJAX 默认使用float32要启用 64 位需在启动时设置jax_enable_x64标志或在环境变量中设置JAX_ENABLE_X64True对应配置项见 jax/_src/config.py。在 TPU 上除jnp.dot、lax.conv等类矩阵乘运算内部临时变量外均默认 32 位这些算子带precision参数可通过三次 bfloat16 通道近似 32 位运算代价是可能变慢。TPU 上的非矩阵乘运算多倾向速度优先的实现因此实际精度通常低于其它后端。部分 dtype 提升语义不同NumPy 中 Python 标量与 NumPy 类型混用的部分提升语义不被保留例如np.add(1, np.array([2], np.float32)).dtype在 JAX 中得到float64而非float32。控制流受限jit等变换会约束 Python 控制流的使用方式出错时会给出响亮的报错。解决方案包括使用jit的static_argnums参数、结构化控制流原语如lax.scan见 jax/_src/lax/init.py 与 docs/jax.lax.rst或把jit用于较小的子函数。七、安装指南支持平台矩阵Linux x86_64Linux aarch64Mac x86_64Mac ARMWindows x86_64Windows WSL2 x86_64CPUyesyesyesyesyesyesNVIDIA GPUyesyesnon/anoexperimentalGoogle TPUyesn/an/an/an/an/aAMD GPUexperimentalnonon/anonoApple GPUn/anoexperimentalexperimentaln/an/a安装命令硬件安装指令CPUpip install -U jaxNVIDIA GPUpip install -U jax[cuda12]Google TPUpip install -U jax[tpu] -f https://storage.googleapis.com/jax-releases/libtpu_releases.htmlAMD GPU使用官方 Docker 镜像或从源码构建Apple GPU按 Apple 官方 Metal/JAX 指引操作其它安装策略从源码编译、Docker、其它 CUDA 版本、社区 conda 构建与 FAQ见 docs/installation.md。从源码构建的详细步骤可参考 docs/developer.mdGPU 内核 wheel 与插件 wheel 的构建脚本位于 jaxlib/tools/如 jaxlib/tools/build_gpu_kernels_wheel.py、jaxlib/tools/build_gpu_plugin_wheel.py。本仓库当前源码版本信息见 jax/version.py_version 0.4.31且与 jaxlib 存在最低版本约束_minimum_jaxlib_version。安装完成后可用以下方式快速验证import jax print(jax.devices()) # 列出当前可见的设备CPU/GPU/TPU print(jax.__version__) # 打印 JAX 版本八、神经网络生态与上手资源JAX 本身不内置训练框架但其生态由多个库补齐Flax功能完整的神经网络训练库含示例与指南推荐其新的NNXAPI 以获得更简化的开发体验。Equinox以 JAX 为基础构建的神经网络库是 JAX 生态中若干上层库的基础。DeepMind 开源生态包括Optax梯度处理与优化、RLax强化学习算法、chex可靠测试工具等。在本仓库内可直接上手与验证的资源快速入门教程docs/quickstart.md、docs/key-concepts.md神经网络与数据加载含 TFDS 数据集docs/notebooks/neural_network_with_tfds_data.ipynbMNIST 分类器完整实现examples/mnist_classifier.py 与从零开始的版本 examples/mnist_classifier_fromscratch.py更多笔记本见 docs/notebooks/ 目录。九、引用与参考文档如需在论文中引用 JAX可使用 README 提供的 BibTeX 条目software{jax2018github, author {James Bradbury and Roy Frostig and Peter Hawkins and Matthew James Johnson 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/google/jax}, version {0.3.13}, year {2018}, }条目中名字按字母序排列version 字段取自 jax/version.pyyear 对应项目开源发布年份。早期仅支持自动微分与 XLA 编译的 JAX 雏形曾在 SysML 2018 会议上以论文形式描述。JAX API 的详细参考文档对应仓库 docs/jax.rst 及其下各子模块文档如 docs/jax.numpy.rst、docs/jax.lax.rst、docs/jax.random.rstJAX 开发者入门文档见 docs/developer.md。想要从零理解 JAX 变换原理的读者强烈推荐 docs/autodidax.md用少量 Python 从零实现 JAX 核心变换的教程。结语JAX 的核心价值在于变换这一统一抽象grad、jit、vmap、pmap四个变换覆盖了求导、编译、向量化与多设备并行四大高频需求且可以任意嵌套组合——正如 README 中的示例所示jit(vmap(grad(loss)))一条表达式就同时获得了梯度、编译与逐样本批处理。理解这四个变换的语义边界纯函数、不可变数组、显式 PRNG、32 位默认精度、控制流约束是写出正确 JAX 程序的关键。本文所有结论均可回溯至仓库源码与测试建议读者结合 tests/ 目录下的lax_test.py、lax_autodiff_test.py、lax_vmap_test.py、pmap_test.py、pjit_test.py等测试文件进一步验证各变换的实际行为与组合语义。【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/gh_mirrors/jax/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考