
JAX 601深入 JAX 内部工作原理——jaxpr 语言、Primitive 机制与从零实现【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax本系列文档docs/601/index.rst及其子文档面向对 JAX 内部实现好奇的读者、贡献者以及希望在最低层级扩展 JAX 的开发者系统讲解 JAX 的核心内部机制tracing 产生的中间表示 jaxpr 语言、作为基本计算单元的 primitive 机制以及如何用纯 Python 从零构建 JAX 核心。读完本系列你将能够阅读并理解 jaxpr 打印输出、掌握定义新 primitive 并为其注册求值/编译/自动微分/批处理规则的全流程从而具备在 JAX 底层进行扩展的能力。正如索引页所强调的使用 JAX 并不需要这些知识但理解这些是深入 JAX 内部的钥匙。系列定位谁需要了解 JAX 的内部JAX 601 系列docs/601/index.rst是一个面向 JAX 内部工作原理 的教程门户包含四篇环环相扣的文档jaxpr 语言——tracing 产生的中间表示IR其文法、以及如何阅读它jax-primitivesPrimitive 机制——primitive 操作如何工作JAX 对 primitive 的要求以及如何用jax.extend.core.Primitive定义新的 primitiveAutodidax从零构建 JAX 核心——用纯 Python 逐层构建 tracing、jaxpr、自动微分和 jitAutodidax2, part 1——反映 JAX 当前内部实现的从零重建。这些内容服务于三类人群好奇者想知道 JAX 为什么这样设计、贡献者需要修改 JAX 源码、以及最底层扩展者为 JAX 添加自定义操作。索引页明确说明 nothing here is needed touseJAX即正常使用 JAX 的开发者无需阅读本系列但它提供了理解 JAX 全部变换能力根基的完整路径。核心思想变换即解释要理解 JAX 内部首先需要抓住一条主线JAX 的变换jax.jit、jax.grad、jax.vmap本质上是以不同方式解释程序。正如 jaxpr 文档 所述JAX 变换在概念上分两步走先把待变换的 Python 函数通过 tracing 特化成一个小而规整的中间形式再用变换专属的解释规则去解释这个中间形式。jaxpr 文档 指出JAX 之所以能用很小的代码量承载如此强大的能力是因为它从一个熟悉而灵活的编程接口Python NumPy出发借助真实的 Python 解释器完成大部分繁重工作把计算精髓蒸馏成一种受限的、显式类型化的表达式语言。这个语言就是 jaxpr。Autodidax2 文档 将这一思想概括得更精炼JAX 是两样东西的集合——(1) 一组 primitive 操作大致对应 NumPy API(2) 一组建立在这些 primitive 之上的解释器编译、自动微分等。在它给出的最小实现里只用加法和乘法两个 primitive通过一个全局上下文变量记录当前解释器用户可见的add、mul函数会分派给当前解释器程序开始时当前解释器就是普通求值解释器。随着逐步叠加新的解释器一个微型 JAX 便诞生了。深入 jaxpr 语言jaxpr 是什么JaxprJAX program是 JAX 程序内部的中间表示IR具有四个关键性质显式类型化explicitly typed每个变量都带有类型标注函数式functional没有副作用结果仅由输入决定一阶first-order大多数 primitive 只接收一个或多个原子表达式作为参数代数范式 / ANFalgebraic normal form所有中间结果都被绑定为具名变量便于后续处理。jaxpr 文档 同时提醒读者并非所有变换都会字面物化出 jaxpr。例如微分和批处理会在 tracing 过程中增量地应用变换但若想理解 JAX 内部或想利用 JAX tracing 的结果比如导出计算图理解 jaxpr 就非常必要。jaxpr 的语法jaxpr 的术语term语法如下jaxpr :: { lambda binder , ... . let eqn ... in ( atom , ... ) } binder :: var:array_type var :: a | b | c | ... atom :: var | literal literal :: int32 | int64 | float32 | float64 eqn :: binder , ... primitive [ params ] atom , ...并非所有 Python 程序都能被这样处理但大量科学计算与机器学习程序都可以。Python 层面的控制流和函数调用会在 tracing 时正常执行并被内联展开因此 jaxpr 中不必然出现控制流或高阶特性。ClosedJaxpr你实际拿到的对象jaxpr 在代码中有两种关联表示jax.core.Jaxpr与jax.core.ClosedJaxpr。用jax.make_jaxpr检查 jaxpr 时得到的是一个ClosedJaxpr——它表示一个部分应用的Jaxpr包含两个字段jaxpr一个jax.core.Jaxpr承载函数实际的计算内容consts一个常量列表。Jaxpr自身的打印文法为jaxpr :: { lambda Var* ; Var. let Eqn* in [Expr] }其中lambda后分号分隔的两组变量第一组constvars是被提升出来的常量所对应的变量在ClosedJaxpr中其值存放在consts字段第二组invars对应被 traced 的 Python 函数的输入。Eqn*是方程列表每个方程定义若干中间变量作为某个 primitive 作用在若干原子表达式上的结果每个方程只使用输入变量与前面方程定义的中间变量。Expr是 jaxpr 的输出原子表达式字面量或变量列表。方程打印为Eqn :: let Var Primitive [ Param* ] Expr其中Var是被 primitive 调用定义的一个或多个中间变量有些 primitive 可返回多值Expr是一个或多个原子表达式变量或字面量常量特殊变量unitvar或字面量unit打印为*表示后续计算不再需要、已被省略的值即占位符Param*是零个或多个具名参数打印在方括号中形式为Name Value。绝大多数 jaxpr primitive 是一阶的Primitive : add | sub | sin | mul | ...最常见的 primitive 在jax.lax模块中有文档说明。用 make_jaxpr 阅读第一个 jaxprjaxpr 文档 给出如下示例源码中的make_jaxpr实现在 jax/_src/api.py它会通过jit(...).trace(...)完成 tracing 后返回ClosedJaxprfrom jax import make_jaxpr import jax.numpy as jnp def func1(first, second): temp first jnp.sin(second) * 3. return jnp.sum(temp) print(make_jaxpr(func1)(jnp.zeros(8), jnp.ones(8)))生成的 jaxpr 中没有 constvarsa和b是输入变量分别对应first与second两个函数参数标量字面量3.0被内联保留在方程中reduce_sumprimitive 除了操作数e外还带有具名参数axes和input_shape。Python 控制流与函数调用会被内联因为 Python 级控制流和函数在 tracing 期间照常执行jaxpr 不会包含它们。例如对func3进行 tracing 时对inner的调用以及if second.shape[0] 4条件都会被内联展开最终产生与func1相同的 jaxprdef func2(inner, first, second): temp first inner(second) * 3. return jnp.sum(temp) def inner(second): if second.shape[0] 4: return jnp.sin(second) else: assert False def func3(first, second): return func2(inner, first, second) print(make_jaxpr(func3)(jnp.zeros(8), jnp.ones(8)))pytrees 的展平jaxpr 中没有元组类型primitive 以多输入多输出的方式工作。当函数输入输出是结构化对象如元组时JAX 会将其展平jaxpr 中以输入/输出列表形式呈现。例如下面的func4产生与之前完全相同的 jaxpr两个输入变量分别对应元组的两个元素def func4(arg): # The arg is a pair. temp arg[0] jnp.sin(arg[1]) * 3. return jnp.sum(temp) print(make_jaxpr(func4)((jnp.zeros(8), jnp.ones(8))))更详细的说明可参考 pytrees 教程。常量变量constvarsjaxpr 中有些值是常量——其值不依赖 jaxpr 的参数。标量常量直接内联在方程里非标量数组常量则被提升hoist到 jaxpr 顶层对应为 constvars。constvars 与其他 jaxpr 参数invars的区别仅是簿记约定上的在ClosedJaxpr中consts字段持有它们的值。高阶 JAX primitivejaxpr 中的子程序普通 primitive 是一阶的但 jaxpr 还包含若干高阶 primitive它们内嵌子 jaxpr因此更复杂。这部分是理解jax.lax控制流算子内部表示的关键。cond primitive条件分支Python 条件在 tracing 时会被直接走通要捕获条件表达式以进行动态执行必须使用jax.lax.switch与jax.lax.cond构造器其签名如下lax.switch(index: int, branches: Sequence[A - B], operand: A) - B lax.cond(pred: bool, true_body: A - B, false_body: A - B, operand: A) - B两者内部都会绑定一个名为cond的 primitive。jaxpr 中的condprimitive 反映的是更一般的lax.switch签名它接收一个表示要执行哪个分支的整数会被截断到合法索引范围内。例如from jax import lax def one_of_three(index, arg): return lax.switch(index, [lambda x: x 1., lambda x: x - 2., lambda x: x 3.], arg) print(make_jaxpr(one_of_three)(1, 5.))condprimitive 有两个关键参数branches对应各分支函数体的 jaxpr。上例中每个函数体接收一个输入变量对应xlinear一个布尔元组供自动微分机制内部使用编码条件中哪些输入参数被线性使用。lax.cond的情形中布尔谓词会被转换为整数索引0 或 1branches按 false、true 顺序对应两个分支的 jaxpr。当分支函数体的输入是元组、且某个分支内包含被提升为 constvar 的常量如jnp.ones(1)时jaxpr 会体现更复杂的情况。while primitive循环与条件一样Python 循环在 tracing 时被内联。要捕获循环以动态执行必须使用jax.lax.while_loop本身是一个 primitive或jax.lax.fori_loop生成 while_loop primitive 的辅助函数lax.while_loop(cond_fun: (C - bool), body_fun: (C - C), init: C) - C lax.fori_loop(start: int, end: int, body: (int - C - C), init: C) - C其中C表示循环 carry携带值的类型。示例import numpy as np def func10(arg, n): ones jnp.ones(arg.shape) # A constant. return lax.fori_loop(0, n, lambda i, carry: carry ones * 3. arg, arg ones) print(make_jaxpr(func10)(np.ones(16), 5))生成的whileprimitive 接收 5 个参数c a 0 b d——其中0是cond_jaxpr的常量个数cond_nconsts为 0c、a是body_jaxpr的 2 个常量b d是 carry 初值的 3 个参数。scan primitive定长循环JAX 支持一种对数组元素形状静态已知进行循环的特殊形式。因为迭代次数固定这种循环很容易做反向微分。用jax.lax.scan构造lax.scan(body_fun: (C - A - (C, B)), init_carry: C, in_arr: Array[A]) - (C, Array[B])这里C是 scan 的 carry 类型A是输入数组元素类型B是输出数组元素类型。示例def func11(arr, extra): ones jnp.ones(arr.shape) # A constant def body(carry, aelems): # carry: running dot-product of the two arrays # aelems: a pair with corresponding elements from the two arrays ae1, ae2 aelems return (carry ae1 * ae2 extra, carry) return lax.scan(body, 0., (arr, ones)) print(make_jaxpr(func11)(np.ones(16), 5.))linear参数描述每个输入变量是否保证在 body 中被线性使用scan 经过线性化后会有更多参数变为线性。scanprimitive 接收 4 个参数b 0.0 a c——一个是 body 的自由变量一个是 carry 的初值另外两个是 scan 操作的数组。(p)jit primitive调用封装调用 primitive 源于 JIT 编译它封装一个子 jaxpr连同指定后端backend与运行设备的参数。示例from jax import jit def func12(arg): jit def inner(x): return x arg * jnp.ones(1) # Include a constant in the inner function. return arg inner(arg - 2.) print(make_jaxpr(func12)(1.))可以看到被jit装饰的内部函数以子 jaxpr 形式嵌套在调用 primitive 中体现了 JAX 变换的可组合性。深入 Primitive 机制什么是 primitive什么是 JAX-traceableJAX primitives 文档 给出的定义是JAX primitive 是 JAX 程序的基本计算单元。例如 multiply-add 既可以用底层jax.lax.*primitive 实现它们类似 XLA 算子包装器也可以用jax.extend.core.Primitive(multiply_add)定义。JAX 之所以能对 Python 函数施加jax.jit、jax.grad、jax.vmap等可组合变换是因为变换以JAX-traceable的方式实现当 Python 函数被执行时它作用于数据的操作只可能是两类——对数据属性的检视如形状shape或类型dtypeJAX primitive 调用即本教程介绍的 JAX 特殊操作。关键点在于JAX primitive 既能处理具体数据值也能处理抽象 JAX 值。例如抽象值ShapedArray(float32[2,2])只捕获值的类型与形状不含具体数据。JAX 可以携带抽象参数来调用一个 JAX-traceable 函数。而被变换后的函数本身必须仍是 JAX-traceable 函数以保证变换可组合例如jax.jit(jax.jacfwd(jax.grad(f)))。JAX 预定义了对应大多数 XLA 操作的 primitiveadd、matmul、sin、cos、索引等并且用 JAX primitive 实现了 NumPy 函数——因此使用 JAX 版 NumPy 编写的 Python 程序天然 JAX-traceable、天然可变换。其他库也可以通过基于 JAX primitive 实现来获得 traceable 能力。更重要的是JAX primitive 的集合是可扩展的你可以定义一个新 primitive 来封装某个函数的行为而不必用既有 primitive 重新实现它。方式一使用现有的 JAX primitive定义新函数最简单的途径是用 JAX primitive 或那些本身基于 primitive 写成的函数如jax.lax模块中的函数来组合from jax._src.lax import lax from jax._src import api def multiply_add_lax(x, y, z): Implementation of multiply-add using the jax.lax primitives. return lax.add(lax.mul(x, y), z) def square_add_lax(a, b): A square-add function using the newly defined multiply-add. return multiply_add_lax(a, a, b) print(square_add_lax , square_add_lax(2., 10.)) # Differentiate w.r.t. the first argument print(grad(square_add_lax) , api.grad(square_add_lax, argnums0)(2.0, 10.))除了直接使用jax.laxprimitive也可以使用已经基于它们写好的函数例如jax.numpyjnp.add(jnp.multiply(x, y), z)。在计算jax.grad的过程中JAX 会用特殊参数ConcreteArray(...)调用这些函数——这说明JAX-traceable 函数必须不仅能处理具体参数还要能处理 JAX 用来抽象函数执行的抽象参数。只要函数基于 JAX primitive 编写traceable 性质就能得到满足。方式二定义新的 JAX primitive为了演示 primitive 的工作机制可以假装要向 JAX 添加一个 multiply-add 的新 primitive尽管正确做法通常是复用现有 primitivefrom jax.extend import core multiply_add_p core.Primitive(multiply_add) # Create the primitive def multiply_add_prim(x, y, z): The JAX-traceable way to use the JAX primitive. return multiply_add_p.bind(x, y, z) def square_add_prim(a, b): A square-add function implemented using the new JAX-primitive. return multiply_add_prim(a, a, b)注意被 trace 的参数必须以位置参数形式传给bind。在源码层面Primitive.bind的实现位于 jax/_src/core.py它会先对每个参数做类型规范化dtypes.canonicalize_value、提取抽象值typeof、校验 tracer 有效性然后调用bind_with_trace最终交给当前 trace 的process_primitive处理。Primitive类jax/_src/core.py还带有multiple_results、call_primitive、ref_primitive、is_effectful等标志位分别表示多输出 primitive、以 final style 处理的调用 primitive、引用类 primitive 与效果属性。刚定义好的 primitive 还不能被调用——因为还没有告诉 JAX 它的任何语义直接调用会得到NotImplementedError。接下来需要逐步注册各条规则。规则一Primal 求值规则def_implprimal 求值规则是 primitive 的具体实现不需要是 JAX-traceable 的只会被具体值调用内部可以使用普通非 JAXNumPyimport numpy as np def multiply_add_impl(x, y, z): Concrete implementation of the primitive. return np.add(np.multiply(x, y), z) # Now, register the primal implementation with JAX: multiply_add_p.def_impl(multiply_add_impl)注册后square_add_prim(2., 10.)就能得到14.。源码中def_impljax/_src/core.py只是把实现赋给self.impl未注册时默认的impl方法会抛出NotImplementedError见 jax/_src/core.py。规则二抽象求值规则def_abstract_eval——JIT 的关键尝试对square_add_prim使用jax.jit会再次遇到NotImplementedError。要 JIT以及支持其他变换JAX 必须先用参数的形状和类型对函数做抽象求值其目的有二得到计算中使用的 JAX primitive 序列——这个序列将被编译计算出计算中所有向量与操作的形状和类型。例如一个 3 元素向量的抽象可以是ShapedArray(float32[3])也可以是ConcreteArray([1., 2., 3.])——后者是 JAX 把实际具体值包装成抽象值。ShapedArray在源码中定义于 jax/_src/core.py包含shape、dtype、weak_type、sharding、memory_space等槽位。抽象求值规则如下from jax import core def multiply_add_abstract_eval(xs, ys, zs): Abstract evaluation of the primitive. assert xs.shape ys.shape assert xs.shape zs.shape return core.ShapedArray(xs.shape, xs.dtype) # Now, register the abstract evaluation with JAX: multiply_add_p.def_abstract_eval(multiply_add_abstract_eval)该函数同样不必是 JAX-traceable 的它接收参数的抽象表示并返回结果的ShapedArray。注册后再次尝试jit会看到抽象求值已能推进但会因缺少 XLA 编译规则而报错。源码中def_abstract_evaljax/_src/core.py会把抽象求值包装成无效果版本_effect_free_abstract_eval如果 primitive 有副作用则需使用def_effectful_abstract_eval等变体。规则三XLA 编译规则loweringJAX 编译的本质是把每个 primitive 编译成一张 XLA 操作图。这是给 JAX 添加新功能的最大门槛——因为 XLA 操作集合有限且 JAX 已为大多数操作预定义了 primitive。不过 XLA 提供了CustomCall操作可用它封装任意用 C 实现的功能。在现代 JAX 中lowering 规则基于 MLIR 编写对应源码中的 jax/_src/interpreters/mlir.py 的register_lowering(prim, rule, platform...)from jax._src.lib.mlir.dialects import hlo def multiply_add_lowering(ctx, xc, yc, zc): The compilation to XLA of the primitive. return [hlo.AddOp(hlo.MulOp(xc, yc), zc).result] # Now, register the lowering rule with JAX. from jax.interpreters import mlir mlir.register_lowering(multiply_add_p, multiply_add_lowering, platformcpu)lowering 规则接收每个参数的mlir.ir.Value返回结果的mlir.ir.Value同样无需是 JAX-traceable 函数。注册之后jax.jit即可成功JAX 先抽象求值触发multiply_add_abstract_eval再编译遇到的 primitive 集合触发multiply_add_lowering。还有一个有趣的细节用jit只对第一个参数编译static_argnums1时square_add_prim的第二个参数是具体的导致multiply_add_abstract_eval收到的第三个参数是ConcreteArray——可见抽象求值规则可以同时接受ShapedArray与ConcreteArray。规则四前向微分JVPJAX 以 Jacobian-Vector ProductJVP形式实现前向微分概念细节可参考 自定义 JVP/VJP 指南。未注册微分规则前jax.jvp会报错。JVP 规则的形式是给定各参数的值与切向量tangent计算 primal 输出与输出切向量。该规则必须 JAX-traceable因为 JAX 可能以抽象值调用它from jax.interpreters import ad def multiply_add_value_and_jvp(arg_values, arg_tangents): Evaluates the primal output and the tangents (Jacobian-vector product). x, y, z arg_values xt, yt, zt arg_tangents # Now, you have a JAX-traceable computation of the output. primal_out multiply_add_prim(x, y, z) # You must use a JAX-traceable way to compute the tangent. # The output tangent is (xt * y x * yt zt), implemented with # the same multiply_add_prim primitive. def make_zero(tan): return lax.full_like(x, 0) if type(tan) is ad.Zero else tan output_tangent multiply_add_prim(make_zero(xt), y, multiply_add_prim(x, make_zero(yt), make_zero(zt))) return (primal_out, output_tangent) # Register the forward differentiation rule with JAX: ad.primitive_jvps[multiply_add_p] multiply_add_value_and_jvp注意arg_tangents中某些切向量可能是特殊值ad.Zero表示零切向量需要特殊处理如make_zero将其转成同形状的 0 张量或者做代数化简。注册后# Tangent is: xt*y x*yt zt 1.*2. 2.*1. 1. 5. assert api.jvp(square_add_prim, (2., 10.), (1., 1.)) (14., 5.)对 JVP 再套jit也是可行的JAX 会先抽象求值multiply_add_value_and_jvp它会抽象求值 primal 与 tangent 两条计算共 3 次调用 multiply_add primitive然后编译这 3 处 primitive。源码中 JVP 规则的注册模式形如ad.primitive_jvps[primitive] rule可在 jax/_src/ad_checkpoint.py 等处看到内置 primitive 的同类用法。规则五反向微分transposition使用jax.grad反向微分时JAX 会先用multiply_add_value_and_jvp对抽象值做前向微分得到一段计算输出切向量的 primitive 轨迹然后JAX 会把这轨迹抽象地反向解释对每个 primitive 应用一条转置规则transposition rule。此时会因缺少转置规则而报NotImplementedError。转置的含义可通过简单例子理解。对f(x, y) x * y y在(2., 4.)处微分JVP 切向计算为a xt * 4. b 2. * yt c a b ft c yt按构造切向计算对输入切向量总是线性的切向计算中可能出现的唯一非线性算子是乘法且其中一个操作数必为常量。JAX 通过逆序处理 JVP 计算来产生反向微分计算——对切向计算中的每个操作用其结果余切cotangent累加该操作所用变量的余切# Initialize cotangents of inputs and intermediate variables: xct yct act bct cct 0. # Initialize cotangent of the output: fct 1. # Process ft c yt: cct fct yct fct # Process c a b: act cct bct cct # Process b 2. * yt: yct 2. * bct # Process a xt * 4.: xct act * 4.可验证该计算得到xct 4.、yct 3.正是f的两个偏导数。概念上若 primitivep(x, y, z)对参数y、z线性x视为常量即p(x, y, z) y*cy z*cz则其转置为p_transpose(out_ct, x, _, _) (None, out_ct*cy, out_ct*cz)p_transpose接收 primitive 输出的余切以及每个参数的对应值线性参数得到未定义值_其余参数得到实际常量返回每个参数的余切常量参数对应位置返回None。典型例子add_transpose(out_ct, _, _) (out_ct, out_ct) mult_transpose(out_ct, x, _) (None, x * out_ct) mult_transpose(out_ct, _, y) (out_ct * y, None)对于本教程的 multiply_add它本身不是线性 primitive但在multiply_add_value_and_jvp中相对于切向量是线性使用的output_tangent(xt, yt, zt) multiply_add_prim(xt, y, multiply_add_prim(x, yt, zt))两个乘法参数中总有一个是常量。转置规则如下from jax.interpreters import ad def multiply_add_transpose(ct, x, y, z): Evaluates the transpose of a linear primitive. if not ad.is_undefined_primal(x): # This use of multiply_add is with a constant x. assert ad.is_undefined_primal(y) ct_y ad.Zero(y.aval) if type(ct) is ad.Zero else multiply_add_prim(x, ct, lax.full_like(x, 0)) res None, ct_y, ct else: # This use of multiply_add is with a constant y. assert ad.is_undefined_primal(x) ct_x ad.Zero(x.aval) if type(ct) is ad.Zero else multiply_add_prim(ct, y, lax.full_like(y, 0)) res ct_x, None, ct return res ad.primitive_transposes[multiply_add_p] multiply_add_transpose这里线性参数收到ad.UndefinedPrimal值常量参数收到实际常量值。注册转置后api.grad(square_add_prim)(2., 10.) 4.即可通过。注意grad运行中multiply_add_transpose被调用两次对应multiply_add_value_and_jvp中output_tangent计算对multiply_add_prim的两次使用先转置最后一次multiply_add_prim(xt, y, ...)其中y是常量2.0。对grad再套jit同样可行且此时multiply_add_value_and_jvp的抽象求值只用抽象值而非无 jit 时的ConcreteArray。规则六批处理batchingjax.vmap变换把一个逐点计算变成向量上的计算。未注册规则时vmap报NotImplementedError。对于 multiply_add 这类本身逐点操作任意维度张量的 primitive批处理版本可以复用其自身实现要求输入同维、且沿相同轴批处理from jax.interpreters import batching def multiply_add_batch(vector_arg_values, batch_axes): Computes the batched version of the primitive. assert batch_axes[0] batch_axes[1] assert batch_axes[0] batch_axes[2] res multiply_add_prim(*vector_arg_values) return res, batch_axes[0] batching.primitive_batchers[multiply_add_p] multiply_add_batch批处理规则必须是 JAX-traceable 函数返回(结果, 被批处理的结果轴)。注册后assert np.allclose(api.vmap(square_add_prim, in_axes0, out_axes0)( np.array([2., 3.]), np.array([10., 20.])), [14., 29.])对vmap套jitapi.jit(api.vmap(...))同样能正确工作。源码中batching.primitive_batchers[prim] batcher是批处理规则的注册方式内置 primitive 的同类用法可参考 jax/_src/ad_checkpoint.pyfancy batcher与 jax/_src/ad_checkpoint.py普通 batcher。从零构建 JAXAutodidax 系列Autodidax逐层重建核心Autodidax 文档 的目标是让读者通过动手实现学到 JAX 核心系统的每一个大思想。它以 变换即解释器 开篇把sin及中缀运算符背后的mul、add、neg视为 primitive 操作原子处理单元而非组合然后通过拦截 primitive 的应用、让不同的值流过程序来实现不同解释。例如把每个 primitive 的应用替换为它的 JVP 规则让 primal-tangent 对流过程序多个变换还可以组合成解释器栈。文中用NamedTuple定义了Primitive含name字段与add_p、mul_p、neg_p、sin_p等 primitive以及bind1(prim, *args, **params)绑定函数——这正是真实 JAX 中Primitive.bind的最小原型。该文档声明为进行中的草稿部分第 5、6 部分内容尚缺但对理解 JAX 核心的 tracing、jaxpr、autodiff、jit 四件套极有价值。Autodidax2, part 1反映当前内部实现的再构建Autodidax2 文档 是反映 JAX 当前内部实现的从零重建理念是去掉杂乱代码的精简版 JAX。它的主线是上下文敏感解释context-sensitive interpretationJAX 是 (1) 一组 primitive 操作大致是 NumPy API与 (2) 一组基于这些 primitive 的解释器编译、自动微分等的集合。在最小实现中只从加法和乘法两个 primitive 起步逐个添加解释器为每种解释定义一个带各 primitive 处理规则的Interpreter对象用全局上下文变量记录当前解释器用户可见的add、mul函数分派给当前解释器程序初始时当前解释器为普通求值解释器。由此同一个用户函数如foo(x) mul(x, add(x, 3.0))无需修改实现就能被求值、微分、转成 IR、编译——这正是 JAX 设计的精髓。从文档到源码一条完整的扩展路径把本系列与仓库源码对照可以勾勒出为 JAX 添加自定义操作如 multiply-add的完整路径定义core.Primitive(multiply_add)创建 primitive用户函数通过bind调用它jax/_src/core.py 的Primitive类定义了bind、def_impl、def_abstract_eval、def_effectful_abstract_eval等接口求值def_impl注册具体实现形状推断def_abstract_eval注册抽象求值返回ShapedArrayjax/_src/core.py编译mlir.register_lowering(prim, rule, platform...)注册到 MLIR/XLAjax/_src/interpreters/mlir.py微分ad.primitive_jvps[prim]注册 JVP、ad.primitive_transposes[prim]注册转置批处理batching.primitive_batchers[prim]注册 vmap 规则。每一步缺失都会以NotImplementedError显式报出这正是 JAX 缺什么补什么 的设计哲学。如果你希望以更直观的方式验证这些概念可以对照 Autodidax 与 Autodidax2 的纯 Python 实现逐行演练相关 notebook 版本位于 docs/autodidax.ipynb 与 docs/autodidax2_part1.ipynb。从Python NumPy 程序的可组合变换到tiny 中间语言 可插拔解释器JAX 的内部设计在 601 系列中一览无余。【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考