JAX 601:深入 JAX 内部工作原理——jaxpr 语言、Primitive 机制与从零实现
【免费下载链接】jaxComposable transformations of Python+NumPy 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-primitives:Primitive 机制——primitive 操作如何工作:JAX 对 primitive 的要求,以及如何用
jax.extend.core.Primitive定义新的 primitive; - Autodidax:从零构建 JAX 核心——用纯 Python 逐层构建 tracing、jaxpr、自动微分和 jit;
- Autodidax2, 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 是什么
Jaxpr(JAX + program)是 JAX 程序内部的中间表示(IR),具有四个关键性质:
- 显式类型化(explicitly typed):每个变量都带有类型标注;
- 函数式(functional):没有副作用,结果仅由输入决定;
- 一阶(first-order):大多数 primitive 只接收一个或多个原子表达式作为参数;
- 代数范式 / ANF(algebraic 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 阅读第一个 jaxpr
jaxpr 文档 给出如下示例(源码中的make_jaxpr实现在 jax/_src/api.py,它会通过jit(...).trace(...)完成 tracing 后返回ClosedJaxpr):
from 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 中没有 constvars;a和b是输入变量,分别对应first与second两个函数参数;标量字面量3.0被内联保留在方程中;reduce_sumprimitive 除了操作数e外,还带有具名参数axes和input_shape。
Python 控制流与函数调用会被内联
因为 Python 级控制流和函数在 tracing 期间照常执行,jaxpr 不会包含它们。例如对func3进行 tracing 时,对inner的调用以及if second.shape[0] > 4条件都会被内联展开,最终产生与func1相同的 jaxpr:
def 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 教程。
常量变量(constvars)
jaxpr 中有些值是常量——其值不依赖 jaxpr 的参数。标量常量直接内联在方程里;非标量数组常量则被提升(hoist)到 jaxpr 顶层,对应为 constvars。constvars 与其他 jaxpr 参数(invars)的区别仅是簿记约定上的,在ClosedJaxpr中consts字段持有它们的值。
高阶 JAX primitive:jaxpr 中的子程序
普通 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。上例中每个函数体接收一个输入变量,对应x;linear:一个布尔元组,供自动微分机制内部使用,编码条件中哪些输入参数被线性使用。
lax.cond的情形中,布尔谓词会被转换为整数索引(0 或 1),branches按 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为 0),c、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-traceable
JAX 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)或类型(dtype);
- JAX primitive 调用:即本教程介绍的 JAX 特殊操作。
关键点在于:JAX primitive 既能处理具体数据值,也能处理抽象 JAX 值。例如抽象值ShapedArray(float32[2,2])只捕获值的类型与形状,不含具体数据。JAX 可以携带抽象参数来调用一个 JAX-traceable 函数。而被变换后的函数本身必须仍是 JAX-traceable 函数,以保证变换可组合,例如jax.jit(jax.jacfwd(jax.grad(f)))。
JAX 预定义了对应大多数 XLA 操作的 primitive(add、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, argnums=0)(2.0, 10.))除了直接使用jax.laxprimitive,也可以使用已经基于它们写好的函数,例如jax.numpy(jnp.add(jnp.multiply(x, y), z))。在计算jax.grad的过程中,JAX 会用特殊参数ConcreteArray(...)调用这些函数——这说明JAX-traceable 函数必须不仅能处理具体参数,还要能处理 JAX 用来抽象函数执行的抽象参数。只要函数基于 JAX primitive 编写,traceable 性质就能得到满足。
方式二:定义新的 JAX primitive
为了演示 primitive 的工作机制,可以假装要向 JAX 添加一个 multiply-add 的新 primitive(尽管正确做法通常是复用现有 primitive):
from 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_impl)
primal 求值规则是 primitive 的具体实现,不需要是 JAX-traceable 的,只会被具体值调用,内部可以使用普通(非 JAX)NumPy:
import 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_impl(jax/_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_eval(jax/_src/core.py)会把抽象求值包装成无效果版本(_effect_free_abstract_eval);如果 primitive 有副作用,则需使用def_effectful_abstract_eval等变体。
规则三:XLA 编译规则(lowering)
JAX 编译的本质是把每个 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, platform='cpu')lowering 规则接收每个参数的mlir.ir.Value,返回结果的mlir.ir.Value,同样无需是 JAX-traceable 函数。注册之后jax.jit即可成功:JAX 先抽象求值(触发multiply_add_abstract_eval),再编译遇到的 primitive 集合(触发multiply_add_lowering)。
还有一个有趣的细节:用jit只对第一个参数编译(static_argnums=1)时,square_add_prim的第二个参数是具体的,导致multiply_add_abstract_eval收到的第三个参数是ConcreteArray——可见抽象求值规则可以同时接受ShapedArray与ConcreteArray。
规则四:前向微分(JVP)
JAX 以 Jacobian-Vector Product(JVP)形式实现前向微分(概念细节可参考 自定义 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)。
规则六:批处理(batching)
jax.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_axes=0, out_axes=0)( np.array([2., 3.]), np.array([10., 20.])), [14., 29.])对vmap套jit(api.jit(api.vmap(...)))同样能正确工作。源码中batching.primitive_batchers[prim] = batcher是批处理规则的注册方式,内置 primitive 的同类用法可参考 jax/_src/ad_checkpoint.py(fancy batcher)与 jax/_src/ad_checkpoint.py(普通 batcher)。
从零构建 JAX:Autodidax 系列
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 interpretation):JAX 是 (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注册抽象求值,返回ShapedArray(jax/_src/core.py); - 编译:
mlir.register_lowering(prim, rule, platform=...)注册到 MLIR/XLA(jax/_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 Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考