
机器学习深度学习【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址https://gitcode.com/gh_mirrors/jax/jax点击查看免费下载JAX 是一个面向数组数值计算à la NumPy的库内置自动微分与 JIT 编译能力用于支撑高性能的机器学习研究。本文基于仓库中的 docs/quickstart.md 展开系统介绍 JAX 的安装方式、NumPy 风格编程接口以及三个最核心的变换原语——jax.jit即时编译、jax.grad自动微分与jax.vmap自动向量化并深入讲解它们如何任意组合帮助你快速上手 JAX 的高性能计算开发。阅读完本文你将能够在自己的机器上安装 JAXCPU 或 NVIDIA GPU用熟悉的 NumPy 语法编写并运行 JAX 程序通过jit把逐算子执行的开销合成为一次 XLA 编译用grad/jacfwd/jacrev等变换高效求导并自行组合出 Hessian用vmap把逐样本循环自动改写为原生向量化实现。安装 JAXJAX 可以按平台直接通过 Python 包索引安装。在 Linux、Windows 与 macOS 上安装 CPU 版本pip install jax如需在 NVIDIA GPU 上运行则安装 CUDA 12 插件版本pip install -U jax[cuda12]更详细的平台专属安装说明包括 TPU、ROCm、源码构建等参见仓库中的 安装指南。注意从仓库的jax_plugins/目录结构可以看出CUDA 与 ROCm 支持是通过独立的插件包jax_plugins/cuda/ 与 jax_plugins/rocm/提供的pip install jax只安装 CPU 后端GPU 后端按需通过 extra 依赖单独安装。JAX 作为 NumPy 使用JAX 的大部分用法都集中在jax.numpy这个与 NumPy 兼容的 API 上通常以jnp为别名导入import jax.numpy as jnp导入之后就可以像编写 NumPy 程序一样使用 JAX支持 NumPy 风格的数组创建函数、Python 函数与运算符以及数组的属性与方法。例如定义一个带参数的 SELU 激活函数def selu(x, alpha1.67, lmbda1.05): return lmbda * jnp.where(x 0, x, alpha * jnp.exp(x) - alpha) x jnp.arange(5.0) print(selu(x))这里jnp.arange(5.0)创建数组[0., 1., 2., 3., 4.]对非负输入 SELU 等价于恒等映射因此输出应为[0., 1.05, 2.1, 3.15, 4.2]。JAX 数组与 NumPy 数组之间也存在一些细微差异如不可变性、随机数的显式 key 管理、类型提升规则等这些差异在 JAX - The Sharp Bits 中有专门梳理建议进阶使用时通读。用jax.jit实现即时编译JAX 可以在 GPU 或 TPU 上透明运行没有对应硬件时回退到 CPU。但在上面的例子中JAX 是逐算子把 kernel 派发到芯片上执行的。如果有一连串操作可以用jax.jit把它们通过 XLA 编译成单个融合计算从而显著降低派发开销。由于 JAX 采用动态派发异步分发模型基准测试时需要使用block_until_ready()阻塞等待实际计算结果才能得到真实的执行时间参见 异步分发说明from jax import random key random.key(1701) x random.normal(key, (1_000_000,)) %timeit selu(x).block_until_ready()这里用到了jax.random生成随机数。JAX 随机数采用显式 key 管理random.key(seed)根据整数种子创建 PRNG key源码见 jax/_src/random.py#L197random.normal(key, shape)以该 key 采样标准正态分布源码见 jax/_src/random.py#L676。关于 JAX 伪随机数生成的设计细节可参考 随机数文档。接下来用jax.jit变换加速selu的执行——第一次调用时完成编译之后的结果会被缓存复用from jax import jit selu_jit jit(selu) _ selu_jit(x) # 首次调用触发编译 %timeit selu_jit(x).block_until_ready()上述计时是在 CPU 上执行的同样的代码在 GPU 或 TPU 上运行通常能获得更大的加速。jax.jit的完整签名与关键参数从源码看jit的完整签名位于 jax/_src/api.py#L142除了函数本身外还支持以下常用参数参数作用默认值in_shardings指定输入分片Sharding用于分布式场景缺省时从参数的分片推断未指定out_shardings指定输出分片等价于对输出施加lax.with_sharding_constraint未指定static_argnums/static_argnames将若干位置参数/关键字参数标记为静态编译期常量参数值变化会触发重编译非数组参数必须标记为静态Nonedonate_argnums/donate_argnames标记可捐赠buffer 复用的输入参数减少内存分配Nonekeep_unused是否保留未使用的输入Falsedevice指定编译运行的目标设备Nonebackend指定后端名称Noneinline是否内联展开False理解static_argnums很重要静态参数会作为编译缓存 key 的一部分参与缓存匹配因此必须是可哈希实现__hash__与__eq__且不可变的对象用不同值调用被 jit 的函数会触发重新编译。更深入的 JIT 编译机制缓存、trace 流程、AOT 等参见 jit 编译专题。用jax.grad求导除了 JIT 编译JAX 还提供其他函数变换其中最典型的是jax.grad——它执行自动微分autodiff。看一个对 logistic 函数之和求导的例子from jax import grad def sum_logistic(x): return jnp.sum(1.0 / (1.0 jnp.exp(-x))) x_small jnp.arange(3.) derivative_fn grad(sum_logistic) print(derivative_fn(x_small))grad要求被求导函数返回标量包括 shape 为()的数组但不包括(1,)等形状默认对第一个位置参数求导。可以用有限差分验证结果正确性def first_finite_differences(f, x, eps1E-3): return jnp.array([(f(x eps * v) - f(x - eps * v)) / (2 * eps) for v in jnp.eye(len(x))]) print(first_finite_differences(sum_logistic, x_small))grad与jit可以任意组合、任意嵌套。上面的例子先 jit 了sum_logistic再求导还能继续叠加print(grad(jit(grad(jit(grad(sum_logistic)))))(1.0))即对sum_logistic求三阶导数其中每层求导之间都穿插了 JIT 编译。grad的关键参数grad的完整签名在 jax/_src/api.py#L566argnums指定对哪个些位置参数求导默认0传入整数元组时返回对应参数的梯度元组has_auxfun返回(输出, 辅助数据)二元组时置为True梯度与辅助数据一起返回holomorphic承诺被求导函数是全纯函数输入输出为复数时置为Trueallow_int是否允许对整数输入求导此时梯度为 float0 类型的平凡向量空间。从源码实现看grad内部实际是基于value_and_grad构造的jax/_src/api.py#L609value_and_grad一次性同时返回函数值与梯度而grad只取出其中的梯度分量。向量值函数jacobian、jvp、vjp与jacfwd/jacrev对于输出为向量的函数jax.jacobian变换可以计算完整的 Jacobian 矩阵from jax import jacobian print(jacobian(jnp.exp)(x_small))更进阶的自动微分原语包括jax.vjp反向模式向量-Jacobian 乘积reverse-mode vector-Jacobian productsjax.jvp与jax.linearize前向模式 Jacobian-向量乘积forward-mode Jacobian-vector products。这些原语可以彼此任意组合也可以与其他 JAX 变换组合。例如jax.jvp与jax.vjp分别被用来定义前向模式的jax.jacfwd与反向模式的jax.jacrev用于在相应模式下计算 Jacobian。下面是用它们组合出高效 Hessian 计算函数的一个示例from jax import jacfwd, jacrev def hessian(fun): return jit(jacfwd(jacrev(fun))) print(hessian(sum_logistic)(x_small))这种嵌套组合在实践中能生成高效代码——它本质上就是 JAX 内置jax.hessian的实现方式内置版本源码见 jax/_src/api.py#L918同样支持argnums、has_aux、holomorphic等参数。更系统的自动微分讲解参见 自动微分文档。用jax.vmap自动向量化另一个常用的变换是jax.vmapvectorizing map。它的语义与在数组轴上显式映射函数一致但并非在 Python 层循环调用函数而是把函数变换为原生向量化版本以获得更好性能。当与jit组合时其性能可以媲美手写批处理维度。下面用一个具体例子说明把矩阵-向量乘积提升为矩阵-矩阵乘积。key1, key2 random.split(key) mat random.normal(key1, (150, 100)) batched_x random.normal(key2, (10, 100)) def apply_matrix(x): return jnp.dot(mat, x)apply_matrix把向量映射到向量但我们想逐行作用在一个 batch 矩阵上。最直接的做法是在 Python 层对 batch 维循环但通常性能不佳def naively_batched_apply_matrix(v_batched): return jnp.stack([apply_matrix(v) for v in v_batched]) print(Naively batched) %timeit naively_batched_apply_matrix(batched_x).block_until_ready()熟悉jnp.dot的程序员可能会手动改写利用jnp.dot内置的批处理语义避免显式循环import numpy as np jit def batched_apply_matrix(batched_x): return jnp.dot(batched_x, mat.T) np.testing.assert_allclose(naively_batched_apply_matrix(batched_x), batched_apply_matrix(batched_x), atol1E-4, rtol1E-4) print(Manually batched) %timeit batched_apply_matrix(batched_x).block_until_ready()但随着函数越来越复杂这种手工批处理会愈发困难且易错。jax.vmap的设计目标就是把函数自动变换为感知批处理维度的版本from jax import vmap jit def vmap_batched_apply_matrix(batched_x): return vmap(apply_matrix)(batched_x) np.testing.assert_allclose(naively_batched_apply_matrix(batched_x), vmap_batched_apply_matrix(batched_x), atol1E-4, rtol1E-4) print(Auto-vectorized with vmap) %timeit vmap_batched_apply_matrix(batched_x).block_until_ready()三种实现的数值结果通过np.testing.assert_allclose校验一致而vmap版本免去了手工重写函数的负担。vmap的关键参数vmap的完整签名位于 jax/_src/api.py#L1035in_axes整数、None或序列指定对哪些输入的哪些轴做映射。整数必须落在对应输入数组的[-ndim, ndim)范围内None表示该参数不参与映射对 pytree 容器输入则要求是参数的 tree prefix。关键字参数总是映射其首轴轴索引 0。默认0out_axes整数、None或嵌套容器指明映射轴出现在输出的哪个位置取值范围同样是[-ndim, ndim)默认0axis_name可哈希对象用于标识被映射的轴以便在映射轴之上应用并行集合通信如lax.pmeanaxis_size可选整数指定映射轴的尺寸未提供时从输入参数推断此时至少有一个位置参数的非None映射轴且所有映射轴的尺寸必须一致。与预期一致vmap可以与jit、grad以及任何其他 JAX 变换任意组合。更深入的向量化变换机制in_axes/out_axes的完整规则、与jit的协同等参见 自动向量化文档。三种变换的任意组合JAX 的设计精髓在于变换之间的可组合性jit编译优化、grad自动微分、vmap自动向量化可以以任意顺序、任意深度嵌套。例如grad(jit(grad(jit(grad(sum_logistic)))))(1.0)对函数反复求导并穿插编译jit(jacfwd(jacrev(fun)))前向与反向模式组合求 Hessianjit(vmap(apply_matrix))自动向量化后再整体编译。这种组合能力意味着你可以先用自然、直观的方式写出标量/逐样本逻辑再通过变换声明式地获得编译优化、梯度与批处理能力而无需手工改写算法结构。进一步阅读本快速入门只是 JAX 能力的冰山一角建议按需深入仓库中的以下文档JIT 编译专题编译缓存、trace 机制与 AOT 编译自动微分jvp/vjp、jacfwd/jacrev、自定义导数规则自动向量化vmap的轴映射规则与复杂用法异步分发理解block_until_ready()与动态派发模型随机数生成PRNG key 的生成、拆分与复用规则Common GotchasJAX 数组与 NumPy 数组的关键差异。赞分享机器学习深度学习【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址https://gitcode.com/gh_mirrors/jax/jax点击查看免费下载相关推荐Flax NNX 变换系统全解析grad、jit、vmap、scan 等 16 个核心变换 API 实战指南Flax NNX 变换系统全解析grad、jit、vmap、scan 等 16 个核心变换 API 实战指南 本篇技术指南以 Flax NNX 的官方 API人工智能深度学习机器学习JAX 完全指南PythonNumPy 程序的可组合变换grad / jit / vmap / pmap与安装实战JAX 完全指南PythonNumPy 程序的可组合变换grad / jit / vmap / pmap与安装实战 导读 JAX 是一个面向加速器GP机器学习深度学习JAX 原语Primitives完全指南从自定义 primitive 到 jit/grad/vmap 全变换支持JAX 原语Primitives完全指南从自定义 primitive 到 jit/grad/vmap 全变换支持 导读 本指南基于 JAX 官方教程 do机器学习深度学习上一篇WrenAI 的 wren 包开发指南WrenEngine 语义层、CLI 与 wren-core-py 绑定的完整解读下一篇VMulti 开源项目使用教程创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考