JAX Shape Polymorphism 完全指南:用符号维度实现一次导出、多形状复用
【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax
导读
本文围绕 JAX 官方文档 docs/501/shape-polymorphism.md 展开,系统讲解Shape Polymorphism(形状多态):这是 JAX export 提供的一项能力,允许把函数只追踪(trace)和降级(lower)一次,导出的Exported对象就能面向一整族输入形状编译执行,从而在没有 Python 源码的另一套系统上也能按需适配不同 batch 大小、序列长度等维度。读完本文,你将掌握jax.export.symbolic_shape/symbolic_args_specs的完整用法、符号维度表达式与约束的编写规则、InconclusiveDimensionOperation的成因与化解策略,以及用symbolic_dim_bounds、JAX_DUMP_IR_TO等工具调试导出代码的实战手段。
1. 为什么需要 Shape Polymorphism:从"每形状一编译"到"一次导出,多形状复用"
在 JIT 模式下,JAX 会对每一种输入类型与形状的组合分别执行追踪、降级到 StableHLO 并编译:
>>> import jax >>> from jax import export >>> from jax import numpy as jnp >>> def f(x): # f: f32[a, b] ... return jnp.concatenate([x, x], axis=1)问题在于:当函数被export导出、序列化并在另一台机器上反序列化之后,Python 源码已经不在现场,无法重新追踪(re-trace)和重新降级(re-lower)。此时若输入形状发生变化(例如 batch 维度从 8 变成 16),传统方案就会失效。
Shape polymorphism 正是为这一场景而生:在导出阶段用"符号形状"(symbolic shapes)进行追踪与降级,把维度写成一个或多个维度变量(dimension variables),而Exported对象中保存了足以在多种具体输入形状下编译执行的全部信息。函数在调用时仍会按需重新编译,但只有编译这一步发生在调用现场,追踪与降级只发生一次、随导出而固化。这一点在官方文档中明确强调,也是理解该特性性能与语义的关键。
2. 核心 API 一:jax.export.symbolic_shape与symbolic_args_specs
2.1 用symbolic_shape构造符号形状
jax.export.symbolic_shape(shape_spec, *, constraints=(), scope=None, like=None)接受一个字符串形式的形状规格,返回维度表达式对象的元组(类型为_DimExpr,可直接替代整数常量用于构造形状):
>>> # 我们构造符号维度变量。 >>> a, b = export.symbolic_shape("a, b") >>> # 符号维度可以直接用来构造形状。 >>> x_shape = (a, b) >>> x_shape (a, b) >>> # 然后用符号形状导出: >>> exp: export.Exported = export.export(jax.jit(f))( ... jax.ShapeDtypeStruct(x_shape, jnp.int32)) >>> exp.in_avals (ShapedArray(int32[a,b]),) >>> exp.out_avals (ShapedArray(int32[a,2*b]),) >>> # 之后可以用具体形状调用(这里 a=3, b=4),无需重新追踪 `f`。 >>> res = exp.call(np.ones((3, 4), dtype=np.int32)) >>> res.shape (3, 8)观察out_avals中的int32[a,2*b]:jnp.concatenate([x, x], axis=1)的输出形状被 JAX 自动计算为2*b这一符号维度表达式。
维度表达式对象重载了绝大多数整数运算符,因此在大多数场景下可以像使用整数常量一样参与算术、切片与形状计算(详见 docs/501/shape-polymorphism.md 与第 4 节)。
2.2 用symbolic_args_specs从真实参数构造规格 pytree
jax.export.symbolic_args_specs(args, shapes_specs, *, constraints=(), scope=None)的用途是:基于实际(具体形状)参数构造出与之一一对应的jax.ShapeDtypeStructpytree,其中被占位符覆盖的维度替换为符号维度,dtype 则沿用真实参数。看文档中的完整示例:
>>> def f1(x, y): # x: f32[a, 1], y : f32[a, 4] ... return x + y >>> # 假设你已有具体形状的真实参数 >>> x = np.ones((3, 1), dtype=np.int32) >>> y = np.ones((3, 4), dtype=np.int32) >>> args_specs = export.symbolic_args_specs((x, y), "a, ...") >>> exp = export.export(jax.jit(f1))(* args_specs) >>> exp.in_avals (ShapedArray(int32[a,1]), ShapedArray(int32[a,4]))这里的规格字符串"a, ..."中:
...占位符代表0 个或多个维度,其取值由真实参数的具体形状填充;_占位符代表恰好一个维度;- 规格可以是pytree 前缀,即一条规格可同时应用于多个参数(如上例中
x、y同时共享a这个符号维度)。
从源码看,其实现位于 jax/_src/export/shape_poly.py:先用tree_util.tree_flatten拍平参数,用tree_util.broadcast_prefix将规格广播到每个参数,再对每个维度规格调用symbolic_shape(spec, like=s, scope=scope)完成占位符填充,最终用args_tree.unflatten还原为与args结构一致的ShapeDtypeStructpytree。因此它天然支持任意嵌套的 pytree 参数结构。
2.3 形状规格的常见写法
官方文档给出了几种典型规格:
("(b, _, _)", None):适用于两个参数的函数。第一个参数是 3D 数组,b是符号化的 batch 引导维度,其余维度按真实参数特化;None表示第二个参数完全非符号化(等价于写...)。由于规格是 pytree 前缀,若第一个参数本身是"多个 3D 数组组成的 pytree",该规格同样适用——只要它们共享同一个引导维度b。("(batch, ...)", "(batch,)"):约束两个参数的引导维度相同,第一个参数秩至少为 1,第二个参数秩恰好为 1。
3. 正确性契约:何时可以相信导出的程序
Shape polymorphism 的正确性定义如下(文档原文要点):
对任意 JAX 函数
f与任意含符号形状的参数规格arg_spec,以及任意形状匹配arg_spec的具体参数arg:
- 若 JAX 原生执行成功:
res = f(arg);- 且符号形状导出成功:
exp = export.export(f)(arg_spec);- 则编译并运行导出结果必然成功且结果一致:
res == exp.call(arg)。
需要强调的是:f(arg)对每种不同的具体形状都会重新调用 JAX 追踪机制;而exp.call(arg)的执行不再依赖任何追踪能力——它甚至可能运行在根本没有f源码的环境中。
要保证这种正确性并不容易,在最棘手的场景下导出会直接失败。本文后续章节(第 5、6 节)即围绕这些失败的处理方法展开,这也是官方文档"Errors in presence of shape polymorphism"与调试部分的主题。
4. 用符号维度做计算:表达式语义与隐式转数组规则
JAX 会跟踪所有中间结果的形状。当形状依赖维度变量时,JAX 把它们计算为符号维度表达式(symbolic dimension expressions)。文档明确了两条基础语义:
- 维度变量表示大于等于 1 的整数值;
- 符号表达式支持在维度表达式与整数(
int、np.int,或任何可用operator.index转换的值)之间应用算术运算符(加、减、乘、整除floordiv、取模mod,以及 NumPy 变体np.sum、np.prod等),结果可继续用于jnp.reshape、jnp.arange、切片索引等形状参数。
典型的扁平化示例:x.shape[0] * x.shape[1]被计算为符号表达式4 * b:
>>> f = lambda x: jnp.reshape(x, (x.shape[0] * x.shape[1],)) >>> arg_spec = jax.ShapeDtypeStruct(export.symbolic_shape("b, 4"), jnp.int32) >>> exp = export.export(jax.jit(f))(arg_spec) >>> exp.out_avals (ShapedArray(int32[4*b]),)4.1 显式转成 JAX 数组:jnp.array(x.shape[0])
可以用jnp.array(x.shape[0])甚至jnp.array(x.shape)把维度表达式显式转为 JAX 数组。得到的数组可作为普通 JAX 数组参与运算,但不能再当作形状维度使用(例如用于reshape):
>>> exp = export.export(jax.jit(lambda x: jnp.array(x.shape[0]) + x))( ... jax.ShapeDtypeStruct(export.symbolic_shape("b"), np.int32)) >>> exp.call(jnp.arange(3, dtype=np.int32)) Array([3, 4, 5], dtype=int32) >>> exp = export.export(jax.jit(lambda x: x.reshape(jnp.array(x.shape[0]) + 2)))( ... jax.ShapeDtypeStruct(export.symbolic_shape("b"), np.int32)) # doctest: +IGNORE_EXCEPTION_DETAIL Traceback (most recent call last): TypeError: Shapes must be 1D sequences of concrete values of integer type, got [Traced<ShapedArray(int32[], weak_type=True)>with<DynamicJaxprTrace(level=1/0)>].4.2 与非整数运算时自动转数组
当符号维度与非整数(float、np.float、np.ndarray、JAX 数组)进行算术运算时,JAX 会隐式地用jnp.array(x.shape[0])把它转成 JAX 数组。文档示例中,5. + x.shape[0]、x.shape[0] - np.arange(5, dtype=jnp.int32)、x + x.shape[0] + jnp.sin(x.shape[0])三处的x.shape[0]都被自动转换:
>>> exp = export.export(jax.jit( ... lambda x: (5. + x.shape[0], ... x.shape[0] - np.arange(5, dtype=jnp.int32), ... x + x.shape[0] + jnp.sin(x.shape[0]))))( ... jax.ShapeDtypeStruct(export.symbolic_shape("b"), jnp.int32)) >>> exp.out_avals (ShapedArray(float32[], weak_type=True), ShapedArray(int32[5]), ShapedArray(float32[b], weak_type=True)) >>> exp.call(jnp.ones((3,), jnp.int32)) (Array(8., dtype=float32, weak_type=True), Array([ 3, 2, 1, 0, -1], dtype=int32), Array([4.14112, 4.14112, 4.14112], dtype=float32, weak_type=True))另一个典型场景是求平均:jnp.sum(x, axis=0) / x.shape[0]中x.shape[0]同样被自动转成数组参与除法,得到正确结果Array([4., 5., 6., 7.], dtype=float32)。
该自动转换机制的底层实现可见 jax/_src/export/shape_poly.py:_DimExpr实现了__jax_array__,为多项式到 JAX 数组的隐式强制转换提供了入口,最终通过dim_as_value_p原语(_dim_as_value)在降级阶段用mlir.eval_dynamic_shape计算维度值。
4.3 符号形状下的常见错误
大多数 JAX 代码假定数组形状是整数元组;引入符号维度后,形状检查会照常触发,但报错信息中会出现符号表达式。例如:
>>> v, = export.symbolic_shape("v,") >>> export.export(jax.jit(lambda x, y: x + y))( # doctest: +IGNORE_EXCEPTION_DETAIL ... jax.ShapeDtypeStruct((v,), dtype=np.int32), # doctest: +IGNORE_EXCEPTION_DETAIL ... jax.ShapeDtypeStruct((4,), dtype=np.int32)) # doctest: +IGNORE_EXCEPTION_DETAIL Traceback (most recent call last): TypeError: add got incompatible shapes for broadcasting: (v,), (4,). >>> export.export(jax.jit(lambda x: jnp.matmul(x, x)))( # doctest: +IGNORE_EXCEPTION_DETAIL ... jax.ShapeDtypeStruct((v, 4), dtype=np.int32)) # doctest: +IGNORE_EXCEPTION_DETAIL Traceback (most recent call last): TypeError: dot_general requires contracting dimensions to have the same shape, got (4,) and (v,).修复方式通常很简单:把 matmul 的参数形状规格改为(v, v),让收缩维度一致即可。
5. 符号维度的比较:部分支持与InconclusiveDimensionOperation
JAX 内部存在大量涉及形状的相等/不等比较,用于形状检查甚至选择某些原语的实现。符号维度下的比较规则如下:
- 相等比较(带一个警示):若两个符号维度在所有维度变量取值下都相等,则结果为
True(如b + b == 2*b);否则一律为False。该行为的深远影响见第 5.4 节。 - 不等比较:恒为相等的否定。
- 不等式比较:部分支持,且会利用"维度变量取值于严格正整数"这一事实。例如
b >= 1、b >= 0、2 * a + b >= 3判定为True;而b >= 2、a >= b、a - b >= 0无法判定,抛出异常。
5.1InconclusiveDimensionOperation的典型触发
当比较无法归结为布尔值时,JAX 抛出jax.errors.InconclusiveDimensionOperation(源码定义于 jax/_src/export/shape_poly.py,是core.InconclusiveDimensionOperation的子类):
import jax >>> export.export(jax.jit(lambda x: 0 if x.shape[0] + 1 >= x.shape[1] else 1))( ... jax.ShapeDtypeStruct(export.symbolic_shape("a, b"), dtype=np.int32)) # doctest: +IGNORE_EXCEPTION_DETAIL Traceback (most recent call last): jax._src.export.shape_poly.InconclusiveDimensionOperation: Symbolic dimension comparison 'a + 1' >= 'b' is inconclusive. This error arises for comparison operations with shapes that are non-constant, and the result of the operation cannot be represented as a boolean value for all values of the symbolic dimensions involved.另一个文档中的例子:JAX 需要证明切片大小mod(b, 3)不超过轴大小b(对所有严格正整数b均成立,但 JAX 的符号比较规则证不出来),于是lax.slice_in_dim(x, 0, x.shape[0] % 3)报错。其解决方案见第 5.3 节。
5.2 应对策略一览
文档给出四条可操作的策略:
- 用
core.max_dim/core.min_dim替代内建max/min(含np.max/np.min):把不等式比较推迟到编译期——此时形状已具体化,比较自然可解。 - 重写条件表达式:例如把
d if d > 0 else 0改写为core.max_dim(d, 0)。 - 降低对"维度是整数"的依赖:符号维度对大多数算术运算是"鸭子类型"的整数,例如把
int(d) + 5写成d + 5。 - 指定符号约束(见下一节)。
5.3 用户自定义符号约束:隐式与显式
默认情况下,JAX 假定所有维度变量取值 >= 1,并从中推导简单不等式,例如a + 2 >= 3、a * 2 >= 1、a + b + c >= 3、a // 4 >= 0、a**2 >= 1等。
隐式约束:通过改变规格本身给维度"加码"。例如用2*b作为维度规格,即约束该维度为偶数且 >= 2;用b + 15约束该维度至少为 16。文档示例:若不写+ 15,JAX 无法证明切片大小 16 不超过轴大小b,导出会失败:
>>> _ = export.export(jax.jit(lambda x: x[0:16]))( ... jax.ShapeDtypeStruct(export.symbolic_shape("b + 15"), dtype=np.int32))显式约束:通过symbolic_shape的constraints参数指定,支持>=、<=、==,并与隐式约束构成合取:
>>> # 引入带约束的维度变量。 >>> a, b = export.symbolic_shape("a, b", ... constraints=("a >= b", "b >= 16")) >>> _ = export.export(jax.jit(lambda x: x[:x.shape[1], :16]))( ... jax.ShapeDtypeStruct((a, b), dtype=np.int32))JAX 目前对符号约束的推理能力有限(源码中的约束类_SymbolicConstraint与归一化规则见 jax/_src/export/shape_poly.py):
- 形式为"变量与常量比较"(
>=/<=)的约束收益最大:由a >= 16、b >= 8可推出a + 2*b >= 32; - 复杂表达式约束能力有限:由
a >= b + 8能推出a - b >= 8,但推不出a >= 9(该领域未来可能改进); - 相等约束被当作重写规则:遇到
==左侧的符号表达式时,改写为右侧表达式。例如floordiv(a, b) == c会把所有floordiv(a, b)替换为c。注意:相等约束的左侧顶层不能是加法或减法,合法示例包括a * b、4 * a、floordiv(a + c, b)。
>>> # 引入带相等约束的维度变量。 >>> a, b, c, d = export.symbolic_shape("a, b, c, d", ... constraints=("a * b == c + d",)) >>> 2 * b * a 2*d + 2*c >>> a * b * b b*d + b*c回到 5.1 的mod例子:要么把轴大小规格改为3*b(此时mod(3*b, 3)可化简为0),要么把 JAX 试图证明的那个不等式原样写成显式约束:
>>> b, = export.symbolic_shape("b", ... constraints=["b >= mod(b, 3)"]) >>> f = lambda x: lax.slice_in_dim(x, 0, x.shape[0] % 3) >>> _ = export.export(jax.jit(f))( ... jax.ShapeDtypeStruct((b,), dtype=np.int32))与隐式约束一样,显式约束也会在编译期被检查,机制见第 7 节"Shape assertion errors"。
5.4 相等比较的警示:刻意为之的不完备语义
相等比较对b + 1 == b、b == 0返回False(确定不同),但对b == 1、a == b也返回False——这在不完备的意义上是不健全(unsound)的:某些估值下真、某些估值下假,按理应抛InconclusiveDimensionOperation。JAX 之所以选择让相等**全函数化(total)**并容忍这种不完备,是为了避免哈希碰撞场景下的误报——维度表达式(及其包含对象:形状、core.AbstractValue、core.Jaxpr)参与哈希,若相等语义部分化,会在b == a or b == b、b in [a, b]这类表达式上产生与书写顺序相关的偶发错误。
实践建议(文档原话):if x.shape[0] != 1: raise NiceErrorMessage这种"先断言后报错"的写法是健全的;而if x.shape[0] != 1: return 1这种依赖比较结果的写法则是不健全的。
5.5 符号维度作用域(SymbolicScope)
符号约束存储于jax.export.SymbolicScope对象中,每次调用symbolic_shape都会隐式创建一个新作用域。来自不同作用域的符号表达式严禁混用,否则报错:
>>> a1, = export.symbolic_shape("a,") >>> a2, = export.symbolic_shape("a,", constraints=("a >= 8",)) >>> a1 + a2 # doctest: +IGNORE_EXCEPTION_DETAIL Traceback (most recent call last): ValueError: Invalid mixing of symbolic scopes for linear combination. Expected scope 4776451856 created at <doctest shape_poly.md[31]>:1:6 (<module>) and found for 'a' (unknown) scope 4776979920 created at <doctest shape_poly.md[32]>:1:6 (<module>) with constraints: a >= 8同一作用域内的表达式(含算术结果)共享作用域、可自由混合;JAX 的追踪缓存以形状为部分键,打印相同的符号形状若来自不同作用域也会被视为不同。
复用作用域的两种方式:通过scope=a.scope复用已有维度变量的作用域(此时不能再附加新约束),或显式创建SymbolicScope:
>>> a, = export.symbolic_shape("a,", constraints=("a >= 8",)) >>> b, = export.symbolic_shape("b,", scope=a.scope) # 复用 a 的作用域 >>> a + b # 允许 b + a >>> my_scope = export.SymbolicScope() >>> c, = export.symbolic_shape("c", scope=my_scope) >>> d, = export.symbolic_shape("d", scope=my_scope) >>> c + d # 允许 d + c5.6 用symbolic_dim_bounds检查可证明的边界
jax.export.symbolic_dim_bounds(dimension)返回 JAX 能为某符号维度(或派生表达式)证明的包含式(inclusive)上下界。文档强调:边界是保守的、可能不紧的;无限边界只表示 JAX 未能建立有限界,并不证明维度数学上无界。
>>> batch, free = export.symbolic_shape( ... "batch, free", constraints=("batch >= 128", "batch <= 1024")) >>> export.symbolic_dim_bounds(batch) (128, 1024) >>> export.symbolic_dim_bounds(2 * batch + 1) (257, 2049) >>> export.symbolic_dim_bounds(free) (1, inf)对应的测试用例见 tests/shape_poly_test.py:symbolic_dim_bounds(np.int32(7)) == (7, 7)、symbolic_dim_bounds(m * n + 1) == (7, 81)、symbolic_dim_bounds(1.5)抛TypeError,而对无法证明有定义的表达式(如a // (b - 1),除数可能为 0)则传播InconclusiveDimensionOperation。实现位于 jax/_src/export/shape_poly.py,内部调用core.concrete_dim_or_error与决策过程_bounds_decision。
6. 维度变量必须能从输入形状解出
目前向已导出的对象传递维度变量值的唯一途径,是经由数组参数的形状间接推导。例如b的值可在调用点从第一个参数的类型f32[b]中读出。这镜像了 JIT 函数的调用约定,适用于绝大多数场景。
但若想导出一个由整数参数决定形状的函数,就会撞上限制。看文档中的my_top_k例子:k决定输出形状,却不出现在输入x: i32[4, 10]的形状里,导出会失败:
>>> def my_top_k(k, x): # x: i32[4, 10], k <= 10 ... return lax.top_k(x, k)[0] # : i32[4, 3] >>> x = np.arange(40, dtype=np.int32).reshape((4, 10)) >>> # 用静态 k=3 导出。由于 k 出现在形状中,必须放进 static_argnums。 >>> exp_static_k = export.export(jax.jit(my_top_k, static_argnums=0))(3, x) >>> exp_static_k.in_avals[0] ShapedArray(int32[4,10]) >>> exp_static_k.out_avals[0] ShapedArray(int32[4,3]) >>> # 调用导出函数时只传非静态参数 >>> exp_static_k.call(x) Array([[ 9, 8, 7], [19, 18, 17], [29, 28, 27], [39, 38, 37]], dtype=int32) >>> # 现在尝试用符号 k 导出,以便导出后再选 k。 >>> k, = export.symbolic_shape("k", constraints=["k <= 10"]) >>> export.export(jax.jit(my_top_k, static_argnums=0))(k, x) # doctest: +IGNORE_EXCEPTION_DETAIL Traceback (most recent call last): UnexpectedDimVar: "Encountered dimension variable 'k' that is not appearing in the shapes of the function arguments文档给出的绕过方案:把函数参数k替换为形状(0, k)的数组,使k可从输入形状推导。首维取 0 保证数组为空、调用时零性能开销:
>>> def my_top_k_with_dimensions(dimensions, x): # dimensions: i32[0, k], x: i32[4, 10] ... return my_top_k(dimensions.shape[1], x) >>> exp = export.export(jax.jit(my_top_k_with_dimensions))( ... jax.ShapeDtypeStruct((0, k), dtype=np.int32), ... x) >>> exp.in_avals (ShapedArray(int32[0,k]), ShapedArray(int32[4,10])) >>> exp.out_avals[0] ShapedArray(int32[4,k]) >>> # 调用 exp 时必须构造并传入形状为 (0, k) 的数组 >>> exp.call(np.zeros((0, 3), dtype=np.int32), x) Array([[ 9, 8, 7], [19, 18, 17], [29, 28, 27], [39, 38, 37]], dtype=int32)另一种报错场景是:维度变量虽然出现在输入形状中,但以 JAX当前无法求解的非线性表达式出现(线性求解逻辑见 jax/_src/export/shape_poly.py 附近的_solve_dim_equations,目前仅支持线性单变量约束):
>>> a, = export.symbolic_shape("a") >>> export.export(jax.jit(lambda x: x.shape[0]))( ... jax.ShapeDtypeStruct((a * a,), dtype=np.int32)) # doctest: +IGNORE_EXCEPTION_DETAIL Traceback (most recent call last): ValueError: Cannot solve for values of dimension variables {'a'}. We can only solve linear uni-variate constraints. Using the following polymorphic shapes specifications: args[0].shape = (a^2,). Unprocessed specifications: 'a^2' for dimension size args[0].shape[0].7. Shape assertion errors:编译期检查维度变量约束
JAX 假定维度变量取严格正整数,并在针对具体输入形状编译时校验这一假定。例如对符号输入形状(b, b, 2*d),用实际参数arg调用时会生成如下断言:
arg.shape[0] >= 1arg.shape[1] == arg.shape[0]arg.shape[2] % 2 == 0arg.shape[2] // 2 >= 1
用形状(3, 3, 5)调用会得到:
>>> def f(x): # x: f32[b, b, 2*d] ... return x >>> exp = export.export(jax.jit(f))( ... jax.ShapeDtypeStruct(export.symbolic_shape("b, b, 2*d"), dtype=np.int32)) >>> exp.call(np.ones((3, 3, 5), dtype=np.int32)) # doctest: +IGNORE_EXCEPTION_DETAIL Traceback (most recent call last): ValueError: Input shapes do not match the polymorphic shapes specification. Division had remainder 1 when computing the value of 'd'. Using the following polymorphic shapes specifications: args[0].shape = (b, b, 2*d). Obtained dimension variables: 'b' = 3 from specification 'b' for dimension args[0].shape[0] (= 3), .这些错误发生在编译前的预处理步骤。其底层实现是形状断言原语shape_assertion_p(源码见 jax/_src/export/shape_poly.py):一个带ShapeAssertionEffect副作用、无返回值的原语,被降级为"shape_assertion"自定义调用(custom_call),并把错误消息以属性形式嵌入,供形状细化(shape refinement)阶段求值后触发报错。
8. 调试:定位形状细化失败
Exported模块在编译期对含维度变量或多平台支持的模块执行形状细化(shape refinement)。调试分两步:
- 先参考导出调试文档docs/export/export.md 中的相关章节;
- 若形状细化阶段报错,可设置
JAX_DUMP_IR_TO环境变量,把形状细化之前的 HLO 模块 dump 出来,文件名为..._before_refine_polymorphic_shapes.mlir,此模块应已具有静态输入形状,便于对比细化前后差异。
若要记录形状细化的所有阶段日志,可设置TF_CPP_VMODULE=refine_polymorphic_shapes=3(OSS 环境;Google 内部用--vmodule=refine_polymorphic_shapes=3)。文档给出的完整命令示例:
# Log from python JAX_DUMP_IR_TO=/tmp/export.dumps/ TF_CPP_VMODULE=refine_polymorphic_shapes=3 python tests/shape_poly_test.py ShapePolyTest.test_simple_unary -v=3该命令同时演示了如何运行官方测试套件中的用例ShapePolyTest.test_simple_unary(完整测试集见 tests/shape_poly_test.py,覆盖解析、求值、边界算术、比较决策与错误传播等维度)。
9. 综合示例与仓库佐证
把本文知识点串成一个可运行的导出流程:
import jax import numpy as np from jax import export from jax import numpy as jnp # 1) 定义符号形状与约束 batch, channels = export.symbolic_shape( "batch, channels", constraints=("batch >= 16", "channels >= 1")) # 2) 导出归一化函数(输出形状含符号表达式) def normalize(x): return (x - jnp.mean(x, axis=0)) / x.shape[0] exp = export.export(jax.jit(normalize))( jax.ShapeDtypeStruct((batch, channels), jnp.float32)) print(exp.in_avals) # ShapedArray(float32[batch,channels]) print(exp.out_avals) # 输出形状含符号表达式 # 3) 用不同具体形状调用,无需重新追踪 for n in (16, 32, 64): r = exp.call(np.ones((n, 3), dtype=np.float32)) assert r.shape == (n, 3)仓库中可交叉验证的素材包括:
- 核心实现:jax/_src/export/shape_poly.py(
symbolic_shape、symbolic_args_specs、symbolic_dim_bounds、SymbolicScope、shape_assertion_p等); - 导出主流程与
Exported.call:jax/_src/export/_export.py; - 公开 API 出口:
jax.export命名空间(SymbolicScope、symbolic_dim_bounds、symbolic_shape、symbolic_args_specs等)与jax.errors.InconclusiveDimensionOperation; - 官方测试:tests/shape_poly_test.py 与 jax/experimental/jax2tf/tests/shape_poly_test.py;
- 性能基准:benchmarks/shape_poly_benchmark.py(覆盖
symbolic_shape解析、算术构造、min/max 操作与约束加载等子场景)。
从源码结构看,符号维度的比较与边界决策通过可替换的决策过程(_bounds_decision、_geq_decision等,见 jax/_src/export/shape_poly.py)实现,并经由shape_poly_decision.py注入具体策略,这为未来增强推理能力预留了扩展点。
结语
Shape polymorphism 把 JAX 的"按形状编译"模型扩展为"按形状族编译",是 JAX export 在多平台、无源码环境下实现可移植推理与部署的关键机制。掌握符号形状规格、维度表达式、显式/隐式约束、作用域管理以及InconclusiveDimensionOperation的处理套路,即可在 batch 化、序列化与模型服务等场景中写出"一次导出、多形状复用"的健壮代码。
【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考