JAX 前向与反向自动微分完全指南:JVP、VJP 与 Hessian-vector products 的原理与实战

发布时间:2026/9/10 15:40:38
JAX 前向与反向自动微分完全指南:JVP、VJP 与 Hessian-vector products 的原理与实战 JAX 前向与反向自动微分完全指南JVP、VJP 与 Hessian-vector products 的原理与实战【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jaxJAX 同时内置了前向模式forward-mode与反向模式reverse-mode两种自动微分automatic differentiation, AD实现本文以官方进阶指南 docs/jacobian-vector-products.md 为主体系统讲解 Jacobian-vector productsJVP与 Vector-Jacobian productsVJP的数学定义、类型签名、计算复杂度以及如何用它们组合出 Hessian-vector products、矩阵-Jacobian 乘积并最终还原jax.jacfwd/jax.jacrev的底层实现。读完本文你将理解jax.grad为什么能高效训练百万、十亿级参数的神经网络掌握jax.jvp、jax.vjp、jax.jacfwd、jax.jacrev、jax.hessian的选型逻辑并能用jax.vmap组合它们写出高性能的自定义微分算子。一、两种自动微分模式为什么需要一点数学背景熟悉的jax.grad构建在反向模式之上但要讲清楚两种模式的差异、以及各自在什么场景下更有用需要先铺垫一些数学背景。核心问题可以概括为给定函数 $f : \mathbb{R}^n \to \mathbb{R}^m$我们想要沿着某个方向求导数。前向模式直接计算 Jacobian-vector productJVP即 $\partial f(x) v$代价约为一次函数求值的 3 倍且内存与计算深度无关反向模式计算 vector-Jacobian productVJP即 $v^\mathsf{T} \partial f(x)$一次调用就能得到标量损失函数的梯度但内存随计算深度线性增长。JAX 对两种模式都有高效且通用的实现。相关 API 的完整定义集中在 jax/_src/api.py底层 AD 解释器实现在 jax/_src/interpreters/ad.py本文后续会逐一给出源码级证据。二、前向模式Jacobian-vector productsJVP2.1 数学定义从 Jacobian 矩阵到 pushforward 线性映射给定函数 $f : \mathbb{R}^n \to \mathbb{R}^m$$f$ 在输入点 $x \in \mathbb{R}^n$ 处的 Jacobian 记为 $\partial f(x)$通常被看作一个 $m \times n$ 矩阵$\qquad \partial f(x) \in \mathbb{R}^{m \times n}$。但也可以把 $\partial f(x)$ 看作一个线性映射它把 $f$ 的定义域在 $x$ 处的切空间即另一份 $\mathbb{R}^n$映到 $f$ 的值域在 $f(x)$ 处的切空间即一份 $\mathbb{R}^m$$\qquad \partial f(x) : \mathbb{R}^n \to \mathbb{R}^m$。这个映射被称为 $f$ 在 $x$ 处的 pushforward前推映射Jacobian 矩阵只是该线性映射在标准基下的矩阵表示。如果不固定具体的输入点 $x$可以把 $\partial f$ 看成一个先接收输入点、再返回该点处 Jacobian 线性映射的函数$\qquad \partial f : \mathbb{R}^n \to \mathbb{R}^n \to \mathbb{R}^m$。对输入点 $x \in \mathbb{R}^n$ 和一个切向量 $v \in \mathbb{R}^n$我们得到输出切向量 $\in \mathbb{R}^m$。这个从 $(x, v)$ 对到输出切向量的映射就是Jacobian-vector product$\qquad (x, v) \mapsto \partial f(x) v$。2.2 JAX 代码中的 JVPjax.jvp回到 Python 代码JAX 的jax.jvp正是对这个变换的建模给定一个计算 $f$ 的 Python 函数jax.jvp返回一个计算 $(x, v) \mapsto (f(x), \partial f(x) v)$ 的 Python 函数。下面用官方指南中的玩具模型sigmoid 二分类预测演示import jax import jax.numpy as jnp key jax.random.key(0) # 初始化随机模型系数 key, W_key, b_key jax.random.split(key, 3) W jax.random.normal(W_key, (3,)) b jax.random.normal(b_key, ()) # 定义 sigmoid 函数 def sigmoid(x): return 0.5 * (jnp.tanh(x / 2) 1) # 输出标签为真的概率 def predict(W, b, inputs): return sigmoid(jnp.dot(inputs, W) b) # 构造玩具数据集 inputs jnp.array([[0.52, 1.12, 0.77], [0.88, -1.08, 0.15], [0.52, 0.06, -1.30], [0.74, -2.49, 1.39]]) # 隔离出从权重矩阵到预测值的函数 f lambda W: predict(W, b, inputs) key, subkey jax.random.split(key) v jax.random.normal(subkey, W.shape) # 沿 f 在 W 处前推向量 v y, u jax.jvp(f, (W,), (v,))这里primals与tangents都要求是 tuple 或 listjax/_src/api.py 中会强制校验类型并检查二者的树结构与形状、dtype 是否匹配返回值是一个(primals_out, tangents_out)对y f(W)u ∂f(W)·v。从源码看jax.jvp实际调用 jax/_src/interpreters/ad.py 中的ad.jvp它会创建JVPTrace把每个 primal 值配上对应的 tangent 值封装成JVPTracer重新执行fun。每遇到一个原始数值运算就同时执行该运算的 JVP 规则——既在 primal 上求值也在该 primal 点应用其 JVP——这正是下文复杂度结论的实现基础。2.3 类型签名视角借用 Haskell 风格的签名可以更精确地描述jvp :: (a - b) - a - T a - (b, T b)其中T a表示a的切空间类型。也就是说jvp接收一个类型为a - b的函数、一个类型为a的值、一个类型为T a的切向量返回一个由类型为b的值和类型为T b的输出切向量组成的对。被jvp变换后的函数其求值方式与原函数几乎一样只是每个类型为a的 primal 值旁边都带上了类型为T a的 tangent 值。原函数每应用一个基本数值运算jvp变换后的函数就执行该基本运算的 JVP 规则既在 primal 上求值该运算又在该 primal 值处应用该运算的 JVP。2.4 计算复杂度约 3 倍 FLOPs 与与深度无关的内存这种边算边推的求值策略直接决定了复杂度特征内存与计算深度无关由于 JVP 是随求值过程即时推进的不需要为后续存储任何中间结果因此内存成本不随计算的深度增长FLOPs 约为原函数的 3 倍一部分工作量用于求值原函数例如sin(x)一部分用于线性化例如cos(x)一部分用于把线性化函数作用到向量上例如cos_x * v。换句话说固定 primal 点 $x$ 后以约等于一次f求值的边际代价就可以对任意方向 $v$ 计算 $v \mapsto \partial f(x) \cdot v$。2.5 为什么机器学习中很少单独使用前向模式内存优势听起来很有吸引力那为什么前向模式在机器学习里不常见关键在于如何用 JVP 拼出完整 Jacobian 矩阵如果把 JVP 作用在 one-hot 切向量上它就揭示 Jacobian 矩阵中与该非零分量对应的一列。因此可以逐列构建完整的 Jacobian而每一列的成本都约等于一次函数求值。这对高瘦tall的 Jacobian 高效但对宽扁wide的 Jacobian 低效。而基于梯度的机器学习优化目标函数是从参数空间 $\mathbb{R}^n$ 到标量损失 $\mathbb{R}$ 的映射其 Jacobian 是一个极宽的矩阵$\partial f(x) \in \mathbb{R}^{1 \times n}$通常与梯度向量 $\nabla f(x) \in \mathbb{R}^n$ 等同。逐列构建这个矩阵、每列都花一次函数求值的 FLOPs显然低效——尤其当 $f$ 是训练损失函数、$n$ 高达百万甚至十亿时这种方法根本无法扩展。要做得更好就需要反向模式。三、反向模式Vector-Jacobian productsVJP前向模式给出计算 Jacobian-vector product 的函数可逐列构建 Jacobian反向模式则给出计算 vector-Jacobian product等价地Jacobian 转置-向量乘积的函数可逐行构建 Jacobian。3.1 数学定义pullback 与转置仍考虑函数 $f : \mathbb{R}^n \to \mathbb{R}^m$。沿用 JVP 的记号VJP 的记号非常简洁$\qquad (x, v) \mapsto v \partial f(x)$其中 $v$ 是 $f$ 在 $x$ 处余切空间cotangent space与另一份 $\mathbb{R}^m$ 同构中的元素。严格地说应把 $v$ 看作线性映射 $v : \mathbb{R}^m \to \mathbb{R}$把 $v \partial f(x)$ 理解为函数复合 $v \circ \partial f(x)$——类型之所以成立是因为 $\partial f(x) : \mathbb{R}^n \to \mathbb{R}^m$。但在常见情况下可以把 $v$ 等同于 $\mathbb{R}^m$ 中的向量两者几乎可以互换使用就像我们有时会在列向量和行向量之间切换而不加说明一样。有了这个等同也可以把 VJP 的线性部分看作 JVP 线性部分的转置或伴随共轭$\qquad (x, v) \mapsto \partial f(x)^\mathsf{T} v$。对给定点 $x$签名可以写作$\qquad \partial f(x)^\mathsf{T} : \mathbb{R}^m \to \mathbb{R}^n$。余切空间上的这个对应映射常被称为 $f$ 在 $x$ 处的 pullback拉回。对我们的目的而言关键在于它从看起来像 $f$ 输出的量出发得到看起来像 $f$ 输入的量——正如我们对转置线性映射的预期。3.2 JAX 代码中的 VJPjax.vjp从数学回到 PythonJAX 的vjp接收一个计算 $f$ 的 Python 函数返回一个计算 $(x, v) \mapsto (f(x), v^\mathsf{T} \partial f(x))$ 的 Python 函数from jax import vjp # 隔离出从权重矩阵到预测值的函数 f lambda W: predict(W, b, inputs) y, vjp_fun vjp(f, W) key, subkey jax.random.split(key) u jax.random.normal(subkey, y.shape) # 沿 f 在 W 处拉回余向量 u v vjp_fun(u)vjp返回的vjp_fun是一个线性函数它接收与输出同形状的余切向量cotangent返回与每个输入同形状的余切向量。值得注意的是vjp在这里是分两步走的——先调用vjp(f, W)完成正向求值与线性化、拿到vjp_fun之后再多次调用vjp_fun(u)传入不同的余切向量做拉回。从源码看jax/_src/api.pyvjp通过ad.linearize对函数做线性化生成切线方向的 jaxpr 与残差residuals再封装出可调用的VJP对象其 docstring 明确指出jax.grad是jax.vjp的特例gradis implemented as a special case ofvjp。3.3 类型签名视角同样可以用 Haskell 风格签名表示vjp :: (a - b) - a - (b, CT b - CT a)其中CT a表示a的余切空间类型。vjp接收一个类型为a - b的函数和一个类型为a的点返回一个由类型为b的值和类型为CT b - CT a的线性映射组成的对。3.4 复杂度与jax.grad的效率来源VJP 让我们可以逐行构建 Jacobian 矩阵而计算 $(x, v) \mapsto (f(x), v^\mathsf{T} \partial f(x))$ 的 FLOPs 代价同样只有求值 $f$ 的约 3 倍。特别是如果要求 $f : \mathbb{R}^n \to \mathbb{R}$ 的梯度只需一次调用即可完成。这就是jax.grad对基于梯度的优化高效的原因——即使目标函数是参数以百万、十亿计的神经网络训练损失。不过反向模式也有代价虽然 FLOPs 友好但内存随计算深度增长需要保存正向传播中的中间值供反向使用且其实现传统上比前向模式复杂得多——不过 JAX 对此有一些技巧例如通过部分求值partial evaluation把线性化后的计算压缩成 tangent jaxpr 再执行反向传播相关机制可参见 jax/_src/interpreters/ad.py 中的linearize/linearize_jaxpr/backward_pass系列函数。3.5 用 VJP 实现向量值梯度如果你需要的是向量值梯度类似tf.gradients可以用 VJP 这样实现def vgrad(f, x): y, vjp_fn jax.vjp(f, x) return vjp_fn(jnp.ones(y.shape))[0] print(vgrad(lambda x: 3*x**2, jnp.ones((2, 2))))其思路是把全 1 向量作为输出余切传入vjp_fn一次性拉回得到对输入每个分量的梯度。四、Hessian-vector products前向与反向的协同4.1 纯反向模式的 HVP 基线在前一节的基础上先用纯反向模式实现一个 Hessian-vector product假设二阶导数连续def hvp(f, x, v): return jax.grad(lambda x: jnp.vdot(jax.grad(f)(x), v))(x)这个实现是高效的但还可以做得更好——把前向模式和反向模式组合起来能进一步节省内存。4.2 数学推导对梯度做 JVP给定要微分的函数 $f : \mathbb{R}^n \to \mathbb{R}$、线性化点 $x \in \mathbb{R}^n$ 和向量 $v \in \mathbb{R}^n$我们想要的 Hessian-vector product 是$(x, v) \mapsto \partial^2 f(x) v$。考虑辅助函数 $g : \mathbb{R}^n \to \mathbb{R}^n$它是 $f$ 的导数梯度即 $g(x) \partial f(x)$。我们只需要它的 JVP因为$(x, v) \mapsto \partial g(x) v \partial^2 f(x) v$。这个推导几乎可以一字不差地翻译成代码# forward-over-reverse def hvp(f, primals, tangents): return jax.jvp(jax.grad(f), primals, tangents)[1]这里jax.grad(f)是反向模式外层是前向的jax.jvp因此这种写法被称为forward-over-reverse前向套反向。更妙的是由于不需要直接调用jnp.dot这个hvp函数适用于任意形状的数组适用于任意容器类型如嵌套的 list / dict / tuple 存储的向量甚至不依赖jax.numpy模块。4.3 用jax.jacfwd(jax.jacrev)验证正确性下面用官方指南的示例验证用jax.hessian的朴素实现物化完整 Hessian 张量与hvp对比。def f(X): return jnp.sum(jnp.tanh(X)**2) key, subkey1, subkey2 jax.random.split(key, 3) X jax.random.normal(subkey1, (30, 40)) V jax.random.normal(subkey2, (30, 40)) def hessian(f): return jax.jacfwd(jax.jacrev(f)) ans1 hvp(f, (X,), (V,)) ans2 jnp.tensordot(hessian(f)(X), V, 2) print(jnp.allclose(ans1, ans2, 1e-4, 1e-4))如果输出True说明 forward-over-reverse 的hvp与显式物化 Hessian 再与 $V$ 做二阶张量缩并的结果一致。这里的hessian jacfwd(jacrev(f))也正是 JAX 内置jax.hessian的实现方式——见 jax/_src/api.py其源码就是return jacfwd(jacrev(fun, ...), ...)即前向套反向forward-over-reverse。4.4 反向套前向与反向套反向除了 forward-over-reverse还有另外两种组合方式# Reverse-over-forward def hvp_revfwd(f, primals, tangents): g lambda primals: jax.jvp(f, primals, tangents)[1] return jax.grad(g)(primals)# Reverse-over-reverse仅适用于单一参数 def hvp_revrev(f, primals, tangents): x, primals v, tangents return jax.grad(lambda x: jnp.vdot(jax.grad(f)(x), v))(x)官方指南给出的结论是reverse-over-forward 不如 forward-over-reverse。原因在于前向模式的开销小于反向模式而此处外层微分算子要微分的计算比内层更大因此把开销更小的前向模式放在外层效果最好。4.5 三种 HVP 的性能对比下面用 IPython 的%timeit魔法对三种 HVP 及朴素完整 Hessian 物化做基准对比%timeit为 IPython/Jupyter 内置魔法脚本运行请改用timeit模块print(Forward over reverse) %timeit -n10 -r3 hvp(f, (X,), (V,)) print(Reverse over forward) %timeit -n10 -r3 hvp_revfwd(f, (X,), (V,)) print(Reverse over reverse) %timeit -n10 -r3 hvp_revrev(f, (X,), (V,)) print(Naive full Hessian materialization) %timeit -n10 -r3 jnp.tensordot(jax.hessian(f)(X), V, 2)一般规律是hvpforward-over-reverse最快hvp_revfwd次之hvp_revrev再次而显式物化完整 Hessian 的做法最慢——因为它把 $30 \times 40 \times 30 \times 40$ 的完整 Hessian 张量都算了出来而 HVP 根本不需要物化它。这一对比也解释了为何在优化算法如共轭梯度、L-BFGS 的隐式求解中HVP 是标准的高效原语。五、组合 VJP、JVP 与jax.vmapjax.jvp和jax.vjp每次只前推/拉回单个向量。要同时前推/拉回一整批向量可以借助 JAX 的jax.vmap变换自动向量化参见 docs/automatic-vectorization.md用它写出快速的矩阵-Jacobian 与 Jacobian-矩阵乘积。5.1 矩阵-Jacobian 乘积Matrix-Jacobian Product, MJP需求把矩阵 $M$ 的每一行 $m_i$ 作为余切向量沿 $f$ 在 $W$ 处拉回。先看朴素的 Python 循环版本# 隔离出从权重矩阵到预测值的函数 f lambda W: predict(W, b, inputs) # 沿 f 在 W 处拉回余向量 m_i对 M 的所有行 i # 先用列表推导式在矩阵 M 的行上循环 def loop_mjp(f, x, M): y, vjp_fun jax.vjp(f, x) return jnp.vstack([jnp.asarray(vjp_fun(mi)) for mi in M]) # 再用 vmap 构造单次快速的矩阵-矩阵乘法 # 而不是外层循环若干次向量-矩阵乘法 def vmap_mjp(f, x, M): y, vjp_fun jax.vjp(f, x) outs, jax.vmap(vjp_fun)(M) return outs key jax.random.key(0) num_covecs 128 U jax.random.normal(key, (num_covecs,) y.shape) loop_vs loop_mjp(f, W, MU) print(Non-vmapped Matrix-Jacobian product) %timeit -n10 -r3 loop_mjp(f, W, MU) print(\nVmapped Matrix-Jacobian product) vmap_vs vmap_mjp(f, W, MU) %timeit -n10 -r3 vmap_mjp(f, W, MU) assert jnp.allclose(loop_vs, vmap_vs), Vmap and non-vmapped Matrix-Jacobian Products should be identical两者的结果必须一致代码末尾的assert会校验但vmap_mjp内部把 128 次独立的向量-矩阵乘法合并成一次矩阵-矩阵乘法通常能获得数量级的加速。5.2 Jacobian-矩阵乘积Jacobian-Matrix Product, JMP对称地可以把矩阵 $M$ 的每一行作为切向量前推def loop_jmp(f, W, M): # jvp 会立即以元组形式返回 primal 与 tangent 值 # 因此在列表推导式中计算并选取 tangent 部分 return jnp.vstack([jax.jvp(f, (W,), (mi,))[1] for mi in M]) def vmap_jmp(f, W, M): _jvp lambda s: jax.jvp(f, (W,), (s,))[1] return jax.vmap(_jvp)(M) num_vecs 128 S jax.random.normal(key, (num_vecs,) W.shape) loop_vs loop_jmp(f, W, MS) print(Non-vmapped Jacobian-Matrix product) %timeit -n10 -r3 loop_jmp(f, W, MS) vmap_vs vmap_jmp(f, W, MS) print(\nVmapped Jacobian-Matrix product) %timeit -n10 -r3 vmap_jmp(f, W, MS) assert jnp.allclose(loop_vs, vmap_vs), Vmap and non-vmapped Jacobian-Matrix products should be identical注意jax.jvp的返回值是(primals_out, tangents_out)对所以循环版本里要用[1]取 tangent 部分vmap_jmp则把对每行 $s$ 做 JVP 再取 tangent封装成_jvp交给jax.vmap批量执行。六、jax.jacfwd与jax.jacrev的实现原理有了快速 Jacobian-矩阵与矩阵-Jacobian 乘积jax.jacfwd与jax.jacrev的实现思路就呼之欲出了用同样的技巧一次性前推或拉回整个标准基与单位矩阵同构。6.1 用vjpvmap实现jacrevfrom jax import jacrev as builtin_jacrev def our_jacrev(f): def jacfun(x): y, vjp_fun jax.vjp(f, x) # 用 vmap 做矩阵-Jacobian 乘积。 # 这里的矩阵是欧氏基因此一次得到 Jacobian 的全部元素。 J, jax.vmap(vjp_fun, in_axes0)(jnp.eye(len(y))) return J return jacfun assert jnp.allclose(builtin_jacrev(f)(W), our_jacrev(f)(W)), Incorrect reverse-mode Jacobian results!思路是vjp先拉回单个余切向量得到 Jacobian 的一行把单位矩阵的每一行作为余切批量喂给vjp_fun就同时得到所有行即完整 Jacobian。这正对应 jax/_src/api.py 中jacrev的真实实现模式y, pullback, ... vjp(f_partial, *dyn_args)之后再jac vmap(pullback)(_std_basis(y))。6.2 用jvpvmap实现jacfwdfrom jax import jacfwd as builtin_jacfwd def our_jacfwd(f): def jacfun(x): _jvp lambda s: jax.jvp(f, (x,), (s,))[1] Jt jax.vmap(_jvp, in_axes1)(jnp.eye(len(x))) return jnp.transpose(Jt) return jacfun assert jnp.allclose(builtin_jacfwd(f)(W), our_jacfwd(f)(W)), Incorrect forward-mode Jacobian results!jax.jvp一次前推单个切向量得到 Jacobian 的一列。为了高效地把单位矩阵的所有列一起前推这里用in_axes1把jnp.eye(len(x))的列映射为vmap的批量维度由于vmap输出的每一行对应一列 Jacobian即转置后的结果最后用jnp.transpose还原。JAX 内置jacfwd的实现与之同理jax/_src/api.py它通过vmap(pushfwd, out_axes(None, -1))(_std_basis(dyn_args))一次推过整个标准基。补充源码中构造标准基的工具是_std_basisjax/_src/api.py它把 pytree 展平后调用jnp.eye(ndim)生成单位矩阵作为基。另外jax.jacobian只是jax.jacrev的别名jax/_src/api.py需要前向模式请显式使用jax.jacfwd。6.3 为什么 Autograd 做不到这些有趣的是Autograd 库做不到上述实现。Autograd 的反向模式jacobian只能通过外层循环的map逐次拉回单个向量一次只把一个向量推过整个计算远比用jax.vmap把整个批次合并起来计算低效。这正是jax.vmap与 JAX 的jaxpr追踪机制带来的优势。6.4 微分计算的线性部分可以被 JITAutograd 做不到的另一件事是jax.jit。有趣的是无论被微分的函数里使用了多少 Python 动态行为我们总能对计算的线性部分使用jax.jit。例如def f(x): try: if x 3: return 2 * x ** 3 else: raise ValueError except ValueError: return jnp.pi * x y, f_vjp jax.vjp(f, 4.) print(jax.jit(f_vjp)(1.))f内部有try/except和条件分支追踪阶段就已经把 Python 控制流消化掉了f_vjp是纯线性函数因此可以安全地交给jax.jit编译执行。对vjp函数做 JIT 的机制在 JAX 的测试套件如 tests/jax_jit_test.py、tests/lax_autodiff_test.py中都有覆盖可用作进一步验证。七、小结与延伸阅读两种模式的核心取舍可以浓缩为一张对照表维度前向模式JVPjax.jvp反向模式VJPjax.vjp基本运算$(x, v) \mapsto (f(x), \partial f(x) v)$$(x, v) \mapsto (f(x), v^\mathsf{T}\partial f(x))$构建完整 Jacobian逐列one-hot 切向量逐行one-hot 余切向量FLOPs约 3 倍于函数求值约 3 倍于函数求值内存与计算深度无关随计算深度增长典型场景输出维度远小于输入维度高瘦 Jacobian、HVP标量损失梯度宽扁 Jacobian、jax.grad实现复杂度相对简单边算边推相对复杂需保存中间值做反向传播jax.jacfwd/jax.jacrev/jax.hessian都是上述原语的组合hessian即jacfwd(jacrev(f))前向套反向自定义 Hessian-vector product 的最佳实践是jax.jvp(jax.grad(f), primals, tangents)[1]需要批量前推/拉回时用jax.vmap与jvp/vjp组合成矩阵-Jacobian / Jacobian-矩阵乘积。如果想继续深入仓库内还有大量相关资源更进阶的自动微分讨论见 docs/advanced_autodiff.md完整的 autodiff 实践手册见 docs/autodiff_cookbook.md从零实现 JAX 风格自动微分的教程见 docs/autodidax.md底层 AD 解释器源码在 jax/_src/interpreters/ad.pyjax.grad的 docstring 也明确指出它是jax.vjp的特例jax/_src/api.py。【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考